Compare commits
165
Commits
master
..
bcada300e3
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bcada300e3 | ||
|
|
9756174676 | ||
|
|
6dd8461a46 | ||
|
|
12cc2eb0e9 | ||
|
|
38dd717aa6 | ||
|
|
6640c902de | ||
|
|
583c343d08 | ||
|
|
dfc48f7a05 | ||
|
|
cd9d854595 | ||
|
|
766cbd17e5 | ||
|
|
6df95bf981 | ||
|
|
6398ca0893 | ||
|
|
62372a48cc | ||
|
|
733632509a | ||
|
|
e1c13ec314 | ||
|
|
928ff0eabe | ||
|
|
406559b13d | ||
|
|
20c16aa6fd | ||
|
|
44ba5fd6d4 | ||
|
|
745c6adbf2 | ||
|
|
a9aa09636f | ||
|
|
6945986dd1 | ||
|
|
80c1f48f0e | ||
|
|
31e18205f0 | ||
|
|
e84a9d3f9b | ||
|
|
133feb8c76 | ||
|
|
c5fd9c01e5 | ||
|
|
a7bf5ceac3 | ||
|
|
74139aeb7e | ||
|
|
0cae4fd05c | ||
|
|
4b3b4fda61 | ||
|
|
adb684a6bf | ||
|
|
8493472983 | ||
|
|
862eeb7add | ||
|
|
4d9b211d69 | ||
|
|
ebb272324c | ||
|
|
c0290512b3 | ||
|
|
4ec56fe41e | ||
|
|
f8a7c46cf9 | ||
|
|
2bab8a9bb6 | ||
|
|
89e6a6215a | ||
|
|
88be87e03e | ||
|
|
22867faa9c | ||
|
|
16c0fc704d | ||
|
|
32fdd076bf | ||
|
|
402ae0d466 | ||
|
|
acb3c6d68b | ||
|
|
40ac83e632 | ||
|
|
58da395941 | ||
|
|
0e3ef94c9e | ||
|
|
1f68dfc2b5 | ||
|
|
84977a464c | ||
|
|
62ada5eaa4 | ||
|
|
f5ff0b7c13 | ||
|
|
3337cafcdf | ||
|
|
e87784118b | ||
|
|
40fada28ea | ||
|
|
58cc94d4b7 | ||
|
|
ccabea59c9 | ||
|
|
8cc1dc042d | ||
|
|
183c37446e | ||
|
|
7aa06afc45 | ||
|
|
1515a2fb86 | ||
|
|
ec798c58d7 | ||
|
|
e365189276 | ||
|
|
116d610ad0 | ||
|
|
75c570962d | ||
|
|
cae8ac1799 | ||
|
|
917cc222a3 | ||
|
|
7edc588202 | ||
|
|
c83461508b | ||
|
|
651d64f34d | ||
|
|
b98d4b59f5 | ||
|
|
374449e663 | ||
|
|
4c876a201b | ||
|
|
21b3dd1da1 | ||
|
|
5ca0ea9228 | ||
|
|
d5c3a68a37 | ||
|
|
2b33b9158d | ||
|
|
b31642e284 | ||
|
|
3a7a3307ef | ||
|
|
0496cd907b | ||
|
|
975b4fa700 | ||
|
|
060f280fdf | ||
|
|
7aaf189247 | ||
|
|
df6d99c07d | ||
|
|
63306cf017 | ||
|
|
08be5e85e4 | ||
|
|
9843510e1f | ||
|
|
83bda3dfb2 | ||
|
|
0ab15aa227 | ||
|
|
29c2fb8e06 | ||
|
|
4ebc465e8d | ||
|
|
4b132a21e9 | ||
|
|
b644971d45 | ||
|
|
df34533765 | ||
|
|
6f4efb36bb | ||
|
|
108d5b14d7 | ||
|
|
f633b86b35 | ||
|
|
1873e18f8e | ||
|
|
3cdcbb47bf | ||
|
|
aaa9c7987c | ||
|
|
048007a042 | ||
|
|
cb35b40b9d | ||
|
|
3b71fe03b4 | ||
|
|
471db64bcc | ||
|
|
5e2234763c | ||
|
|
1d140be715 | ||
|
|
ffb2a34ae5 | ||
|
|
ccf7de1a55 | ||
|
|
ccf3c80d29 | ||
|
|
3a3c89e0b4 | ||
|
|
17c629136a | ||
|
|
52a5c4141f | ||
|
|
c9ba27c333 | ||
|
|
46f6e2c58b | ||
|
|
d1f47e5a22 | ||
|
|
65de94bad3 | ||
|
|
4e935c6203 | ||
|
|
eac07cf5a8 | ||
|
|
260259d461 | ||
|
|
a9fb092834 | ||
|
|
3a21a68792 | ||
|
|
9d572d18bc | ||
|
|
f7852e8034 | ||
|
|
8c075de147 | ||
|
|
864367f4f5 | ||
|
|
33db2ea7f4 | ||
|
|
8396d09891 | ||
|
|
bf7171924d | ||
|
|
097c363fbc | ||
|
|
7dd8809e38 | ||
|
|
1749757036 | ||
|
|
5857e6121c | ||
|
|
a41147916b | ||
|
|
cabe38db1d | ||
|
|
f079479160 | ||
|
|
87e160a01a | ||
|
|
f94d829bf8 | ||
|
|
9a05bfa0c3 | ||
|
|
c1d46859a3 | ||
|
|
87ecbcb113 | ||
|
|
1fb2949561 | ||
|
|
11be777fc0 | ||
|
|
ff94161fc0 | ||
|
|
c2ab9a950f | ||
|
|
8e26a0f5a8 | ||
|
|
d1c15ee295 | ||
|
|
00a96234c6 | ||
|
|
fc3b663510 | ||
|
|
e4e045d059 | ||
|
|
436feaf33d | ||
|
|
29f450b962 | ||
|
|
db343893c8 | ||
|
|
554906ec02 | ||
|
|
18c37f4842 | ||
|
|
83382b824a | ||
|
|
6203316aa1 | ||
|
|
d57b4d1d5e | ||
|
|
379ae214fc | ||
|
|
163a403636 | ||
|
|
53edaadc3a | ||
|
|
3c2664c3ce | ||
|
|
d9048954a5 | ||
|
|
4f84dfd73f |
Generated
+548
-4
@@ -26,6 +26,16 @@ dependencies = [
|
||||
"pom",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aead"
|
||||
version = "0.5.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0"
|
||||
dependencies = [
|
||||
"crypto-common 0.1.7",
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aes"
|
||||
version = "0.8.4"
|
||||
@@ -37,6 +47,20 @@ dependencies = [
|
||||
"cpufeatures 0.2.17",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aes-gcm"
|
||||
version = "0.10.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1"
|
||||
dependencies = [
|
||||
"aead",
|
||||
"aes",
|
||||
"cipher",
|
||||
"ctr",
|
||||
"ghash",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "agen"
|
||||
version = "0.2.1"
|
||||
@@ -326,6 +350,12 @@ dependencies = [
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "base16ct"
|
||||
version = "0.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4c7f02d4ea65f2c1853089ffd8d2787bdbc63de2f0d29dedbcf8ccdfa0ccd4cf"
|
||||
|
||||
[[package]]
|
||||
name = "base64"
|
||||
version = "0.21.7"
|
||||
@@ -338,6 +368,12 @@ version = "0.22.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6"
|
||||
|
||||
[[package]]
|
||||
name = "base64ct"
|
||||
version = "1.8.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06"
|
||||
|
||||
[[package]]
|
||||
name = "base64urlsafedata"
|
||||
version = "0.5.5"
|
||||
@@ -349,6 +385,17 @@ dependencies = [
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "bcrypt-pbkdf"
|
||||
version = "0.10.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6aeac2e1fe888769f34f05ac343bbef98b14d1ffb292ab69d4608b3abc86f2a2"
|
||||
dependencies = [
|
||||
"blowfish",
|
||||
"pbkdf2",
|
||||
"sha2 0.10.9",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "bit-set"
|
||||
version = "0.5.3"
|
||||
@@ -403,6 +450,16 @@ dependencies = [
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "blowfish"
|
||||
version = "0.9.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e412e2cd0f2b2d93e02543ceae7917b3c70331573df19ee046bcbc35e45e87d7"
|
||||
dependencies = [
|
||||
"byteorder",
|
||||
"cipher",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "bstr"
|
||||
version = "1.12.1"
|
||||
@@ -435,6 +492,12 @@ version = "1.25.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec"
|
||||
|
||||
[[package]]
|
||||
name = "byteorder"
|
||||
version = "1.5.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b"
|
||||
|
||||
[[package]]
|
||||
name = "bytes"
|
||||
version = "1.11.1"
|
||||
@@ -495,6 +558,17 @@ version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
|
||||
|
||||
[[package]]
|
||||
name = "chacha20"
|
||||
version = "0.9.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cipher",
|
||||
"cpufeatures 0.2.17",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "chrono"
|
||||
version = "0.4.44"
|
||||
@@ -563,8 +637,8 @@ checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9"
|
||||
name = "client"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"chrono",
|
||||
"futures",
|
||||
"manifest",
|
||||
"protocol",
|
||||
"reqwest",
|
||||
"serde",
|
||||
@@ -654,6 +728,12 @@ dependencies = [
|
||||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "const-oid"
|
||||
version = "0.9.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8"
|
||||
|
||||
[[package]]
|
||||
name = "const-oid"
|
||||
version = "0.10.2"
|
||||
@@ -937,6 +1017,18 @@ version = "0.2.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5"
|
||||
|
||||
[[package]]
|
||||
name = "crypto-bigint"
|
||||
version = "0.5.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0dc92fb57ca44df6db8059111ab3af99a63d5d0f8375d9972e319a379c6bab76"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
"rand_core 0.6.4",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "crypto-common"
|
||||
version = "0.1.7"
|
||||
@@ -966,6 +1058,41 @@ dependencies = [
|
||||
"phf 0.11.3",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ctr"
|
||||
version = "0.9.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835"
|
||||
dependencies = [
|
||||
"cipher",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "curve25519-dalek"
|
||||
version = "4.1.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.2.17",
|
||||
"curve25519-dalek-derive",
|
||||
"digest 0.10.7",
|
||||
"fiat-crypto",
|
||||
"rustc_version",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "curve25519-dalek-derive"
|
||||
version = "0.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f46882e17999c6cc590af592290432be3bce0428cb0d5f8b6715e4dc7b383eb3"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "darling"
|
||||
version = "0.23.0"
|
||||
@@ -1056,6 +1183,16 @@ version = "0.3.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5729f5117e208430e437df2f4843f5e5952997175992d1414f94c57d61e270b4"
|
||||
|
||||
[[package]]
|
||||
name = "der"
|
||||
version = "0.7.10"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb"
|
||||
dependencies = [
|
||||
"const-oid 0.9.6",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "der-parser"
|
||||
version = "9.0.0"
|
||||
@@ -1114,7 +1251,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292"
|
||||
dependencies = [
|
||||
"block-buffer 0.10.4",
|
||||
"const-oid 0.9.6",
|
||||
"crypto-common 0.1.7",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -1124,7 +1263,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4850db49bf08e663084f7fb5c87d202ef91a3907271aff24a94eb97ff039153c"
|
||||
dependencies = [
|
||||
"block-buffer 0.12.0",
|
||||
"const-oid",
|
||||
"const-oid 0.10.2",
|
||||
"crypto-common 0.2.1",
|
||||
]
|
||||
|
||||
@@ -1175,12 +1314,66 @@ dependencies = [
|
||||
"cipher",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ecdsa"
|
||||
version = "0.16.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ee27f32b5c5292967d2d4a9d7f1e0b0aed2c15daded5a60300e4abb9d8020bca"
|
||||
dependencies = [
|
||||
"der",
|
||||
"digest 0.10.7",
|
||||
"elliptic-curve",
|
||||
"rfc6979",
|
||||
"signature",
|
||||
"spki",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ed25519"
|
||||
version = "2.2.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "115531babc129696a58c64a4fef0a8bf9e9698629fb97e9e40767d235cfbcd53"
|
||||
dependencies = [
|
||||
"signature",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ed25519-dalek"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "70e796c081cee67dc755e1a36a0a172b897fab85fc3f6bc48307991f64e4eca9"
|
||||
dependencies = [
|
||||
"curve25519-dalek",
|
||||
"ed25519",
|
||||
"sha2 0.10.9",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "either"
|
||||
version = "1.15.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719"
|
||||
|
||||
[[package]]
|
||||
name = "elliptic-curve"
|
||||
version = "0.13.8"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b5e6043086bf7973472e0c7dff2142ea0b680d30e18d9cc40f267efbf222bd47"
|
||||
dependencies = [
|
||||
"base16ct",
|
||||
"crypto-bigint",
|
||||
"digest 0.10.7",
|
||||
"ff",
|
||||
"generic-array",
|
||||
"group",
|
||||
"pkcs8",
|
||||
"rand_core 0.6.4",
|
||||
"sec1",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "embedded-io"
|
||||
version = "0.4.0"
|
||||
@@ -1284,6 +1477,22 @@ version = "2.3.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be"
|
||||
|
||||
[[package]]
|
||||
name = "ff"
|
||||
version = "0.13.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c0b50bfb653653f9ca9095b427bed08ab8d75a137839d9ad64eb11810d5b6393"
|
||||
dependencies = [
|
||||
"rand_core 0.6.4",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "fiat-crypto"
|
||||
version = "0.2.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d"
|
||||
|
||||
[[package]]
|
||||
name = "filedescriptor"
|
||||
version = "0.8.3"
|
||||
@@ -1526,6 +1735,7 @@ checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a"
|
||||
dependencies = [
|
||||
"typenum",
|
||||
"version_check",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -1568,6 +1778,16 @@ dependencies = [
|
||||
"wasip3",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ghash"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1"
|
||||
dependencies = [
|
||||
"opaque-debug",
|
||||
"polyval",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "gimli"
|
||||
version = "0.33.0"
|
||||
@@ -1636,6 +1856,17 @@ dependencies = [
|
||||
"memmap2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "group"
|
||||
version = "0.13.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f0f9ef7462f7c099f518d754361858f86d8a07af53ba9af0fe635bbccb151a63"
|
||||
dependencies = [
|
||||
"ff",
|
||||
"rand_core 0.6.4",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "h2"
|
||||
version = "0.4.13"
|
||||
@@ -1724,6 +1955,15 @@ version = "0.4.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70"
|
||||
|
||||
[[package]]
|
||||
name = "hmac"
|
||||
version = "0.12.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e"
|
||||
dependencies = [
|
||||
"digest 0.10.7",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "html5ever"
|
||||
version = "0.26.0"
|
||||
@@ -2212,6 +2452,9 @@ name = "lazy_static"
|
||||
version = "1.5.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe"
|
||||
dependencies = [
|
||||
"spin",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "leb128fmt"
|
||||
@@ -2386,6 +2629,7 @@ version = "0.1.0"
|
||||
dependencies = [
|
||||
"agen",
|
||||
"arc-swap",
|
||||
"decodal",
|
||||
"protocol",
|
||||
"secrets",
|
||||
"serde",
|
||||
@@ -2672,6 +2916,22 @@ dependencies = [
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-bigint-dig"
|
||||
version = "0.8.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e661dda6640fad38e827a6d4a310ff4763082116fe217f279885c97f511bb0b7"
|
||||
dependencies = [
|
||||
"lazy_static",
|
||||
"libm",
|
||||
"num-integer",
|
||||
"num-iter",
|
||||
"num-traits",
|
||||
"rand 0.8.5",
|
||||
"smallvec",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-conv"
|
||||
version = "0.2.1"
|
||||
@@ -2698,6 +2958,16 @@ dependencies = [
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-iter"
|
||||
version = "0.1.46"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b"
|
||||
dependencies = [
|
||||
"num-integer",
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "num-traits"
|
||||
version = "0.2.19"
|
||||
@@ -2705,6 +2975,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841"
|
||||
dependencies = [
|
||||
"autocfg",
|
||||
"libm",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -2765,6 +3036,12 @@ version = "1.70.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
|
||||
|
||||
[[package]]
|
||||
name = "opaque-debug"
|
||||
version = "0.3.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381"
|
||||
|
||||
[[package]]
|
||||
name = "openssl"
|
||||
version = "0.10.76"
|
||||
@@ -2818,6 +3095,44 @@ dependencies = [
|
||||
"num-traits",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "p256"
|
||||
version = "0.13.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c9863ad85fa8f4460f9c48cb909d38a0d689dba1f6f6988a5e3e0d31071bcd4b"
|
||||
dependencies = [
|
||||
"ecdsa",
|
||||
"elliptic-curve",
|
||||
"primeorder",
|
||||
"sha2 0.10.9",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "p384"
|
||||
version = "0.13.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fe42f1670a52a47d448f14b6a5c61dd78fce51856e68edaa38f7ae3a46b8d6b6"
|
||||
dependencies = [
|
||||
"ecdsa",
|
||||
"elliptic-curve",
|
||||
"primeorder",
|
||||
"sha2 0.10.9",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "p521"
|
||||
version = "0.13.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0fc9e2161f1f215afdfce23677034ae137bbd45016a880c2eb3ba8eb95f085b2"
|
||||
dependencies = [
|
||||
"base16ct",
|
||||
"ecdsa",
|
||||
"elliptic-curve",
|
||||
"primeorder",
|
||||
"rand_core 0.6.4",
|
||||
"sha2 0.10.9",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "parking_lot"
|
||||
version = "0.12.5"
|
||||
@@ -2847,6 +3162,15 @@ version = "0.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "35fb2e5f958ec131621fdd531e9fc186ed768cbe395337403ae56c17a74c68ec"
|
||||
|
||||
[[package]]
|
||||
name = "pbkdf2"
|
||||
version = "0.12.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f8ed6a7761f76e3b9f92dfb0a60a6a6477c61024b775147ff0973a02653abaf2"
|
||||
dependencies = [
|
||||
"digest 0.10.7",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pdf-extract"
|
||||
version = "0.10.0"
|
||||
@@ -2864,6 +3188,15 @@ dependencies = [
|
||||
"unicode-normalization",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pem-rfc7468"
|
||||
version = "0.7.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "88b39c9bfcfc231068454382784bb460aae594343fb030d46e9f50a645418412"
|
||||
dependencies = [
|
||||
"base64ct",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "percent-encoding"
|
||||
version = "2.3.2"
|
||||
@@ -3009,6 +3342,27 @@ version = "0.2.17"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd"
|
||||
|
||||
[[package]]
|
||||
name = "pkcs1"
|
||||
version = "0.7.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c8ffb9f10fa047879315e6625af03c164b16962a5368d724ed16323b68ace47f"
|
||||
dependencies = [
|
||||
"der",
|
||||
"pkcs8",
|
||||
"spki",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pkcs8"
|
||||
version = "0.10.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7"
|
||||
dependencies = [
|
||||
"der",
|
||||
"spki",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pkg-config"
|
||||
version = "0.3.32"
|
||||
@@ -3021,6 +3375,29 @@ version = "0.2.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b4596b6d070b27117e987119b4dac604f3c58cfb0b191112e24771b2faeac1a6"
|
||||
|
||||
[[package]]
|
||||
name = "poly1305"
|
||||
version = "0.8.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf"
|
||||
dependencies = [
|
||||
"cpufeatures 0.2.17",
|
||||
"opaque-debug",
|
||||
"universal-hash",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "polyval"
|
||||
version = "0.6.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.2.17",
|
||||
"opaque-debug",
|
||||
"universal-hash",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pom"
|
||||
version = "1.1.0"
|
||||
@@ -3101,6 +3478,15 @@ dependencies = [
|
||||
"syn 2.0.117",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "primeorder"
|
||||
version = "0.13.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "353e1ca18966c16d9deb1c69278edbc5f194139612772bd9537af60ac231e1e6"
|
||||
dependencies = [
|
||||
"elliptic-curve",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "proc-macro2"
|
||||
version = "1.0.106"
|
||||
@@ -3525,6 +3911,16 @@ dependencies = [
|
||||
"web-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rfc6979"
|
||||
version = "0.4.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f8dd2a808d456c4a54e300a23e9f5a67e122c3024119acbfd73e3bf664491cb2"
|
||||
dependencies = [
|
||||
"hmac",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ring"
|
||||
version = "0.17.14"
|
||||
@@ -3539,6 +3935,27 @@ dependencies = [
|
||||
"windows-sys 0.52.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rsa"
|
||||
version = "0.9.10"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b8573f03f5883dcaebdfcf4725caa1ecb9c15b2ef50c43a07b816e06799bb12d"
|
||||
dependencies = [
|
||||
"const-oid 0.9.6",
|
||||
"digest 0.10.7",
|
||||
"num-bigint-dig",
|
||||
"num-integer",
|
||||
"num-traits",
|
||||
"pkcs1",
|
||||
"pkcs8",
|
||||
"rand_core 0.6.4",
|
||||
"sha2 0.10.9",
|
||||
"signature",
|
||||
"spki",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rusqlite"
|
||||
version = "0.37.0"
|
||||
@@ -3745,6 +4162,20 @@ version = "1.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49"
|
||||
|
||||
[[package]]
|
||||
name = "sec1"
|
||||
version = "0.7.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d3e97a565f76233a6003f9f5c54be1d9c5bdfa3eccfb189469f11ec4901c47dc"
|
||||
dependencies = [
|
||||
"base16ct",
|
||||
"der",
|
||||
"generic-array",
|
||||
"pkcs8",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "secrets"
|
||||
version = "0.1.0"
|
||||
@@ -4060,6 +4491,16 @@ dependencies = [
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "signature"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de"
|
||||
dependencies = [
|
||||
"digest 0.10.7",
|
||||
"rand_core 0.6.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "simd-adler32"
|
||||
version = "0.3.9"
|
||||
@@ -4103,12 +4544,98 @@ dependencies = [
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "spin"
|
||||
version = "0.9.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3763264f6b73151db08c50ff20d7d8a0b8796e021cdea7ceedad07b80155fa0e"
|
||||
|
||||
[[package]]
|
||||
name = "spki"
|
||||
version = "0.7.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d"
|
||||
dependencies = [
|
||||
"base64ct",
|
||||
"der",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ssh-cipher"
|
||||
version = "0.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "caac132742f0d33c3af65bfcde7f6aa8f62f0e991d80db99149eb9d44708784f"
|
||||
dependencies = [
|
||||
"aes",
|
||||
"aes-gcm",
|
||||
"cbc",
|
||||
"chacha20",
|
||||
"cipher",
|
||||
"ctr",
|
||||
"poly1305",
|
||||
"ssh-encoding",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ssh-encoding"
|
||||
version = "0.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "eb9242b9ef4108a78e8cd1a2c98e193ef372437f8c22be363075233321dd4a15"
|
||||
dependencies = [
|
||||
"base64ct",
|
||||
"pem-rfc7468",
|
||||
"sha2 0.10.9",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ssh-key"
|
||||
version = "0.6.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3b86f5297f0f04d08cabaa0f6bff7cb6aec4d9c3b49d87990d63da9d9156a8c3"
|
||||
dependencies = [
|
||||
"bcrypt-pbkdf",
|
||||
"ed25519-dalek",
|
||||
"p256",
|
||||
"p384",
|
||||
"p521",
|
||||
"rand_core 0.6.4",
|
||||
"rsa",
|
||||
"sec1",
|
||||
"sha2 0.10.9",
|
||||
"signature",
|
||||
"ssh-cipher",
|
||||
"ssh-encoding",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "stable_deref_trait"
|
||||
version = "1.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596"
|
||||
|
||||
[[package]]
|
||||
name = "standalone"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"agen",
|
||||
"async-trait",
|
||||
"fs4",
|
||||
"futures",
|
||||
"manifest",
|
||||
"protocol",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"session-store",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"uuid",
|
||||
"worker",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "static_assertions"
|
||||
version = "1.1.0"
|
||||
@@ -4798,7 +5325,6 @@ dependencies = [
|
||||
"base64 0.22.1",
|
||||
"client",
|
||||
"crossterm 0.28.1",
|
||||
"fs4",
|
||||
"manifest",
|
||||
"protocol",
|
||||
"pulldown-cmark",
|
||||
@@ -4807,13 +5333,14 @@ dependencies = [
|
||||
"serde",
|
||||
"serde_json",
|
||||
"session-store",
|
||||
"standalone",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"ticket",
|
||||
"tokio",
|
||||
"toml",
|
||||
"unicode-width",
|
||||
"uuid",
|
||||
"worker",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -4934,6 +5461,16 @@ version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853"
|
||||
|
||||
[[package]]
|
||||
name = "universal-hash"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea"
|
||||
dependencies = [
|
||||
"crypto-common 0.1.7",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "unsafe-libyaml"
|
||||
version = "0.2.11"
|
||||
@@ -6101,6 +6638,7 @@ dependencies = [
|
||||
"wasmtime",
|
||||
"wat",
|
||||
"workdir",
|
||||
"workspace-api",
|
||||
"yoi-plugin-pdk",
|
||||
]
|
||||
|
||||
@@ -6131,9 +6669,12 @@ dependencies = [
|
||||
"tokio-tungstenite 0.29.0",
|
||||
"toml",
|
||||
"tower",
|
||||
"url",
|
||||
"uuid",
|
||||
"workdir",
|
||||
"worker",
|
||||
"workspace-api",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -6257,11 +6798,13 @@ dependencies = [
|
||||
"project-record",
|
||||
"protocol",
|
||||
"reqwest",
|
||||
"ring",
|
||||
"rusqlite",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serde_yaml",
|
||||
"sha2 0.11.0",
|
||||
"ssh-key",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"ticket",
|
||||
@@ -6278,6 +6821,7 @@ dependencies = [
|
||||
"worker",
|
||||
"worker-runtime",
|
||||
"workspace-api",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
@@ -5,6 +5,7 @@ members = [
|
||||
"crates/agen",
|
||||
"crates/agen-macros",
|
||||
"crates/session-store",
|
||||
"crates/standalone",
|
||||
"crates/secrets",
|
||||
"crates/manifest",
|
||||
"crates/mcp",
|
||||
@@ -36,6 +37,7 @@ default-members = [
|
||||
"crates/agen",
|
||||
"crates/agen-macros",
|
||||
"crates/session-store",
|
||||
"crates/standalone",
|
||||
"crates/secrets",
|
||||
"crates/manifest",
|
||||
"crates/mcp",
|
||||
@@ -87,6 +89,7 @@ protocol = { path = "crates/protocol" }
|
||||
session-metrics = { path = "crates/session-metrics" }
|
||||
session-analytics = { path = "crates/session-analytics" }
|
||||
session-store = { path = "crates/session-store" }
|
||||
standalone = { path = "crates/standalone" }
|
||||
secrets = { path = "crates/secrets" }
|
||||
tools = { path = "crates/tools" }
|
||||
config-source = { path = "crates/config-source" }
|
||||
@@ -115,6 +118,7 @@ tar = "0.4"
|
||||
rusqlite = { version = "0.37", features = ["backup", "bundled"] }
|
||||
ring = "0.17.14"
|
||||
sha2 = "0.11"
|
||||
ssh-key = { version = "0.6.7", features = ["ed25519", "encryption"] }
|
||||
tempfile = "3.27"
|
||||
thiserror = "2.0"
|
||||
tokio = "1.52"
|
||||
@@ -124,4 +128,5 @@ toml = "1.1"
|
||||
tracing = "0.1"
|
||||
url = "2.5"
|
||||
uuid = "1.23"
|
||||
zeroize = "1"
|
||||
webauthn-rs = { version = "0.5.2", features = ["danger-allow-state-serialisation", "danger-credential-internals"] }
|
||||
|
||||
+1
-1
@@ -21,7 +21,7 @@ services:
|
||||
- "8787"
|
||||
volumes:
|
||||
- server-data:/server-data
|
||||
- ./docker/workspace:/workspace:ro
|
||||
- /etc/yoi/server.toml:/server-config/server.toml:ro
|
||||
|
||||
webui:
|
||||
image: yoi-webui:latest
|
||||
|
||||
@@ -21,20 +21,21 @@ agen = { version = "0.2.1", features = ["codex"] }
|
||||
|
||||
## Quick start
|
||||
|
||||
Supply an implementation of [`LlmClient`](https://docs.rs/agen/latest/agen/llm_client/trait.LlmClient.html), then run a turn. The first call consumes the mutable engine and returns a cache-locked engine for later turns.
|
||||
Supply an implementation of [`LlmClient`](https://docs.rs/agen/latest/agen/llm_client/trait.LlmClient.html), keep conversation history in your application, then run a turn. The first call consumes the mutable engine and returns a cache-locked engine for later turns.
|
||||
|
||||
```no_run
|
||||
use agen::{Engine, EngineError};
|
||||
use agen::{Engine, EngineError, History};
|
||||
use agen::llm_client::LlmClient;
|
||||
|
||||
async fn conversation<C: LlmClient>(client: C) -> Result<(), EngineError> {
|
||||
let mut history = History::new();
|
||||
let output = Engine::new(client)
|
||||
.system_prompt("You are a concise assistant.")
|
||||
.run("Explain typed state in one sentence.")
|
||||
.await?;
|
||||
.run(&mut history, "Explain typed state in one sentence.")
|
||||
.await;
|
||||
|
||||
let mut engine = output.engine;
|
||||
let _result = engine.run("Give a Rust example.").await?;
|
||||
let _result = engine.run(&mut history, "Give a Rust example.").await;
|
||||
Ok(())
|
||||
}
|
||||
```
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
|
||||
use agen::llm_client::scheme::{Scheme, anthropic::AnthropicScheme};
|
||||
use agen::llm_client::transport::{HttpTransport, ResolvedAuth};
|
||||
use agen::{Engine, EngineResult};
|
||||
use agen::{Engine, EngineRunExit, StopReason};
|
||||
use std::time::Duration;
|
||||
|
||||
#[tokio::main]
|
||||
@@ -29,6 +29,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
let base_url = scheme.default_base_url().to_string();
|
||||
let client = HttpTransport::new(scheme, model, base_url, ResolvedAuth::ApiKey(api_key), cap);
|
||||
let engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
println!("🚀 Starting Engine...");
|
||||
println!("💡 Will cancel after 2 seconds\n");
|
||||
@@ -45,16 +46,15 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
|
||||
println!("📡 Sending request to LLM...");
|
||||
|
||||
match engine.run("Tell me a very long story about a brave knight. Make it as detailed as possible with many paragraphs.").await {
|
||||
Ok(out) => match out.result {
|
||||
EngineResult::Finished => println!("✅ Task completed normally"),
|
||||
EngineResult::Paused => println!("⏸️ Task paused"),
|
||||
EngineResult::LimitReached => println!("🔒 Turn limit reached"),
|
||||
EngineResult::Yielded => println!("↩️ Task yielded"),
|
||||
},
|
||||
Err(e) => {
|
||||
println!("❌ Task error: {}", e);
|
||||
let output = engine.run(&mut history, "Tell me a very long story about a brave knight. Make it as detailed as possible with many paragraphs.").await;
|
||||
match output.result {
|
||||
EngineRunExit::Finished => println!("✅ Task completed normally"),
|
||||
EngineRunExit::Paused => println!("⏸️ Task paused"),
|
||||
EngineRunExit::Yielded => println!("↩️ Task yielded"),
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached) => {
|
||||
println!("🔒 Turn limit reached")
|
||||
}
|
||||
EngineRunExit::Interrupted(reason) => println!("❌ Task interrupted: {reason:?}"),
|
||||
}
|
||||
|
||||
println!("\n✨ Demo complete!");
|
||||
|
||||
@@ -39,7 +39,7 @@ use tracing::info;
|
||||
use tracing_subscriber::EnvFilter;
|
||||
|
||||
use agen::{
|
||||
Engine,
|
||||
Engine, EngineRunExit, StopReason,
|
||||
interceptor::{Interceptor, PostToolAction, ToolResultInfo},
|
||||
llm_client::{
|
||||
LlmClient,
|
||||
@@ -451,6 +451,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
|
||||
// Create Engine
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
let tool_call_names = Arc::new(Mutex::new(HashMap::new()));
|
||||
|
||||
@@ -476,12 +477,9 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
|
||||
// One-shot mode
|
||||
if let Some(prompt) = args.prompt {
|
||||
match engine.run(&prompt).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
eprintln!("\n❌ Error: {}", e);
|
||||
std::process::exit(1);
|
||||
}
|
||||
let output = engine.run(&mut history, &prompt).await;
|
||||
if let EngineRunExit::Interrupted(StopReason::Unexpected(error)) = output.result {
|
||||
eprintln!("\n❌ Error: {error}");
|
||||
}
|
||||
|
||||
return Ok(());
|
||||
@@ -500,13 +498,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let mut locked = match engine.run(first_input).await {
|
||||
Ok(out) => out.engine,
|
||||
Err(e) => {
|
||||
eprintln!("\n❌ Error: {}", e);
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
let output = engine.run(&mut history, first_input).await;
|
||||
let mut locked = output.engine;
|
||||
|
||||
loop {
|
||||
print!("\n👤 You: ");
|
||||
@@ -525,11 +518,10 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
break;
|
||||
}
|
||||
|
||||
match locked.run(input).await {
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
eprintln!("\n❌ Error: {}", e);
|
||||
}
|
||||
if let EngineRunExit::Interrupted(StopReason::Unexpected(error)) =
|
||||
locked.run(&mut history, input).await
|
||||
{
|
||||
eprintln!("\n❌ Error: {error}");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+886
-283
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,199 @@
|
||||
//! Typed conversation history containers.
|
||||
//!
|
||||
//! Agen keeps provider-visible [`Item`](crate::Item) values separate from any
|
||||
//! host-domain provenance. The host chooses the annotation type `A`, while Agen
|
||||
//! preserves each item and annotation as one entry for clone/truncate/restore
|
||||
//! style history operations.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::Item;
|
||||
|
||||
/// One conversation-history entry with host-owned annotation.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
pub struct HistoryEntry<A = ()> {
|
||||
/// Provider/model-visible conversation item.
|
||||
pub item: Item,
|
||||
/// Host-domain metadata kept with the item and never projected to providers.
|
||||
pub annotation: A,
|
||||
}
|
||||
|
||||
impl<A> HistoryEntry<A> {
|
||||
/// Build an entry from an item and its annotation.
|
||||
pub fn new(item: Item, annotation: A) -> Self {
|
||||
Self { item, annotation }
|
||||
}
|
||||
|
||||
/// Split the entry into its item and annotation.
|
||||
pub fn into_parts(self) -> (Item, A) {
|
||||
(self.item, self.annotation)
|
||||
}
|
||||
}
|
||||
|
||||
impl HistoryEntry<()> {
|
||||
/// Build a unit-annotated entry.
|
||||
pub fn from_item(item: Item) -> Self {
|
||||
Self {
|
||||
item,
|
||||
annotation: (),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Conversation history with one annotation per item.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
|
||||
pub struct History<A = ()> {
|
||||
entries: Vec<HistoryEntry<A>>,
|
||||
}
|
||||
|
||||
impl<A> History<A> {
|
||||
/// Create an empty history.
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
entries: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Build history from already annotated entries, preserving order.
|
||||
pub fn from_entries(entries: Vec<HistoryEntry<A>>) -> Self {
|
||||
Self { entries }
|
||||
}
|
||||
|
||||
/// Replace all entries as one restore/rebuild operation and return the old entries.
|
||||
pub fn replace_entries(&mut self, entries: Vec<HistoryEntry<A>>) -> Vec<HistoryEntry<A>> {
|
||||
std::mem::replace(&mut self.entries, entries)
|
||||
}
|
||||
|
||||
/// Borrow annotated entries.
|
||||
pub fn entries(&self) -> &[HistoryEntry<A>] {
|
||||
&self.entries
|
||||
}
|
||||
|
||||
/// Mutably borrow annotated entries for host-owned rebuild operations.
|
||||
pub fn entries_mut(&mut self) -> &mut [HistoryEntry<A>] {
|
||||
&mut self.entries
|
||||
}
|
||||
|
||||
/// Consume the history into annotated entries.
|
||||
pub fn into_entries(self) -> Vec<HistoryEntry<A>> {
|
||||
self.entries
|
||||
}
|
||||
|
||||
/// Number of entries.
|
||||
pub fn len(&self) -> usize {
|
||||
self.entries.len()
|
||||
}
|
||||
|
||||
/// Whether the history is empty.
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.entries.is_empty()
|
||||
}
|
||||
|
||||
/// Iterate over annotated entries.
|
||||
pub fn iter(&self) -> impl ExactSizeIterator<Item = &HistoryEntry<A>> {
|
||||
self.entries.iter()
|
||||
}
|
||||
|
||||
/// Iterate over provider-visible items only.
|
||||
pub fn items(&self) -> impl ExactSizeIterator<Item = &Item> {
|
||||
self.entries.iter().map(|entry| &entry.item)
|
||||
}
|
||||
|
||||
/// Clone provider-visible items into a request-local projection.
|
||||
pub fn items_cloned(&self) -> Vec<Item> {
|
||||
self.items().cloned().collect()
|
||||
}
|
||||
|
||||
/// Append an already annotated entry.
|
||||
pub fn push_entry(&mut self, entry: HistoryEntry<A>) {
|
||||
self.entries.push(entry);
|
||||
}
|
||||
|
||||
/// Append many already annotated entries.
|
||||
pub fn extend_entries(&mut self, entries: impl IntoIterator<Item = HistoryEntry<A>>) {
|
||||
self.entries.extend(entries);
|
||||
}
|
||||
|
||||
/// Commit one item through a trusted annotation callback before it becomes live.
|
||||
///
|
||||
/// The callback may durably persist the item and returns the annotation that
|
||||
/// must be stored with it. If the callback fails, the history is left unchanged.
|
||||
pub fn append_with(
|
||||
&mut self,
|
||||
item: Item,
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
) -> Result<(), String> {
|
||||
let annotation = annotate(&item)?;
|
||||
self.entries.push(HistoryEntry { item, annotation });
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Commit items through a trusted annotation callback before they become live.
|
||||
///
|
||||
/// Items before a failure remain appended; the failing item and later items do
|
||||
/// not enter history. This mirrors append-only durable logs where each accepted
|
||||
/// item is already committed before the next item is attempted.
|
||||
pub fn extend_with(
|
||||
&mut self,
|
||||
items: impl IntoIterator<Item = Item>,
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
) -> Result<(), String> {
|
||||
for item in items {
|
||||
self.append_with(item, annotate)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Truncate entries, preserving item+annotation pairing for retained entries.
|
||||
pub fn truncate(&mut self, len: usize) {
|
||||
self.entries.truncate(len);
|
||||
}
|
||||
|
||||
/// Clear all entries.
|
||||
pub fn clear(&mut self) {
|
||||
self.entries.clear();
|
||||
}
|
||||
}
|
||||
|
||||
impl History<()> {
|
||||
/// Build unit-annotated history from provider-visible items.
|
||||
pub fn from_items(items: Vec<Item>) -> Self {
|
||||
Self {
|
||||
entries: items.into_iter().map(HistoryEntry::from_item).collect(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Replace history from provider-visible items using unit annotations.
|
||||
pub fn replace_items(&mut self, items: Vec<Item>) -> Vec<HistoryEntry<()>> {
|
||||
self.replace_entries(items.into_iter().map(HistoryEntry::from_item).collect())
|
||||
}
|
||||
|
||||
/// Append one item with unit annotation.
|
||||
pub fn push(&mut self, item: Item) {
|
||||
self.entries.push(HistoryEntry::from_item(item));
|
||||
}
|
||||
|
||||
/// Append items with unit annotations.
|
||||
pub fn extend_items(&mut self, items: impl IntoIterator<Item = Item>) {
|
||||
self.entries
|
||||
.extend(items.into_iter().map(HistoryEntry::from_item));
|
||||
}
|
||||
}
|
||||
|
||||
impl<A> IntoIterator for History<A> {
|
||||
type Item = HistoryEntry<A>;
|
||||
type IntoIter = std::vec::IntoIter<HistoryEntry<A>>;
|
||||
|
||||
fn into_iter(self) -> Self::IntoIter {
|
||||
self.entries.into_iter()
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a, A> IntoIterator for &'a History<A> {
|
||||
type Item = &'a HistoryEntry<A>;
|
||||
type IntoIter = std::slice::Iter<'a, HistoryEntry<A>>;
|
||||
|
||||
fn into_iter(self) -> Self::IntoIter {
|
||||
self.entries.iter()
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
mod engine;
|
||||
mod handler;
|
||||
mod history;
|
||||
mod message;
|
||||
|
||||
pub(crate) mod callback;
|
||||
@@ -20,13 +21,18 @@ pub mod usage_record;
|
||||
pub use agen_macros::{description, tool, tool_registry};
|
||||
pub use callback::{TextBlockScope, ThinkingBlockScope, ToolUseBlockScope};
|
||||
pub use engine::{
|
||||
Engine, EngineConfig, EngineError, EngineResult, EngineRunOutput, LlmRetryNotice,
|
||||
ToolRegistryError,
|
||||
Engine, EngineConfig, EngineError, EngineResult, EngineRunExit, EngineRunOutput,
|
||||
LlmRetryNotice, StopReason, ToolRegistryError,
|
||||
};
|
||||
pub use handler::ToolUseBlockStart;
|
||||
pub use history::{History, HistoryEntry};
|
||||
pub use interceptor::Interceptor;
|
||||
pub use message::{ContentPart, Item, Message, Role};
|
||||
pub use tool::{ToolCall, ToolExecutionContext, ToolOutputLimits, ToolResult};
|
||||
pub use tool::{
|
||||
ToolCall, ToolExecutionContext, ToolExecutionHandle, ToolExecutionPolicy,
|
||||
ToolExecutionTerminal, ToolExecutionTerminalFuture, ToolOutputLimits, ToolResult,
|
||||
ToolResultDisposition,
|
||||
};
|
||||
pub use usage_record::UsageRecord;
|
||||
|
||||
/// Implementation dependencies used by code generated from `agen` macros.
|
||||
|
||||
@@ -18,6 +18,9 @@ pub enum ClientError {
|
||||
message: String,
|
||||
retry_after: Option<Duration>,
|
||||
},
|
||||
/// The provider rejected the request because it exceeded the model context window.
|
||||
/// Classified only from a structured provider error code, never message text.
|
||||
ContextWindowExceeded,
|
||||
/// A request lifecycle phase exceeded its hard timeout.
|
||||
Timeout {
|
||||
phase: &'static str,
|
||||
@@ -48,6 +51,7 @@ impl fmt::Display for ClientError {
|
||||
}
|
||||
write!(f, ": {}", message)
|
||||
}
|
||||
ClientError::ContextWindowExceeded => write!(f, "Model context window reached"),
|
||||
ClientError::Timeout { phase, timeout } => {
|
||||
write!(f, "{phase} timed out after {}s", timeout.as_secs())
|
||||
}
|
||||
@@ -112,7 +116,10 @@ pub fn is_retryable(error: &ClientError) -> bool {
|
||||
ClientError::Api { status: None, .. } => false,
|
||||
ClientError::Timeout { .. } => true,
|
||||
ClientError::Http(e) => e.is_connect() || e.is_timeout(),
|
||||
ClientError::Json(_) | ClientError::Sse(_) | ClientError::Config(_) => false,
|
||||
ClientError::ContextWindowExceeded
|
||||
| ClientError::Json(_)
|
||||
| ClientError::Sse(_)
|
||||
| ClientError::Config(_) => false,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -431,13 +431,7 @@ fn api_error_code(error: &ClientError) -> Option<&str> {
|
||||
}
|
||||
|
||||
fn is_context_length_exceeded(error: &ClientError) -> bool {
|
||||
match error {
|
||||
ClientError::Api { code, message, .. } => {
|
||||
code.as_deref() == Some("context_length_exceeded")
|
||||
|| message.contains("context_length_exceeded")
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
matches!(error, ClientError::ContextWindowExceeded)
|
||||
}
|
||||
|
||||
async fn response_with_timeout(
|
||||
@@ -487,6 +481,9 @@ async fn classify_error_response(resp: reqwest::Response) -> ClientError {
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or(&text)
|
||||
.to_string();
|
||||
if code.as_deref() == Some("context_length_exceeded") {
|
||||
return ClientError::ContextWindowExceeded;
|
||||
}
|
||||
ClientError::Api {
|
||||
status: Some(status),
|
||||
code,
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
use std::{fmt, sync::Arc};
|
||||
|
||||
use crate::tool::Attachment;
|
||||
use crate::tool::{Attachment, ToolResultDisposition};
|
||||
use base64::Engine as _;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
@@ -121,6 +121,9 @@ pub enum Item {
|
||||
/// Detailed output (removed by pruning when old enough)
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
content: Option<String>,
|
||||
/// Typed terminal state used for replay and recovery.
|
||||
#[serde(default, skip_serializing_if = "ToolResultDisposition::is_success")]
|
||||
disposition: ToolResultDisposition,
|
||||
/// Whether the tool result represents an execution error.
|
||||
#[serde(default, skip_serializing_if = "is_false")]
|
||||
is_error: bool,
|
||||
@@ -261,7 +264,17 @@ impl Item {
|
||||
content: Option<String>,
|
||||
is_error: bool,
|
||||
) -> Self {
|
||||
Self::tool_result_item_with_attachments(call_id, summary, content, is_error, Vec::new())
|
||||
Self::tool_result_item_with_disposition_and_attachments(
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
if is_error {
|
||||
ToolResultDisposition::Error
|
||||
} else {
|
||||
ToolResultDisposition::Success
|
||||
},
|
||||
Vec::new(),
|
||||
)
|
||||
}
|
||||
|
||||
/// Create a tool result item with durable, prunable structured attachments.
|
||||
@@ -272,11 +285,33 @@ impl Item {
|
||||
is_error: bool,
|
||||
attachments: Vec<Attachment>,
|
||||
) -> Self {
|
||||
Self::tool_result_item_with_disposition_and_attachments(
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
if is_error {
|
||||
ToolResultDisposition::Error
|
||||
} else {
|
||||
ToolResultDisposition::Success
|
||||
},
|
||||
attachments,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn tool_result_item_with_disposition_and_attachments(
|
||||
call_id: impl Into<String>,
|
||||
summary: impl Into<String>,
|
||||
content: Option<String>,
|
||||
disposition: ToolResultDisposition,
|
||||
attachments: Vec<Attachment>,
|
||||
) -> Self {
|
||||
let is_error = !disposition.is_success();
|
||||
Self::ToolResult {
|
||||
id: None,
|
||||
call_id: call_id.into(),
|
||||
summary: summary.into(),
|
||||
content,
|
||||
disposition,
|
||||
is_error,
|
||||
attachments,
|
||||
}
|
||||
|
||||
@@ -19,7 +19,7 @@ mod private {
|
||||
/// - Editing message history (add, delete, clear)
|
||||
/// - Registering tools and hooks
|
||||
///
|
||||
/// Can transition to [`Locked`] state via `Engine::lock()`.
|
||||
/// Can transition to [`Locked`] state via `Engine::lock(&history)`.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
|
||||
+227
-2
@@ -3,7 +3,14 @@
|
||||
//! Traits for defining tools callable by LLM.
|
||||
//! Usually auto-implemented using the `#[tool]` macro.
|
||||
|
||||
use std::{collections::HashMap, fmt, sync::Arc};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
fmt,
|
||||
future::Future,
|
||||
pin::Pin,
|
||||
sync::Arc,
|
||||
task::{Context, Poll},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use base64::{Engine as _, engine::general_purpose::STANDARD};
|
||||
@@ -23,6 +30,12 @@ pub enum ToolError {
|
||||
/// Internal error
|
||||
#[error("Internal error: {0}")]
|
||||
Internal(String),
|
||||
/// Cooperative cancellation completed with bounded terminal output.
|
||||
#[error("Tool execution cancelled")]
|
||||
Cancelled(ToolOutput),
|
||||
/// Execution was interrupted with a confirmed bounded terminal output.
|
||||
#[error("Tool execution interrupted")]
|
||||
Interrupted(ToolOutput),
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -158,6 +171,28 @@ pub enum Attachment {
|
||||
Image(ImageAttachment),
|
||||
}
|
||||
|
||||
/// Terminal disposition of one started tool call.
|
||||
///
|
||||
/// `Cancelled` means the tool confirmed cancellation. `OutcomeUnknown` means
|
||||
/// execution stopped without confirmation, so neither completion nor side
|
||||
/// effects may be inferred.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ToolResultDisposition {
|
||||
#[default]
|
||||
Success,
|
||||
Error,
|
||||
Interrupted,
|
||||
Cancelled,
|
||||
OutcomeUnknown,
|
||||
}
|
||||
|
||||
impl ToolResultDisposition {
|
||||
pub const fn is_success(&self) -> bool {
|
||||
matches!(self, Self::Success)
|
||||
}
|
||||
}
|
||||
|
||||
/// Tool execution result.
|
||||
///
|
||||
/// Every output has a mandatory `summary` (1-2 lines) that persists in
|
||||
@@ -322,6 +357,12 @@ impl ToolExecutionContext {
|
||||
}
|
||||
}
|
||||
|
||||
/// Identifies one live execution attempt without making the batch id a durable
|
||||
/// replay or idempotency authority.
|
||||
pub fn execution_id(&self) -> String {
|
||||
format!("{}:{}", self.batch_id, self.call_id)
|
||||
}
|
||||
|
||||
/// Context for direct, non-engine calls in unit tests and low-level callers.
|
||||
pub fn direct() -> Self {
|
||||
Self::new("direct", "direct", 0)
|
||||
@@ -334,6 +375,142 @@ impl Default for ToolExecutionContext {
|
||||
}
|
||||
}
|
||||
|
||||
/// The provider-confirmed terminal result of one started tool execution.
|
||||
///
|
||||
/// `OutcomeUnknown` is reserved for an execution task that had to be force-closed
|
||||
/// or failed before the provider could confirm its terminal result.
|
||||
#[derive(Debug)]
|
||||
pub enum ToolExecutionTerminal {
|
||||
Confirmed(Result<ToolOutput, ToolError>),
|
||||
OutcomeUnknown,
|
||||
}
|
||||
|
||||
/// The completion future paired with a [`ToolExecutionHandle`]. Dropping this
|
||||
/// future does not drop the provider execution: the spawned execution remains
|
||||
/// owned by its handle until it completes or is explicitly force-closed.
|
||||
pub struct ToolExecutionTerminalFuture {
|
||||
task: tokio::task::JoinHandle<Result<ToolOutput, ToolError>>,
|
||||
}
|
||||
|
||||
impl Future for ToolExecutionTerminalFuture {
|
||||
type Output = ToolExecutionTerminal;
|
||||
|
||||
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
|
||||
match Pin::new(&mut self.task).poll(cx) {
|
||||
Poll::Ready(Ok(result)) => Poll::Ready(ToolExecutionTerminal::Confirmed(result)),
|
||||
Poll::Ready(Err(_)) => Poll::Ready(ToolExecutionTerminal::OutcomeUnknown),
|
||||
Poll::Pending => Poll::Pending,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Live ownership and control for one started tool execution.
|
||||
///
|
||||
/// Execution, cancellation, and terminal confirmation remain provider-owned:
|
||||
/// this handle starts `Tool::execute`, delegates cooperative cancellation to
|
||||
/// `Tool::cancel_execution`, and treats execution-future completion as the
|
||||
/// provider's terminal confirmation. Agen may force-close only after its caller's
|
||||
/// deadline expires, at which point the outcome is necessarily unknown.
|
||||
#[derive(Clone)]
|
||||
pub struct ToolExecutionHandle {
|
||||
inner: Arc<ToolExecutionHandleInner>,
|
||||
}
|
||||
|
||||
struct ToolExecutionHandleInner {
|
||||
tool: Arc<dyn Tool>,
|
||||
context: ToolExecutionContext,
|
||||
abort: tokio::task::AbortHandle,
|
||||
}
|
||||
|
||||
impl Drop for ToolExecutionHandleInner {
|
||||
fn drop(&mut self) {
|
||||
// Losing the final live owner is an explicit forced close, never a
|
||||
// best-effort detached provider future.
|
||||
self.abort.abort();
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for ToolExecutionHandle {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("ToolExecutionHandle")
|
||||
.field("call_id", &self.inner.context.call_id)
|
||||
.field("batch_id", &self.inner.context.batch_id)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl ToolExecutionHandle {
|
||||
pub fn start(
|
||||
tool: Arc<dyn Tool>,
|
||||
input_json: String,
|
||||
context: ToolExecutionContext,
|
||||
) -> (Self, ToolExecutionTerminalFuture) {
|
||||
let execution_tool = Arc::clone(&tool);
|
||||
let execution_context = context.clone();
|
||||
let task =
|
||||
tokio::spawn(
|
||||
async move { execution_tool.execute(&input_json, execution_context).await },
|
||||
);
|
||||
let abort = task.abort_handle();
|
||||
(
|
||||
Self {
|
||||
inner: Arc::new(ToolExecutionHandleInner {
|
||||
tool,
|
||||
context,
|
||||
abort,
|
||||
}),
|
||||
},
|
||||
ToolExecutionTerminalFuture { task },
|
||||
)
|
||||
}
|
||||
|
||||
pub fn context(&self) -> &ToolExecutionContext {
|
||||
&self.inner.context
|
||||
}
|
||||
|
||||
pub async fn cancel_before(&self, deadline: tokio::time::Instant) -> Result<(), ToolError> {
|
||||
match tokio::time::timeout_at(
|
||||
deadline,
|
||||
self.inner.tool.cancel_execution(&self.inner.context),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(_) => Err(ToolError::Internal(format!(
|
||||
"tool cancellation request exceeded its deadline for call {}",
|
||||
self.inner.context.call_id
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn force_close(&self) {
|
||||
self.inner.abort.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct ToolExecutionPolicy {
|
||||
/// Time a pause waits for already-started providers to reach a natural safe
|
||||
/// boundary before escalating to explicit cooperative cancellation.
|
||||
pub pause_safe_boundary_timeout: std::time::Duration,
|
||||
/// Maximum time allowed for a provider to accept one cooperative
|
||||
/// cancellation request.
|
||||
pub cancellation_request_timeout: std::time::Duration,
|
||||
/// Maximum time allowed for all providers to confirm terminal results after
|
||||
/// cancellation has been requested.
|
||||
pub terminal_confirmation_timeout: std::time::Duration,
|
||||
}
|
||||
|
||||
impl Default for ToolExecutionPolicy {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
pause_safe_boundary_timeout: std::time::Duration::from_millis(100),
|
||||
cancellation_request_timeout: std::time::Duration::from_millis(100),
|
||||
terminal_confirmation_timeout: std::time::Duration::from_millis(500),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Tool trait
|
||||
// =============================================================================
|
||||
@@ -402,6 +579,26 @@ pub trait Tool: Send + Sync {
|
||||
input_json: &str,
|
||||
ctx: ToolExecutionContext,
|
||||
) -> Result<ToolOutput, ToolError>;
|
||||
|
||||
/// Request cooperative cancellation for one started call.
|
||||
///
|
||||
/// Implementations that own cancellable provider operations should signal
|
||||
/// every live execution identified by `call_id`, then let `execute` return
|
||||
/// the confirmed bounded terminal output. Direct callers may use this
|
||||
/// compatibility surface; Agen uses [`Tool::cancel_execution`] so providers
|
||||
/// can bind cancellation to one exact live attempt.
|
||||
async fn cancel(&self, _call_id: &str) -> Result<(), ToolError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Request cooperative cancellation for one exact started execution.
|
||||
///
|
||||
/// The default preserves existing tools by delegating to `cancel(call_id)`.
|
||||
/// Providers with their own execution registry should override this method
|
||||
/// and key cancellation by [`ToolExecutionContext::execution_id`].
|
||||
async fn cancel_execution(&self, ctx: &ToolExecutionContext) -> Result<(), ToolError> {
|
||||
self.cancel(&ctx.call_id).await
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -429,6 +626,9 @@ pub struct ToolCall {
|
||||
pub struct ToolResult {
|
||||
/// Corresponding tool call ID
|
||||
pub tool_use_id: String,
|
||||
/// Typed terminal state.
|
||||
#[serde(default, skip_serializing_if = "ToolResultDisposition::is_success")]
|
||||
pub disposition: ToolResultDisposition,
|
||||
/// Short summary (always kept in history)
|
||||
pub summary: String,
|
||||
/// Detailed output (prunable)
|
||||
@@ -445,11 +645,20 @@ pub struct ToolResult {
|
||||
impl ToolResult {
|
||||
/// Create a success result from a [`ToolOutput`].
|
||||
pub fn from_output(tool_use_id: impl Into<String>, output: ToolOutput) -> Self {
|
||||
Self::from_output_with_disposition(tool_use_id, output, ToolResultDisposition::Success)
|
||||
}
|
||||
|
||||
pub fn from_output_with_disposition(
|
||||
tool_use_id: impl Into<String>,
|
||||
output: ToolOutput,
|
||||
disposition: ToolResultDisposition,
|
||||
) -> Self {
|
||||
Self {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
disposition,
|
||||
summary: output.summary,
|
||||
content: output.content,
|
||||
is_error: false,
|
||||
is_error: !disposition.is_success(),
|
||||
attachments: output.attachments,
|
||||
}
|
||||
}
|
||||
@@ -458,12 +667,28 @@ impl ToolResult {
|
||||
pub fn error(tool_use_id: impl Into<String>, message: impl Into<String>) -> Self {
|
||||
Self {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
disposition: ToolResultDisposition::Error,
|
||||
summary: message.into(),
|
||||
content: None,
|
||||
is_error: true,
|
||||
attachments: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Close an execution whose completion and side effects cannot be confirmed.
|
||||
pub fn outcome_unknown(tool_use_id: impl Into<String>) -> Self {
|
||||
Self {
|
||||
tool_use_id: tool_use_id.into(),
|
||||
disposition: ToolResultDisposition::OutcomeUnknown,
|
||||
summary: "Tool execution outcome unknown".to_string(),
|
||||
content: Some(
|
||||
"Execution was interrupted before completion could be confirmed. Completion and side effects are unknown."
|
||||
.to_string(),
|
||||
),
|
||||
is_error: true,
|
||||
attachments: Vec::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
mod common;
|
||||
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::{Engine, EngineError, History, HistoryEntry, Item, Role};
|
||||
use common::MockLlmClient;
|
||||
|
||||
fn completed_text_events(text: &str) -> Vec<Event> {
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, text),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn run_preserves_item_annotations_without_projecting_them() {
|
||||
let client = MockLlmClient::new(completed_text_events("assistant reply"));
|
||||
let engine = Engine::<_, agen::state::Mutable, String>::new_annotated(client);
|
||||
let mut history = History::<String>::new();
|
||||
let mut next = 0usize;
|
||||
let mut annotate = |item: &Item| {
|
||||
next += 1;
|
||||
let kind = match item {
|
||||
Item::Message { role, .. } => match role {
|
||||
Role::User => "user",
|
||||
Role::Assistant => "assistant",
|
||||
Role::System => "system",
|
||||
},
|
||||
Item::ToolCall { .. } => "tool_call",
|
||||
Item::ToolResult { .. } => "tool_result",
|
||||
Item::Reasoning { .. } => "reasoning",
|
||||
};
|
||||
Ok(format!("{next}:{kind}"))
|
||||
};
|
||||
|
||||
let output = engine
|
||||
.run_with_annotation(&mut history, "hello", &mut annotate)
|
||||
.await;
|
||||
|
||||
assert!(matches!(output.result, agen::EngineRunExit::Finished));
|
||||
assert_eq!(history.len(), 2);
|
||||
assert_eq!(history.entries()[0].annotation, "1:user");
|
||||
assert_eq!(history.entries()[1].annotation, "2:assistant");
|
||||
assert_eq!(history.items_cloned().len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn append_failure_does_not_make_item_live() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut engine = Engine::<_, agen::state::Mutable, usize>::new_annotated(client);
|
||||
let mut history = History::<usize>::new();
|
||||
let mut fail = |_item: &Item| Err("commit failed".to_string());
|
||||
|
||||
let err = engine
|
||||
.append_history_with(&mut history, [Item::user_message("uncommitted")], &mut fail)
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(err, EngineError::HistoryAppend(message) if message == "commit failed"));
|
||||
assert!(history.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replacement_keeps_items_and_annotations_together() {
|
||||
let mut history = History::from_entries(vec![
|
||||
HistoryEntry::new(Item::user_message("old"), "old-ann".to_string()),
|
||||
HistoryEntry::new(Item::user_message("second"), "second-ann".to_string()),
|
||||
]);
|
||||
|
||||
history.truncate(1);
|
||||
assert_eq!(history.entries()[0].item.as_text(), Some("old"));
|
||||
assert_eq!(history.entries()[0].annotation, "old-ann");
|
||||
|
||||
let previous = history.replace_entries(vec![HistoryEntry::new(
|
||||
Item::user_message("restored"),
|
||||
"restored-ann".to_string(),
|
||||
)]);
|
||||
|
||||
assert_eq!(previous.len(), 1);
|
||||
assert_eq!(history.entries()[0].item.as_text(), Some("restored"));
|
||||
assert_eq!(history.entries()[0].annotation, "restored-ann");
|
||||
}
|
||||
@@ -58,6 +58,7 @@ async fn test_callback_llm_retry_event() {
|
||||
max_attempts: 2,
|
||||
total_timeout: Duration::from_secs(1),
|
||||
});
|
||||
let mut history = agen::History::new();
|
||||
|
||||
let notices = Arc::new(Mutex::new(Vec::new()));
|
||||
let sink = notices.clone();
|
||||
@@ -65,8 +66,11 @@ async fn test_callback_llm_retry_event() {
|
||||
sink.lock().unwrap().push((llm_call, notice.clone()));
|
||||
});
|
||||
|
||||
let result = engine.run("retry once").await;
|
||||
assert!(result.is_ok(), "engine should succeed after one retry");
|
||||
let result = engine.run(&mut history, "retry once").await;
|
||||
assert!(
|
||||
matches!(result.result, agen::EngineRunExit::Finished),
|
||||
"engine should succeed after one retry"
|
||||
);
|
||||
|
||||
let notices = notices.lock().unwrap();
|
||||
assert_eq!(notices.len(), 1);
|
||||
@@ -91,6 +95,7 @@ async fn test_callback_text_block_events() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
let text_deltas = Arc::new(Mutex::new(Vec::new()));
|
||||
let text_completes = Arc::new(Mutex::new(Vec::new()));
|
||||
@@ -108,9 +113,12 @@ async fn test_callback_text_block_events() {
|
||||
});
|
||||
});
|
||||
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let result = engine.run("Greet me").await;
|
||||
assert!(result.is_ok(), "Engine should complete");
|
||||
// Mutable::run consumes self, returns (Locked, EngineRunExit)
|
||||
let result = engine.run(&mut history, "Greet me").await;
|
||||
assert!(
|
||||
matches!(result.result, agen::EngineRunExit::Finished),
|
||||
"Engine should complete"
|
||||
);
|
||||
|
||||
let deltas = text_deltas.lock().unwrap();
|
||||
assert_eq!(deltas.len(), 2);
|
||||
@@ -137,6 +145,7 @@ async fn test_callback_tool_call_complete() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
let tool_starts = Arc::new(Mutex::new(Vec::<(String, String)>::new()));
|
||||
let tool_completes = Arc::new(Mutex::new(Vec::new()));
|
||||
@@ -154,8 +163,8 @@ async fn test_callback_tool_call_complete() {
|
||||
});
|
||||
});
|
||||
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let _ = engine.run("Weather please").await;
|
||||
// Mutable::run consumes self, returns (Locked, EngineRunExit)
|
||||
let _ = engine.run(&mut history, "Weather please").await;
|
||||
|
||||
let starts = tool_starts.lock().unwrap();
|
||||
assert_eq!(starts.len(), 1);
|
||||
@@ -183,6 +192,7 @@ async fn test_callback_turn_events() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
let turn_starts = Arc::new(Mutex::new(Vec::new()));
|
||||
let turn_ends = Arc::new(Mutex::new(Vec::new()));
|
||||
@@ -197,9 +207,9 @@ async fn test_callback_turn_events() {
|
||||
ends.lock().unwrap().push(turn);
|
||||
});
|
||||
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let result = engine.run("Do something").await;
|
||||
assert!(result.is_ok());
|
||||
// Mutable::run consumes self, returns (Locked, EngineRunExit)
|
||||
let result = engine.run(&mut history, "Do something").await;
|
||||
assert!(matches!(result.result, agen::EngineRunExit::Finished));
|
||||
|
||||
let starts = turn_starts.lock().unwrap();
|
||||
let ends = turn_ends.lock().unwrap();
|
||||
@@ -254,6 +264,7 @@ async fn test_callback_tool_result_events() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
engine.register_tool(fixed_tool(
|
||||
"fixed",
|
||||
@@ -276,7 +287,7 @@ async fn test_callback_tool_result_events() {
|
||||
));
|
||||
});
|
||||
|
||||
let _ = engine.run("call it").await;
|
||||
let _ = engine.run(&mut history, "call it").await;
|
||||
|
||||
let observed = captured.lock().unwrap();
|
||||
assert_eq!(observed.len(), 1);
|
||||
@@ -330,6 +341,7 @@ async fn test_callback_tool_result_error_path() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
engine.register_tool(erroring_tool("erroring", "boom"));
|
||||
|
||||
@@ -345,7 +357,7 @@ async fn test_callback_tool_result_error_path() {
|
||||
));
|
||||
});
|
||||
|
||||
let _ = engine.run("fail it").await;
|
||||
let _ = engine.run(&mut history, "fail it").await;
|
||||
|
||||
let observed = captured.lock().unwrap();
|
||||
assert_eq!(observed.len(), 1);
|
||||
@@ -374,6 +386,7 @@ async fn test_callback_usage_events() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
let usage_events = Arc::new(Mutex::new(Vec::new()));
|
||||
|
||||
@@ -382,8 +395,8 @@ async fn test_callback_usage_events() {
|
||||
usages.lock().unwrap().push(event.clone());
|
||||
});
|
||||
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let _ = engine.run("Hello").await;
|
||||
// Mutable::run consumes self, returns (Locked, EngineRunExit)
|
||||
let _ = engine.run(&mut history, "Hello").await;
|
||||
|
||||
let usages = usage_events.lock().unwrap();
|
||||
assert_eq!(usages.len(), 1);
|
||||
|
||||
@@ -19,6 +19,7 @@ use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
pub struct MockLlmClient {
|
||||
responses: Arc<Vec<Vec<Event>>>,
|
||||
call_count: Arc<AtomicUsize>,
|
||||
requests: Arc<Mutex<Vec<Request>>>,
|
||||
}
|
||||
|
||||
impl MockLlmClient {
|
||||
@@ -30,6 +31,7 @@ impl MockLlmClient {
|
||||
Self {
|
||||
responses: Arc::new(responses),
|
||||
call_count: Arc::new(AtomicUsize::new(0)),
|
||||
requests: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -41,6 +43,10 @@ impl MockLlmClient {
|
||||
pub fn event_count(&self) -> usize {
|
||||
self.responses.iter().map(|v| v.len()).sum()
|
||||
}
|
||||
|
||||
pub fn requests(&self) -> Vec<Request> {
|
||||
self.requests.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -51,8 +57,9 @@ impl LlmClient for MockLlmClient {
|
||||
|
||||
async fn stream(
|
||||
&self,
|
||||
_request: Request,
|
||||
request: Request,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>, ClientError> {
|
||||
self.requests.lock().unwrap().push(request);
|
||||
let count = self.call_count.fetch_add(1, Ordering::SeqCst);
|
||||
if count >= self.responses.len() {
|
||||
return Err(ClientError::Api {
|
||||
|
||||
@@ -134,11 +134,15 @@ async fn test_engine_simple_text_response() {
|
||||
|
||||
let client = MockLlmClient::from_fixture(&fixture_path).unwrap();
|
||||
let engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
// Send a simple message (Mutable::run consumes self, returns tuple)
|
||||
let result = engine.run("Hello").await;
|
||||
let result = engine.run(&mut history, "Hello").await;
|
||||
|
||||
assert!(result.is_ok(), "Engine should complete successfully");
|
||||
assert!(
|
||||
matches!(result.result, agen::EngineRunExit::Finished),
|
||||
"Engine should complete successfully"
|
||||
);
|
||||
}
|
||||
|
||||
/// Verify that Engine can correctly process responses containing tool calls
|
||||
@@ -156,6 +160,7 @@ async fn test_engine_tool_call() {
|
||||
|
||||
let client = MockLlmClient::from_fixture(&fixture_path).unwrap();
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
// Register tool
|
||||
let weather_tool = MockWeatherTool::new();
|
||||
@@ -163,7 +168,9 @@ async fn test_engine_tool_call() {
|
||||
engine.register_tool(weather_tool.definition());
|
||||
|
||||
// Send message (Mutable::run consumes self, returns tuple)
|
||||
let _result = engine.run("What's the weather in Tokyo?").await;
|
||||
let _result = engine
|
||||
.run(&mut history, "What's the weather in Tokyo?")
|
||||
.await;
|
||||
|
||||
// Verify tool was called
|
||||
// Note: max_turns=1 so no request is sent after tool result
|
||||
@@ -195,11 +202,15 @@ async fn test_engine_with_programmatic_events() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let engine = Engine::new(client);
|
||||
let mut history = agen::History::new();
|
||||
|
||||
// Mutable::run consumes self, returns tuple
|
||||
let result = engine.run("Greet me").await;
|
||||
let result = engine.run(&mut history, "Greet me").await;
|
||||
|
||||
assert!(result.is_ok(), "Engine should complete successfully");
|
||||
assert!(
|
||||
matches!(result.result, agen::EngineRunExit::Finished),
|
||||
"Engine should complete successfully"
|
||||
);
|
||||
}
|
||||
|
||||
/// Verify that ToolCallCollector correctly collects ToolCall from ToolUse block events
|
||||
|
||||
@@ -9,9 +9,12 @@ use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use agen::Item;
|
||||
use agen::interceptor::{
|
||||
Interceptor, PreRequestAction, PreToolAction, ToolCallInfo, TurnEndAction,
|
||||
};
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput};
|
||||
use agen::{Engine, EngineError};
|
||||
use agen::{Engine, EngineError, EngineRunExit, History, StopReason};
|
||||
use async_trait::async_trait;
|
||||
use common::MockLlmClient;
|
||||
|
||||
@@ -39,36 +42,37 @@ fn test_mutable_set_system_prompt() {
|
||||
fn test_mutable_history_manipulation() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
// Initial state is empty
|
||||
assert!(engine.history().is_empty());
|
||||
assert!(history.is_empty());
|
||||
|
||||
// Add to history
|
||||
engine
|
||||
.append_history(vec![Item::user_message("Hello")])
|
||||
.append_history(&mut history, vec![Item::user_message("Hello")])
|
||||
.unwrap();
|
||||
engine
|
||||
.append_history(vec![Item::assistant_message("Hi there!")])
|
||||
.append_history(&mut history, vec![Item::assistant_message("Hi there!")])
|
||||
.unwrap();
|
||||
assert_eq!(engine.history().len(), 2);
|
||||
assert_eq!(history.len(), 2);
|
||||
|
||||
// Append to history via the callback-aware API.
|
||||
engine
|
||||
.append_history(vec![Item::user_message("How are you?")])
|
||||
.append_history(&mut history, vec![Item::user_message("How are you?")])
|
||||
.unwrap();
|
||||
assert_eq!(engine.history().len(), 3);
|
||||
assert_eq!(history.len(), 3);
|
||||
|
||||
// Clear history
|
||||
engine.clear_history();
|
||||
assert!(engine.history().is_empty());
|
||||
engine.clear_history(&mut history);
|
||||
assert!(history.is_empty());
|
||||
|
||||
// Set history
|
||||
let items = vec![
|
||||
Item::user_message("Test"),
|
||||
Item::assistant_message("Response"),
|
||||
];
|
||||
engine.set_history(items);
|
||||
assert_eq!(engine.history().len(), 2);
|
||||
engine.set_history(&mut history, items);
|
||||
assert_eq!(history.len(), 2);
|
||||
}
|
||||
|
||||
/// Verify that Engine can be constructed using builder pattern
|
||||
@@ -76,9 +80,10 @@ fn test_mutable_history_manipulation() {
|
||||
fn test_mutable_builder_pattern() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let engine = Engine::new(client).system_prompt("System prompt");
|
||||
let history: History = History::new();
|
||||
|
||||
assert_eq!(engine.get_system_prompt(), Some("System prompt"));
|
||||
assert!(engine.history().is_empty());
|
||||
assert!(history.is_empty());
|
||||
}
|
||||
|
||||
/// Verify that multiple items can be added with append_history and callbacks fire.
|
||||
@@ -88,6 +93,7 @@ fn test_mutable_append_history() {
|
||||
let observed = Arc::new(Mutex::new(Vec::new()));
|
||||
let observed_for_callback = Arc::clone(&observed);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
engine.on_history_append(move |item| {
|
||||
if let Some(text) = item.as_text() {
|
||||
observed_for_callback.lock().unwrap().push(text.to_string());
|
||||
@@ -96,18 +102,21 @@ fn test_mutable_append_history() {
|
||||
});
|
||||
|
||||
engine
|
||||
.append_history(vec![Item::user_message("First")])
|
||||
.append_history(&mut history, vec![Item::user_message("First")])
|
||||
.unwrap();
|
||||
|
||||
engine
|
||||
.append_history(vec![
|
||||
Item::assistant_message("Response 1"),
|
||||
Item::user_message("Second"),
|
||||
Item::assistant_message("Response 2"),
|
||||
])
|
||||
.append_history(
|
||||
&mut history,
|
||||
vec![
|
||||
Item::assistant_message("Response 1"),
|
||||
Item::user_message("Second"),
|
||||
Item::assistant_message("Response 2"),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(engine.history().len(), 4);
|
||||
assert_eq!(history.len(), 4);
|
||||
assert_eq!(
|
||||
observed.lock().unwrap().as_slice(),
|
||||
["First", "Response 1", "Second", "Response 2"]
|
||||
@@ -182,6 +191,7 @@ async fn history_append_failure_stops_before_tool_execution() {
|
||||
]);
|
||||
let tool = CountingTool::new("count_tool");
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
engine.register_tool(tool.definition());
|
||||
engine.on_history_append(|item| {
|
||||
if item.is_tool_call() {
|
||||
@@ -191,15 +201,15 @@ async fn history_append_failure_stops_before_tool_execution() {
|
||||
}
|
||||
});
|
||||
|
||||
let mut engine = engine.lock();
|
||||
let error = engine.run("use the tool").await.unwrap_err();
|
||||
let mut engine = engine.lock(&history);
|
||||
let exit = engine.run(&mut history, "use the tool").await;
|
||||
|
||||
assert!(
|
||||
matches!(error, EngineError::HistoryAppend(ref message) if message == "simulated ENOSPC")
|
||||
matches!(exit, EngineRunExit::Interrupted(StopReason::Unexpected(EngineError::HistoryAppend(ref message))) if message == "simulated ENOSPC")
|
||||
);
|
||||
assert_eq!(tool.call_count(), 0);
|
||||
assert_eq!(engine.history().len(), 1);
|
||||
assert_eq!(engine.history()[0].as_text(), Some("use the tool"));
|
||||
assert_eq!(history.len(), 1);
|
||||
assert_eq!(history.entries()[0].item.as_text(), Some("use the tool"));
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -211,21 +221,22 @@ async fn history_append_failure_stops_before_tool_execution() {
|
||||
fn test_lock_transition() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
engine.set_system_prompt("System");
|
||||
engine
|
||||
.append_history(vec![Item::user_message("Hello")])
|
||||
.append_history(&mut history, vec![Item::user_message("Hello")])
|
||||
.unwrap();
|
||||
engine
|
||||
.append_history(vec![Item::assistant_message("Hi")])
|
||||
.append_history(&mut history, vec![Item::assistant_message("Hi")])
|
||||
.unwrap();
|
||||
|
||||
// Lock
|
||||
let locked_engine = engine.lock();
|
||||
let locked_engine = engine.lock(&history);
|
||||
|
||||
// History and system prompt are still accessible in Locked state
|
||||
assert_eq!(locked_engine.get_system_prompt(), Some("System"));
|
||||
assert_eq!(locked_engine.history().len(), 2);
|
||||
assert_eq!(history.len(), 2);
|
||||
assert_eq!(locked_engine.locked_prefix_len(), 2);
|
||||
}
|
||||
|
||||
@@ -234,21 +245,22 @@ fn test_lock_transition() {
|
||||
fn test_unlock_transition() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
engine
|
||||
.append_history(vec![Item::user_message("Hello")])
|
||||
.append_history(&mut history, vec![Item::user_message("Hello")])
|
||||
.unwrap();
|
||||
let locked_engine = engine.lock();
|
||||
let locked_engine = engine.lock(&history);
|
||||
|
||||
// Unlock
|
||||
let mut engine = locked_engine.unlock();
|
||||
|
||||
// History operations are available again in Mutable state
|
||||
engine
|
||||
.append_history(vec![Item::assistant_message("Hi")])
|
||||
.append_history(&mut history, vec![Item::assistant_message("Hi")])
|
||||
.unwrap();
|
||||
engine.clear_history();
|
||||
assert!(engine.history().is_empty());
|
||||
engine.clear_history(&mut history);
|
||||
assert!(history.is_empty());
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -269,20 +281,20 @@ async fn test_mutable_run_updates_history() -> Result<(), EngineError> {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
// Execute (Mutable::run consumes self, returns EngineRunOutput)
|
||||
let out = engine.run("Hi there").await?;
|
||||
let engine = out.engine;
|
||||
let _out = engine.run(&mut history, "Hi there").await;
|
||||
|
||||
// History is updated
|
||||
let history = engine.history();
|
||||
let entries = history.entries();
|
||||
assert_eq!(history.len(), 2); // user + assistant
|
||||
|
||||
// User message
|
||||
assert_eq!(history[0].as_text(), Some("Hi there"));
|
||||
assert_eq!(entries[0].item.as_text(), Some("Hi there"));
|
||||
|
||||
// Assistant message
|
||||
assert_eq!(history[1].as_text(), Some("Hello, I'm an assistant!"));
|
||||
assert_eq!(entries[1].item.as_text(), Some("Hello, I'm an assistant!"));
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -313,35 +325,36 @@ async fn test_locked_multi_turn_history_accumulation() {
|
||||
]);
|
||||
|
||||
let engine = Engine::new(client).system_prompt("You are helpful.");
|
||||
let mut history: History = History::new();
|
||||
|
||||
// Lock (after setting system prompt)
|
||||
let mut locked_engine = engine.lock();
|
||||
let mut locked_engine = engine.lock(&history);
|
||||
assert_eq!(locked_engine.locked_prefix_len(), 0); // No items yet
|
||||
|
||||
// Turn 1
|
||||
let result1 = locked_engine.run("Hello!").await;
|
||||
assert!(result1.is_ok());
|
||||
assert_eq!(locked_engine.history().len(), 2); // user + assistant
|
||||
let result1 = locked_engine.run(&mut history, "Hello!").await;
|
||||
assert!(matches!(result1, EngineRunExit::Finished));
|
||||
assert_eq!(history.len(), 2); // user + assistant
|
||||
|
||||
// Turn 2
|
||||
let result2 = locked_engine.run("Can you help me?").await;
|
||||
assert!(result2.is_ok());
|
||||
assert_eq!(locked_engine.history().len(), 4); // 2 * (user + assistant)
|
||||
let result2 = locked_engine.run(&mut history, "Can you help me?").await;
|
||||
assert!(matches!(result2, EngineRunExit::Finished));
|
||||
assert_eq!(history.len(), 4); // 2 * (user + assistant)
|
||||
|
||||
// Verify history contents
|
||||
let history = locked_engine.history();
|
||||
let entries = history.entries();
|
||||
|
||||
// Turn 1 user message
|
||||
assert_eq!(history[0].as_text(), Some("Hello!"));
|
||||
assert_eq!(entries[0].item.as_text(), Some("Hello!"));
|
||||
|
||||
// Turn 1 assistant message
|
||||
assert_eq!(history[1].as_text(), Some("Nice to meet you!"));
|
||||
assert_eq!(entries[1].item.as_text(), Some("Nice to meet you!"));
|
||||
|
||||
// Turn 2 user message
|
||||
assert_eq!(history[2].as_text(), Some("Can you help me?"));
|
||||
assert_eq!(entries[2].item.as_text(), Some("Can you help me?"));
|
||||
|
||||
// Turn 2 assistant message
|
||||
assert_eq!(history[3].as_text(), Some("I can help with that."));
|
||||
assert_eq!(entries[3].item.as_text(), Some("I can help with that."));
|
||||
}
|
||||
|
||||
/// Verify that locked_prefix_len correctly records history length at lock time
|
||||
@@ -367,26 +380,33 @@ async fn test_locked_prefix_len_tracking() {
|
||||
]);
|
||||
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
// Add items beforehand
|
||||
engine
|
||||
.append_history(vec![Item::user_message("Pre-existing message 1")])
|
||||
.append_history(
|
||||
&mut history,
|
||||
vec![Item::user_message("Pre-existing message 1")],
|
||||
)
|
||||
.unwrap();
|
||||
engine
|
||||
.append_history(vec![Item::assistant_message("Pre-existing response 1")])
|
||||
.append_history(
|
||||
&mut history,
|
||||
vec![Item::assistant_message("Pre-existing response 1")],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(engine.history().len(), 2);
|
||||
assert_eq!(history.len(), 2);
|
||||
|
||||
// Lock
|
||||
let mut locked_engine = engine.lock();
|
||||
let mut locked_engine = engine.lock(&history);
|
||||
assert_eq!(locked_engine.locked_prefix_len(), 2); // 2 items at lock time
|
||||
|
||||
// Execute turn
|
||||
locked_engine.run("New message").await.unwrap();
|
||||
locked_engine.run(&mut history, "New message").await;
|
||||
|
||||
// History grows but locked_prefix_len remains unchanged
|
||||
assert_eq!(locked_engine.history().len(), 4); // 2 + 2
|
||||
assert_eq!(history.len(), 4); // 2 + 2
|
||||
assert_eq!(locked_engine.locked_prefix_len(), 2); // Unchanged
|
||||
}
|
||||
|
||||
@@ -413,18 +433,22 @@ async fn test_turn_count_increment() -> Result<(), EngineError> {
|
||||
]);
|
||||
|
||||
let engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
assert_eq!(engine.turn_count(), 0);
|
||||
assert_eq!(engine.llm_call_count(), 0);
|
||||
|
||||
// First run consumes Mutable, returns EngineRunOutput
|
||||
let mut engine = engine.run("First").await?.engine;
|
||||
let mut engine = engine.run(&mut history, "First").await.engine;
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
// Retry not yet implemented → AgentTurn:LlmCall is 1:1.
|
||||
assert_eq!(engine.llm_call_count(), 1);
|
||||
|
||||
// Subsequent runs on Locked take &mut self
|
||||
engine.run("Second").await?;
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "Second").await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 2);
|
||||
assert_eq!(engine.llm_call_count(), 2);
|
||||
|
||||
@@ -444,28 +468,29 @@ async fn test_unlock_edit_relock() {
|
||||
]]);
|
||||
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
engine
|
||||
.append_history(vec![
|
||||
Item::user_message("Hello"),
|
||||
Item::assistant_message("Hi"),
|
||||
])
|
||||
.append_history(
|
||||
&mut history,
|
||||
vec![Item::user_message("Hello"), Item::assistant_message("Hi")],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
// Lock -> Unlock
|
||||
let locked = engine.lock();
|
||||
let locked = engine.lock(&history);
|
||||
assert_eq!(locked.locked_prefix_len(), 2);
|
||||
|
||||
let mut unlocked = locked.unlock();
|
||||
|
||||
// Edit history
|
||||
unlocked.clear_history();
|
||||
unlocked.clear_history(&mut history);
|
||||
unlocked
|
||||
.append_history(vec![Item::user_message("Fresh start")])
|
||||
.append_history(&mut history, vec![Item::user_message("Fresh start")])
|
||||
.unwrap();
|
||||
|
||||
// Re-lock
|
||||
let relocked = unlocked.lock();
|
||||
assert_eq!(relocked.history().len(), 1);
|
||||
let relocked = unlocked.lock(&history);
|
||||
assert_eq!(history.len(), 1);
|
||||
assert_eq!(relocked.locked_prefix_len(), 1);
|
||||
}
|
||||
|
||||
@@ -508,19 +533,26 @@ async fn test_lock_unlock_relock_tools_remain_effective() {
|
||||
]);
|
||||
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
let tool_a = CountingTool::new("tool_a");
|
||||
engine.register_tool(tool_a.definition());
|
||||
|
||||
let mut locked = engine.lock();
|
||||
locked.run("first").await.expect("first run");
|
||||
let mut locked = engine.lock(&history);
|
||||
assert!(matches!(
|
||||
locked.run(&mut history, "first").await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(tool_a.call_count(), 1, "tool_a should be called once");
|
||||
|
||||
let mut unlocked = locked.unlock();
|
||||
let tool_b = CountingTool::new("tool_b");
|
||||
unlocked.register_tool(tool_b.definition());
|
||||
|
||||
let mut relocked = unlocked.lock();
|
||||
relocked.run("second").await.expect("second run");
|
||||
let mut relocked = unlocked.lock(&history);
|
||||
assert!(matches!(
|
||||
relocked.run(&mut history, "second").await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
|
||||
assert_eq!(tool_a.call_count(), 1, "tool_a should not be called again");
|
||||
assert_eq!(tool_b.call_count(), 1, "tool_b should be called once");
|
||||
@@ -535,8 +567,9 @@ async fn test_lock_unlock_relock_tools_remain_effective() {
|
||||
fn test_system_prompt_preserved_in_locked_state() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let engine = Engine::new(client).system_prompt("Important system prompt");
|
||||
let history: History = History::new();
|
||||
|
||||
let locked = engine.lock();
|
||||
let locked = engine.lock(&history);
|
||||
assert_eq!(locked.get_system_prompt(), Some("Important system prompt"));
|
||||
|
||||
let unlocked = locked.unlock();
|
||||
@@ -551,13 +584,228 @@ fn test_system_prompt_preserved_in_locked_state() {
|
||||
fn test_system_prompt_change_after_unlock() {
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let engine = Engine::new(client).system_prompt("Original prompt");
|
||||
let history: History = History::new();
|
||||
|
||||
let locked = engine.lock();
|
||||
let locked = engine.lock(&history);
|
||||
let mut unlocked = locked.unlock();
|
||||
|
||||
unlocked.set_system_prompt("New prompt");
|
||||
assert_eq!(unlocked.get_system_prompt(), Some("New prompt"));
|
||||
|
||||
let relocked = unlocked.lock();
|
||||
let relocked = unlocked.lock(&history);
|
||||
assert_eq!(relocked.get_system_prompt(), Some("New prompt"));
|
||||
}
|
||||
|
||||
fn completed_text_events() -> Vec<Event> {
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "done"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]
|
||||
}
|
||||
|
||||
struct YieldOnce {
|
||||
calls: AtomicUsize,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for YieldOnce {
|
||||
async fn pre_llm_request(&self, _context: &mut Vec<Item>) -> PreRequestAction {
|
||||
if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
PreRequestAction::Yield
|
||||
} else {
|
||||
PreRequestAction::Continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct PauseToolOnce {
|
||||
calls: AtomicUsize,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for PauseToolOnce {
|
||||
async fn pre_tool_call(&self, _info: &mut ToolCallInfo) -> PreToolAction {
|
||||
if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
PreToolAction::Pause
|
||||
} else {
|
||||
PreToolAction::Continue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct ContinueTurnOnce {
|
||||
calls: AtomicUsize,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Interceptor for ContinueTurnOnce {
|
||||
async fn on_turn_end(&self, _history: &[Item]) -> TurnEndAction {
|
||||
if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
|
||||
TurnEndAction::ContinueWithMessages(vec![Item::system_message("continue")])
|
||||
} else {
|
||||
TurnEndAction::Finish
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn max_turns_is_scoped_to_each_fresh_run() {
|
||||
let mut history: History = History::new();
|
||||
let responses = vec![completed_text_events(), completed_text_events()];
|
||||
let mut engine = Engine::new(MockLlmClient::with_responses(responses));
|
||||
engine.set_max_turns(Some(1));
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "first").await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "second").await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 2);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn yielded_resume_keeps_the_same_unspent_turn_budget() {
|
||||
let mut history: History = History::new();
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_max_turns(Some(1));
|
||||
engine.set_interceptor(YieldOnce {
|
||||
calls: AtomicUsize::new(0),
|
||||
});
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "start").await,
|
||||
EngineRunExit::Yielded
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 0);
|
||||
assert_eq!(engine.active_run_turn_count(), Some(0));
|
||||
|
||||
assert!(matches!(
|
||||
engine.resume(&mut history).await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn paused_tool_resume_does_not_reset_the_consumed_turn_budget() {
|
||||
let mut history: History = History::new();
|
||||
let events = vec![
|
||||
Event::tool_use_start(0, "call_1", "count_tool"),
|
||||
Event::tool_input_delta(0, "{}"),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
];
|
||||
let tool = CountingTool::new("count_tool");
|
||||
let mut engine = Engine::new(MockLlmClient::new(events));
|
||||
engine.set_max_turns(Some(1));
|
||||
engine.register_tool(tool.definition());
|
||||
engine.set_interceptor(PauseToolOnce {
|
||||
calls: AtomicUsize::new(0),
|
||||
});
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "call it").await,
|
||||
EngineRunExit::Paused
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
assert_eq!(engine.active_run_turn_count(), Some(1));
|
||||
assert_eq!(tool.call_count(), 0);
|
||||
|
||||
assert!(matches!(
|
||||
engine.resume(&mut history).await,
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached)
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
assert_eq!(tool.call_count(), 1, "the consumed turn's tool still runs");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn fresh_input_abandons_a_paused_run_and_starts_a_new_budget() {
|
||||
let mut history: History = History::new();
|
||||
let tool_events = vec![
|
||||
Event::tool_use_start(0, "call_1", "count_tool"),
|
||||
Event::tool_input_delta(0, "{}"),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
];
|
||||
let client = MockLlmClient::with_responses(vec![tool_events, completed_text_events()]);
|
||||
let tool = CountingTool::new("count_tool");
|
||||
let mut engine = Engine::new(client);
|
||||
engine.set_max_turns(Some(1));
|
||||
engine.register_tool(tool.definition());
|
||||
engine.set_interceptor(PauseToolOnce {
|
||||
calls: AtomicUsize::new(0),
|
||||
});
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "pause").await,
|
||||
EngineRunExit::Paused
|
||||
));
|
||||
assert_eq!(engine.active_run_turn_count(), Some(1));
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "replace").await,
|
||||
EngineRunExit::Finished
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 2);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
assert_eq!(tool.call_count(), 1, "pending-tool semantics are unchanged");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn interceptor_continuation_consumes_the_logical_run_budget() {
|
||||
let mut history: History = History::new();
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_max_turns(Some(1));
|
||||
engine.set_interceptor(ContinueTurnOnce {
|
||||
calls: AtomicUsize::new(0),
|
||||
});
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
assert!(matches!(
|
||||
engine.run(&mut history, "start").await,
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached)
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 1);
|
||||
assert_eq!(engine.llm_call_count(), 1);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn restored_active_run_budget_is_enforced_before_another_llm_call() {
|
||||
let mut history: History = History::new();
|
||||
let mut engine = Engine::new(MockLlmClient::new(completed_text_events()));
|
||||
engine.set_max_turns(Some(1));
|
||||
engine.set_turn_count(7);
|
||||
engine.set_active_run_turn_count(Some(1));
|
||||
let mut engine = engine.lock(&history);
|
||||
|
||||
assert!(matches!(
|
||||
engine.resume(&mut history).await,
|
||||
EngineRunExit::Interrupted(StopReason::LimitReached)
|
||||
));
|
||||
assert_eq!(engine.turn_count(), 7);
|
||||
assert_eq!(engine.llm_call_count(), 0);
|
||||
assert_eq!(engine.active_run_turn_count(), None);
|
||||
}
|
||||
|
||||
@@ -6,12 +6,13 @@ use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use agen::Engine;
|
||||
use agen::interceptor::{Interceptor, PostToolAction, PreToolAction, ToolCallInfo, ToolResultInfo};
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::tool::{
|
||||
Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult,
|
||||
ToolResultDisposition,
|
||||
};
|
||||
use agen::{Engine, History, Item, ToolExecutionPolicy};
|
||||
use async_trait::async_trait;
|
||||
|
||||
mod common;
|
||||
@@ -70,6 +71,144 @@ impl Tool for SlowTool {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct FirstAttemptHangsTool {
|
||||
calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
impl FirstAttemptHangsTool {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
calls: Arc::new(AtomicUsize::new(0)),
|
||||
}
|
||||
}
|
||||
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
let tool = self.clone();
|
||||
Arc::new(move || {
|
||||
let meta = ToolMeta::new("hang_once")
|
||||
.description("Hangs on the first execution attempt")
|
||||
.input_schema(serde_json::json!({"type": "object"}));
|
||||
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
|
||||
})
|
||||
}
|
||||
|
||||
fn call_count(&self) -> usize {
|
||||
self.calls.load(Ordering::SeqCst)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for FirstAttemptHangsTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
_input_json: &str,
|
||||
_ctx: ToolExecutionContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let attempt = self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
if attempt == 0 {
|
||||
std::future::pending::<()>().await;
|
||||
}
|
||||
Ok("completed on retry".to_string().into())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct CooperativeCancelTool {
|
||||
calls: Arc<AtomicUsize>,
|
||||
cancelled: Arc<tokio::sync::Notify>,
|
||||
}
|
||||
|
||||
impl CooperativeCancelTool {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
calls: Arc::new(AtomicUsize::new(0)),
|
||||
cancelled: Arc::new(tokio::sync::Notify::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
let tool = self.clone();
|
||||
Arc::new(move || {
|
||||
let meta = ToolMeta::new("cooperative")
|
||||
.description("Returns bounded progress after cancellation")
|
||||
.input_schema(serde_json::json!({"type": "object"}));
|
||||
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for CooperativeCancelTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
_input_json: &str,
|
||||
_ctx: ToolExecutionContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.cancelled.notified().await;
|
||||
Err(ToolError::Cancelled(ToolOutput {
|
||||
summary: "cooperative command cancelled".to_string(),
|
||||
content: Some("stdout before cancellation\nstderr before cancellation".to_string()),
|
||||
attachments: Vec::new(),
|
||||
}))
|
||||
}
|
||||
|
||||
async fn cancel(&self, _call_id: &str) -> Result<(), ToolError> {
|
||||
self.cancelled.notify_one();
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct SafePauseTool {
|
||||
calls: Arc<AtomicUsize>,
|
||||
cancellations: Arc<AtomicUsize>,
|
||||
release: Arc<tokio::sync::Notify>,
|
||||
}
|
||||
|
||||
impl SafePauseTool {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
calls: Arc::new(AtomicUsize::new(0)),
|
||||
cancellations: Arc::new(AtomicUsize::new(0)),
|
||||
release: Arc::new(tokio::sync::Notify::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
let tool = self.clone();
|
||||
Arc::new(move || {
|
||||
let meta = ToolMeta::new("safe_pause")
|
||||
.description("Waits for a safe-boundary release")
|
||||
.input_schema(serde_json::json!({"type": "object"}));
|
||||
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for SafePauseTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
_input_json: &str,
|
||||
_ctx: ToolExecutionContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.release.notified().await;
|
||||
Ok(ToolOutput {
|
||||
summary: "safe-boundary complete".to_string(),
|
||||
content: Some("safe-boundary complete".to_string()),
|
||||
attachments: Vec::new(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn cancel(&self, _call_id: &str) -> Result<(), ToolError> {
|
||||
self.cancellations.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ContextRecordingTool {
|
||||
name: String,
|
||||
@@ -145,6 +284,7 @@ async fn test_parallel_tool_execution() {
|
||||
],
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
let tool1 = SlowTool::new("slow_tool_1", 100);
|
||||
let tool2 = SlowTool::new("slow_tool_2", 100);
|
||||
let tool3 = SlowTool::new("slow_tool_3", 100);
|
||||
@@ -159,7 +299,7 @@ async fn test_parallel_tool_execution() {
|
||||
|
||||
let start = Instant::now();
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let _result = engine.run("Run all tools").await;
|
||||
let _result = engine.run(&mut history, "Run all tools").await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
// Verify all tools were called
|
||||
@@ -178,6 +318,450 @@ async fn test_parallel_tool_execution() {
|
||||
println!("Parallel execution completed in {:?}", elapsed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn completed_results_commit_before_publish_without_waiting_for_siblings() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
vec![
|
||||
Event::tool_use_start(0, "call_slow", "slow_first"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::tool_use_start(1, "call_fast", "fast_second"),
|
||||
Event::tool_input_delta(1, r#"{}"#),
|
||||
Event::tool_use_stop(1),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Done"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
]);
|
||||
let client_probe = client.clone();
|
||||
let mut engine = Engine::new(client);
|
||||
engine.register_tool(SlowTool::new("slow_first", 100).definition());
|
||||
engine.register_tool(SlowTool::new("fast_second", 5).definition());
|
||||
|
||||
let observed = Arc::new(Mutex::new(Vec::<String>::new()));
|
||||
let published = observed.clone();
|
||||
engine.on_tool_result(move |result| {
|
||||
published
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(format!("publish:{}", result.tool_use_id));
|
||||
});
|
||||
|
||||
let committed = observed.clone();
|
||||
let mut annotate = move |item: &Item| {
|
||||
if let Item::ToolResult { call_id, .. } = item {
|
||||
committed.lock().unwrap().push(format!("commit:{call_id}"));
|
||||
}
|
||||
Ok(())
|
||||
};
|
||||
let mut history = History::new();
|
||||
let _ = engine
|
||||
.run_with_annotation(&mut history, "run both", &mut annotate)
|
||||
.await;
|
||||
observed.lock().unwrap().push("run-returned".to_string());
|
||||
|
||||
assert_eq!(
|
||||
observed.lock().unwrap().as_slice(),
|
||||
[
|
||||
"commit:call_fast",
|
||||
"publish:call_fast",
|
||||
"commit:call_slow",
|
||||
"publish:call_slow",
|
||||
"run-returned",
|
||||
]
|
||||
);
|
||||
|
||||
let committed_order: Vec<_> = history
|
||||
.iter()
|
||||
.filter_map(|entry| match &entry.item {
|
||||
Item::ToolResult { call_id, .. } => Some(call_id.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(committed_order, ["call_fast", "call_slow"]);
|
||||
|
||||
let requests = client_probe.requests();
|
||||
let projected_order: Vec<_> = requests[1]
|
||||
.items
|
||||
.iter()
|
||||
.filter_map(|item| match item {
|
||||
Item::ToolResult { call_id, .. } => Some(call_id.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(projected_order, ["call_slow", "call_fast"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_preserves_completed_results_and_resume_skips_them() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
vec![
|
||||
Event::tool_use_start(0, "call_hang", "hang_once"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::tool_use_start(1, "call_fast_a", "fast_a"),
|
||||
Event::tool_input_delta(1, r#"{}"#),
|
||||
Event::tool_use_stop(1),
|
||||
Event::tool_use_start(2, "call_fast_b", "fast_b"),
|
||||
Event::tool_input_delta(2, r#"{}"#),
|
||||
Event::tool_use_stop(2),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
Event::text_delta(0, "Recovered"),
|
||||
Event::text_block_stop(0, None),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let hanging = FirstAttemptHangsTool::new();
|
||||
let fast_a = SlowTool::new("fast_a", 1);
|
||||
let fast_b = SlowTool::new("fast_b", 2);
|
||||
engine.register_tool(hanging.definition());
|
||||
engine.register_tool(fast_a.definition());
|
||||
engine.register_tool(fast_b.definition());
|
||||
|
||||
let cancel = engine.cancel_sender();
|
||||
let cancel_task = tokio::spawn(async move {
|
||||
tokio::time::sleep(Duration::from_millis(30)).await;
|
||||
cancel.send(()).await.unwrap();
|
||||
});
|
||||
let mut history = History::new();
|
||||
let output = engine.run(&mut history, "start").await;
|
||||
let mut engine = output.engine;
|
||||
cancel_task.await.unwrap();
|
||||
|
||||
let completed_before_resume = history
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult { call_id, .. }
|
||||
if call_id == "call_fast_a" || call_id == "call_fast_b"
|
||||
)
|
||||
})
|
||||
.count();
|
||||
let unknown_before_resume = history
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition: ToolResultDisposition::OutcomeUnknown,
|
||||
..
|
||||
} if call_id == "call_hang"
|
||||
)
|
||||
})
|
||||
.count();
|
||||
assert_eq!(completed_before_resume, 2);
|
||||
assert_eq!(unknown_before_resume, 1);
|
||||
assert_eq!(fast_a.call_count(), 1);
|
||||
assert_eq!(fast_b.call_count(), 1);
|
||||
assert_eq!(hanging.call_count(), 1);
|
||||
|
||||
let _ = engine.resume(&mut history).await;
|
||||
|
||||
assert_eq!(
|
||||
fast_a.call_count(),
|
||||
1,
|
||||
"completed call must not be re-executed"
|
||||
);
|
||||
assert_eq!(
|
||||
fast_b.call_count(),
|
||||
1,
|
||||
"completed call must not be re-executed"
|
||||
);
|
||||
assert_eq!(
|
||||
hanging.call_count(),
|
||||
1,
|
||||
"OutcomeUnknown is terminal and must not be re-executed"
|
||||
);
|
||||
let completed_after_resume = history
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult { call_id, .. }
|
||||
if call_id == "call_fast_a" || call_id == "call_fast_b"
|
||||
)
|
||||
})
|
||||
.count();
|
||||
assert_eq!(completed_after_resume, 2);
|
||||
assert_eq!(
|
||||
history
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition: ToolResultDisposition::OutcomeUnknown,
|
||||
..
|
||||
} if call_id == "call_hang"
|
||||
)
|
||||
})
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cooperative_cancellation_commits_bounded_terminal_output() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
Event::tool_use_start(0, "call_cooperative", "cooperative"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]]);
|
||||
let mut engine = Engine::new(client);
|
||||
let tool = CooperativeCancelTool::new();
|
||||
engine.register_tool(tool.definition());
|
||||
let observed = Arc::new(Mutex::new(Vec::<&'static str>::new()));
|
||||
let published = observed.clone();
|
||||
engine.on_tool_result(move |_| published.lock().unwrap().push("published"));
|
||||
let committed = observed.clone();
|
||||
let mut annotate = move |item: &Item| {
|
||||
if matches!(item, Item::ToolResult { .. }) {
|
||||
committed.lock().unwrap().push("committed");
|
||||
}
|
||||
Ok(())
|
||||
};
|
||||
|
||||
let cancel = engine.cancel_sender();
|
||||
let cancel_task = tokio::spawn(async move {
|
||||
tokio::time::sleep(Duration::from_millis(30)).await;
|
||||
cancel.send(()).await.unwrap();
|
||||
});
|
||||
let mut history = History::new();
|
||||
let output = engine
|
||||
.run_with_annotation(&mut history, "start", &mut annotate)
|
||||
.await;
|
||||
observed.lock().unwrap().push("run-returned");
|
||||
cancel_task.await.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
observed.lock().unwrap().as_slice(),
|
||||
["committed", "published", "run-returned"]
|
||||
);
|
||||
assert_eq!(tool.calls.load(Ordering::SeqCst), 1);
|
||||
let terminal: Vec<_> = history
|
||||
.iter()
|
||||
.filter_map(|entry| match &entry.item {
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition,
|
||||
content,
|
||||
..
|
||||
} if call_id == "call_cooperative" => Some((*disposition, content.as_deref())),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(terminal.len(), 1);
|
||||
assert_eq!(terminal[0].0, ToolResultDisposition::Cancelled);
|
||||
assert_eq!(
|
||||
terminal[0].1,
|
||||
Some("stdout before cancellation\nstderr before cancellation")
|
||||
);
|
||||
assert!(matches!(
|
||||
output.result,
|
||||
agen::EngineRunExit::Interrupted(agen::StopReason::Cancelled)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pause_waits_for_started_tool_terminal_without_cancelling_provider() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
Event::tool_use_start(0, "call_safe_pause", "safe_pause"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]]);
|
||||
let mut engine = Engine::new(client);
|
||||
let tool = SafePauseTool::new();
|
||||
engine.register_tool(tool.definition());
|
||||
|
||||
let pause = engine.pause_sender();
|
||||
let calls = Arc::clone(&tool.calls);
|
||||
let release = Arc::clone(&tool.release);
|
||||
let control = tokio::spawn(async move {
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while calls.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("tool execution starts");
|
||||
pause.send(()).await.unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
release.notify_one();
|
||||
});
|
||||
|
||||
let started_at = std::time::Instant::now();
|
||||
let mut history = History::new();
|
||||
let output = engine.run(&mut history, "pause safely").await;
|
||||
control.await.unwrap();
|
||||
|
||||
assert!(started_at.elapsed() >= Duration::from_millis(50));
|
||||
assert_eq!(tool.calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(tool.cancellations.load(Ordering::SeqCst), 0);
|
||||
assert!(matches!(output.result, agen::EngineRunExit::Paused));
|
||||
assert!(history.iter().any(|entry| matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition: ToolResultDisposition::Success,
|
||||
..
|
||||
} if call_id == "call_safe_pause"
|
||||
)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pause_escalates_to_explicit_cancel_and_confirm_after_safe_boundary_deadline() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
Event::tool_use_start(0, "call_pause_cancel", "cooperative"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]]);
|
||||
let mut engine = Engine::new(client);
|
||||
engine.set_tool_execution_policy(ToolExecutionPolicy {
|
||||
pause_safe_boundary_timeout: Duration::from_millis(20),
|
||||
cancellation_request_timeout: Duration::from_millis(50),
|
||||
terminal_confirmation_timeout: Duration::from_millis(100),
|
||||
});
|
||||
let tool = CooperativeCancelTool::new();
|
||||
engine.register_tool(tool.definition());
|
||||
|
||||
let pause = engine.pause_sender();
|
||||
let calls = Arc::clone(&tool.calls);
|
||||
let control = tokio::spawn(async move {
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while calls.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("tool execution starts");
|
||||
pause.send(()).await.unwrap();
|
||||
});
|
||||
|
||||
let mut history = History::new();
|
||||
let output = engine.run(&mut history, "pause with escalation").await;
|
||||
control.await.unwrap();
|
||||
|
||||
assert!(matches!(output.result, agen::EngineRunExit::Paused));
|
||||
assert!(history.iter().any(|entry| matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition: ToolResultDisposition::Cancelled,
|
||||
..
|
||||
} if call_id == "call_pause_cancel"
|
||||
)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancellation_completion_race_commits_one_terminal_output() {
|
||||
for iteration in 0..24u64 {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
vec![
|
||||
Event::tool_use_start(0, "call_racy", "racy"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
],
|
||||
vec![Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
})],
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let delay = 2 + iteration % 3;
|
||||
let tool = SlowTool::new("racy", delay);
|
||||
engine.register_tool(tool.definition());
|
||||
let cancel = engine.cancel_sender();
|
||||
let cancel_task = tokio::spawn(async move {
|
||||
tokio::time::sleep(Duration::from_millis(delay)).await;
|
||||
let _ = cancel.send(()).await;
|
||||
});
|
||||
|
||||
let mut history = History::new();
|
||||
let _ = engine.run(&mut history, "race").await;
|
||||
cancel_task.await.unwrap();
|
||||
let terminal_count = history
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult { call_id, .. } if call_id == "call_racy"
|
||||
)
|
||||
})
|
||||
.count();
|
||||
assert_eq!(terminal_count, 1, "iteration {iteration}");
|
||||
assert_eq!(tool.call_count(), 1, "iteration {iteration}");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tool_result_commit_failure_prevents_publication() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
Event::tool_use_start(0, "call_fast", "fast"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]]);
|
||||
let mut engine = Engine::new(client);
|
||||
engine.register_tool(SlowTool::new("fast", 1).definition());
|
||||
|
||||
let published = Arc::new(AtomicUsize::new(0));
|
||||
let published_probe = published.clone();
|
||||
engine.on_tool_result(move |_| {
|
||||
published_probe.fetch_add(1, Ordering::SeqCst);
|
||||
});
|
||||
|
||||
let mut history = History::new();
|
||||
let mut reject_tool_result = |item: &Item| {
|
||||
if matches!(item, Item::ToolResult { .. }) {
|
||||
Err("session log unavailable".to_string())
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
};
|
||||
let _ = engine
|
||||
.run_with_annotation(&mut history, "start", &mut reject_tool_result)
|
||||
.await;
|
||||
|
||||
assert_eq!(published.load(Ordering::SeqCst), 0);
|
||||
assert!(
|
||||
history
|
||||
.iter()
|
||||
.all(|entry| !matches!(entry.item, Item::ToolResult { .. }))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_tool_execution_context_order_and_batch_id() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
@@ -205,13 +789,14 @@ async fn test_tool_execution_context_order_and_batch_id() {
|
||||
],
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
let contexts = Arc::new(Mutex::new(Vec::new()));
|
||||
|
||||
engine.register_tool(ContextRecordingTool::new("record_a", contexts.clone()).definition());
|
||||
engine.register_tool(ContextRecordingTool::new("record_b", contexts.clone()).definition());
|
||||
engine.register_tool(ContextRecordingTool::new("record_c", contexts.clone()).definition());
|
||||
|
||||
let _ = engine.run("record contexts").await;
|
||||
let _ = engine.run(&mut history, "record contexts").await;
|
||||
|
||||
let mut contexts = contexts.lock().unwrap().clone();
|
||||
contexts.sort_by_key(|ctx| ctx.call_index);
|
||||
@@ -256,11 +841,12 @@ async fn test_tool_execution_context_batch_id_changes_between_batches() {
|
||||
],
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
let contexts = Arc::new(Mutex::new(Vec::new()));
|
||||
|
||||
engine.register_tool(ContextRecordingTool::new("record", contexts.clone()).definition());
|
||||
|
||||
let _ = engine.run("record batches").await;
|
||||
let _ = engine.run(&mut history, "record batches").await;
|
||||
|
||||
let contexts = contexts.lock().unwrap().clone();
|
||||
assert_eq!(contexts.len(), 2);
|
||||
@@ -298,6 +884,7 @@ async fn test_tool_execution_context_for_skipped_and_synthetic_paths() {
|
||||
],
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
let executed_contexts = Arc::new(Mutex::new(Vec::new()));
|
||||
let pre_contexts = Arc::new(Mutex::new(Vec::new()));
|
||||
let post_contexts = Arc::new(Mutex::new(Vec::new()));
|
||||
@@ -344,7 +931,9 @@ async fn test_tool_execution_context_for_skipped_and_synthetic_paths() {
|
||||
post_contexts: post_contexts.clone(),
|
||||
});
|
||||
|
||||
let _ = engine.run("record skipped and synthetic contexts").await;
|
||||
let _ = engine
|
||||
.run(&mut history, "record skipped and synthetic contexts")
|
||||
.await;
|
||||
|
||||
let mut pre_contexts = pre_contexts.lock().unwrap().clone();
|
||||
pre_contexts.sort_by_key(|ctx| ctx.call_index);
|
||||
@@ -389,6 +978,7 @@ async fn test_before_tool_call_skip() {
|
||||
|
||||
let client = MockLlmClient::new(events);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
let allowed_tool = SlowTool::new("allowed_tool", 10);
|
||||
let blocked_tool = SlowTool::new("blocked_tool", 10);
|
||||
@@ -416,7 +1006,7 @@ async fn test_before_tool_call_skip() {
|
||||
engine.set_interceptor(BlockingPolicy);
|
||||
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let _result = engine.run("Test hook").await;
|
||||
let _result = engine.run(&mut history, "Test hook").await;
|
||||
|
||||
// allowed_tool is called, but blocked_tool is not
|
||||
assert_eq!(
|
||||
@@ -457,6 +1047,7 @@ async fn test_post_tool_call_modification() {
|
||||
]);
|
||||
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
|
||||
#[derive(Clone)]
|
||||
struct SimpleTool;
|
||||
@@ -503,9 +1094,12 @@ async fn test_post_tool_call_modification() {
|
||||
});
|
||||
|
||||
// Mutable::run consumes self, returns (Locked, EngineResult)
|
||||
let result = engine.run("Test modification").await;
|
||||
let result = engine.run(&mut history, "Test modification").await;
|
||||
|
||||
assert!(result.is_ok(), "Engine should complete");
|
||||
assert!(
|
||||
matches!(result.result, agen::EngineRunExit::Finished),
|
||||
"Engine should complete"
|
||||
);
|
||||
|
||||
// Verify hook was called and content was modified
|
||||
let content = modified_content.lock().unwrap().clone();
|
||||
@@ -540,6 +1134,7 @@ async fn test_before_tool_call_synthetic_result_committed() {
|
||||
],
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
let blocked_tool = SlowTool::new("blocked_tool", 10);
|
||||
let blocked_clone = blocked_tool.clone();
|
||||
engine.register_tool(blocked_tool.definition());
|
||||
@@ -558,10 +1153,10 @@ async fn test_before_tool_call_synthetic_result_committed() {
|
||||
|
||||
engine.set_interceptor(SyntheticPolicy);
|
||||
|
||||
let result = engine.run("Test synthetic result").await.unwrap();
|
||||
let _result = engine.run(&mut history, "Test synthetic result").await;
|
||||
|
||||
assert_eq!(blocked_clone.call_count(), 0, "Blocked tool should not run");
|
||||
assert!(result.engine.history().iter().any(|item| matches!(
|
||||
assert!(history.items().any(|item| matches!(
|
||||
item,
|
||||
agen::Item::ToolResult {
|
||||
call_id,
|
||||
@@ -571,3 +1166,76 @@ async fn test_before_tool_call_synthetic_result_committed() {
|
||||
} if call_id == "call_1" && summary == "permission denied"
|
||||
)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn post_tool_abort_commits_confirmed_result_before_stopping_run() {
|
||||
let client = MockLlmClient::new(vec![
|
||||
Event::tool_use_start(0, "call_confirmed", "confirmed"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]);
|
||||
let mut engine = Engine::new(client);
|
||||
let tool = SlowTool::new("confirmed", 1);
|
||||
engine.register_tool(tool.definition());
|
||||
|
||||
struct AbortAfterResult;
|
||||
#[async_trait]
|
||||
impl Interceptor for AbortAfterResult {
|
||||
async fn post_tool_call(&self, _info: &mut ToolResultInfo) -> PostToolAction {
|
||||
PostToolAction::Abort("policy stopped the run".to_string())
|
||||
}
|
||||
}
|
||||
engine.set_interceptor(AbortAfterResult);
|
||||
|
||||
let observed = Arc::new(Mutex::new(Vec::<&'static str>::new()));
|
||||
let published = observed.clone();
|
||||
engine.on_tool_result(move |_| published.lock().unwrap().push("published"));
|
||||
let committed = observed.clone();
|
||||
let mut annotate = move |item: &Item| {
|
||||
if matches!(item, Item::ToolResult { .. }) {
|
||||
committed.lock().unwrap().push("committed");
|
||||
}
|
||||
Ok(())
|
||||
};
|
||||
|
||||
let mut history = History::new();
|
||||
let output = engine
|
||||
.run_with_annotation(&mut history, "run confirmed tool", &mut annotate)
|
||||
.await;
|
||||
observed.lock().unwrap().push("run-returned");
|
||||
|
||||
assert_eq!(tool.call_count(), 1);
|
||||
assert_eq!(
|
||||
observed.lock().unwrap().as_slice(),
|
||||
["committed", "published", "run-returned"]
|
||||
);
|
||||
assert!(matches!(
|
||||
output.result,
|
||||
agen::EngineRunExit::Interrupted(agen::StopReason::Unexpected(
|
||||
agen::EngineError::Aborted(ref reason)
|
||||
)) if reason == "policy stopped the run"
|
||||
));
|
||||
let terminal: Vec<_> = history
|
||||
.iter()
|
||||
.filter_map(|entry| match &entry.item {
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition,
|
||||
..
|
||||
} if call_id == "call_confirmed" => Some(*disposition),
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(terminal, [ToolResultDisposition::Success]);
|
||||
assert!(!history.iter().any(|entry| matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition: ToolResultDisposition::OutcomeUnknown,
|
||||
..
|
||||
} if call_id == "call_confirmed"
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -13,12 +13,12 @@
|
||||
|
||||
mod common;
|
||||
|
||||
use agen::Engine;
|
||||
use agen::Item;
|
||||
use agen::llm_client::event::{
|
||||
BlockMetadata, BlockStart, BlockStop, BlockType, Event, ReasoningBlockData, ResponseStatus,
|
||||
StatusEvent,
|
||||
};
|
||||
use agen::{Engine, History};
|
||||
use common::MockLlmClient;
|
||||
|
||||
fn reasoning_block(text: impl Into<String>, data: ReasoningBlockData) -> Vec<Event> {
|
||||
@@ -65,15 +65,15 @@ async fn anthropic_thinking_round_trips_signature_into_history() {
|
||||
]);
|
||||
let client = MockLlmClient::new(events);
|
||||
let engine = Engine::new(client);
|
||||
let out = engine.run("question?").await.expect("run ok");
|
||||
let engine = out.engine;
|
||||
let mut history: History = History::new();
|
||||
let _out = engine.run(&mut history, "question?").await;
|
||||
|
||||
let history = engine.history();
|
||||
let entries = history.entries();
|
||||
// user / reasoning / assistant_message
|
||||
assert_eq!(history.len(), 3, "history: {history:?}");
|
||||
|
||||
assert!(matches!(history[0], Item::Message { .. }));
|
||||
match &history[1] {
|
||||
assert!(matches!(entries[0].item, Item::Message { .. }));
|
||||
match &entries[1].item {
|
||||
Item::Reasoning {
|
||||
text, signature, ..
|
||||
} => {
|
||||
@@ -82,7 +82,7 @@ async fn anthropic_thinking_round_trips_signature_into_history() {
|
||||
}
|
||||
other => panic!("expected Reasoning, got {other:?}"),
|
||||
}
|
||||
assert_eq!(history[2].as_text(), Some("Here's the answer"));
|
||||
assert_eq!(entries[2].item.as_text(), Some("Here's the answer"));
|
||||
}
|
||||
|
||||
/// OpenAI Responses 風: encrypted_content + summary を持った reasoning が
|
||||
@@ -109,11 +109,11 @@ async fn openai_reasoning_round_trips_encrypted_and_summary() {
|
||||
]);
|
||||
let client = MockLlmClient::new(events);
|
||||
let engine = Engine::new(client);
|
||||
let out = engine.run("q").await.expect("run ok");
|
||||
let engine = out.engine;
|
||||
let mut history: History = History::new();
|
||||
let _out = engine.run(&mut history, "q").await;
|
||||
|
||||
let history = engine.history();
|
||||
match &history[1] {
|
||||
let entries = history.entries();
|
||||
match &entries[1].item {
|
||||
Item::Reasoning {
|
||||
text,
|
||||
summary,
|
||||
@@ -155,13 +155,13 @@ async fn reasoning_precedes_text_in_assistant_burst() {
|
||||
}));
|
||||
let client = MockLlmClient::new(events);
|
||||
let engine = Engine::new(client);
|
||||
let out = engine.run("q").await.expect("run ok");
|
||||
let engine = out.engine;
|
||||
let mut history: History = History::new();
|
||||
let _out = engine.run(&mut history, "q").await;
|
||||
|
||||
let history = engine.history();
|
||||
let entries = history.entries();
|
||||
// user / reasoning(先頭) / assistant_message
|
||||
assert!(matches!(history[1], Item::Reasoning { .. }));
|
||||
assert_eq!(history[2].as_text(), Some("intermediate"));
|
||||
assert!(matches!(entries[1].item, Item::Reasoning { .. }));
|
||||
assert_eq!(entries[2].item.as_text(), Some("intermediate"));
|
||||
}
|
||||
|
||||
/// resume シナリオ: history.json 由来の Item::Reasoning(signature) を Engine に
|
||||
@@ -207,14 +207,18 @@ async fn injected_reasoning_survives_into_outgoing_request() {
|
||||
};
|
||||
|
||||
let mut engine = Engine::new(client);
|
||||
let mut history: History = History::new();
|
||||
// resume: 既存 history を流し込む
|
||||
engine.set_history(vec![
|
||||
Item::user_message("prior question"),
|
||||
Item::reasoning("prior thinking").with_signature("SIG-PRIOR"),
|
||||
Item::assistant_message("prior answer"),
|
||||
]);
|
||||
engine.set_history(
|
||||
&mut history,
|
||||
vec![
|
||||
Item::user_message("prior question"),
|
||||
Item::reasoning("prior thinking").with_signature("SIG-PRIOR"),
|
||||
Item::assistant_message("prior answer"),
|
||||
],
|
||||
);
|
||||
|
||||
let _ = engine.run("follow up").await.expect("run ok");
|
||||
let _ = engine.run(&mut history, "follow up").await;
|
||||
|
||||
let req = captured
|
||||
.lock()
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use agen::Engine;
|
||||
use agen::{Engine, History};
|
||||
use agen::llm_client::capability::{
|
||||
CacheStrategy, ModelCapability, StructuredOutput, ToolCallingSupport,
|
||||
};
|
||||
@@ -22,7 +22,8 @@ fn main() {
|
||||
cap,
|
||||
);
|
||||
let engine = Engine::new(client);
|
||||
let mut locked = engine.lock();
|
||||
let history = History::new();
|
||||
let mut locked = engine.lock(&history);
|
||||
let def: agen::tool::ToolDefinition = Arc::new(|| panic!("unused"));
|
||||
let _ = locked.register_tool(def);
|
||||
}
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
error[E0599]: no method named `register_tool` found for struct `Engine<HttpTransport<AnthropicScheme>, Locked>` in the current scope
|
||||
--> tests/ui/locked_register_tool.rs:27:20
|
||||
--> tests/ui/locked_register_tool.rs:28:20
|
||||
|
|
||||
27 | let _ = locked.register_tool(def);
|
||||
28 | let _ = locked.register_tool(def);
|
||||
| ^^^^^^^^^^^^^ method not found in `Engine<HttpTransport<AnthropicScheme>, Locked>`
|
||||
|
|
||||
= note: the method was found for
|
||||
- `Engine<C>`
|
||||
- `Engine<C, Mutable, A>`
|
||||
|
||||
@@ -5,15 +5,15 @@ edition.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
chrono = { version = "0.4", default-features = false, features = ["clock"] }
|
||||
protocol = { workspace = true }
|
||||
manifest = { workspace = true }
|
||||
ticket = { workspace = true }
|
||||
futures = { workspace = true }
|
||||
reqwest = { version = "0.13", default-features = false, features = ["blocking", "json", "native-tls"] }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
thiserror = { workspace = true }
|
||||
tokio = { workspace = true, features = ["rt", "macros", "net", "io-util", "sync", "time", "process", "fs"] }
|
||||
tokio = { workspace = true, features = ["rt", "macros", "net", "io-util", "sync", "time"] }
|
||||
tokio-tungstenite = { workspace = true }
|
||||
uuid = { workspace = true }
|
||||
workspace-api.workspace = true
|
||||
|
||||
@@ -0,0 +1,768 @@
|
||||
use chrono::{DateTime, Utc};
|
||||
use reqwest::{Method, StatusCode, Url, redirect};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::BTreeMap;
|
||||
use std::env;
|
||||
use std::fmt;
|
||||
use std::fs::{self, OpenOptions};
|
||||
use std::io::Write as _;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
const TOKEN_FILE_NAME: &str = "backend-tokens.json";
|
||||
const MAX_REDIRECTS: usize = 10;
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
|
||||
pub struct BackendOrigin(String);
|
||||
|
||||
impl BackendOrigin {
|
||||
pub fn parse(input: &str) -> Result<Self, BackendApiClientError> {
|
||||
let url = Url::parse(input.trim()).map_err(|error| {
|
||||
BackendApiClientError::InvalidBackendOrigin(format!(
|
||||
"Backend URL is not a valid absolute URL: {error}"
|
||||
))
|
||||
})?;
|
||||
if !url.path().bytes().all(|byte| byte == b'/')
|
||||
|| url.query().is_some()
|
||||
|| url.fragment().is_some()
|
||||
{
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL must contain only an origin, without a path, query, or fragment"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
Self::from_url(url)
|
||||
}
|
||||
|
||||
fn from_url(mut url: Url) -> Result<Self, BackendApiClientError> {
|
||||
if !matches!(url.scheme(), "http" | "https") {
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL scheme must be http or https".to_string(),
|
||||
));
|
||||
}
|
||||
if !url.username().is_empty() || url.password().is_some() {
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL must not contain user information".to_string(),
|
||||
));
|
||||
}
|
||||
if url.host().is_none() {
|
||||
return Err(BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL must contain a host".to_string(),
|
||||
));
|
||||
}
|
||||
let default_port = match url.scheme() {
|
||||
"http" => 80,
|
||||
"https" => 443,
|
||||
_ => unreachable!("validated Backend URL scheme"),
|
||||
};
|
||||
if url.port() == Some(default_port) {
|
||||
url.set_port(None).map_err(|()| {
|
||||
BackendApiClientError::InvalidBackendOrigin(
|
||||
"Backend URL contains an invalid port".to_string(),
|
||||
)
|
||||
})?;
|
||||
}
|
||||
url.set_path("");
|
||||
url.set_query(None);
|
||||
url.set_fragment(None);
|
||||
let normalized = url.as_str().trim_end_matches('/').to_string();
|
||||
Ok(Self(normalized))
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
|
||||
fn url(&self, path_and_query: &str) -> Result<Url, BackendApiClientError> {
|
||||
if !path_and_query.starts_with('/') || path_and_query.starts_with("//") {
|
||||
return Err(BackendApiClientError::InvalidRequestPath(
|
||||
"Backend API request path must start with one `/`".to_string(),
|
||||
));
|
||||
}
|
||||
Url::parse(&format!("{}{path_and_query}", self.0)).map_err(|error| {
|
||||
BackendApiClientError::InvalidRequestPath(format!(
|
||||
"Backend API request path is invalid: {error}"
|
||||
))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for BackendOrigin {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_tuple("BackendOrigin").field(&self.0).finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for BackendOrigin {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str(&self.0)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct BackendAccessToken(String);
|
||||
|
||||
impl fmt::Debug for BackendAccessToken {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.write_str("BackendAccessToken([REDACTED])")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct BackendApiClient {
|
||||
origin: BackendOrigin,
|
||||
access_token: BackendAccessToken,
|
||||
asynchronous: reqwest::Client,
|
||||
}
|
||||
|
||||
impl fmt::Debug for BackendApiClient {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("BackendApiClient")
|
||||
.field("origin", &self.origin)
|
||||
.field("access_token", &self.access_token)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl BackendApiClient {
|
||||
pub fn from_stored_token(base_url: &str) -> Result<Self, BackendApiClientError> {
|
||||
let path = backend_token_file_path()?;
|
||||
Self::from_token_file(base_url, &path)
|
||||
}
|
||||
|
||||
fn from_token_file(base_url: &str, path: &Path) -> Result<Self, BackendApiClientError> {
|
||||
let origin = BackendOrigin::parse(base_url)?;
|
||||
let token_file = read_token_file(path)?;
|
||||
let entry = token_file.tokens.get(origin.as_str()).ok_or_else(|| {
|
||||
BackendApiClientError::TokenEntryMissing {
|
||||
origin: origin.clone(),
|
||||
path: path.to_path_buf(),
|
||||
}
|
||||
})?;
|
||||
validate_token_entry(entry, &origin, path)?;
|
||||
Self::new(origin, BackendAccessToken(entry.access_token.clone()))
|
||||
}
|
||||
|
||||
fn new(
|
||||
origin: BackendOrigin,
|
||||
access_token: BackendAccessToken,
|
||||
) -> Result<Self, BackendApiClientError> {
|
||||
let asynchronous = reqwest::Client::builder()
|
||||
.redirect(redirect_policy(origin.clone()))
|
||||
.build()
|
||||
.map_err(BackendApiClientError::Http)?;
|
||||
Ok(Self {
|
||||
origin,
|
||||
access_token,
|
||||
asynchronous,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn origin(&self) -> &BackendOrigin {
|
||||
&self.origin
|
||||
}
|
||||
|
||||
pub fn request(
|
||||
&self,
|
||||
method: Method,
|
||||
path_and_query: &str,
|
||||
) -> Result<reqwest::RequestBuilder, BackendApiClientError> {
|
||||
let url = self.origin.url(path_and_query)?;
|
||||
Ok(self
|
||||
.asynchronous
|
||||
.request(method, url)
|
||||
.bearer_auth(&self.access_token.0))
|
||||
}
|
||||
|
||||
pub fn blocking_request(
|
||||
&self,
|
||||
method: Method,
|
||||
path_and_query: &str,
|
||||
) -> Result<reqwest::blocking::RequestBuilder, BackendApiClientError> {
|
||||
let url = self.origin.url(path_and_query)?;
|
||||
let client = reqwest::blocking::Client::builder()
|
||||
.redirect(redirect_policy(self.origin.clone()))
|
||||
.build()
|
||||
.map_err(BackendApiClientError::Http)?;
|
||||
Ok(client
|
||||
.request(method, url)
|
||||
.bearer_auth(&self.access_token.0))
|
||||
}
|
||||
|
||||
pub(crate) fn authorization_header_value(&self) -> String {
|
||||
format!("Bearer {}", self.access_token.0)
|
||||
}
|
||||
|
||||
pub fn check_status(&self, status: StatusCode) -> Result<(), BackendApiClientError> {
|
||||
match status {
|
||||
StatusCode::UNAUTHORIZED => Err(BackendApiClientError::Unauthorized {
|
||||
origin: self.origin.clone(),
|
||||
}),
|
||||
StatusCode::FORBIDDEN => Err(BackendApiClientError::Forbidden {
|
||||
origin: self.origin.clone(),
|
||||
}),
|
||||
status if !status.is_success() => Err(BackendApiClientError::BackendStatus {
|
||||
origin: self.origin.clone(),
|
||||
status: status.as_u16(),
|
||||
}),
|
||||
_ => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn from_access_token_for_test(
|
||||
base_url: &str,
|
||||
access_token: &str,
|
||||
) -> Result<Self, BackendApiClientError> {
|
||||
Self::new(
|
||||
BackendOrigin::parse(base_url)?,
|
||||
BackendAccessToken(access_token.to_string()),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn redirect_policy(origin: BackendOrigin) -> redirect::Policy {
|
||||
redirect::Policy::custom(move |attempt| {
|
||||
if attempt.previous().len() >= MAX_REDIRECTS {
|
||||
return attempt.error("Backend request exceeded the redirect limit");
|
||||
}
|
||||
match BackendOrigin::from_url(attempt.url().clone()) {
|
||||
Ok(target_origin) if target_origin == origin => attempt.follow(),
|
||||
Ok(target_origin) => attempt.error(format!(
|
||||
"Backend request refused a cross-origin redirect from {origin} to {target_origin}"
|
||||
)),
|
||||
Err(error) => attempt.error(error.to_string()),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum BackendApiClientError {
|
||||
InvalidBackendOrigin(String),
|
||||
InvalidRequestPath(String),
|
||||
ConfigDirectoryUnavailable,
|
||||
TokenFileMissing {
|
||||
path: PathBuf,
|
||||
},
|
||||
TokenFileMalformed {
|
||||
path: PathBuf,
|
||||
message: String,
|
||||
},
|
||||
TokenEntryMissing {
|
||||
origin: BackendOrigin,
|
||||
path: PathBuf,
|
||||
},
|
||||
TokenExpired {
|
||||
origin: BackendOrigin,
|
||||
expired_at: String,
|
||||
},
|
||||
Http(reqwest::Error),
|
||||
Unauthorized {
|
||||
origin: BackendOrigin,
|
||||
},
|
||||
Forbidden {
|
||||
origin: BackendOrigin,
|
||||
},
|
||||
BackendStatus {
|
||||
origin: BackendOrigin,
|
||||
status: u16,
|
||||
},
|
||||
Io {
|
||||
path: PathBuf,
|
||||
source: std::io::Error,
|
||||
},
|
||||
}
|
||||
|
||||
impl fmt::Display for BackendApiClientError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::InvalidBackendOrigin(message) | Self::InvalidRequestPath(message) => {
|
||||
f.write_str(message)
|
||||
}
|
||||
Self::ConfigDirectoryUnavailable => f.write_str(
|
||||
"cannot locate the client configuration directory for backend-tokens.json",
|
||||
),
|
||||
Self::TokenFileMissing { path } => write!(
|
||||
f,
|
||||
"Backend token file {} is missing; run `yoi login --backend <BACKEND>` first",
|
||||
path.display()
|
||||
),
|
||||
Self::TokenFileMalformed { path, message } => write!(
|
||||
f,
|
||||
"Backend token file {} is malformed: {message}; run `yoi login --backend <BACKEND>` again",
|
||||
path.display()
|
||||
),
|
||||
Self::TokenEntryMissing { origin, path } => write!(
|
||||
f,
|
||||
"no Backend token for {origin} exists in {}; login URLs are matched by normalized origin, so run `yoi login --backend {origin}`",
|
||||
path.display()
|
||||
),
|
||||
Self::TokenExpired { origin, expired_at } => write!(
|
||||
f,
|
||||
"Backend token for {origin} expired at {expired_at}; run `yoi login --backend {origin}` again"
|
||||
),
|
||||
Self::Http(error) => write!(f, "Backend request failed: {error}"),
|
||||
Self::Unauthorized { origin } => write!(
|
||||
f,
|
||||
"Backend {origin} returned HTTP 401 for the saved token; it may be expired or revoked, so run `yoi login --backend {origin}` again"
|
||||
),
|
||||
Self::Forbidden { origin } => write!(
|
||||
f,
|
||||
"Backend {origin} returned HTTP 403; the saved token is authenticated but is not authorized for this operation"
|
||||
),
|
||||
Self::BackendStatus { origin, status } => {
|
||||
write!(f, "Backend {origin} returned HTTP {status}")
|
||||
}
|
||||
Self::Io { path, source } => {
|
||||
write!(f, "failed to access {}: {source}", path.display())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for BackendApiClientError {
|
||||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||||
match self {
|
||||
Self::Http(error) => Some(error),
|
||||
Self::Io { source, .. } => Some(source),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize, Serialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct BackendTokenFile {
|
||||
tokens: BTreeMap<String, BackendTokenEntry>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize, Serialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct BackendTokenEntry {
|
||||
token_type: String,
|
||||
access_token: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
expires_at: Option<String>,
|
||||
}
|
||||
|
||||
pub fn save_backend_token(
|
||||
base_url: &str,
|
||||
token_type: &str,
|
||||
access_token: &str,
|
||||
) -> Result<PathBuf, BackendApiClientError> {
|
||||
save_backend_token_with_expiry(base_url, token_type, access_token, None)
|
||||
}
|
||||
|
||||
fn save_backend_token_with_expiry(
|
||||
base_url: &str,
|
||||
token_type: &str,
|
||||
access_token: &str,
|
||||
expires_at: Option<String>,
|
||||
) -> Result<PathBuf, BackendApiClientError> {
|
||||
let path = backend_token_file_path()?;
|
||||
save_backend_token_to_file(base_url, token_type, access_token, expires_at, &path)?;
|
||||
Ok(path)
|
||||
}
|
||||
|
||||
fn save_backend_token_to_file(
|
||||
base_url: &str,
|
||||
token_type: &str,
|
||||
access_token: &str,
|
||||
expires_at: Option<String>,
|
||||
path: &Path,
|
||||
) -> Result<(), BackendApiClientError> {
|
||||
let origin = BackendOrigin::parse(base_url)?;
|
||||
let mut token_file = if path.exists() {
|
||||
read_token_file(&path)?
|
||||
} else {
|
||||
BackendTokenFile {
|
||||
tokens: BTreeMap::new(),
|
||||
}
|
||||
};
|
||||
let entry = BackendTokenEntry {
|
||||
token_type: token_type.to_string(),
|
||||
access_token: access_token.to_string(),
|
||||
expires_at,
|
||||
};
|
||||
validate_token_entry(&entry, &origin, path)?;
|
||||
token_file.tokens.insert(origin.to_string(), entry);
|
||||
write_token_file(path, &token_file)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn backend_token_file_path() -> Result<PathBuf, BackendApiClientError> {
|
||||
if let Some(config_home) = env::var_os("XDG_CONFIG_HOME") {
|
||||
return Ok(PathBuf::from(config_home).join("yoi").join(TOKEN_FILE_NAME));
|
||||
}
|
||||
let Some(home) = env::var_os("HOME") else {
|
||||
return Err(BackendApiClientError::ConfigDirectoryUnavailable);
|
||||
};
|
||||
Ok(PathBuf::from(home)
|
||||
.join(".config")
|
||||
.join("yoi")
|
||||
.join(TOKEN_FILE_NAME))
|
||||
}
|
||||
|
||||
fn read_token_file(path: &Path) -> Result<BackendTokenFile, BackendApiClientError> {
|
||||
let bytes = fs::read(path).map_err(|source| {
|
||||
if source.kind() == std::io::ErrorKind::NotFound {
|
||||
BackendApiClientError::TokenFileMissing {
|
||||
path: path.to_path_buf(),
|
||||
}
|
||||
} else {
|
||||
BackendApiClientError::Io {
|
||||
path: path.to_path_buf(),
|
||||
source,
|
||||
}
|
||||
}
|
||||
})?;
|
||||
let raw: BackendTokenFile = serde_json::from_slice(&bytes).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: error.to_string(),
|
||||
}
|
||||
})?;
|
||||
normalize_token_file(raw, path)
|
||||
}
|
||||
|
||||
fn normalize_token_file(
|
||||
token_file: BackendTokenFile,
|
||||
path: &Path,
|
||||
) -> Result<BackendTokenFile, BackendApiClientError> {
|
||||
let mut normalized = BTreeMap::new();
|
||||
for (raw_origin, entry) in token_file.tokens {
|
||||
let origin = BackendOrigin::parse(&raw_origin).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token key `{raw_origin}` is invalid: {error}"),
|
||||
}
|
||||
})?;
|
||||
if normalized.insert(origin.to_string(), entry).is_some() {
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("more than one token entry normalizes to `{origin}`"),
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(BackendTokenFile { tokens: normalized })
|
||||
}
|
||||
|
||||
fn validate_token_entry(
|
||||
entry: &BackendTokenEntry,
|
||||
origin: &BackendOrigin,
|
||||
path: &Path,
|
||||
) -> Result<(), BackendApiClientError> {
|
||||
if !entry.token_type.eq_ignore_ascii_case("Bearer") {
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` does not use the Bearer token type"),
|
||||
});
|
||||
}
|
||||
if entry.access_token.trim().is_empty()
|
||||
|| entry.access_token.contains('\r')
|
||||
|| entry.access_token.contains('\n')
|
||||
{
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` is empty or contains an invalid line break"),
|
||||
});
|
||||
}
|
||||
if reqwest::header::HeaderValue::from_str(&format!("Bearer {}", entry.access_token)).is_err() {
|
||||
return Err(BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` cannot be represented as an HTTP header"),
|
||||
});
|
||||
}
|
||||
if let Some(expires_at) = entry.expires_at.as_deref() {
|
||||
let expiration = DateTime::parse_from_rfc3339(expires_at).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: format!("token for `{origin}` has invalid expires_at: {error}"),
|
||||
}
|
||||
})?;
|
||||
if expiration <= Utc::now() {
|
||||
return Err(BackendApiClientError::TokenExpired {
|
||||
origin: origin.clone(),
|
||||
expired_at: expires_at.to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn write_token_file(
|
||||
path: &Path,
|
||||
token_file: &BackendTokenFile,
|
||||
) -> Result<(), BackendApiClientError> {
|
||||
let parent = path
|
||||
.parent()
|
||||
.ok_or(BackendApiClientError::ConfigDirectoryUnavailable)?;
|
||||
fs::create_dir_all(parent).map_err(|source| BackendApiClientError::Io {
|
||||
path: parent.to_path_buf(),
|
||||
source,
|
||||
})?;
|
||||
let payload = serde_json::to_vec_pretty(token_file).map_err(|error| {
|
||||
BackendApiClientError::TokenFileMalformed {
|
||||
path: path.to_path_buf(),
|
||||
message: error.to_string(),
|
||||
}
|
||||
})?;
|
||||
let temp_path = parent.join(format!(".{TOKEN_FILE_NAME}.tmp-{}", std::process::id()));
|
||||
let mut options = OpenOptions::new();
|
||||
options.write(true).create(true).truncate(true);
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::OpenOptionsExt;
|
||||
options.mode(0o600);
|
||||
}
|
||||
let mut file = options
|
||||
.open(&temp_path)
|
||||
.map_err(|source| BackendApiClientError::Io {
|
||||
path: temp_path.clone(),
|
||||
source,
|
||||
})?;
|
||||
file.write_all(&payload)
|
||||
.and_then(|()| file.write_all(b"\n"))
|
||||
.and_then(|()| file.sync_all())
|
||||
.map_err(|source| BackendApiClientError::Io {
|
||||
path: temp_path.clone(),
|
||||
source,
|
||||
})?;
|
||||
fs::rename(&temp_path, path).map_err(|source| BackendApiClientError::Io {
|
||||
path: path.to_path_buf(),
|
||||
source,
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::io::{Read, Write};
|
||||
use std::net::TcpListener;
|
||||
use std::thread;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
fn temp_path(label: &str) -> PathBuf {
|
||||
let nonce = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_nanos();
|
||||
env::temp_dir().join(format!(
|
||||
"yoi-client-{label}-{}-{nonce}.json",
|
||||
std::process::id()
|
||||
))
|
||||
}
|
||||
|
||||
fn write_fixture(path: &Path, value: serde_json::Value) {
|
||||
fs::write(path, serde_json::to_vec(&value).unwrap()).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_origin_normalizes_safe_equivalents() {
|
||||
let variants = [
|
||||
"HTTP://Example.COM",
|
||||
"http://example.com/",
|
||||
"http://example.com:80////",
|
||||
];
|
||||
for variant in variants {
|
||||
assert_eq!(
|
||||
BackendOrigin::parse(variant).unwrap().as_str(),
|
||||
"http://example.com"
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
BackendOrigin::parse("https://EXAMPLE.com:443/")
|
||||
.unwrap()
|
||||
.as_str(),
|
||||
"https://example.com"
|
||||
);
|
||||
assert_eq!(
|
||||
BackendOrigin::parse("https://example.com:8443/")
|
||||
.unwrap()
|
||||
.as_str(),
|
||||
"https://example.com:8443"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_origin_rejects_unsafe_authority_changes() {
|
||||
for invalid in [
|
||||
"ftp://example.com",
|
||||
"https://user@example.com",
|
||||
"https://example.com/api",
|
||||
"https://example.com/?query=1",
|
||||
"https://example.com/#fragment",
|
||||
] {
|
||||
assert!(BackendOrigin::parse(invalid).is_err(), "accepted {invalid}");
|
||||
}
|
||||
assert_ne!(
|
||||
BackendOrigin::parse("http://localhost:8787").unwrap(),
|
||||
BackendOrigin::parse("http://127.0.0.1:8787").unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn token_lookup_distinguishes_missing_malformed_mismatch_and_expired() {
|
||||
let missing = temp_path("missing");
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://localhost:8787", &missing),
|
||||
Err(BackendApiClientError::TokenFileMissing { .. })
|
||||
));
|
||||
|
||||
let malformed = temp_path("malformed");
|
||||
fs::write(&malformed, b"not json").unwrap();
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://localhost:8787", &malformed),
|
||||
Err(BackendApiClientError::TokenFileMalformed { .. })
|
||||
));
|
||||
|
||||
let mismatch = temp_path("mismatch");
|
||||
write_fixture(
|
||||
&mismatch,
|
||||
serde_json::json!({"tokens": {"http://localhost:8787": {
|
||||
"token_type": "Bearer", "access_token": "secret"
|
||||
}}}),
|
||||
);
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://127.0.0.1:8787", &mismatch),
|
||||
Err(BackendApiClientError::TokenEntryMissing { .. })
|
||||
));
|
||||
|
||||
let expired = temp_path("expired");
|
||||
write_fixture(
|
||||
&expired,
|
||||
serde_json::json!({"tokens": {"http://localhost:8787": {
|
||||
"token_type": "Bearer",
|
||||
"access_token": "secret",
|
||||
"expires_at": "2000-01-01T00:00:00Z"
|
||||
}}}),
|
||||
);
|
||||
assert!(matches!(
|
||||
BackendApiClient::from_token_file("http://localhost:8787", &expired),
|
||||
Err(BackendApiClientError::TokenExpired { .. })
|
||||
));
|
||||
|
||||
for path in [malformed, mismatch, expired] {
|
||||
let _ = fs::remove_file(path);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn token_write_and_lookup_share_origin_normalization() {
|
||||
let path = temp_path("normalized-write");
|
||||
save_backend_token_to_file(
|
||||
"HTTP://Example.COM:80////",
|
||||
"Bearer",
|
||||
"normalized-secret",
|
||||
None,
|
||||
&path,
|
||||
)
|
||||
.unwrap();
|
||||
let contents = fs::read_to_string(&path).unwrap();
|
||||
assert!(contents.contains("\"http://example.com\""));
|
||||
let client = BackendApiClient::from_token_file("http://example.com/", &path).unwrap();
|
||||
assert_eq!(
|
||||
client.authorization_header_value(),
|
||||
"Bearer normalized-secret"
|
||||
);
|
||||
fs::remove_file(path).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn client_debug_and_errors_never_include_token_value() {
|
||||
let client = BackendApiClient::from_access_token_for_test(
|
||||
"http://localhost:8787",
|
||||
"never-print-this-token",
|
||||
)
|
||||
.unwrap();
|
||||
assert!(!format!("{client:?}").contains("never-print-this-token"));
|
||||
assert!(
|
||||
!BackendApiClientError::Unauthorized {
|
||||
origin: client.origin().clone()
|
||||
}
|
||||
.to_string()
|
||||
.contains("never-print-this-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_requests_follow_only_same_origin_redirects() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let origin = format!("http://{}", listener.local_addr().unwrap());
|
||||
let handle = thread::spawn(move || {
|
||||
for response in [
|
||||
"HTTP/1.1 302 Found\r\nLocation: /final\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
||||
"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
||||
] {
|
||||
let (mut stream, _) = listener.accept().unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let read = stream.read(&mut request).unwrap();
|
||||
let request = String::from_utf8_lossy(&request[..read]).to_ascii_lowercase();
|
||||
assert!(request.contains("authorization: bearer redirect-secret\r\n"));
|
||||
stream.write_all(response.as_bytes()).unwrap();
|
||||
}
|
||||
});
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test(&origin, "redirect-secret").unwrap();
|
||||
let response = client
|
||||
.blocking_request(Method::GET, "/start")
|
||||
.unwrap()
|
||||
.send()
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authenticated_requests_reject_cross_origin_redirects_without_leaking_token() {
|
||||
let source = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let target = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
target.set_nonblocking(true).unwrap();
|
||||
let source_origin = format!("http://{}", source.local_addr().unwrap());
|
||||
let target_origin = format!("http://{}", target.local_addr().unwrap());
|
||||
let location = format!("{target_origin}/capture");
|
||||
let handle = thread::spawn(move || {
|
||||
let (mut stream, _) = source.accept().unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let read = stream.read(&mut request).unwrap();
|
||||
let request = String::from_utf8_lossy(&request[..read]).to_ascii_lowercase();
|
||||
assert!(request.contains("authorization: bearer redirect-secret\r\n"));
|
||||
let response = format!(
|
||||
"HTTP/1.1 302 Found\r\nLocation: {location}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
|
||||
);
|
||||
stream.write_all(response.as_bytes()).unwrap();
|
||||
});
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test(&source_origin, "redirect-secret")
|
||||
.unwrap();
|
||||
let error = client
|
||||
.blocking_request(Method::GET, "/start")
|
||||
.unwrap()
|
||||
.send()
|
||||
.unwrap_err();
|
||||
let message = error.to_string();
|
||||
assert!(message.contains("redirect"));
|
||||
assert!(!message.contains("redirect-secret"));
|
||||
handle.join().unwrap();
|
||||
thread::sleep(Duration::from_millis(20));
|
||||
assert!(matches!(
|
||||
target.accept(),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn status_diagnostics_distinguish_unauthorized_and_forbidden() {
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test("http://localhost:8787", "secret")
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
client.check_status(StatusCode::UNAUTHORIZED),
|
||||
Err(BackendApiClientError::Unauthorized { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
client.check_status(StatusCode::FORBIDDEN),
|
||||
Err(BackendApiClientError::Forbidden { .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::BackendOrigin;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
use std::time::Duration;
|
||||
@@ -9,9 +10,11 @@ pub struct BackendAuthTarget {
|
||||
|
||||
impl BackendAuthTarget {
|
||||
pub fn new(base_url: impl Into<String>) -> Self {
|
||||
Self {
|
||||
base_url: base_url.into(),
|
||||
}
|
||||
let base_url = base_url.into();
|
||||
let base_url = BackendOrigin::parse(&base_url)
|
||||
.map(|origin| origin.to_string())
|
||||
.unwrap_or(base_url);
|
||||
Self { base_url }
|
||||
}
|
||||
|
||||
fn api_url(&self, path: &str) -> String {
|
||||
|
||||
@@ -1,11 +1,16 @@
|
||||
use crate::{BackendApiClient, BackendApiClientError};
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use protocol::stream::{decode_event, encode_method};
|
||||
use protocol::{ErrorCode, Event, Method};
|
||||
use reqwest::Method as HttpMethod;
|
||||
use std::collections::VecDeque;
|
||||
use std::fmt;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_tungstenite::connect_async;
|
||||
use tokio_tungstenite::tungstenite::Message as TungsteniteMessage;
|
||||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
use tokio_tungstenite::tungstenite::http::HeaderValue;
|
||||
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
|
||||
pub use workdir::workspace::WorkingDirectorySummary as BackendWorkingDirectorySummary;
|
||||
pub use workspace_api::{
|
||||
Diagnostic as BackendDiagnostic, DiagnosticSeverity as BackendDiagnosticSeverity,
|
||||
@@ -113,6 +118,7 @@ pub struct BackendRuntimeClient {
|
||||
#[derive(Debug)]
|
||||
pub enum BackendRuntimeClientError {
|
||||
InvalidTarget(String),
|
||||
Api(BackendApiClientError),
|
||||
Http(reqwest::Error),
|
||||
}
|
||||
|
||||
@@ -120,6 +126,7 @@ impl fmt::Display for BackendRuntimeClientError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::InvalidTarget(message) => f.write_str(message),
|
||||
Self::Api(error) => write!(f, "{error}"),
|
||||
Self::Http(error) => write!(f, "{error}"),
|
||||
}
|
||||
}
|
||||
@@ -127,6 +134,12 @@ impl fmt::Display for BackendRuntimeClientError {
|
||||
|
||||
impl std::error::Error for BackendRuntimeClientError {}
|
||||
|
||||
impl From<BackendApiClientError> for BackendRuntimeClientError {
|
||||
fn from(error: BackendApiClientError) -> Self {
|
||||
Self::Api(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<reqwest::Error> for BackendRuntimeClientError {
|
||||
fn from(error: reqwest::Error) -> Self {
|
||||
Self::Http(error)
|
||||
@@ -137,7 +150,7 @@ pub async fn list_backend_workers(
|
||||
target: &BackendRuntimeListTarget,
|
||||
) -> Result<BackendRuntimeListResponse<BackendWorkerSummary>, BackendRuntimeClientError> {
|
||||
validate_list_target(target)?;
|
||||
let http = reqwest::Client::new();
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
if let Some(runtime_id) = target.runtime_id.as_deref() {
|
||||
let path = backend_runtime_workers_path(
|
||||
target
|
||||
@@ -146,12 +159,9 @@ pub async fn list_backend_workers(
|
||||
.expect("validated Backend Workspace scope"),
|
||||
runtime_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &path);
|
||||
return Ok(http
|
||||
.get(url)
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
let response = api.request(HttpMethod::GET, &path)?.send().await?;
|
||||
api.check_status(response.status())?;
|
||||
return Ok(response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?);
|
||||
}
|
||||
@@ -162,12 +172,9 @@ pub async fn list_backend_workers(
|
||||
.as_deref()
|
||||
.expect("validated Backend Workspace scope"),
|
||||
);
|
||||
let runtime_url = join_base_and_path(&target.base_url, &runtime_path);
|
||||
let runtimes = http
|
||||
.get(runtime_url)
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
let response = api.request(HttpMethod::GET, &runtime_path)?.send().await?;
|
||||
api.check_status(response.status())?;
|
||||
let runtimes = response
|
||||
.json::<BackendRuntimeListResponse<BackendRuntimeSummary>>()
|
||||
.await?;
|
||||
|
||||
@@ -181,29 +188,43 @@ pub async fn list_backend_workers(
|
||||
.expect("validated Backend Workspace scope"),
|
||||
&runtime.runtime_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &path);
|
||||
match http
|
||||
.get(url)
|
||||
.send()
|
||||
.await
|
||||
.and_then(|response| response.error_for_status())
|
||||
{
|
||||
Ok(response) => {
|
||||
let response = response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?;
|
||||
diagnostics.extend(response.diagnostics);
|
||||
items.extend(response.items);
|
||||
let response = match api.request(HttpMethod::GET, &path)?.send().await {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
diagnostics.push(BackendDiagnostic {
|
||||
code: "runtime_worker_list_failed".to_string(),
|
||||
severity: BackendDiagnosticSeverity::Error,
|
||||
message: format!(
|
||||
"failed to list workers for runtime {}: {error}",
|
||||
runtime.runtime_id
|
||||
),
|
||||
});
|
||||
continue;
|
||||
}
|
||||
Err(error) => diagnostics.push(BackendDiagnostic {
|
||||
};
|
||||
if matches!(
|
||||
response.status(),
|
||||
reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN
|
||||
) {
|
||||
api.check_status(response.status())?;
|
||||
}
|
||||
if !response.status().is_success() {
|
||||
diagnostics.push(BackendDiagnostic {
|
||||
code: "runtime_worker_list_failed".to_string(),
|
||||
severity: BackendDiagnosticSeverity::Error,
|
||||
message: format!(
|
||||
"failed to list workers for runtime {}: {error}",
|
||||
runtime.runtime_id
|
||||
"failed to list workers for runtime {}: Backend returned HTTP {}",
|
||||
runtime.runtime_id,
|
||||
response.status().as_u16()
|
||||
),
|
||||
}),
|
||||
});
|
||||
continue;
|
||||
}
|
||||
let response = response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?;
|
||||
diagnostics.extend(response.diagnostics);
|
||||
items.extend(response.items);
|
||||
}
|
||||
|
||||
Ok(BackendRuntimeListResponse {
|
||||
@@ -224,7 +245,7 @@ pub async fn list_backend_stopped_workers(
|
||||
"stopped worker listing requires a runtime id".to_string(),
|
||||
));
|
||||
};
|
||||
let http = reqwest::Client::new();
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let path = backend_runtime_workers_path(
|
||||
target
|
||||
.workspace_id
|
||||
@@ -232,12 +253,12 @@ pub async fn list_backend_stopped_workers(
|
||||
.expect("validated Backend Workspace scope"),
|
||||
runtime_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &format!("{path}?status=stopped"));
|
||||
Ok(http
|
||||
.get(url)
|
||||
let response = api
|
||||
.request(HttpMethod::GET, &format!("{path}?status=stopped"))?
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
.await?;
|
||||
api.check_status(response.status())?;
|
||||
Ok(response
|
||||
.json::<BackendRuntimeListResponse<BackendWorkerSummary>>()
|
||||
.await?)
|
||||
}
|
||||
@@ -246,33 +267,33 @@ pub async fn restore_backend_worker(
|
||||
target: &BackendRuntimeTarget,
|
||||
) -> Result<BackendWorkerRestoreResponse, BackendRuntimeClientError> {
|
||||
validate_target(target)?;
|
||||
let http = reqwest::Client::new();
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let path = backend_runtime_worker_restore_path(
|
||||
&target.workspace_id,
|
||||
&target.runtime_id,
|
||||
&target.worker_id,
|
||||
);
|
||||
let url = join_base_and_path(&target.base_url, &path);
|
||||
Ok(http
|
||||
.post(url)
|
||||
let response = api
|
||||
.request(HttpMethod::POST, &path)?
|
||||
.json(&serde_json::json!({}))
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
.json::<BackendWorkerRestoreResponse>()
|
||||
.await?)
|
||||
.await?;
|
||||
api.check_status(response.status())?;
|
||||
Ok(response.json::<BackendWorkerRestoreResponse>().await?)
|
||||
}
|
||||
|
||||
impl BackendRuntimeClient {
|
||||
pub async fn connect(target: BackendRuntimeTarget) -> Result<Self, BackendRuntimeClientError> {
|
||||
validate_target(&target)?;
|
||||
let api = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let (event_tx, rx) = mpsc::unbounded_channel();
|
||||
let (command_tx, command_rx) = mpsc::unbounded_channel();
|
||||
|
||||
let protocol_target = target.clone();
|
||||
let protocol_event_tx = event_tx.clone();
|
||||
let protocol_task = tokio::spawn(async move {
|
||||
run_worker_protocol_transport(protocol_target, command_rx, protocol_event_tx).await;
|
||||
run_worker_protocol_transport(protocol_target, api, command_rx, protocol_event_tx)
|
||||
.await;
|
||||
});
|
||||
|
||||
Ok(Self {
|
||||
@@ -317,11 +338,21 @@ impl Drop for BackendRuntimeClient {
|
||||
|
||||
async fn run_worker_protocol_transport(
|
||||
target: BackendRuntimeTarget,
|
||||
api: BackendApiClient,
|
||||
mut commands: mpsc::UnboundedReceiver<Method>,
|
||||
tx: mpsc::UnboundedSender<Event>,
|
||||
) {
|
||||
let url = protocol_ws_url(&target);
|
||||
match connect_async(&url).await {
|
||||
let request = match protocol_ws_request(&target, &api) {
|
||||
Ok(request) => request,
|
||||
Err(error) => {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol request could not be constructed for {}: {error}",
|
||||
target.display_label()
|
||||
)));
|
||||
return;
|
||||
}
|
||||
};
|
||||
match connect_async(request).await {
|
||||
Ok((ws, _)) => {
|
||||
let (mut sink, mut stream) = ws.split();
|
||||
loop {
|
||||
@@ -387,10 +418,8 @@ async fn run_worker_protocol_transport(
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol WebSocket connect failed for {}: {error}",
|
||||
target.display_label()
|
||||
)));
|
||||
let message = protocol_connect_error_message(&target, &api, &error);
|
||||
let _ = tx.send(diagnostic_event(message));
|
||||
while commands.recv().await.is_some() {
|
||||
let _ = tx.send(diagnostic_event(format!(
|
||||
"Backend protocol command was not sent because command stream is unavailable for {}",
|
||||
@@ -401,6 +430,29 @@ async fn run_worker_protocol_transport(
|
||||
}
|
||||
}
|
||||
|
||||
fn protocol_connect_error_message(
|
||||
target: &BackendRuntimeTarget,
|
||||
api: &BackendApiClient,
|
||||
error: &tokio_tungstenite::tungstenite::Error,
|
||||
) -> String {
|
||||
if let tokio_tungstenite::tungstenite::Error::Http(response) = error {
|
||||
if let Ok(status) = reqwest::StatusCode::from_u16(response.status().as_u16()) {
|
||||
if matches!(
|
||||
status,
|
||||
reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN
|
||||
) {
|
||||
if let Err(error) = api.check_status(status) {
|
||||
return error.to_string();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
format!(
|
||||
"Backend protocol WebSocket connect failed for {}: {error}",
|
||||
target.display_label()
|
||||
)
|
||||
}
|
||||
|
||||
fn diagnostic_event(message: impl Into<String>) -> Event {
|
||||
Event::Error {
|
||||
code: ErrorCode::Internal,
|
||||
@@ -496,6 +548,19 @@ fn backend_runtime_worker_restore_path(
|
||||
)
|
||||
}
|
||||
|
||||
fn protocol_ws_request(
|
||||
target: &BackendRuntimeTarget,
|
||||
api: &BackendApiClient,
|
||||
) -> Result<tokio_tungstenite::tungstenite::http::Request<()>, String> {
|
||||
let mut request = protocol_ws_url(target)
|
||||
.into_client_request()
|
||||
.map_err(|error| error.to_string())?;
|
||||
let value = HeaderValue::from_str(&api.authorization_header_value())
|
||||
.map_err(|_| "saved Backend token is not a valid Authorization header".to_string())?;
|
||||
request.headers_mut().insert(AUTHORIZATION, value);
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
fn protocol_ws_url(target: &BackendRuntimeTarget) -> String {
|
||||
let path = format!(
|
||||
"/api/w/{}/runtimes/{}/workers/{}/protocol/ws",
|
||||
@@ -557,6 +622,26 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn protocol_request_attaches_saved_bearer_authorization() {
|
||||
let target = BackendRuntimeTarget::new(
|
||||
"http://127.0.0.1:8787/",
|
||||
"workspace alpha",
|
||||
"runtime/one",
|
||||
"worker one",
|
||||
);
|
||||
let api = BackendApiClient::from_access_token_for_test(
|
||||
"http://127.0.0.1:8787",
|
||||
"websocket-secret",
|
||||
)
|
||||
.unwrap();
|
||||
let request = protocol_ws_request(&target, &api).unwrap();
|
||||
assert_eq!(
|
||||
request.headers().get(AUTHORIZATION).unwrap(),
|
||||
"Bearer websocket-secret"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_worker_summary_decodes_current_occupied_workdir_contract() {
|
||||
let payload = serde_json::json!({
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
use crate::{BackendApiClient, BackendApiClientError};
|
||||
use reqwest::Method;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
use workspace_api::{RepositoryObservedStatus, RepositorySource};
|
||||
|
||||
const DEFAULT_WORKSPACE_LIMIT: usize = 200;
|
||||
|
||||
@@ -44,8 +47,13 @@ pub struct CreateBackendWorkspaceRepositoryRecord {
|
||||
pub repository_id: String,
|
||||
pub name: String,
|
||||
pub kind: String,
|
||||
pub uri: String,
|
||||
pub provider: Option<String>,
|
||||
pub source: RepositorySource,
|
||||
pub default_ref: Option<String>,
|
||||
pub source_revision: u64,
|
||||
pub source_fingerprint: String,
|
||||
pub observed_status: RepositoryObservedStatus,
|
||||
pub observed_at: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -64,7 +72,7 @@ impl BackendWorkspaceCatalogTarget {
|
||||
#[derive(Debug)]
|
||||
pub enum BackendWorkspaceClientError {
|
||||
InvalidTarget(String),
|
||||
RequestFailed { status: u16, message: String },
|
||||
Api(BackendApiClientError),
|
||||
Http(reqwest::Error),
|
||||
}
|
||||
|
||||
@@ -72,9 +80,7 @@ impl fmt::Display for BackendWorkspaceClientError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::InvalidTarget(message) => f.write_str(message),
|
||||
Self::RequestFailed { status, message } => {
|
||||
write!(f, "Backend request failed with HTTP {status}: {message}")
|
||||
}
|
||||
Self::Api(error) => write!(f, "{error}"),
|
||||
Self::Http(error) => write!(f, "{error}"),
|
||||
}
|
||||
}
|
||||
@@ -82,6 +88,12 @@ impl fmt::Display for BackendWorkspaceClientError {
|
||||
|
||||
impl std::error::Error for BackendWorkspaceClientError {}
|
||||
|
||||
impl From<BackendApiClientError> for BackendWorkspaceClientError {
|
||||
fn from(error: BackendApiClientError) -> Self {
|
||||
Self::Api(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<reqwest::Error> for BackendWorkspaceClientError {
|
||||
fn from(error: reqwest::Error) -> Self {
|
||||
Self::Http(error)
|
||||
@@ -91,13 +103,21 @@ impl From<reqwest::Error> for BackendWorkspaceClientError {
|
||||
pub async fn list_backend_workspaces(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
) -> Result<Vec<BackendWorkspace>, BackendWorkspaceClientError> {
|
||||
validate_target(target)?;
|
||||
let url = format!(
|
||||
"{}/api/workspaces?limit={DEFAULT_WORKSPACE_LIMIT}",
|
||||
target.base_url.trim_end_matches('/')
|
||||
);
|
||||
let response = reqwest::Client::new().get(url).send().await?;
|
||||
let response = require_success(response).await?;
|
||||
let client = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
list_backend_workspaces_with_client(&client).await
|
||||
}
|
||||
|
||||
async fn list_backend_workspaces_with_client(
|
||||
client: &BackendApiClient,
|
||||
) -> Result<Vec<BackendWorkspace>, BackendWorkspaceClientError> {
|
||||
let response = client
|
||||
.request(
|
||||
Method::GET,
|
||||
&format!("/api/workspaces?limit={DEFAULT_WORKSPACE_LIMIT}"),
|
||||
)?
|
||||
.send()
|
||||
.await?;
|
||||
client.check_status(response.status())?;
|
||||
Ok(response.json::<Vec<BackendWorkspace>>().await?)
|
||||
}
|
||||
|
||||
@@ -105,42 +125,50 @@ pub async fn create_backend_workspace(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
request: &CreateBackendWorkspaceRequest,
|
||||
) -> Result<CreateBackendWorkspaceResponse, BackendWorkspaceClientError> {
|
||||
validate_target(target)?;
|
||||
let url = format!("{}/api/workspaces", target.base_url.trim_end_matches('/'));
|
||||
let response = reqwest::Client::new()
|
||||
.post(url)
|
||||
let client = BackendApiClient::from_stored_token(&target.base_url)?;
|
||||
let response = client
|
||||
.request(Method::POST, "/api/workspaces")?
|
||||
.json(request)
|
||||
.send()
|
||||
.await?;
|
||||
let response = require_success(response).await?;
|
||||
client.check_status(response.status())?;
|
||||
Ok(response.json::<CreateBackendWorkspaceResponse>().await?)
|
||||
}
|
||||
|
||||
async fn require_success(
|
||||
response: reqwest::Response,
|
||||
) -> Result<reqwest::Response, BackendWorkspaceClientError> {
|
||||
if response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
let status = response.status().as_u16();
|
||||
let message = response.text().await.unwrap_or_default();
|
||||
Err(BackendWorkspaceClientError::RequestFailed { status, message })
|
||||
}
|
||||
|
||||
fn validate_target(
|
||||
target: &BackendWorkspaceCatalogTarget,
|
||||
) -> Result<(), BackendWorkspaceClientError> {
|
||||
if !(target.base_url.starts_with("http://") || target.base_url.starts_with("https://")) {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend API base URL must start with http:// or https://".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::io::{Read, Write};
|
||||
use std::net::TcpListener;
|
||||
use std::thread;
|
||||
|
||||
#[tokio::test]
|
||||
async fn workspace_catalog_request_uses_shared_bearer_client() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let base_url = format!("http://{}", listener.local_addr().unwrap());
|
||||
let handle = thread::spawn(move || {
|
||||
let (mut stream, _) = listener.accept().unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let read = stream.read(&mut request).unwrap();
|
||||
let request = String::from_utf8_lossy(&request[..read]).to_ascii_lowercase();
|
||||
assert!(request.starts_with("get /api/workspaces?limit=200 "));
|
||||
assert!(request.contains("authorization: bearer catalog-secret\r\n"));
|
||||
stream
|
||||
.write_all(
|
||||
b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 2\r\nConnection: close\r\n\r\n[]",
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
let client =
|
||||
BackendApiClient::from_access_token_for_test(&base_url, "catalog-secret").unwrap();
|
||||
assert!(
|
||||
list_backend_workspaces_with_client(&client)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn create_request_keeps_operation_key_for_exact_retry() {
|
||||
|
||||
+12
-27
@@ -1,23 +1,20 @@
|
||||
//! Worker プロトコルを喋るクライアント。
|
||||
//! Backend Workspace/Runtime と既存 Worker protocol へ接続するクライアント。
|
||||
//!
|
||||
//! - [`WorkerClient`]: 既存 worker の Unix ソケットへ接続して `Method` を送り、
|
||||
//! `Event` を受け取る低レベル接続。
|
||||
//! - [`spawn`]: worker バイナリをサブプロセスとして起動し、`YOI-READY`
|
||||
//! ハンドシェイクが終わるまで待つフロー。subprocess を立ち上げる必要が
|
||||
//! ない呼び出し側 (=既存 worker に attach する場合) は使わなくてよい。
|
||||
//!
|
||||
//! TUI / GUI / E2E ハーネスはこの crate に依存して protocol を喋る。
|
||||
//! Standalone execution is owned by the `standalone` crate and does not spawn
|
||||
//! a Worker subprocess through this crate.
|
||||
|
||||
pub mod backend_auth;
|
||||
pub mod backend_api;
|
||||
mod backend_auth;
|
||||
pub mod backend_runtime;
|
||||
pub mod backend_workspace;
|
||||
pub mod runtime_command;
|
||||
pub mod spawn;
|
||||
pub mod target;
|
||||
pub mod ticket_role;
|
||||
mod worker_client;
|
||||
mod workspace_product;
|
||||
|
||||
pub use backend_api::{
|
||||
BackendApiClient, BackendApiClientError, BackendOrigin, backend_token_file_path,
|
||||
save_backend_token,
|
||||
};
|
||||
pub use backend_auth::{
|
||||
BackendAuthClientError, BackendAuthTarget, DeviceLoginPollResponse, DeviceLoginStartResponse,
|
||||
poll_device_login, start_device_login, wait_for_device_login,
|
||||
@@ -35,22 +32,10 @@ pub use backend_workspace::{
|
||||
CreateBackendWorkspaceRepository, CreateBackendWorkspaceRequest,
|
||||
CreateBackendWorkspaceResponse, create_backend_workspace, list_backend_workspaces,
|
||||
};
|
||||
pub use runtime_command::WorkerRuntimeCommand;
|
||||
pub use target::{
|
||||
BackendTarget, Dashboard, LocalTarget, ResolvedTarget, Target, TargetError, TargetKind,
|
||||
WorkerByName, WorkerConnection, WorkerConnectionSelector, WorkerList, WorkerListRequest,
|
||||
WorkerResume, WorkerSpawn,
|
||||
};
|
||||
|
||||
pub use spawn::{
|
||||
SpawnConfig, SpawnError, SpawnReady, WorkerProcessLaunchConfig, WorkerProcessLaunchOptions,
|
||||
spawn_worker, spawn_worker_with_options,
|
||||
};
|
||||
pub use ticket_role::{
|
||||
TicketRef, TicketRoleLaunchContext, TicketRoleLaunchError, TicketRoleLaunchOptions,
|
||||
TicketRoleLaunchPlan, TicketRoleLaunchResult, TicketRolePreRunWarning,
|
||||
launch_ticket_role_worker, launch_ticket_role_worker_with_options, plan_ticket_role_launch,
|
||||
plan_ticket_role_launch_with_config,
|
||||
BackendTarget, Dashboard, ResolvedTarget, StandaloneSessionListIntent,
|
||||
StandaloneSessionResumeIntent, StandaloneTarget, Target, TargetError, TargetKind,
|
||||
WorkerConnection, WorkerConnectionSelector, WorkerList, WorkerListRequest, WorkerSpawn,
|
||||
};
|
||||
pub use worker_client::WorkerClient;
|
||||
pub use workspace_api::{ObjectiveDetail, ObjectiveSummary};
|
||||
|
||||
@@ -1,435 +0,0 @@
|
||||
//! Worker runtime command をサブプロセスとして立ち上げ、`YOI-READY` を待つ
|
||||
//! ハンドシェイク。
|
||||
//!
|
||||
//! - 親プロセス (TUI / GUI / E2E) は profile/default/typed restore flags を
|
||||
//! 指定してこの関数に渡す。worker はそれを受けて socket を bind し、stderr に
|
||||
//! `YOI-READY\t<name>\t<socket>` を吐く。
|
||||
//! - 待機中の stderr 行は `progress` コールバック越しに呼び出し側へ流す。
|
||||
//! UI の進捗表示や E2E のログ収集はここで賄う。
|
||||
//! - `kill_on_drop = false` + `process_group(0)` により、親プロセス
|
||||
//! ライフサイクルから切り離した detached worker を作る。ready 後の lifecycle
|
||||
//! 管理は runtime ディレクトリ / socket を介して行う。
|
||||
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::process::Stdio;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::WorkerRuntimeCommand;
|
||||
use tokio::process::Command;
|
||||
use uuid::Uuid;
|
||||
|
||||
const READY_PREFIX: &str = "YOI-READY\t";
|
||||
const READY_TIMEOUT: Duration = Duration::from_secs(20);
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerProcessLaunchConfig {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
/// `worker.name` として使う識別子。runtime ディレクトリ
|
||||
/// (`manifest::paths::worker_runtime_dir`) の解決と、ready 行に乗る
|
||||
/// 名前との突き合わせに使う。
|
||||
pub worker_name: String,
|
||||
/// Optional reusable Profile selector. Worker identity is always supplied
|
||||
/// separately with `--worker`; profile selection must not imply a name.
|
||||
pub profile: Option<String>,
|
||||
/// Explicit runtime workspace root. The child receives it via
|
||||
/// `--workspace` so startup does not infer workspace identity from the
|
||||
/// parent process cwd.
|
||||
pub workspace_root: PathBuf,
|
||||
/// Optional child process cwd. This is not runtime workspace identity and
|
||||
/// is not passed as a CLI argument; the child observes it as its ordinary
|
||||
/// process current directory.
|
||||
pub cwd: Option<PathBuf>,
|
||||
/// `Some(id)` のとき `--session <id>` を付与し、当該セッションから
|
||||
/// resume させる。
|
||||
pub resume_from: Option<Uuid>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default)]
|
||||
pub struct WorkerProcessLaunchOptions {
|
||||
/// Extra child CLI arguments supplied by an upper resolver layer. The
|
||||
/// low-level launch config intentionally does not model Ticket IDs,
|
||||
/// Ticket roles, orchestration roles, executable authority, or raw
|
||||
/// browser-provided profile/cwd/workspace inputs.
|
||||
pub extra_args: Vec<String>,
|
||||
}
|
||||
|
||||
impl WorkerProcessLaunchOptions {
|
||||
pub fn with_hidden_arg(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
|
||||
self.extra_args.extend([name.into(), value.into()]);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.extra_args.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
pub type SpawnConfig = WorkerProcessLaunchConfig;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct SpawnReady {
|
||||
pub worker_name: String,
|
||||
pub socket_path: PathBuf,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum SpawnError {
|
||||
Io(io::Error),
|
||||
/// runtime ディレクトリが解決できなかった (環境変数未設定等)。
|
||||
RuntimeDirUnavailable,
|
||||
WorkerLaunchFailed {
|
||||
command: WorkerRuntimeCommand,
|
||||
source: io::Error,
|
||||
},
|
||||
WorkerExitedEarly {
|
||||
stderr_tail: String,
|
||||
},
|
||||
Timeout,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for SpawnError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(e) => write!(f, "io error: {e}"),
|
||||
Self::RuntimeDirUnavailable => write!(
|
||||
f,
|
||||
"could not resolve runtime directory (set YOI_HOME, YOI_RUNTIME_DIR, XDG_RUNTIME_DIR, or HOME)"
|
||||
),
|
||||
Self::WorkerLaunchFailed { command, source } => write!(
|
||||
f,
|
||||
"failed to launch worker runtime command `{command}`: {source}"
|
||||
),
|
||||
Self::WorkerExitedEarly { stderr_tail } => {
|
||||
if stderr_tail.is_empty() {
|
||||
write!(f, "worker exited before becoming ready")
|
||||
} else {
|
||||
write!(f, "worker exited before becoming ready: {stderr_tail}")
|
||||
}
|
||||
}
|
||||
Self::Timeout => write!(
|
||||
f,
|
||||
"worker did not become ready within {}s",
|
||||
READY_TIMEOUT.as_secs()
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for SpawnError {
|
||||
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
|
||||
match self {
|
||||
Self::Io(error) | Self::WorkerLaunchFailed { source: error, .. } => Some(error),
|
||||
Self::RuntimeDirUnavailable | Self::WorkerExitedEarly { .. } | Self::Timeout => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<io::Error> for SpawnError {
|
||||
fn from(e: io::Error) -> Self {
|
||||
Self::Io(e)
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_args(
|
||||
config: &WorkerProcessLaunchConfig,
|
||||
options: &WorkerProcessLaunchOptions,
|
||||
) -> Vec<String> {
|
||||
let mut args = vec![
|
||||
"--workspace".to_string(),
|
||||
config.workspace_root.display().to_string(),
|
||||
];
|
||||
if let Some(id) = config.resume_from {
|
||||
args.extend([
|
||||
"--session".to_string(),
|
||||
id.to_string(),
|
||||
"--worker".to_string(),
|
||||
config.worker_name.clone(),
|
||||
]);
|
||||
} else {
|
||||
args.extend(["--worker".to_string(), config.worker_name.clone()]);
|
||||
if let Some(profile) = &config.profile {
|
||||
args.extend(["--profile".to_string(), profile.clone()]);
|
||||
}
|
||||
}
|
||||
args.extend(options.extra_args.clone());
|
||||
args
|
||||
}
|
||||
|
||||
/// worker を spawn し、`YOI-READY` ハンドシェイクが終わるまで待つ。
|
||||
///
|
||||
/// `progress` は ready 行を見つけるまでに観測した stderr の各行で呼ばれる
|
||||
/// (ready 行自体は除外される)。UI の表示更新や E2E ログ取得に使う。
|
||||
pub async fn spawn_worker<F>(
|
||||
config: WorkerProcessLaunchConfig,
|
||||
progress: F,
|
||||
) -> Result<SpawnReady, SpawnError>
|
||||
where
|
||||
F: FnMut(&str),
|
||||
{
|
||||
spawn_worker_with_options(config, WorkerProcessLaunchOptions::default(), progress).await
|
||||
}
|
||||
|
||||
pub async fn spawn_worker_with_options<F>(
|
||||
config: WorkerProcessLaunchConfig,
|
||||
options: WorkerProcessLaunchOptions,
|
||||
mut progress: F,
|
||||
) -> Result<SpawnReady, SpawnError>
|
||||
where
|
||||
F: FnMut(&str),
|
||||
{
|
||||
let worker_runtime_dir = manifest::paths::worker_runtime_dir(&config.worker_name)
|
||||
.ok_or(SpawnError::RuntimeDirUnavailable)?;
|
||||
std::fs::create_dir_all(&worker_runtime_dir).map_err(SpawnError::Io)?;
|
||||
let stderr_path = worker_runtime_dir.join("stderr.log");
|
||||
let stderr_file = std::fs::File::create(&stderr_path).map_err(SpawnError::Io)?;
|
||||
|
||||
let mut command = Command::new(config.runtime_command.program());
|
||||
command
|
||||
.args(config.runtime_command.prefix_args())
|
||||
.current_dir(config.cwd.as_ref().unwrap_or(&config.workspace_root))
|
||||
.stdin(Stdio::null())
|
||||
.stdout(Stdio::null())
|
||||
.stderr(Stdio::from(stderr_file))
|
||||
.process_group(0);
|
||||
for arg in runtime_args(&config, &options) {
|
||||
command.arg(arg);
|
||||
}
|
||||
let mut child = command
|
||||
.spawn()
|
||||
.map_err(|source| SpawnError::WorkerLaunchFailed {
|
||||
command: config.runtime_command.clone(),
|
||||
source,
|
||||
})?;
|
||||
|
||||
// Default `kill_on_drop = false` plus `process_group(0)` makes this
|
||||
// a detached Worker once startup succeeds: dropping the handle does not
|
||||
// terminate it, and terminal-generated signals for the parent's
|
||||
// process group do not hit the Worker. Runtime state/socket files are
|
||||
// the source of truth after that point.
|
||||
let ready = match wait_for_ready_file(&mut progress, &stderr_path, &mut child).await {
|
||||
Ok(ready) => ready,
|
||||
Err(e) => {
|
||||
let _ = child.start_kill();
|
||||
let _ = child.wait().await;
|
||||
return Err(e);
|
||||
}
|
||||
};
|
||||
tokio::spawn(async move {
|
||||
let _ = child.wait().await;
|
||||
});
|
||||
Ok(ready)
|
||||
}
|
||||
|
||||
async fn wait_for_ready_file<F>(
|
||||
progress: &mut F,
|
||||
stderr_path: &Path,
|
||||
child: &mut tokio::process::Child,
|
||||
) -> Result<SpawnReady, SpawnError>
|
||||
where
|
||||
F: FnMut(&str),
|
||||
{
|
||||
let mut tail = StderrTail::new();
|
||||
let deadline = tokio::time::Instant::now() + READY_TIMEOUT;
|
||||
let mut offset = 0usize;
|
||||
|
||||
loop {
|
||||
let content = match tokio::fs::read_to_string(stderr_path).await {
|
||||
Ok(content) => content,
|
||||
Err(e) if e.kind() == io::ErrorKind::NotFound => String::new(),
|
||||
Err(e) => return Err(SpawnError::Io(e)),
|
||||
};
|
||||
if content.len() > offset {
|
||||
for line in content[offset..].lines() {
|
||||
if let Some(rest) = line.strip_prefix(READY_PREFIX) {
|
||||
let mut parts = rest.splitn(2, '\t');
|
||||
let worker_name = parts.next().unwrap_or("").to_string();
|
||||
let socket_str = parts.next().unwrap_or("").to_string();
|
||||
if worker_name.is_empty() || socket_str.is_empty() {
|
||||
return Err(SpawnError::WorkerExitedEarly {
|
||||
stderr_tail: format!("malformed ready line: {line}"),
|
||||
});
|
||||
}
|
||||
let socket_path = PathBuf::from(socket_str);
|
||||
wait_for_socket(
|
||||
&socket_path,
|
||||
deadline,
|
||||
child,
|
||||
stderr_path,
|
||||
&mut tail,
|
||||
&mut offset,
|
||||
)
|
||||
.await?;
|
||||
return Ok(SpawnReady {
|
||||
worker_name,
|
||||
socket_path,
|
||||
});
|
||||
}
|
||||
tail.push(line);
|
||||
progress(line);
|
||||
}
|
||||
offset = content.len();
|
||||
}
|
||||
|
||||
if tokio::time::Instant::now() >= deadline {
|
||||
return Err(SpawnError::Timeout);
|
||||
}
|
||||
tokio::select! {
|
||||
status = child.wait() => {
|
||||
let _ = status;
|
||||
// Worker は exit 直前に最終 stderr 行を flush することがある。
|
||||
// child.wait() が解決した後に再読みして、原因行を取りこ
|
||||
// ぼさず WorkerExitedEarly に載せる。
|
||||
drain_stderr_into_tail(stderr_path, &mut tail, &mut offset).await;
|
||||
return Err(SpawnError::WorkerExitedEarly {
|
||||
stderr_tail: tail.into_string(),
|
||||
});
|
||||
}
|
||||
_ = tokio::time::sleep(Duration::from_millis(100)) => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn wait_for_socket(
|
||||
socket_path: &Path,
|
||||
deadline: tokio::time::Instant,
|
||||
child: &mut tokio::process::Child,
|
||||
stderr_path: &Path,
|
||||
tail: &mut StderrTail,
|
||||
offset: &mut usize,
|
||||
) -> Result<(), SpawnError> {
|
||||
loop {
|
||||
match tokio::net::UnixStream::connect(socket_path).await {
|
||||
Ok(_) => return Ok(()),
|
||||
Err(e)
|
||||
if e.kind() == io::ErrorKind::NotFound
|
||||
|| e.kind() == io::ErrorKind::ConnectionRefused => {}
|
||||
Err(e) => return Err(SpawnError::Io(e)),
|
||||
}
|
||||
if tokio::time::Instant::now() >= deadline {
|
||||
return Err(SpawnError::Timeout);
|
||||
}
|
||||
tokio::select! {
|
||||
status = child.wait() => {
|
||||
let _ = status;
|
||||
drain_stderr_into_tail(stderr_path, tail, offset).await;
|
||||
return Err(SpawnError::WorkerExitedEarly {
|
||||
stderr_tail: tail.as_string(),
|
||||
});
|
||||
}
|
||||
_ = tokio::time::sleep(Duration::from_millis(50)) => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn drain_stderr_into_tail(stderr_path: &Path, tail: &mut StderrTail, offset: &mut usize) {
|
||||
let Ok(content) = tokio::fs::read_to_string(stderr_path).await else {
|
||||
return;
|
||||
};
|
||||
if content.len() <= *offset {
|
||||
return;
|
||||
}
|
||||
for line in content[*offset..].lines() {
|
||||
if !line.starts_with(READY_PREFIX) {
|
||||
tail.push(line);
|
||||
}
|
||||
}
|
||||
*offset = content.len();
|
||||
}
|
||||
|
||||
struct StderrTail {
|
||||
lines: std::collections::VecDeque<String>,
|
||||
}
|
||||
|
||||
impl StderrTail {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
lines: std::collections::VecDeque::with_capacity(8),
|
||||
}
|
||||
}
|
||||
fn push(&mut self, line: &str) {
|
||||
if self.lines.len() == 8 {
|
||||
self.lines.pop_front();
|
||||
}
|
||||
self.lines.push_back(line.to_string());
|
||||
}
|
||||
fn as_string(&self) -> String {
|
||||
self.lines.iter().cloned().collect::<Vec<_>>().join(" | ")
|
||||
}
|
||||
fn into_string(self) -> String {
|
||||
self.lines.into_iter().collect::<Vec<_>>().join(" | ")
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::ffi::OsString;
|
||||
|
||||
fn base_config() -> WorkerProcessLaunchConfig {
|
||||
WorkerProcessLaunchConfig {
|
||||
runtime_command: WorkerRuntimeCommand::new("/bin/yoi", vec![OsString::from("worker")]),
|
||||
worker_name: "explicit-worker".to_string(),
|
||||
profile: Some("project:companion".to_string()),
|
||||
workspace_root: PathBuf::from("/work/other-project"),
|
||||
cwd: None,
|
||||
resume_from: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_args_keep_workspace_worker_and_profile_separate() {
|
||||
assert_eq!(
|
||||
runtime_args(&base_config(), &WorkerProcessLaunchOptions::default()),
|
||||
vec![
|
||||
"--workspace",
|
||||
"/work/other-project",
|
||||
"--worker",
|
||||
"explicit-worker",
|
||||
"--profile",
|
||||
"project:companion",
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_args_use_session_mode_without_profile_identity_alias() {
|
||||
let mut config = base_config();
|
||||
config.resume_from = Some(Uuid::nil());
|
||||
assert_eq!(
|
||||
runtime_args(&config, &WorkerProcessLaunchOptions::default()),
|
||||
vec![
|
||||
"--workspace",
|
||||
"/work/other-project",
|
||||
"--session",
|
||||
"00000000-0000-0000-0000-000000000000",
|
||||
"--worker",
|
||||
"explicit-worker",
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_args_include_upper_resolver_extra_args_without_child_cwd() {
|
||||
let mut config = base_config();
|
||||
config.cwd = Some(PathBuf::from("/work/main/.worktree/orchestration/yoi"));
|
||||
|
||||
assert_eq!(
|
||||
runtime_args(
|
||||
&config,
|
||||
&WorkerProcessLaunchOptions::default()
|
||||
.with_hidden_arg("--ticket-role", "orchestrator"),
|
||||
),
|
||||
vec![
|
||||
"--workspace",
|
||||
"/work/other-project",
|
||||
"--worker",
|
||||
"explicit-worker",
|
||||
"--profile",
|
||||
"project:companion",
|
||||
"--ticket-role",
|
||||
"orchestrator",
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
+161
-194
@@ -1,16 +1,20 @@
|
||||
use std::fmt;
|
||||
use std::{fmt, path::PathBuf};
|
||||
|
||||
use crate::{BackendRuntimeListTarget, BackendRuntimeTarget, WorkerRuntimeCommand};
|
||||
use crate::{
|
||||
BackendApiClient, BackendApiClientError, BackendOrigin, BackendRuntimeListTarget,
|
||||
BackendRuntimeTarget,
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum TargetKind {
|
||||
Local,
|
||||
/// One-process Standalone authority with no Runtime or Workspace backend.
|
||||
Standalone,
|
||||
Backend,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum ResolvedTarget {
|
||||
Local,
|
||||
Standalone,
|
||||
Backend {
|
||||
base_url: String,
|
||||
workspace_id: String,
|
||||
@@ -20,7 +24,7 @@ pub enum ResolvedTarget {
|
||||
impl ResolvedTarget {
|
||||
pub fn kind(&self) -> TargetKind {
|
||||
match self {
|
||||
Self::Local => TargetKind::Local,
|
||||
Self::Standalone => TargetKind::Standalone,
|
||||
Self::Backend { .. } => TargetKind::Backend,
|
||||
}
|
||||
}
|
||||
@@ -29,31 +33,12 @@ impl ResolvedTarget {
|
||||
impl fmt::Display for TargetKind {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Local => f.write_str("local"),
|
||||
Self::Standalone => f.write_str("Standalone"),
|
||||
Self::Backend => f.write_str("Backend"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct LocalTarget;
|
||||
|
||||
impl LocalTarget {
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
fn runtime_command(&self) -> Result<WorkerRuntimeCommand, TargetError> {
|
||||
WorkerRuntimeCommand::resolve().map_err(TargetError::local_runtime_command)
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for LocalTarget {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BackendTarget {
|
||||
pub base_url: String,
|
||||
@@ -62,11 +47,19 @@ pub struct BackendTarget {
|
||||
|
||||
impl BackendTarget {
|
||||
pub fn new(base_url: impl Into<String>, workspace_id: Option<impl Into<String>>) -> Self {
|
||||
let base_url = base_url.into();
|
||||
let base_url = BackendOrigin::parse(&base_url)
|
||||
.map(|origin| origin.to_string())
|
||||
.unwrap_or(base_url);
|
||||
Self {
|
||||
base_url: base_url.into(),
|
||||
base_url,
|
||||
workspace_id: workspace_id.map(Into::into),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn authenticated_client(&self) -> Result<BackendApiClient, BackendApiClientError> {
|
||||
BackendApiClient::from_stored_token(&self.base_url)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -108,34 +101,31 @@ impl WorkerConnectionSelector {
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerSpawn {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
pub state_dir: PathBuf,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerByName {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
pub struct StandaloneSessionListIntent {
|
||||
pub state_dir: PathBuf,
|
||||
pub cwd: PathBuf,
|
||||
pub include_all: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerResume {
|
||||
pub runtime_command: WorkerRuntimeCommand,
|
||||
pub struct StandaloneSessionResumeIntent {
|
||||
pub state_dir: PathBuf,
|
||||
pub session_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum Dashboard {
|
||||
Local {
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
},
|
||||
Backend {
|
||||
base_url: String,
|
||||
workspace_id: String,
|
||||
},
|
||||
pub struct Dashboard {
|
||||
pub base_url: String,
|
||||
pub workspace_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct WorkerList {
|
||||
pub local_runtime_command: Option<WorkerRuntimeCommand>,
|
||||
pub backend_target: Option<BackendRuntimeListTarget>,
|
||||
pub backend_target: BackendRuntimeListTarget,
|
||||
pub include_stopped: bool,
|
||||
}
|
||||
|
||||
@@ -161,12 +151,6 @@ impl TargetError {
|
||||
message: format!("invalid {target} target: {}", message.into()),
|
||||
}
|
||||
}
|
||||
|
||||
fn local_runtime_command(error: std::io::Error) -> Self {
|
||||
Self {
|
||||
message: format!("failed to resolve local Worker runtime command: {error}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for TargetError {
|
||||
@@ -183,71 +167,40 @@ pub trait Target: fmt::Debug + Send + Sync {
|
||||
/// Resolve the target once for Workspace product-state operations.
|
||||
///
|
||||
/// Backend targets must carry an explicit Workspace identity. Callers use
|
||||
/// this value instead of rediscovering Backend/local authority from cwd or
|
||||
/// process configuration after command dispatch.
|
||||
/// this value instead of rediscovering authority from cwd or process
|
||||
/// configuration after command dispatch.
|
||||
fn resolve(&self) -> Result<ResolvedTarget, TargetError>;
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError>;
|
||||
|
||||
fn worker_by_name(&self) -> Result<WorkerByName, TargetError>;
|
||||
|
||||
fn resume_worker(&self) -> Result<WorkerResume, TargetError>;
|
||||
|
||||
fn dashboard(&self) -> Result<Dashboard, TargetError>;
|
||||
|
||||
fn list_workers(&self, request: WorkerListRequest) -> Result<WorkerList, TargetError>;
|
||||
|
||||
fn connect_worker(
|
||||
&self,
|
||||
selector: WorkerConnectionSelector,
|
||||
) -> Result<WorkerConnection, TargetError>;
|
||||
}
|
||||
|
||||
impl Target for LocalTarget {
|
||||
fn kind(&self) -> TargetKind {
|
||||
TargetKind::Local
|
||||
}
|
||||
|
||||
fn resolve(&self) -> Result<ResolvedTarget, TargetError> {
|
||||
Ok(ResolvedTarget::Local)
|
||||
}
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError> {
|
||||
Ok(WorkerSpawn {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
Err(TargetError::unsupported("Worker spawn", self.kind()))
|
||||
}
|
||||
|
||||
fn worker_by_name(&self) -> Result<WorkerByName, TargetError> {
|
||||
Ok(WorkerByName {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
fn standalone_session_list(
|
||||
&self,
|
||||
_include_all: bool,
|
||||
) -> Result<StandaloneSessionListIntent, TargetError> {
|
||||
Err(TargetError::unsupported(
|
||||
"standalone session listing",
|
||||
self.kind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn resume_worker(&self) -> Result<WorkerResume, TargetError> {
|
||||
Ok(WorkerResume {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
fn standalone_session_resume(
|
||||
&self,
|
||||
_session_id: String,
|
||||
) -> Result<StandaloneSessionResumeIntent, TargetError> {
|
||||
Err(TargetError::unsupported(
|
||||
"standalone session restore",
|
||||
self.kind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn dashboard(&self) -> Result<Dashboard, TargetError> {
|
||||
Ok(Dashboard::Local {
|
||||
runtime_command: self.runtime_command()?,
|
||||
})
|
||||
Err(TargetError::unsupported("Worker dashboard", self.kind()))
|
||||
}
|
||||
|
||||
fn list_workers(&self, request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
if request.runtime_id.is_some() {
|
||||
return Err(TargetError::unsupported(
|
||||
"Explicit runtime id for local worker listing",
|
||||
self.kind(),
|
||||
));
|
||||
}
|
||||
Ok(WorkerList {
|
||||
local_runtime_command: Some(self.runtime_command()?),
|
||||
backend_target: None,
|
||||
include_stopped: request.include_stopped,
|
||||
})
|
||||
fn list_workers(&self, _request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
Err(TargetError::unsupported("Worker listing", self.kind()))
|
||||
}
|
||||
|
||||
fn connect_worker(
|
||||
@@ -261,6 +214,59 @@ impl Target for LocalTarget {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct StandaloneTarget {
|
||||
state_dir: PathBuf,
|
||||
}
|
||||
|
||||
impl StandaloneTarget {
|
||||
#[must_use]
|
||||
pub fn new(state_dir: impl Into<PathBuf>) -> Self {
|
||||
Self {
|
||||
state_dir: state_dir.into(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Target for StandaloneTarget {
|
||||
fn kind(&self) -> TargetKind {
|
||||
TargetKind::Standalone
|
||||
}
|
||||
|
||||
fn resolve(&self) -> Result<ResolvedTarget, TargetError> {
|
||||
Ok(ResolvedTarget::Standalone)
|
||||
}
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError> {
|
||||
Ok(WorkerSpawn {
|
||||
state_dir: self.state_dir.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
fn standalone_session_list(
|
||||
&self,
|
||||
include_all: bool,
|
||||
) -> Result<StandaloneSessionListIntent, TargetError> {
|
||||
let cwd = std::env::current_dir()
|
||||
.map_err(|error| TargetError::invalid(self.kind(), error.to_string()))?;
|
||||
Ok(StandaloneSessionListIntent {
|
||||
state_dir: self.state_dir.clone(),
|
||||
cwd,
|
||||
include_all,
|
||||
})
|
||||
}
|
||||
|
||||
fn standalone_session_resume(
|
||||
&self,
|
||||
session_id: String,
|
||||
) -> Result<StandaloneSessionResumeIntent, TargetError> {
|
||||
Ok(StandaloneSessionResumeIntent {
|
||||
state_dir: self.state_dir.clone(),
|
||||
session_id,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Target for BackendTarget {
|
||||
fn kind(&self) -> TargetKind {
|
||||
TargetKind::Backend
|
||||
@@ -279,42 +285,27 @@ impl Target for BackendTarget {
|
||||
})
|
||||
}
|
||||
|
||||
fn spawn_worker(&self) -> Result<WorkerSpawn, TargetError> {
|
||||
Err(TargetError::unsupported("Worker spawn", self.kind()))
|
||||
}
|
||||
|
||||
fn worker_by_name(&self) -> Result<WorkerByName, TargetError> {
|
||||
Err(TargetError::unsupported(
|
||||
"Worker name attachment",
|
||||
self.kind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn resume_worker(&self) -> Result<WorkerResume, TargetError> {
|
||||
Err(TargetError::unsupported("Worker resume", self.kind()))
|
||||
}
|
||||
|
||||
fn dashboard(&self) -> Result<Dashboard, TargetError> {
|
||||
match self.resolve()? {
|
||||
ResolvedTarget::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
} => Ok(Dashboard::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
}),
|
||||
ResolvedTarget::Local => unreachable!("BackendTarget cannot resolve as Local"),
|
||||
}
|
||||
let ResolvedTarget::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
} = self.resolve()?
|
||||
else {
|
||||
unreachable!("BackendTarget resolves only Backend authority")
|
||||
};
|
||||
Ok(Dashboard {
|
||||
base_url,
|
||||
workspace_id,
|
||||
})
|
||||
}
|
||||
|
||||
fn list_workers(&self, request: WorkerListRequest) -> Result<WorkerList, TargetError> {
|
||||
Ok(WorkerList {
|
||||
local_runtime_command: None,
|
||||
backend_target: Some(BackendRuntimeListTarget::new(
|
||||
backend_target: BackendRuntimeListTarget::new(
|
||||
self.base_url.clone(),
|
||||
self.workspace_id.clone(),
|
||||
request.runtime_id,
|
||||
)),
|
||||
),
|
||||
include_stopped: request.include_stopped,
|
||||
})
|
||||
}
|
||||
@@ -371,8 +362,34 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_target_resolves_local_product_state_authority() {
|
||||
assert_eq!(LocalTarget::new().resolve().unwrap(), ResolvedTarget::Local);
|
||||
fn standalone_target_carries_in_process_state_without_runtime_command() {
|
||||
let target = StandaloneTarget::new("/tmp/yoi-standalone-state");
|
||||
|
||||
assert_eq!(target.kind(), TargetKind::Standalone);
|
||||
assert_eq!(target.resolve().unwrap(), ResolvedTarget::Standalone);
|
||||
assert_eq!(
|
||||
target.spawn_worker().unwrap(),
|
||||
WorkerSpawn {
|
||||
state_dir: PathBuf::from("/tmp/yoi-standalone-state"),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standalone_target_never_exposes_workspace_worker_operations() {
|
||||
let target = StandaloneTarget::new("/tmp/yoi-standalone-state");
|
||||
|
||||
assert_eq!(
|
||||
target
|
||||
.list_workers(WorkerListRequest::new(None))
|
||||
.unwrap_err()
|
||||
.to_string(),
|
||||
"Worker listing is not supported by Standalone target"
|
||||
);
|
||||
assert_eq!(
|
||||
target.dashboard().unwrap_err().to_string(),
|
||||
"Worker dashboard is not supported by Standalone target"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -381,26 +398,13 @@ mod tests {
|
||||
|
||||
assert_eq!(
|
||||
target.dashboard().unwrap(),
|
||||
Dashboard::Backend {
|
||||
Dashboard {
|
||||
base_url: "http://127.0.0.1:8787".to_string(),
|
||||
workspace_id: "workspace-a".to_string(),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_rejects_dashboard_without_workspace_selection() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", None::<String>);
|
||||
|
||||
assert!(
|
||||
target
|
||||
.dashboard()
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("workspace selection is required")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_builds_worker_list() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", Some("workspace-a"));
|
||||
@@ -408,26 +412,13 @@ mod tests {
|
||||
.list_workers(WorkerListRequest::new(Some("runtime-a".to_string())))
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(workers.backend_target.base_url, "http://127.0.0.1:8787");
|
||||
assert_eq!(
|
||||
workers.backend_target.as_ref().unwrap().base_url,
|
||||
"http://127.0.0.1:8787"
|
||||
);
|
||||
assert_eq!(
|
||||
workers
|
||||
.backend_target
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.workspace_id
|
||||
.as_deref(),
|
||||
workers.backend_target.workspace_id.as_deref(),
|
||||
Some("workspace-a")
|
||||
);
|
||||
assert_eq!(
|
||||
workers
|
||||
.backend_target
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.runtime_id
|
||||
.as_deref(),
|
||||
workers.backend_target.runtime_id.as_deref(),
|
||||
Some("runtime-a")
|
||||
);
|
||||
}
|
||||
@@ -446,41 +437,17 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_rejects_worker_connection_before_workspace_selection() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", None::<String>);
|
||||
let error =
|
||||
match target.connect_worker(WorkerConnectionSelector::new("runtime-a", "worker-b")) {
|
||||
Ok(_) => panic!("unscoped connection must fail"),
|
||||
Err(error) => error,
|
||||
};
|
||||
fn standalone_target_builds_explicit_session_intents() {
|
||||
let target = StandaloneTarget::new("/tmp/yoi-client-sessions");
|
||||
let list = target.standalone_session_list(true).unwrap();
|
||||
assert_eq!(list.state_dir, PathBuf::from("/tmp/yoi-client-sessions"));
|
||||
assert!(list.include_all);
|
||||
assert!(list.cwd.is_absolute());
|
||||
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("workspace selection is required")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_target_rejects_local_worker_operations() {
|
||||
let target = BackendTarget::new("http://127.0.0.1:8787", None::<String>);
|
||||
let err = target.spawn_worker().unwrap_err();
|
||||
|
||||
assert_eq!(
|
||||
err.to_string(),
|
||||
"Worker spawn is not supported by Backend target"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_target_builds_local_worker_list() {
|
||||
let target = LocalTarget::new();
|
||||
let workers = target
|
||||
.list_workers(WorkerListRequest::with_stopped(None))
|
||||
let resume = target
|
||||
.standalone_session_resume("019d1234-0000-7000-8000-000000000000".to_string())
|
||||
.unwrap();
|
||||
|
||||
assert!(workers.local_runtime_command.is_some());
|
||||
assert!(workers.backend_target.is_none());
|
||||
assert!(workers.include_stopped);
|
||||
assert_eq!(resume.state_dir, list.state_dir);
|
||||
assert_eq!(resume.session_id, "019d1234-0000-7000-8000-000000000000");
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -14,7 +14,7 @@ use workspace_api::{
|
||||
TICKET_ORCHESTRATION_PLANS_QUERY_PATH, TICKET_RELATIONS_QUERY_PATH,
|
||||
};
|
||||
|
||||
use crate::BackendWorkspaceClientError;
|
||||
use crate::{BackendApiClient, BackendWorkspaceClientError};
|
||||
|
||||
const DEFAULT_PRODUCT_LIST_LIMIT: usize = 1_000;
|
||||
|
||||
@@ -26,7 +26,7 @@ struct BackendWorkerLaunchOptions {
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct BackendWorkerLaunchRuntime {
|
||||
runtime_id: String,
|
||||
can_spawn_worker: bool,
|
||||
worker_creation_available: bool,
|
||||
working_directory_required: bool,
|
||||
}
|
||||
|
||||
@@ -47,9 +47,9 @@ struct BackendWorkspaceOrchestratorResponse {
|
||||
/// Construction requires both the selected Backend URL and Workspace identity.
|
||||
/// Callers should derive these once from `Target::resolve()` and must not retry
|
||||
/// failed requests against repository-local state.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct BackendWorkspaceProductClient {
|
||||
base_url: String,
|
||||
api: BackendApiClient,
|
||||
workspace_id: String,
|
||||
}
|
||||
|
||||
@@ -58,22 +58,32 @@ impl BackendWorkspaceProductClient {
|
||||
base_url: impl Into<String>,
|
||||
workspace_id: impl Into<String>,
|
||||
) -> Result<Self, BackendWorkspaceClientError> {
|
||||
let base_url = base_url.into().trim_end_matches('/').to_string();
|
||||
if base_url.is_empty() {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend base URL must not be empty".into(),
|
||||
));
|
||||
}
|
||||
let base_url = base_url.into();
|
||||
let api = BackendApiClient::from_stored_token(&base_url)?;
|
||||
let workspace_id = workspace_id.into();
|
||||
if workspace_id.trim().is_empty() {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend Workspace identity must not be empty".into(),
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
base_url,
|
||||
workspace_id,
|
||||
})
|
||||
Ok(Self { api, workspace_id })
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn new_with_access_token(
|
||||
base_url: impl Into<String>,
|
||||
workspace_id: impl Into<String>,
|
||||
access_token: &str,
|
||||
) -> Result<Self, BackendWorkspaceClientError> {
|
||||
let base_url = base_url.into();
|
||||
let api = BackendApiClient::from_access_token_for_test(&base_url, access_token)?;
|
||||
let workspace_id = workspace_id.into();
|
||||
if workspace_id.trim().is_empty() {
|
||||
return Err(BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend Workspace identity must not be empty".into(),
|
||||
));
|
||||
}
|
||||
Ok(Self { api, workspace_id })
|
||||
}
|
||||
|
||||
pub fn workspace_id(&self) -> &str {
|
||||
@@ -261,7 +271,7 @@ impl BackendWorkspaceProductClient {
|
||||
let runtime = options
|
||||
.runtimes
|
||||
.iter()
|
||||
.find(|runtime| runtime.can_spawn_worker && !runtime.working_directory_required)
|
||||
.find(|runtime| runtime.worker_creation_available && !runtime.working_directory_required)
|
||||
.ok_or_else(|| {
|
||||
BackendWorkspaceClientError::InvalidTarget(
|
||||
"Backend has no spawn-capable Runtime that supports a Workdir-less Intake Worker"
|
||||
@@ -316,7 +326,7 @@ impl BackendWorkspaceProductClient {
|
||||
body: Option<&B>,
|
||||
) -> Result<R, BackendWorkspaceClientError> {
|
||||
let response = self.request(method, path, body)?.send()?;
|
||||
let response = ensure_success(response)?;
|
||||
self.api.check_status(response.status())?;
|
||||
response.json().map_err(BackendWorkspaceClientError::Http)
|
||||
}
|
||||
|
||||
@@ -326,7 +336,8 @@ impl BackendWorkspaceProductClient {
|
||||
path: &str,
|
||||
body: Option<&B>,
|
||||
) -> Result<(), BackendWorkspaceClientError> {
|
||||
ensure_success(self.request(method, path, body)?.send()?)?;
|
||||
let response = self.request(method, path, body)?.send()?;
|
||||
self.api.check_status(response.status())?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -336,14 +347,12 @@ impl BackendWorkspaceProductClient {
|
||||
path: &str,
|
||||
body: Option<&B>,
|
||||
) -> Result<reqwest::blocking::RequestBuilder, BackendWorkspaceClientError> {
|
||||
let client = reqwest::blocking::Client::builder().build()?;
|
||||
let url = format!(
|
||||
"{}/api/w/{}/{}",
|
||||
self.base_url,
|
||||
let path = format!(
|
||||
"/api/w/{}/{}",
|
||||
encode_path_segment(&self.workspace_id),
|
||||
path.trim_start_matches('/')
|
||||
);
|
||||
let request = client.request(method, url);
|
||||
let request = self.api.blocking_request(method, &path)?;
|
||||
Ok(match body {
|
||||
Some(body) => request.json(body),
|
||||
None => request,
|
||||
@@ -473,8 +482,12 @@ impl TicketBackend for BackendWorkspaceProductClient {
|
||||
.map_err(ticket_client_error)
|
||||
}
|
||||
|
||||
fn queue_ready(&self, id: TicketIdOrSlug, _queued_by: &str) -> ticket::Result<()> {
|
||||
self.send_unit::<()>(
|
||||
fn queue_ready(
|
||||
&self,
|
||||
id: TicketIdOrSlug,
|
||||
_queued_by: &str,
|
||||
) -> ticket::Result<ticket::TicketQueueOutcome> {
|
||||
self.send_json::<(), _>(
|
||||
Method::POST,
|
||||
&format!(
|
||||
"/tickets/{}/workflow/queue",
|
||||
@@ -584,19 +597,6 @@ fn ticket_client_error(error: BackendWorkspaceClientError) -> TicketError {
|
||||
TicketError::Sqlite(format!("Backend request failed: {error}"))
|
||||
}
|
||||
|
||||
fn ensure_success(
|
||||
response: reqwest::blocking::Response,
|
||||
) -> Result<reqwest::blocking::Response, BackendWorkspaceClientError> {
|
||||
if response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
let status = response.status().as_u16();
|
||||
let message = response
|
||||
.text()
|
||||
.unwrap_or_else(|_| "Backend request failed".to_string());
|
||||
Err(BackendWorkspaceClientError::RequestFailed { status, message })
|
||||
}
|
||||
|
||||
fn ticket_reference(id: &TicketIdOrSlug) -> String {
|
||||
match id {
|
||||
TicketIdOrSlug::Id(id) => id.to_string(),
|
||||
@@ -694,24 +694,32 @@ mod tests {
|
||||
fn objective_list_uses_workspace_scoped_backend_route() {
|
||||
let body = r#"{"workspace_id":"workspace-a","limit":1000,"items":[],"source":"sqlite","diagnostics":[]}"#;
|
||||
let (base_url, request, handle) = one_response_server("200 OK", body);
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let response = client.list_objectives(1_000).unwrap();
|
||||
|
||||
assert!(response.items.is_empty());
|
||||
assert!(
|
||||
request
|
||||
.recv()
|
||||
.unwrap()
|
||||
.starts_with("GET /api/w/workspace-a/objectives?limit=1000 ")
|
||||
);
|
||||
let request = request.recv().unwrap();
|
||||
assert!(request.starts_with("GET /api/w/workspace-a/objectives?limit=1000 "));
|
||||
assert!(request.contains("authorization: Bearer test-backend-token\r\n"));
|
||||
handle.join().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backend_mutation_failure_is_returned_without_local_fallback() {
|
||||
let (base_url, request, handle) = one_response_server("403 Forbidden", "denied");
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let (base_url, request, handle) =
|
||||
one_response_server("403 Forbidden", "test-backend-token");
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let error = client
|
||||
.create_objective(&ObjectiveCreateRequest {
|
||||
@@ -723,6 +731,7 @@ mod tests {
|
||||
.unwrap_err();
|
||||
|
||||
assert!(error.to_string().contains("403"));
|
||||
assert!(!error.to_string().contains("test-backend-token"));
|
||||
assert!(
|
||||
request
|
||||
.recv()
|
||||
@@ -735,7 +744,12 @@ mod tests {
|
||||
#[test]
|
||||
fn ticket_relation_query_uses_workspace_scoped_backend_route() {
|
||||
let (base_url, request, handle) = one_response_server("200 OK", "[]");
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let relations = client
|
||||
.query_ticket_relations(
|
||||
@@ -754,7 +768,12 @@ mod tests {
|
||||
#[test]
|
||||
fn orchestration_plan_query_uses_workspace_scoped_backend_route() {
|
||||
let (base_url, request, handle) = one_response_server("200 OK", "[]");
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let records = TicketBackend::query_orchestration_plan_records(&client, None, None).unwrap();
|
||||
|
||||
@@ -773,14 +792,19 @@ mod tests {
|
||||
let (base_url, requests, handle) = response_sequence_server(vec![
|
||||
(
|
||||
"200 OK",
|
||||
r#"{"runtimes":[{"runtime_id":"embedded","can_spawn_worker":true,"working_directory_required":false}]}"#,
|
||||
r#"{"runtimes":[{"runtime_id":"embedded","worker_creation_available":true,"working_directory_required":false}]}"#,
|
||||
),
|
||||
(
|
||||
"200 OK",
|
||||
r#"{"runtime_id":"embedded","worker_id":"worker-1"}"#,
|
||||
),
|
||||
]);
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let status = client.launch_ticket_intake("T-1").unwrap();
|
||||
|
||||
@@ -802,7 +826,12 @@ mod tests {
|
||||
fn workspace_orchestrator_launch_uses_scoped_backend_route() {
|
||||
let body = r#"{"disposition":"created","worker":{"runtime_id":"embedded","worker_id":"worker-2"}}"#;
|
||||
let (base_url, request, handle) = one_response_server("200 OK", body);
|
||||
let client = BackendWorkspaceProductClient::new(base_url, "workspace-a").unwrap();
|
||||
let client = BackendWorkspaceProductClient::new_with_access_token(
|
||||
base_url,
|
||||
"workspace-a",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let status = client.start_workspace_orchestrator().unwrap();
|
||||
|
||||
@@ -818,7 +847,12 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn product_client_requires_workspace_identity() {
|
||||
let error = BackendWorkspaceProductClient::new("http://127.0.0.1:8787", "").unwrap_err();
|
||||
let error = BackendWorkspaceProductClient::new_with_access_token(
|
||||
"http://127.0.0.1:8787",
|
||||
"",
|
||||
"test-backend-token",
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(error.to_string().contains("Workspace identity"));
|
||||
}
|
||||
|
||||
|
||||
@@ -24,7 +24,7 @@ pub fn builtin_flow_source(slug: &str) -> Option<BuiltinFlowSource> {
|
||||
match slug {
|
||||
CODER_REVIEW_FLOW_SLUG => Some(BuiltinFlowSource {
|
||||
slug: CODER_REVIEW_FLOW_SLUG,
|
||||
revision: 3,
|
||||
revision: 4,
|
||||
path: "builtin/flows/coder-review.dcdl",
|
||||
content: CODER_REVIEW_FLOW_SOURCE,
|
||||
}),
|
||||
@@ -35,7 +35,7 @@ pub fn builtin_flow_source(slug: &str) -> Option<BuiltinFlowSource> {
|
||||
pub fn builtin_flow_sources() -> &'static [BuiltinFlowSource] {
|
||||
const SOURCES: &[BuiltinFlowSource] = &[BuiltinFlowSource {
|
||||
slug: CODER_REVIEW_FLOW_SLUG,
|
||||
revision: 3,
|
||||
revision: 4,
|
||||
path: "builtin/flows/coder-review.dcdl",
|
||||
content: CODER_REVIEW_FLOW_SOURCE,
|
||||
}];
|
||||
@@ -46,6 +46,30 @@ pub fn builtin_flow_sources() -> &'static [BuiltinFlowSource] {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn coder_review_flow_uses_current_selector_ref_review_contract() {
|
||||
let source = builtin_flow_source(CODER_REVIEW_FLOW_SLUG).expect("coder review Flow");
|
||||
for required in [
|
||||
"OpenMergeRequest",
|
||||
"ShowMergeRequest",
|
||||
"ReviewMergeRequest",
|
||||
"CompleteMergeRequest",
|
||||
"existing Merge Request `selector_from`",
|
||||
"Target-only movement does not invalidate",
|
||||
] {
|
||||
assert!(source.content.contains(required), "missing {required}");
|
||||
}
|
||||
for stale in [
|
||||
"MergeRequestOpen",
|
||||
"MergeRequestShow",
|
||||
"MergeRequestReview",
|
||||
"MergeRequestComplete",
|
||||
"new immutable revision",
|
||||
] {
|
||||
assert!(!source.content.contains(stale), "stale contract {stale}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_builtin_flow_compiles_and_matches_catalog_identity() {
|
||||
assert!(!builtin_flow_sources().is_empty());
|
||||
|
||||
@@ -157,6 +157,22 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn grep_request(path: &str, pattern: &str) -> GrepRequest {
|
||||
GrepRequest {
|
||||
pattern: pattern.to_string(),
|
||||
path: FsPath::new(path).unwrap(),
|
||||
glob: None,
|
||||
file_type: None,
|
||||
case_insensitive: false,
|
||||
before_context: 0,
|
||||
after_context: 0,
|
||||
multiline: false,
|
||||
output_mode: GrepOutputMode::Content,
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn logical_paths_reject_absolute_parent_and_backslash_forms() {
|
||||
assert!(FsPath::new("src/lib.rs").is_ok());
|
||||
@@ -279,4 +295,313 @@ mod tests {
|
||||
assert_eq!(grep.matched_files, 2);
|
||||
assert!(!grep.output.contains("c.txt"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_accepts_a_direct_file_without_searching_siblings() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let selected = temp.path().join("selected.txt");
|
||||
std::fs::write(&selected, "before\nneedle selected\nafter\n").unwrap();
|
||||
std::fs::write(temp.path().join("sibling.txt"), "needle sibling\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let mut request = grep_request("selected.txt", "needle");
|
||||
request.before_context = 1;
|
||||
request.after_context = 1;
|
||||
let direct = run_grep(&root, selected, request, &readable).unwrap();
|
||||
|
||||
assert_eq!(direct.match_count, 1);
|
||||
assert_eq!(direct.matched_files, 1);
|
||||
assert_eq!(
|
||||
direct.output,
|
||||
concat!(
|
||||
"selected.txt\n",
|
||||
" 1 │ before\n",
|
||||
" > 2 │ needle selected\n",
|
||||
" 3 │ after\n",
|
||||
)
|
||||
);
|
||||
assert!(!direct.output.contains("sibling"));
|
||||
|
||||
let directory = run_grep(
|
||||
&root,
|
||||
root.clone(),
|
||||
GrepRequest {
|
||||
pattern: "needle".to_string(),
|
||||
path: FsPath::root(),
|
||||
glob: None,
|
||||
file_type: None,
|
||||
case_insensitive: false,
|
||||
before_context: 0,
|
||||
after_context: 0,
|
||||
multiline: false,
|
||||
output_mode: GrepOutputMode::Content,
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
},
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(directory.match_count, 2);
|
||||
assert_eq!(directory.matched_files, 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_direct_file_applies_glob_and_type_filters_for_every_output_mode() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let nested = temp.path().join("nested");
|
||||
std::fs::create_dir(&nested).unwrap();
|
||||
let selected = nested.join("selected.rs");
|
||||
std::fs::write(&selected, "needle one\nneedle two\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
for mode in [
|
||||
GrepOutputMode::Content,
|
||||
GrepOutputMode::FilesWithMatches,
|
||||
GrepOutputMode::Count,
|
||||
] {
|
||||
for (glob, file_type) in [(Some("other/*.rs"), None), (None, Some("python"))] {
|
||||
let mut request = grep_request("nested/selected.rs", "needle");
|
||||
request.output_mode = mode;
|
||||
request.glob = glob.map(str::to_string);
|
||||
request.file_type = file_type.map(str::to_string);
|
||||
|
||||
let excluded = run_grep(&root, selected.clone(), request, &readable).unwrap();
|
||||
assert_eq!(excluded.output, "", "mode {mode:?}");
|
||||
assert_eq!(excluded.match_count, 0, "mode {mode:?}");
|
||||
assert_eq!(excluded.matched_files, 0, "mode {mode:?}");
|
||||
assert!(!excluded.truncated, "mode {mode:?}");
|
||||
}
|
||||
|
||||
let mut request = grep_request("nested/selected.rs", "needle");
|
||||
request.output_mode = mode;
|
||||
request.glob = Some("nested/*.rs".to_string());
|
||||
request.file_type = Some("rust".to_string());
|
||||
let matched = run_grep(&root, selected.clone(), request, &readable).unwrap();
|
||||
|
||||
match mode {
|
||||
GrepOutputMode::Content => {
|
||||
assert_eq!(matched.match_count, 2);
|
||||
assert_eq!(matched.matched_files, 1);
|
||||
assert!(matched.output.starts_with("nested/selected.rs\n"));
|
||||
assert!(matched.output.contains("> 1 │ needle one"));
|
||||
assert!(matched.output.contains("> 2 │ needle two"));
|
||||
}
|
||||
GrepOutputMode::FilesWithMatches => {
|
||||
assert_eq!(matched.match_count, 1);
|
||||
assert_eq!(matched.matched_files, 1);
|
||||
assert_eq!(matched.output, "nested/selected.rs\n");
|
||||
}
|
||||
GrepOutputMode::Count => {
|
||||
assert_eq!(matched.match_count, 2);
|
||||
assert_eq!(matched.matched_files, 1);
|
||||
assert_eq!(matched.output, "nested/selected.rs:2\n");
|
||||
}
|
||||
}
|
||||
assert!(!matched.truncated, "mode {mode:?}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_direct_file_preserves_explicit_hidden_and_gitignored_behavior() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let hidden = temp.path().join(".hidden.rs");
|
||||
let ignored = temp.path().join("ignored.rs");
|
||||
std::fs::write(&hidden, "needle hidden\n").unwrap();
|
||||
std::fs::write(&ignored, "needle ignored\n").unwrap();
|
||||
std::fs::write(temp.path().join(".gitignore"), "ignored.rs\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
for (path, expected) in [
|
||||
(".hidden.rs", "needle hidden"),
|
||||
("ignored.rs", "needle ignored"),
|
||||
] {
|
||||
let result = run_grep(
|
||||
&root,
|
||||
root.join(path),
|
||||
grep_request(path, "needle"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(result.match_count, 1, "path {path}");
|
||||
assert!(result.output.contains(expected), "path {path}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_direct_file_preserves_case_multiline_and_bounds() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let selected = temp.path().join("selected.txt");
|
||||
std::fs::write(&selected, "NEEDLE first\nstart\nfinish\nneedle last\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let mut case_request = grep_request("selected.txt", "needle");
|
||||
case_request.case_insensitive = true;
|
||||
case_request.offset = 1;
|
||||
case_request.limit = 1;
|
||||
let bounded = run_grep(&root, selected.clone(), case_request, &readable).unwrap();
|
||||
assert_eq!(bounded.match_count, 1);
|
||||
assert!(!bounded.output.contains("NEEDLE first"));
|
||||
assert!(bounded.output.contains("needle last"));
|
||||
assert!(bounded.truncated);
|
||||
|
||||
let mut multiline_request = grep_request("selected.txt", "start\\nfinish");
|
||||
multiline_request.multiline = true;
|
||||
let multiline = run_grep(&root, selected, multiline_request, &readable).unwrap();
|
||||
assert_eq!(multiline.match_count, 1);
|
||||
assert!(multiline.output.contains("start\nfinish"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_returns_not_found_for_a_missing_direct_path() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let missing = root.join("missing.txt");
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let error = run_grep(
|
||||
&root,
|
||||
missing.clone(),
|
||||
grep_request("missing.txt", "needle"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, FsError::NotFound(path) if path == missing));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn grep_keeps_direct_symlink_directory_and_broken_path_guards() {
|
||||
use std::os::unix::fs::symlink;
|
||||
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
std::fs::create_dir(root.join("target-dir")).unwrap();
|
||||
std::fs::write(root.join("target-file.rs"), "needle file\n").unwrap();
|
||||
symlink(root.join("target-file.rs"), root.join("file-link.rs")).unwrap();
|
||||
symlink(root.join("target-dir"), root.join("directory-link")).unwrap();
|
||||
symlink(root.join("missing-target"), root.join("broken-link")).unwrap();
|
||||
|
||||
let request = |path: &str| grep_request(path, "needle");
|
||||
|
||||
let file_result = run_grep(
|
||||
&root,
|
||||
root.join("file-link.rs"),
|
||||
request("file-link.rs"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(file_result.match_count, 1);
|
||||
assert!(file_result.output.starts_with("file-link.rs\n"));
|
||||
|
||||
let directory_error = run_grep(
|
||||
&root,
|
||||
root.join("directory-link"),
|
||||
request("directory-link"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
directory_error,
|
||||
FsError::SymlinkDirectoryNotTraversed { tool: "Grep", path, .. }
|
||||
if path == root.join("directory-link")
|
||||
));
|
||||
|
||||
let broken_error = run_grep(
|
||||
&root,
|
||||
root.join("broken-link"),
|
||||
request("broken-link"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
broken_error,
|
||||
FsError::BrokenSymlink { path, .. } if path == root.join("broken-link")
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[test]
|
||||
fn grep_rejects_a_direct_special_file_as_invalid_argument() {
|
||||
use std::os::unix::net::UnixListener;
|
||||
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let socket = temp.path().join("grep.sock");
|
||||
let _listener = UnixListener::bind(&socket).unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let error = run_grep(
|
||||
&root,
|
||||
socket,
|
||||
grep_request("grep.sock", "needle"),
|
||||
&readable,
|
||||
)
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
FsError::InvalidArgument(message)
|
||||
if message.contains("must be a regular file or directory")
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grep_content_groups_lines_by_file_and_marks_matches() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
std::fs::write(
|
||||
temp.path().join("first.txt"),
|
||||
"before\nneedle one\nafter\nomitted one\nomitted two\nbefore distant\nneedle distant\nafter distant\n",
|
||||
)
|
||||
.unwrap();
|
||||
std::fs::write(temp.path().join("second.txt"), "needle two\n").unwrap();
|
||||
let root = temp.path().canonicalize().unwrap();
|
||||
let readable = RootAccess(root.clone());
|
||||
|
||||
let grep = run_grep(
|
||||
&root,
|
||||
root.clone(),
|
||||
GrepRequest {
|
||||
pattern: "needle".to_string(),
|
||||
path: FsPath::root(),
|
||||
glob: Some("*.txt".to_string()),
|
||||
output_mode: GrepOutputMode::Content,
|
||||
case_insensitive: false,
|
||||
before_context: 1,
|
||||
after_context: 1,
|
||||
multiline: false,
|
||||
file_type: None,
|
||||
limit: 20,
|
||||
offset: 0,
|
||||
},
|
||||
&readable,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(grep.match_count, 3);
|
||||
assert_eq!(grep.matched_files, 2);
|
||||
assert_eq!(
|
||||
grep.output,
|
||||
concat!(
|
||||
"first.txt\n",
|
||||
" 1 │ before\n",
|
||||
" > 2 │ needle one\n",
|
||||
" 3 │ after\n",
|
||||
" …\n",
|
||||
" 6 │ before distant\n",
|
||||
" > 7 │ needle distant\n",
|
||||
" 8 │ after distant\n",
|
||||
"\n",
|
||||
"second.txt\n",
|
||||
" > 1 │ needle two\n",
|
||||
)
|
||||
);
|
||||
assert_eq!(grep.output.matches("first.txt").count(), 1);
|
||||
assert_eq!(grep.output.matches("second.txt").count(), 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fmt::Write as _;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use crate::FsAccessPolicy;
|
||||
@@ -5,8 +7,8 @@ use grep_regex::RegexMatcherBuilder;
|
||||
use grep_searcher::sinks::UTF8 as UTF8Sink;
|
||||
use grep_searcher::{BinaryDetection, Searcher, SearcherBuilder, Sink, SinkContext, SinkMatch};
|
||||
use ignore::WalkBuilder;
|
||||
use ignore::overrides::OverrideBuilder;
|
||||
use ignore::types::TypesBuilder;
|
||||
use ignore::overrides::{Override, OverrideBuilder};
|
||||
use ignore::types::{Types, TypesBuilder};
|
||||
|
||||
use crate::{FsError, GrepOutputMode, GrepRequest, GrepResult, direct_symlink};
|
||||
|
||||
@@ -57,20 +59,11 @@ impl GrepReport {
|
||||
}
|
||||
}
|
||||
GrepOutputMode::Content => {
|
||||
for line in &self.lines {
|
||||
let separator = if line.is_match { ':' } else { '-' };
|
||||
let path = logical_display(root, &line.path);
|
||||
if self.show_line_numbers
|
||||
&& let Some(number) = line.line_number
|
||||
{
|
||||
output.push_str(&format!(
|
||||
"{path}{separator}{number}{separator}{}\n",
|
||||
line.text
|
||||
));
|
||||
} else {
|
||||
output.push_str(&format!("{path}{separator}{}\n", line.text));
|
||||
}
|
||||
}
|
||||
output.push_str(&render_content_lines(
|
||||
root,
|
||||
&self.lines,
|
||||
self.show_line_numbers,
|
||||
));
|
||||
}
|
||||
}
|
||||
GrepResult {
|
||||
@@ -82,6 +75,48 @@ impl GrepReport {
|
||||
}
|
||||
}
|
||||
|
||||
fn render_content_lines(root: &Path, lines: &[ContentLine], show_line_numbers: bool) -> String {
|
||||
let mut grouped = BTreeMap::<&Path, Vec<&ContentLine>>::new();
|
||||
for line in lines {
|
||||
grouped.entry(&line.path).or_default().push(line);
|
||||
}
|
||||
|
||||
let mut output = String::new();
|
||||
for (file_index, (path, file_lines)) in grouped.into_iter().enumerate() {
|
||||
if file_index > 0 {
|
||||
output.push('\n');
|
||||
}
|
||||
let _ = writeln!(output, "{}", logical_display(root, path));
|
||||
|
||||
let number_width = file_lines
|
||||
.iter()
|
||||
.filter_map(|line| line.line_number)
|
||||
.map(|number| number.to_string().len())
|
||||
.max()
|
||||
.unwrap_or(1);
|
||||
let mut previous_line_end = None;
|
||||
for line in file_lines {
|
||||
if let (Some(previous_end), Some(number)) = (previous_line_end, line.line_number)
|
||||
&& number > previous_end
|
||||
{
|
||||
let _ = writeln!(output, " …");
|
||||
}
|
||||
|
||||
let marker = if line.is_match { '>' } else { ' ' };
|
||||
if show_line_numbers && let Some(number) = line.line_number {
|
||||
let _ = writeln!(output, " {marker} {number:>number_width$} │ {}", line.text);
|
||||
} else {
|
||||
let _ = writeln!(output, " {marker} │ {}", line.text);
|
||||
}
|
||||
previous_line_end = line
|
||||
.line_number
|
||||
.map(|number| number + line.text.split('\n').count() as u64);
|
||||
}
|
||||
}
|
||||
|
||||
output
|
||||
}
|
||||
|
||||
fn logical_display(root: &Path, path: &Path) -> String {
|
||||
path.strip_prefix(root)
|
||||
.unwrap_or(path)
|
||||
@@ -91,6 +126,38 @@ fn logical_display(root: &Path, path: &Path) -> String {
|
||||
|
||||
const DEFAULT_HEAD_LIMIT: usize = 250;
|
||||
|
||||
fn build_overrides(base: &Path, glob: Option<&str>) -> Result<Option<Override>, FsError> {
|
||||
let Some(glob) = glob else {
|
||||
return Ok(None);
|
||||
};
|
||||
let mut builder = OverrideBuilder::new(base);
|
||||
builder
|
||||
.add(glob)
|
||||
.map_err(|error| FsError::InvalidGlob(error.to_string()))?;
|
||||
builder
|
||||
.build()
|
||||
.map(Some)
|
||||
.map_err(|error| FsError::InvalidGlob(error.to_string()))
|
||||
}
|
||||
|
||||
fn build_types(file_type: Option<&str>) -> Result<Option<Types>, FsError> {
|
||||
let Some(file_type) = file_type else {
|
||||
return Ok(None);
|
||||
};
|
||||
let mut builder = TypesBuilder::new();
|
||||
builder.add_defaults();
|
||||
builder.select(file_type);
|
||||
builder
|
||||
.build()
|
||||
.map(Some)
|
||||
.map_err(|error| FsError::InvalidArgument(format!("invalid type {file_type}: {error}")))
|
||||
}
|
||||
|
||||
fn direct_file_selected(path: &Path, overrides: Option<&Override>, types: Option<&Types>) -> bool {
|
||||
!overrides.is_some_and(|filter| filter.matched(path, false).is_ignore())
|
||||
&& !types.is_some_and(|filter| filter.matched(path, false).is_ignore())
|
||||
}
|
||||
|
||||
struct GrepParams {
|
||||
pattern: String,
|
||||
path: Option<PathBuf>,
|
||||
@@ -186,13 +253,15 @@ pub fn run_grep(
|
||||
std::io::ErrorKind::NotFound => FsError::NotFound(base.clone()),
|
||||
_ => FsError::io(&base, e),
|
||||
})?;
|
||||
if !base_meta.is_dir() {
|
||||
if !base_meta.is_file() && !base_meta.is_dir() {
|
||||
return Err(FsError::InvalidArgument(format!(
|
||||
"grep search path is not a directory: {}",
|
||||
"grep search path must be a regular file or directory: {}",
|
||||
base.display()
|
||||
)));
|
||||
}
|
||||
if let Some(info) = symlink.as_ref() {
|
||||
if base_meta.is_dir()
|
||||
&& let Some(info) = symlink.as_ref()
|
||||
{
|
||||
return Err(FsError::SymlinkDirectoryNotTraversed {
|
||||
tool: "Grep",
|
||||
path: base.clone(),
|
||||
@@ -200,32 +269,9 @@ pub fn run_grep(
|
||||
});
|
||||
}
|
||||
|
||||
let mut wb = WalkBuilder::new(&base);
|
||||
wb.hidden(true)
|
||||
.git_ignore(true)
|
||||
.git_global(true)
|
||||
.git_exclude(true)
|
||||
.ignore(true)
|
||||
.parents(true)
|
||||
.follow_links(false);
|
||||
|
||||
if let Some(t) = p.file_type.as_deref() {
|
||||
let mut tb = TypesBuilder::new();
|
||||
tb.add_defaults();
|
||||
tb.select(t);
|
||||
let types = tb
|
||||
.build()
|
||||
.map_err(|e| FsError::InvalidArgument(format!("invalid type {t}: {e}")))?;
|
||||
wb.types(types);
|
||||
}
|
||||
if let Some(g) = p.glob.as_deref() {
|
||||
let mut ob = OverrideBuilder::new(&base);
|
||||
ob.add(g).map_err(|e| FsError::InvalidGlob(e.to_string()))?;
|
||||
let ov = ob
|
||||
.build()
|
||||
.map_err(|e| FsError::InvalidGlob(e.to_string()))?;
|
||||
wb.overrides(ov);
|
||||
}
|
||||
let filter_base = if base_meta.is_file() { root } else { &base };
|
||||
let types = build_types(p.file_type.as_deref())?;
|
||||
let overrides = build_overrides(filter_base, p.glob.as_deref())?;
|
||||
|
||||
let mode = p.output_mode.unwrap_or_default();
|
||||
let head_limit = p.head_limit.unwrap_or(DEFAULT_HEAD_LIMIT);
|
||||
@@ -240,74 +286,133 @@ pub fn run_grep(
|
||||
lines: Vec::new(),
|
||||
truncated: false,
|
||||
};
|
||||
let mut matching_files_seen = 0;
|
||||
let mut matches_seen = 0;
|
||||
|
||||
// Per-mode walker state.
|
||||
let mut matching_files_seen: usize = 0;
|
||||
let mut matches_seen: usize = 0;
|
||||
if base_meta.is_file() {
|
||||
if direct_file_selected(&base, overrides.as_ref(), types.as_ref()) {
|
||||
scan_path(
|
||||
&mut searcher,
|
||||
&matcher,
|
||||
&base,
|
||||
mode,
|
||||
&mut report,
|
||||
&mut matching_files_seen,
|
||||
&mut matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
)?;
|
||||
}
|
||||
return Ok(report.into_result(root));
|
||||
}
|
||||
|
||||
'walker: for entry in wb.build().flatten() {
|
||||
if !entry.file_type().map(|t| t.is_file()).unwrap_or(false) {
|
||||
let mut walker = WalkBuilder::new(&base);
|
||||
walker
|
||||
.hidden(true)
|
||||
.git_ignore(true)
|
||||
.git_global(true)
|
||||
.git_exclude(true)
|
||||
.ignore(true)
|
||||
.parents(true)
|
||||
.follow_links(false);
|
||||
if let Some(types) = types {
|
||||
walker.types(types);
|
||||
}
|
||||
if let Some(overrides) = overrides {
|
||||
walker.overrides(overrides);
|
||||
}
|
||||
|
||||
for entry in walker.build().flatten() {
|
||||
if !entry
|
||||
.file_type()
|
||||
.map(|kind| kind.is_file())
|
||||
.unwrap_or(false)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let path = entry.path();
|
||||
if !access.is_readable(path) {
|
||||
continue;
|
||||
}
|
||||
|
||||
match mode {
|
||||
GrepOutputMode::FilesWithMatches => {
|
||||
let hit = scan_any_match(&mut searcher, &matcher, path)?;
|
||||
if !hit {
|
||||
continue;
|
||||
}
|
||||
if matching_files_seen >= offset {
|
||||
report.files.push(path.to_path_buf());
|
||||
if report.files.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
break 'walker;
|
||||
}
|
||||
}
|
||||
matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Count => {
|
||||
let count = scan_count(&mut searcher, &matcher, path)?;
|
||||
if count == 0 {
|
||||
continue;
|
||||
}
|
||||
if matching_files_seen >= offset {
|
||||
report.counts.push((path.to_path_buf(), count));
|
||||
if report.counts.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
break 'walker;
|
||||
}
|
||||
}
|
||||
matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Content => {
|
||||
let before_count = matches_seen;
|
||||
let mut sink = ContentSink {
|
||||
path: path.to_path_buf(),
|
||||
lines: &mut report.lines,
|
||||
matches_seen: &mut matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
};
|
||||
searcher
|
||||
.search_path(&matcher, path, &mut sink)
|
||||
.map_err(|e| FsError::io(path, e))?;
|
||||
// If we hit head_limit during this file, stop walking.
|
||||
if matches_seen >= offset.saturating_add(head_limit) && matches_seen > before_count
|
||||
{
|
||||
report.truncated = true;
|
||||
break 'walker;
|
||||
}
|
||||
}
|
||||
if scan_path(
|
||||
&mut searcher,
|
||||
&matcher,
|
||||
path,
|
||||
mode,
|
||||
&mut report,
|
||||
&mut matching_files_seen,
|
||||
&mut matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
)? {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(report.into_result(root))
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn scan_path(
|
||||
searcher: &mut Searcher,
|
||||
matcher: &grep_regex::RegexMatcher,
|
||||
path: &Path,
|
||||
mode: GrepOutputMode,
|
||||
report: &mut GrepReport,
|
||||
matching_files_seen: &mut usize,
|
||||
matches_seen: &mut usize,
|
||||
offset: usize,
|
||||
head_limit: usize,
|
||||
) -> Result<bool, FsError> {
|
||||
match mode {
|
||||
GrepOutputMode::FilesWithMatches => {
|
||||
if !scan_any_match(searcher, matcher, path)? {
|
||||
return Ok(false);
|
||||
}
|
||||
if *matching_files_seen >= offset {
|
||||
report.files.push(path.to_path_buf());
|
||||
if report.files.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
*matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Count => {
|
||||
let count = scan_count(searcher, matcher, path)?;
|
||||
if count == 0 {
|
||||
return Ok(false);
|
||||
}
|
||||
if *matching_files_seen >= offset {
|
||||
report.counts.push((path.to_path_buf(), count));
|
||||
if report.counts.len() >= head_limit {
|
||||
report.truncated = true;
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
*matching_files_seen += 1;
|
||||
}
|
||||
GrepOutputMode::Content => {
|
||||
let before_count = *matches_seen;
|
||||
let mut sink = ContentSink {
|
||||
path: path.to_path_buf(),
|
||||
lines: &mut report.lines,
|
||||
matches_seen,
|
||||
offset,
|
||||
head_limit,
|
||||
};
|
||||
searcher
|
||||
.search_path(matcher, path, &mut sink)
|
||||
.map_err(|error| FsError::io(path, error))?;
|
||||
if *matches_seen >= offset.saturating_add(head_limit) && *matches_seen > before_count {
|
||||
report.truncated = true;
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
fn scan_any_match(
|
||||
searcher: &mut Searcher,
|
||||
matcher: &grep_regex::RegexMatcher,
|
||||
|
||||
@@ -7,6 +7,7 @@ license.workspace = true
|
||||
[dependencies]
|
||||
arc-swap = "1"
|
||||
agen = { workspace = true }
|
||||
decodal.workspace = true
|
||||
protocol = { workspace = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
serde_json = { workspace = true }
|
||||
|
||||
@@ -0,0 +1,318 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use decodal::{Data, Engine, ImportLoader, LoadedImport};
|
||||
use serde_json::{Map, Number, Value};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use crate::profile::ProfileError;
|
||||
|
||||
pub const BUILTIN_PROFILE_CATALOG_ID: &str = "builtin-profiles-v2";
|
||||
pub const BUILTIN_DEFAULT_PROFILE: &str = "builtin:default";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct BuiltinProfileImport {
|
||||
pub specifier: &'static str,
|
||||
pub resolved_path: &'static str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct BuiltinProfileResource {
|
||||
pub selector: Option<&'static str>,
|
||||
pub path: &'static str,
|
||||
pub source: &'static str,
|
||||
pub description: &'static str,
|
||||
pub imports: &'static [BuiltinProfileImport],
|
||||
}
|
||||
|
||||
const BASE_PATH: &str = "profiles/base.dcdl";
|
||||
const BASE_IMPORT: &[BuiltinProfileImport] = &[BuiltinProfileImport {
|
||||
specifier: "./base.dcdl",
|
||||
resolved_path: BASE_PATH,
|
||||
}];
|
||||
const NO_IMPORTS: &[BuiltinProfileImport] = &[];
|
||||
|
||||
pub const BUILTIN_PROFILE_RESOURCES: &[BuiltinProfileResource] = &[
|
||||
BuiltinProfileResource {
|
||||
selector: None,
|
||||
path: BASE_PATH,
|
||||
source: include_str!("../../../resources/profiles/base.dcdl"),
|
||||
description: "Shared built-in Profile defaults.",
|
||||
imports: NO_IMPORTS,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some(BUILTIN_DEFAULT_PROFILE),
|
||||
path: "profiles/default.dcdl",
|
||||
source: include_str!("../../../resources/profiles/default.dcdl"),
|
||||
description: "Standalone Yoi coding profile.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:coder"),
|
||||
path: "profiles/coder.dcdl",
|
||||
source: include_str!("../../../resources/profiles/coder.dcdl"),
|
||||
description: "Ticket implementation with direct Reviewer SubWorkers.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:companion"),
|
||||
path: "profiles/companion.dcdl",
|
||||
source: include_str!("../../../resources/profiles/companion.dcdl"),
|
||||
description: "General assistance with Workspace tools.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:intake"),
|
||||
path: "profiles/intake.dcdl",
|
||||
source: include_str!("../../../resources/profiles/intake.dcdl"),
|
||||
description: "Read-only intake and planning.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:reviewer"),
|
||||
path: "profiles/reviewer.dcdl",
|
||||
source: include_str!("../../../resources/profiles/reviewer.dcdl"),
|
||||
description: "Independent review of a published Merge Request source.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:orchestrator"),
|
||||
path: "profiles/orchestrator.dcdl",
|
||||
source: include_str!("../../../resources/profiles/orchestrator.dcdl"),
|
||||
description: "Workspace orchestration and Worker control.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
BuiltinProfileResource {
|
||||
selector: Some("builtin:memory-consolidation"),
|
||||
path: "profiles/memory-consolidation.dcdl",
|
||||
source: include_str!("../../../resources/profiles/memory-consolidation.dcdl"),
|
||||
description: "Internal Memory consolidation service.",
|
||||
imports: BASE_IMPORT,
|
||||
},
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BuiltinProfileCatalogSnapshot {
|
||||
pub id: &'static str,
|
||||
pub sources: BTreeMap<String, String>,
|
||||
pub entrypoints: BTreeMap<String, String>,
|
||||
pub imports: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl BuiltinProfileCatalogSnapshot {
|
||||
pub fn digest(&self) -> String {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(self.id.as_bytes());
|
||||
for (path, source) in &self.sources {
|
||||
hasher.update((path.len() as u64).to_le_bytes());
|
||||
hasher.update(path.as_bytes());
|
||||
hasher.update((source.len() as u64).to_le_bytes());
|
||||
hasher.update(source.as_bytes());
|
||||
}
|
||||
for (selector, path) in &self.entrypoints {
|
||||
hasher.update((selector.len() as u64).to_le_bytes());
|
||||
hasher.update(selector.as_bytes());
|
||||
hasher.update((path.len() as u64).to_le_bytes());
|
||||
hasher.update(path.as_bytes());
|
||||
}
|
||||
for (request, resolved_path) in &self.imports {
|
||||
hasher.update((request.len() as u64).to_le_bytes());
|
||||
hasher.update(request.as_bytes());
|
||||
hasher.update((resolved_path.len() as u64).to_le_bytes());
|
||||
hasher.update(resolved_path.as_bytes());
|
||||
}
|
||||
format!("sha256:{:x}", hasher.finalize())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn builtin_profile_catalog_snapshot() -> BuiltinProfileCatalogSnapshot {
|
||||
let mut sources = BTreeMap::new();
|
||||
let mut entrypoints = BTreeMap::new();
|
||||
let mut imports = BTreeMap::new();
|
||||
|
||||
for resource in BUILTIN_PROFILE_RESOURCES {
|
||||
sources.insert(resource.path.to_owned(), resource.source.to_owned());
|
||||
for import in resource.imports {
|
||||
imports.insert(
|
||||
format!("{}\0{}", resource.path, import.specifier),
|
||||
import.resolved_path.to_owned(),
|
||||
);
|
||||
}
|
||||
if let Some(selector) = resource.selector {
|
||||
entrypoints.insert(selector.to_owned(), resource.path.to_owned());
|
||||
}
|
||||
}
|
||||
|
||||
BuiltinProfileCatalogSnapshot {
|
||||
id: BUILTIN_PROFILE_CATALOG_ID,
|
||||
sources,
|
||||
entrypoints,
|
||||
imports,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn builtin_profile_entrypoints() -> impl Iterator<Item = &'static BuiltinProfileResource> {
|
||||
BUILTIN_PROFILE_RESOURCES
|
||||
.iter()
|
||||
.filter(|resource| resource.selector.is_some())
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_builtin_profile_artifact(
|
||||
selector: &str,
|
||||
) -> Result<Option<Value>, ProfileError> {
|
||||
let catalog = builtin_profile_catalog_snapshot();
|
||||
let Some(entrypoint) = catalog.entrypoints.get(selector) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let source = catalog
|
||||
.sources
|
||||
.get(entrypoint)
|
||||
.expect("built-in Profile entrypoint must name a source")
|
||||
.clone();
|
||||
let mut engine = Engine::new(BuiltinProfileImportLoader {
|
||||
sources: catalog.sources,
|
||||
});
|
||||
let module = engine
|
||||
.add_root_source(entrypoint, entrypoint, &source)
|
||||
.map_err(|error| ProfileError::BuiltinProfileEvaluation {
|
||||
selector: selector.to_owned(),
|
||||
message: format!("{error:?}"),
|
||||
})?;
|
||||
let value =
|
||||
engine
|
||||
.eval_module(module)
|
||||
.map_err(|error| ProfileError::BuiltinProfileEvaluation {
|
||||
selector: selector.to_owned(),
|
||||
message: format!("{error:?}"),
|
||||
})?;
|
||||
let data =
|
||||
engine
|
||||
.materialize(&value)
|
||||
.map_err(|error| ProfileError::BuiltinProfileEvaluation {
|
||||
selector: selector.to_owned(),
|
||||
message: format!("{error:?}"),
|
||||
})?;
|
||||
Ok(Some(data_to_json(&data)))
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct BuiltinProfileImportLoader {
|
||||
sources: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl ImportLoader for BuiltinProfileImportLoader {
|
||||
fn load(
|
||||
&mut self,
|
||||
current_key: Option<&str>,
|
||||
specifier: &str,
|
||||
) -> decodal::Result<LoadedImport> {
|
||||
let current_key = current_key.ok_or_else(|| {
|
||||
decodal::Diagnostic::new(
|
||||
decodal::DiagnosticKind::Import,
|
||||
decodal::Span::default(),
|
||||
format!("built-in Profile import `{specifier}` has no source context"),
|
||||
)
|
||||
})?;
|
||||
let resolved = resolve_import_path(current_key, specifier).ok_or_else(|| {
|
||||
decodal::Diagnostic::new(
|
||||
decodal::DiagnosticKind::Import,
|
||||
decodal::Span::default(),
|
||||
format!("built-in Profile import `{specifier}` from `{current_key}` is invalid"),
|
||||
)
|
||||
})?;
|
||||
let source = self.sources.get(&resolved).ok_or_else(|| {
|
||||
decodal::Diagnostic::new(
|
||||
decodal::DiagnosticKind::Import,
|
||||
decodal::Span::default(),
|
||||
format!("built-in Profile import `{specifier}` from `{current_key}` was not found"),
|
||||
)
|
||||
})?;
|
||||
Ok(LoadedImport::source(
|
||||
resolved.clone(),
|
||||
resolved,
|
||||
source.clone(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_import_path(current_key: &str, specifier: &str) -> Option<String> {
|
||||
let current_parent = current_key
|
||||
.rsplit_once('/')
|
||||
.map_or("", |(parent, _)| parent);
|
||||
let joined = if let Some(relative) = specifier.strip_prefix("./") {
|
||||
format!("{current_parent}/{relative}")
|
||||
} else {
|
||||
return None;
|
||||
};
|
||||
if joined
|
||||
.split('/')
|
||||
.any(|segment| segment.is_empty() || segment == "." || segment == "..")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(joined)
|
||||
}
|
||||
|
||||
fn data_to_json(data: &Data) -> Value {
|
||||
match data {
|
||||
Data::Bool(value) => Value::Bool(*value),
|
||||
Data::Int(value) => Value::Number(Number::from(*value)),
|
||||
Data::Float(value) => Number::from_f64(*value)
|
||||
.map(Value::Number)
|
||||
.unwrap_or(Value::Null),
|
||||
Data::String(value) => Value::String(value.clone()),
|
||||
Data::Array(values) => Value::Array(values.iter().map(data_to_json).collect()),
|
||||
Data::Object(fields) => Value::Object(
|
||||
fields
|
||||
.iter()
|
||||
.map(|field| (field.name.clone(), data_to_json(&field.value)))
|
||||
.collect::<Map<_, _>>(),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn catalog_has_one_explicit_entrypoint_for_each_builtin_profile() {
|
||||
let catalog = builtin_profile_catalog_snapshot();
|
||||
assert_eq!(catalog.sources.len(), BUILTIN_PROFILE_RESOURCES.len());
|
||||
assert_eq!(catalog.entrypoints.len() + 1, catalog.sources.len());
|
||||
assert_eq!(
|
||||
catalog.entrypoints.get(BUILTIN_DEFAULT_PROFILE),
|
||||
Some(&"profiles/default.dcdl".to_owned())
|
||||
);
|
||||
assert!(catalog.digest().starts_with("sha256:"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_profile_evaluates_from_the_shared_resource_graph() {
|
||||
let value = resolve_builtin_profile_artifact(BUILTIN_DEFAULT_PROFILE)
|
||||
.expect("evaluate built-in default")
|
||||
.expect("default exists");
|
||||
assert_eq!(value["slug"], "default");
|
||||
assert_eq!(value["feature"]["task"]["enabled"], true);
|
||||
assert_eq!(value["feature"]["sub_worker"]["enabled"], true);
|
||||
assert_eq!(value["feature"]["memory"]["enabled"], false);
|
||||
assert_eq!(value["feature"]["ticket"]["enabled"], false);
|
||||
assert_eq!(value["feature"]["worker"]["enabled"], false);
|
||||
assert_eq!(value["feature"]["manage_workdir"]["enabled"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn imports_cannot_escape_the_builtin_resource_catalog() {
|
||||
assert_eq!(
|
||||
resolve_import_path("profiles/default.dcdl", "./base.dcdl").as_deref(),
|
||||
Some("profiles/base.dcdl")
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_import_path("profiles/default.dcdl", "../outside.dcdl"),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_import_path("profiles/default.dcdl", "/outside.dcdl"),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -92,6 +92,8 @@ pub struct FeatureConfigPartial {
|
||||
#[serde(default)]
|
||||
pub worker: Option<WorkerFeatureConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub workspace_worker_discovery: Option<FeatureFlagConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub objective: Option<FeatureFlagConfigPartial>,
|
||||
#[serde(default)]
|
||||
pub manage_workdir: Option<FeatureFlagConfigPartial>,
|
||||
@@ -119,6 +121,11 @@ impl FeatureConfigPartial {
|
||||
),
|
||||
flow: merge_option(self.flow, other.flow, FeatureFlagConfigPartial::merge),
|
||||
worker: merge_option(self.worker, other.worker, WorkerFeatureConfigPartial::merge),
|
||||
workspace_worker_discovery: merge_option(
|
||||
self.workspace_worker_discovery,
|
||||
other.workspace_worker_discovery,
|
||||
FeatureFlagConfigPartial::merge,
|
||||
),
|
||||
objective: merge_option(
|
||||
self.objective,
|
||||
other.objective,
|
||||
@@ -265,6 +272,10 @@ impl From<FeatureConfigPartial> for FeatureConfig {
|
||||
.worker
|
||||
.map(WorkerFeatureConfig::from)
|
||||
.unwrap_or_default(),
|
||||
workspace_worker_discovery: value
|
||||
.workspace_worker_discovery
|
||||
.map(FeatureFlagConfig::from)
|
||||
.unwrap_or_default(),
|
||||
objective: value
|
||||
.objective
|
||||
.map(FeatureFlagConfig::from)
|
||||
@@ -394,6 +405,7 @@ impl From<FeatureConfig> for FeatureConfigPartial {
|
||||
sub_worker: Some(value.sub_worker.into()),
|
||||
flow: Some(value.flow.into()),
|
||||
worker: Some(value.worker.into()),
|
||||
workspace_worker_discovery: Some(value.workspace_worker_discovery.into()),
|
||||
objective: Some(value.objective.into()),
|
||||
manage_workdir: Some(value.manage_workdir.into()),
|
||||
ticket: Some(value.ticket.into()),
|
||||
@@ -566,15 +578,16 @@ impl WorkerManifestConfig {
|
||||
})
|
||||
}
|
||||
|
||||
/// Base config populated with the in-code defaults listed in
|
||||
/// [`crate::defaults`]. Profile and one-file Manifest resolvers start
|
||||
/// from this layer so every per-field default lives at exactly one
|
||||
/// call site (the `defaults` module).
|
||||
/// Base config populated with the in-code per-field defaults listed in
|
||||
/// [`crate::defaults`]. This is not a selectable Profile and does not
|
||||
/// enable a launch capability surface. Profile and one-file Manifest
|
||||
/// resolvers start from this layer so every per-field default lives at
|
||||
/// exactly one call site (the `defaults` module).
|
||||
///
|
||||
/// `TryFrom<WorkerManifestConfig>` also reads the same constants as a
|
||||
/// belt-and-suspenders fallback, so a manually-constructed config
|
||||
/// that skips this layer still resolves to the same values.
|
||||
pub fn builtin_defaults() -> Self {
|
||||
pub fn resolution_defaults() -> Self {
|
||||
Self {
|
||||
engine: EngineManifestConfig {
|
||||
tool_output: ToolOutputLimitsPartial {
|
||||
@@ -1973,7 +1986,7 @@ enabled = false
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::builtin_defaults()
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::resolution_defaults()
|
||||
.merge(cfg)
|
||||
.merge(WorkerManifestConfig {
|
||||
worker: WorkerMetaConfig {
|
||||
@@ -2074,7 +2087,7 @@ enabled = true
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::builtin_defaults()
|
||||
let manifest: WorkerManifest = WorkerManifestConfig::resolution_defaults()
|
||||
.merge(base)
|
||||
.merge(upper)
|
||||
.merge(WorkerManifestConfig {
|
||||
@@ -2137,7 +2150,7 @@ permission = "write"
|
||||
|
||||
#[test]
|
||||
fn builtin_defaults_populates_worker_limit_defaults() {
|
||||
let cfg = WorkerManifestConfig::builtin_defaults();
|
||||
let cfg = WorkerManifestConfig::resolution_defaults();
|
||||
assert_eq!(
|
||||
cfg.engine.tool_output.default_max_bytes,
|
||||
Some(defaults::TOOL_OUTPUT_MAX_BYTES)
|
||||
@@ -2172,7 +2185,7 @@ permission = "write"
|
||||
},
|
||||
..Default::default()
|
||||
};
|
||||
let merged = WorkerManifestConfig::builtin_defaults().merge(overlay);
|
||||
let merged = WorkerManifestConfig::resolution_defaults().merge(overlay);
|
||||
let manifest: WorkerManifest = merged.try_into().unwrap();
|
||||
assert_eq!(
|
||||
manifest.engine.tool_output.default_max_bytes,
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
mod builtin_profile;
|
||||
mod config;
|
||||
pub mod defaults;
|
||||
mod model;
|
||||
@@ -7,6 +8,11 @@ pub mod plugin;
|
||||
mod profile;
|
||||
mod scope;
|
||||
|
||||
pub use builtin_profile::{
|
||||
BUILTIN_DEFAULT_PROFILE, BUILTIN_PROFILE_CATALOG_ID, BUILTIN_PROFILE_RESOURCES,
|
||||
BuiltinProfileCatalogSnapshot, BuiltinProfileImport, BuiltinProfileResource,
|
||||
builtin_profile_catalog_snapshot, builtin_profile_entrypoints,
|
||||
};
|
||||
pub use config::{
|
||||
CompactionConfigPartial, EngineManifestConfig, FileUploadLimitsPartial,
|
||||
PermissionConfigPartial, ResolveError, SessionConfigPartial, ToolOutputLimitsPartial,
|
||||
@@ -17,10 +23,11 @@ pub use model::{
|
||||
};
|
||||
pub use paths::user_profiles_path;
|
||||
pub use profile::{
|
||||
ProfileDiscovery, ProfileError, ProfileManifestSnapshot, ProfileMetadata, ProfileRegistry,
|
||||
ProfileRegistryEntry, ProfileRegistrySource, ProfileResolveOptions, ProfileResolver,
|
||||
ProfileSelector, ProfileSource, ResolvedProfile, resolve_profile_artifact,
|
||||
resolve_profile_artifact_value,
|
||||
ProfileDiscovery, ProfileError, ProfileExecutionTarget, ProfileManifestSnapshot,
|
||||
ProfileMetadata, ProfileRegistry, ProfileRegistryEntry, ProfileRegistrySource,
|
||||
ProfileResolveOptions, ProfileResolver, ProfileSelector, ProfileSource, ResolvedProfile,
|
||||
WorkspaceAuthorityRequirement, resolve_profile_artifact, resolve_profile_artifact_value,
|
||||
validate_profile_execution_target,
|
||||
};
|
||||
pub use protocol::{Permission, ScopeRule};
|
||||
pub use scope::{DelegationScope, Scope, ScopeError, SharedScope};
|
||||
@@ -118,6 +125,10 @@ pub struct FeatureConfig {
|
||||
pub flow: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
pub worker: WorkerFeatureConfig,
|
||||
/// Privileged read-only discovery of visible Workspace Workers. Backend
|
||||
/// source proof remains required for every listing operation.
|
||||
#[serde(default)]
|
||||
pub workspace_worker_discovery: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
pub objective: FeatureFlagConfig,
|
||||
#[serde(default)]
|
||||
@@ -142,6 +153,7 @@ impl Default for FeatureConfig {
|
||||
sub_worker: FeatureFlagConfig::disabled(),
|
||||
flow: FeatureFlagConfig::disabled(),
|
||||
worker: WorkerFeatureConfig::disabled(),
|
||||
workspace_worker_discovery: FeatureFlagConfig::disabled(),
|
||||
objective: FeatureFlagConfig::disabled(),
|
||||
manage_workdir: FeatureFlagConfig::disabled(),
|
||||
ticket: TicketFeatureConfig::default(),
|
||||
|
||||
+274
-257
@@ -6,9 +6,14 @@
|
||||
//! from launch context.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::BTreeMap;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::fmt;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use crate::builtin_profile::{
|
||||
BUILTIN_DEFAULT_PROFILE, builtin_profile_catalog_snapshot, builtin_profile_entrypoints,
|
||||
resolve_builtin_profile_artifact,
|
||||
};
|
||||
use crate::config::{
|
||||
CompactionConfigPartial, FeatureConfigPartial, PermissionConfigPartial, SessionConfigPartial,
|
||||
};
|
||||
@@ -23,45 +28,6 @@ use crate::{
|
||||
const PROFILE_FORMAT_V1: &str = "yoi.profile.v1";
|
||||
const BUILTIN_MODEL_CATALOG: &str = include_str!("../../../resources/models/builtin.toml");
|
||||
|
||||
struct BuiltinProfile {
|
||||
name: &'static str,
|
||||
label: &'static str,
|
||||
description: &'static str,
|
||||
}
|
||||
|
||||
const BUILTIN_PROFILES: &[BuiltinProfile] = &[
|
||||
BuiltinProfile {
|
||||
name: "companion",
|
||||
label: "builtin:companion",
|
||||
description: "Bundled Companion role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "intake",
|
||||
label: "builtin:intake",
|
||||
description: "Bundled Intake role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "orchestrator",
|
||||
label: "builtin:orchestrator",
|
||||
description: "Bundled Orchestrator role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "coder",
|
||||
label: "builtin:coder",
|
||||
description: "Bundled Coder role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "reviewer",
|
||||
label: "builtin:reviewer",
|
||||
description: "Bundled Reviewer role profile",
|
||||
},
|
||||
BuiltinProfile {
|
||||
name: "memory-consolidation",
|
||||
label: "builtin:memory-consolidation",
|
||||
description: "Bundled Memory staging consolidation profile",
|
||||
},
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ProfileRegistrySource {
|
||||
@@ -159,6 +125,108 @@ impl ProfileSelector {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ProfileExecutionTarget {
|
||||
Workspace,
|
||||
Standalone,
|
||||
}
|
||||
|
||||
impl fmt::Display for ProfileExecutionTarget {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Workspace => formatter.write_str("workspace"),
|
||||
Self::Standalone => formatter.write_str("standalone"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
|
||||
pub enum WorkspaceAuthorityRequirement {
|
||||
Flow,
|
||||
ManageWorkdir,
|
||||
Memory,
|
||||
MergeRequest,
|
||||
Objective,
|
||||
Orchestration,
|
||||
Plugins,
|
||||
Ticket,
|
||||
Worker,
|
||||
}
|
||||
|
||||
impl fmt::Display for WorkspaceAuthorityRequirement {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
Self::Flow => formatter.write_str("feature.flow"),
|
||||
Self::ManageWorkdir => formatter.write_str("feature.manage_workdir"),
|
||||
Self::Memory => formatter.write_str("feature.memory"),
|
||||
Self::MergeRequest => formatter.write_str("feature.merge_request"),
|
||||
Self::Objective => formatter.write_str("feature.objective"),
|
||||
Self::Orchestration => formatter.write_str("feature.orchestration"),
|
||||
Self::Plugins => formatter.write_str("feature.plugins or plugin packages"),
|
||||
Self::Ticket => formatter.write_str("feature.ticket"),
|
||||
Self::Worker => formatter.write_str("feature.worker"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn validate_profile_execution_target(
|
||||
manifest: &WorkerManifest,
|
||||
target: ProfileExecutionTarget,
|
||||
) -> Result<(), ProfileError> {
|
||||
if target == ProfileExecutionTarget::Workspace {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let feature = &manifest.feature;
|
||||
let mut requirements = BTreeSet::new();
|
||||
if feature.flow.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Flow);
|
||||
}
|
||||
if feature.manage_workdir.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::ManageWorkdir);
|
||||
}
|
||||
if feature.memory.enabled || feature.memory.staging {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Memory);
|
||||
}
|
||||
if feature.merge_request.show
|
||||
|| feature.merge_request.open
|
||||
|| feature.merge_request.review
|
||||
|| feature.merge_request.readiness_check
|
||||
|| feature.merge_request.complete
|
||||
{
|
||||
requirements.insert(WorkspaceAuthorityRequirement::MergeRequest);
|
||||
}
|
||||
if feature.objective.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Objective);
|
||||
}
|
||||
if feature.orchestration.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Orchestration);
|
||||
}
|
||||
if feature.plugins.enabled || !manifest.plugins.is_empty() {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Plugins);
|
||||
}
|
||||
if feature.ticket.enabled
|
||||
|| feature.ticket.authoring
|
||||
|| feature.ticket.thread
|
||||
|| feature.ticket.intake
|
||||
|| feature.ticket.workflow
|
||||
{
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Ticket);
|
||||
}
|
||||
if feature.worker.enabled {
|
||||
requirements.insert(WorkspaceAuthorityRequirement::Worker);
|
||||
}
|
||||
|
||||
if requirements.is_empty() {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(ProfileError::UnsupportedExecutionTarget {
|
||||
target,
|
||||
requirements: requirements.into_iter().collect(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum ProfileSource {
|
||||
@@ -217,13 +285,14 @@ impl ProfileRegistryEntry {
|
||||
source: ProfileRegistrySource,
|
||||
name: &'static str,
|
||||
label: &'static str,
|
||||
provenance: String,
|
||||
description: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
source,
|
||||
name: name.to_string(),
|
||||
path: None,
|
||||
provenance: label.to_string(),
|
||||
provenance,
|
||||
description,
|
||||
is_default: false,
|
||||
artifact: ProfileRegistryArtifact::Builtin { label },
|
||||
@@ -321,12 +390,16 @@ pub struct ProfileDiscovery {
|
||||
}
|
||||
|
||||
impl ProfileDiscovery {
|
||||
pub fn for_cwd(_cwd: &Path) -> Self {
|
||||
pub fn user_settings() -> Self {
|
||||
Self {
|
||||
user_config: paths::user_profiles_path(),
|
||||
project_config: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn for_cwd(_cwd: &Path) -> Self {
|
||||
Self::user_settings()
|
||||
}
|
||||
pub fn with_sources(user_config: Option<PathBuf>, project_config: Option<PathBuf>) -> Self {
|
||||
Self {
|
||||
user_config,
|
||||
@@ -412,15 +485,22 @@ impl ProfileResolver {
|
||||
options,
|
||||
),
|
||||
ProfileSelector::Named { .. } | ProfileSelector::Default => {
|
||||
let cwd = std::env::current_dir().map_err(|source| ProfileError::CommandIo {
|
||||
path: PathBuf::from("."),
|
||||
source,
|
||||
})?;
|
||||
let registry = ProfileDiscovery::for_cwd(&cwd).discover()?;
|
||||
let registry = ProfileDiscovery::user_settings().discover()?;
|
||||
self.resolve_from_registry(selector, ®istry, options)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn resolve_for_target(
|
||||
&self,
|
||||
selector: &ProfileSelector,
|
||||
options: ProfileResolveOptions,
|
||||
target: ProfileExecutionTarget,
|
||||
) -> Result<ResolvedProfile, ProfileError> {
|
||||
let resolved = self.resolve(selector, options)?;
|
||||
validate_profile_execution_target(&resolved.manifest, target)?;
|
||||
Ok(resolved)
|
||||
}
|
||||
/// Resolve a registry/default selector against an already-discovered
|
||||
/// registry. Callers such as SubWorkerSpawn use this to bind discovery to the
|
||||
/// Worker's cwd instead of the process current directory.
|
||||
@@ -503,7 +583,7 @@ impl ProfileResolver {
|
||||
.as_deref()
|
||||
.unwrap_or_else(|| Path::new(".")),
|
||||
)?;
|
||||
let raw_artifact = builtin_profile_artifact(label).ok_or_else(|| {
|
||||
let raw_artifact = resolve_builtin_profile_artifact(label)?.ok_or_else(|| {
|
||||
ProfileError::InvalidProfile(format!("unknown builtin profile artifact `{label}`"))
|
||||
})?;
|
||||
resolve_profile_value(
|
||||
@@ -565,7 +645,8 @@ fn resolve_profile_value(
|
||||
memory: profile.memory.map(Into::into),
|
||||
skills: profile.skills,
|
||||
};
|
||||
let config = WorkerManifestConfig::builtin_defaults().merge(config.resolve_paths(profile_dir));
|
||||
let config =
|
||||
WorkerManifestConfig::resolution_defaults().merge(config.resolve_paths(profile_dir));
|
||||
let mut manifest = WorkerManifest::try_from(config).map_err(ProfileError::ManifestResolve)?;
|
||||
manifest.profile = Some(ProfileManifestSnapshot {
|
||||
source: source.clone(),
|
||||
@@ -759,14 +840,30 @@ fn load_profile_registry_file(
|
||||
}
|
||||
|
||||
fn add_builtin_profiles(registry: &mut ProfileRegistry) {
|
||||
for profile in BUILTIN_PROFILES {
|
||||
let catalog = builtin_profile_catalog_snapshot();
|
||||
let digest = catalog.digest();
|
||||
for profile in builtin_profile_entrypoints() {
|
||||
let label = profile
|
||||
.selector
|
||||
.expect("built-in Profile entrypoint must have a selector");
|
||||
let name = label
|
||||
.strip_prefix("builtin:")
|
||||
.expect("built-in Profile selector must be source-qualified");
|
||||
registry.push_entry(ProfileRegistryEntry::embedded(
|
||||
ProfileRegistrySource::Builtin,
|
||||
profile.name,
|
||||
profile.label,
|
||||
name,
|
||||
label,
|
||||
format!("{}#{digest}", profile.path),
|
||||
Some(profile.description.into()),
|
||||
));
|
||||
}
|
||||
registry.set_default(ProfileDefault {
|
||||
source: Some(ProfileRegistrySource::Builtin),
|
||||
name: BUILTIN_DEFAULT_PROFILE
|
||||
.strip_prefix("builtin:")
|
||||
.expect("built-in default selector must be source-qualified")
|
||||
.to_owned(),
|
||||
});
|
||||
}
|
||||
|
||||
fn parse_profile_ref(raw: &str) -> (Option<ProfileRegistrySource>, String) {
|
||||
@@ -804,201 +901,6 @@ fn read_profile_artifact_file(path: &Path) -> Result<serde_json::Value, ProfileE
|
||||
}
|
||||
}
|
||||
|
||||
fn builtin_profile_artifact(label: &str) -> Option<serde_json::Value> {
|
||||
let mut value = builtin_base_profile_artifact();
|
||||
match label {
|
||||
"builtin:companion" | "companion" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"companion",
|
||||
"Workspace companion profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:intake" | "intake" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"intake",
|
||||
"Ticket intake profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:orchestrator" | "orchestrator" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"orchestrator",
|
||||
"Ticket orchestrator profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:coder" | "coder" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"coder",
|
||||
"Ticket implementation coder profile.",
|
||||
"workspace_write",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:reviewer" | "reviewer" => {
|
||||
apply_role_profile(
|
||||
&mut value,
|
||||
"reviewer",
|
||||
"Ticket review profile.",
|
||||
"workspace_read",
|
||||
true,
|
||||
true,
|
||||
true,
|
||||
false,
|
||||
);
|
||||
Some(value)
|
||||
}
|
||||
"builtin:memory-consolidation" | "memory-consolidation" => {
|
||||
value["slug"] = serde_json::Value::String("memory-consolidation".to_string());
|
||||
value["description"] =
|
||||
serde_json::Value::String("Memory staging consolidation profile.".to_string());
|
||||
value["feature"]["task"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["memory"] = serde_json::json!({ "enabled": true, "staging": true });
|
||||
value["feature"]["web"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["sub_worker"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["objective"] = serde_json::json!({ "enabled": false });
|
||||
value["feature"]["ticket"] = serde_json::json!({ "enabled": false, "thread": false });
|
||||
Some(value)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn builtin_base_profile_artifact() -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"slug": "default",
|
||||
"description": "Default Yoi coding profile.",
|
||||
"model": { "ref": "codex-oauth/gpt-5.5" },
|
||||
"session": { "record_event_trace": true },
|
||||
"engine": { "reasoning": "high" },
|
||||
"compaction": {
|
||||
"kind": "tokens",
|
||||
"threshold": 240000,
|
||||
"request_threshold": 270000,
|
||||
"worker_context_max_tokens": 100000
|
||||
},
|
||||
"feature": {
|
||||
"task": { "enabled": true },
|
||||
"memory": { "enabled": true },
|
||||
"web": { "enabled": true },
|
||||
"image": { "enabled": true },
|
||||
"sub_worker": { "enabled": true },
|
||||
"worker": { "enabled": false },
|
||||
"objective": { "enabled": true },
|
||||
"ticket": { "enabled": true, "authoring": true, "thread": true }
|
||||
},
|
||||
"memory": {
|
||||
"extract_threshold": 50000,
|
||||
"consolidation_threshold_files": 5,
|
||||
"consolidation_threshold_bytes": 50000
|
||||
},
|
||||
"web": {
|
||||
"enabled": true,
|
||||
"search": {
|
||||
"provider": "brave",
|
||||
"api_key_secret": "web/brave/default"
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn apply_role_profile(
|
||||
value: &mut serde_json::Value,
|
||||
slug: &str,
|
||||
description: &str,
|
||||
_scope: &str,
|
||||
task: bool,
|
||||
memory: bool,
|
||||
web: bool,
|
||||
sub_worker: bool,
|
||||
) {
|
||||
value["slug"] = serde_json::Value::String(slug.to_string());
|
||||
value["description"] = serde_json::Value::String(description.to_string());
|
||||
value["feature"]["task"] = serde_json::json!({ "enabled": task });
|
||||
value["feature"]["memory"] = serde_json::json!({ "enabled": memory });
|
||||
value["feature"]["web"] = serde_json::json!({ "enabled": web });
|
||||
value["feature"]["image"] = serde_json::json!({ "enabled": true });
|
||||
value["feature"]["sub_worker"] = serde_json::json!({ "enabled": sub_worker });
|
||||
value["feature"]["flow"] = serde_json::json!({ "enabled": slug == "coder" });
|
||||
value["feature"]["worker"] = serde_json::json!({
|
||||
"enabled": matches!(slug, "companion" | "orchestrator"),
|
||||
"direct_spawn": slug != "orchestrator"
|
||||
});
|
||||
value["feature"]["manage_workdir"] = serde_json::json!({
|
||||
"enabled": matches!(slug, "companion" | "orchestrator")
|
||||
});
|
||||
value["feature"]["orchestration"] = serde_json::json!({ "enabled": slug == "orchestrator" });
|
||||
let ticket = match slug {
|
||||
"companion" => serde_json::json!({ "enabled": true, "authoring": true, "thread": true }),
|
||||
"intake" => {
|
||||
serde_json::json!({ "enabled": true, "authoring": true, "thread": true, "intake": true })
|
||||
}
|
||||
"orchestrator" => {
|
||||
serde_json::json!({ "enabled": true, "thread": true, "workflow": true })
|
||||
}
|
||||
"coder" => serde_json::json!({ "enabled": true, "thread": true }),
|
||||
"reviewer" => serde_json::json!({ "enabled": true, "thread": true }),
|
||||
_ => serde_json::json!({ "enabled": true, "authoring": true, "thread": true }),
|
||||
};
|
||||
value["feature"]["ticket"] = ticket;
|
||||
let merge_request = match slug {
|
||||
"coder" => serde_json::json!({
|
||||
"show": true,
|
||||
"open": true,
|
||||
"review": false,
|
||||
"readiness_check": false,
|
||||
"complete": false
|
||||
}),
|
||||
"reviewer" => serde_json::json!({
|
||||
"show": true,
|
||||
"open": false,
|
||||
"review": true,
|
||||
"readiness_check": false,
|
||||
"complete": false
|
||||
}),
|
||||
"orchestrator" => serde_json::json!({
|
||||
"show": true,
|
||||
"open": false,
|
||||
"review": false,
|
||||
"readiness_check": true,
|
||||
"complete": true
|
||||
}),
|
||||
_ => serde_json::json!({
|
||||
"show": false,
|
||||
"open": false,
|
||||
"review": false,
|
||||
"readiness_check": false,
|
||||
"complete": false
|
||||
}),
|
||||
};
|
||||
value["feature"]["merge_request"] = merge_request;
|
||||
}
|
||||
|
||||
fn reject_manifest_shaped_profile(value: &serde_json::Value) -> Result<(), ProfileError> {
|
||||
let Some(map) = value.as_object() else {
|
||||
return Err(ProfileError::InvalidProfile(
|
||||
@@ -1288,6 +1190,13 @@ pub enum ProfileError {
|
||||
#[source]
|
||||
source: toml::de::Error,
|
||||
},
|
||||
#[error("failed to evaluate built-in Profile `{selector}`: {message}")]
|
||||
BuiltinProfileEvaluation { selector: String, message: String },
|
||||
#[error("Profile requires unsupported {target} launch authorities: {requirements:?}")]
|
||||
UnsupportedExecutionTarget {
|
||||
target: ProfileExecutionTarget,
|
||||
requirements: Vec<WorkspaceAuthorityRequirement>,
|
||||
},
|
||||
#[error("no default profile is configured")]
|
||||
NoDefaultProfile,
|
||||
#[error("profile resolution requires an explicit runtime Worker name")]
|
||||
@@ -1341,18 +1250,21 @@ mod tests {
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn builtin_profiles_do_not_define_an_implicit_default() {
|
||||
fn builtin_default_is_explicit_registry_authority() {
|
||||
let registry = ProfileDiscovery::with_sources(None, None)
|
||||
.discover()
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
registry.default_entry(),
|
||||
Err(ProfileError::NoDefaultProfile)
|
||||
));
|
||||
assert!(matches!(
|
||||
registry.select(&ProfileSelector::Default),
|
||||
Err(ProfileError::NoDefaultProfile)
|
||||
));
|
||||
let default = registry.default_entry().unwrap();
|
||||
assert_eq!(default.source, ProfileRegistrySource::Builtin);
|
||||
assert_eq!(default.name, "default");
|
||||
assert_eq!(default.qualified_name(), BUILTIN_DEFAULT_PROFILE);
|
||||
assert!(default.is_default);
|
||||
assert!(
|
||||
default
|
||||
.provenance
|
||||
.starts_with("profiles/default.dcdl#sha256:")
|
||||
);
|
||||
assert_eq!(registry.select(&ProfileSelector::Default).unwrap(), default);
|
||||
}
|
||||
#[test]
|
||||
fn builtin_role_profiles_are_registered_and_resolve() {
|
||||
@@ -1408,7 +1320,108 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_companion_can_manage_workdirs() {
|
||||
fn builtin_default_resolves_as_a_standalone_local_capability_profile() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve_for_target(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "default"),
|
||||
ProfileResolveOptions::with_worker_name("standalone-worker"),
|
||||
ProfileExecutionTarget::Standalone,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(
|
||||
&resolved.source,
|
||||
ProfileSource::Registry {
|
||||
source: ProfileRegistrySource::Builtin,
|
||||
name,
|
||||
path: None,
|
||||
provenance: Some(provenance),
|
||||
..
|
||||
} if name == "default" && provenance.starts_with("profiles/default.dcdl#sha256:")
|
||||
));
|
||||
assert!(resolved.manifest.feature.task.enabled);
|
||||
assert!(resolved.manifest.feature.web.enabled);
|
||||
assert!(resolved.manifest.feature.image.enabled);
|
||||
assert!(resolved.manifest.feature.sub_worker.enabled);
|
||||
assert!(resolved.manifest.scope.allow.iter().any(|rule| {
|
||||
rule.permission == protocol::Permission::Write && rule.target == tmp.path()
|
||||
}));
|
||||
assert!(resolved.manifest.delegation_scope.allow.iter().any(|rule| {
|
||||
rule.permission == protocol::Permission::Write && rule.target == tmp.path()
|
||||
}));
|
||||
assert!(!resolved.manifest.feature.memory.enabled);
|
||||
assert!(!resolved.manifest.feature.ticket.enabled);
|
||||
assert!(!resolved.manifest.feature.objective.enabled);
|
||||
assert!(!resolved.manifest.feature.flow.enabled);
|
||||
assert!(!resolved.manifest.feature.worker.enabled);
|
||||
assert!(!resolved.manifest.feature.manage_workdir.enabled);
|
||||
assert!(!resolved.manifest.feature.plugins.enabled);
|
||||
assert!(resolved.manifest.plugins.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standalone_rejects_profiles_that_require_workspace_authority() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let error = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve_for_target(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "coder"),
|
||||
ProfileResolveOptions::with_worker_name("standalone-worker"),
|
||||
ProfileExecutionTarget::Standalone,
|
||||
)
|
||||
.unwrap_err();
|
||||
let diagnostic = error.to_string();
|
||||
|
||||
let ProfileError::UnsupportedExecutionTarget {
|
||||
target,
|
||||
requirements,
|
||||
} = error
|
||||
else {
|
||||
panic!("unexpected error: {error}");
|
||||
};
|
||||
assert_eq!(target, ProfileExecutionTarget::Standalone);
|
||||
assert!(requirements.contains(&WorkspaceAuthorityRequirement::Memory));
|
||||
assert!(requirements.contains(&WorkspaceAuthorityRequirement::MergeRequest));
|
||||
assert!(requirements.contains(&WorkspaceAuthorityRequirement::Ticket));
|
||||
assert!(!diagnostic.contains(tmp.path().to_string_lossy().as_ref()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repository_markers_do_not_change_builtin_profile_authority() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let nested = tmp.path().join("repository/nested");
|
||||
std::fs::create_dir_all(&nested).unwrap();
|
||||
std::fs::create_dir_all(tmp.path().join("repository/.yoi")).unwrap();
|
||||
std::fs::write(
|
||||
tmp.path().join("repository/.yoi/profiles.toml"),
|
||||
"default = { source = 'project', name = 'shadow' }\n",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let discovery = ProfileDiscovery::for_cwd(&nested);
|
||||
assert_eq!(discovery.user_config, paths::user_profiles_path());
|
||||
assert!(discovery.project_config.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_coder_uses_sub_worker_control_without_worker_control() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
.resolve(
|
||||
&ProfileSelector::source_named(ProfileRegistrySource::Builtin, "coder"),
|
||||
ProfileResolveOptions::with_worker_name("coder-worker"),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(resolved.manifest.feature.sub_worker.enabled);
|
||||
assert!(!resolved.manifest.feature.worker.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_companion_combines_runtime_and_sub_worker_control_with_discovery() {
|
||||
let tmp = TempDir::new().unwrap();
|
||||
let resolved = ProfileResolver::new()
|
||||
.with_workspace_base(tmp.path())
|
||||
@@ -1419,6 +1432,10 @@ mod tests {
|
||||
.unwrap();
|
||||
|
||||
assert!(resolved.manifest.feature.manage_workdir.enabled);
|
||||
assert!(resolved.manifest.feature.sub_worker.enabled);
|
||||
assert!(resolved.manifest.feature.worker.enabled);
|
||||
assert!(!resolved.manifest.feature.worker.direct_spawn);
|
||||
assert!(resolved.manifest.feature.workspace_worker_discovery.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
use schemars::JsonSchema;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::schema::{EvidenceKind, SourceEvidenceRef, SourceRef};
|
||||
use crate::schema::{EvidenceKind, EvidenceOrigin, SourceEvidenceRef, SourceRef};
|
||||
|
||||
/// Current flat staging schema version.
|
||||
pub const STAGING_SCHEMA_VERSION: u32 = 2;
|
||||
@@ -80,6 +80,8 @@ pub struct StagingEvidence {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub entry_range: Option<[u64; 2]>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub origin: Option<EvidenceOrigin>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub excerpt: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub summary: Option<String>,
|
||||
@@ -159,6 +161,7 @@ mod tests {
|
||||
id: "E001".into(),
|
||||
kind: EvidenceKind::new(EvidenceKind::MESSAGE),
|
||||
entry_range: Some([10, 12]),
|
||||
origin: None,
|
||||
excerpt: Some("extract candidate taxonomy".into()),
|
||||
summary: Some("User and assistant discussed staging kinds".into()),
|
||||
};
|
||||
|
||||
@@ -67,6 +67,40 @@ impl EvidenceKind {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum EvidenceOriginKind {
|
||||
HumanInput,
|
||||
WorkerInput,
|
||||
FlowInstruction,
|
||||
BackendInstruction,
|
||||
ModelOutput,
|
||||
ToolOutput,
|
||||
DerivedSummary,
|
||||
LegacyUnknown,
|
||||
}
|
||||
|
||||
/// Bounded origin snapshot attached to extraction evidence. This is audit
|
||||
/// metadata only and cannot authorize Workspace operations.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
|
||||
pub struct EvidenceOrigin {
|
||||
pub kind: EvidenceOriginKind,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub account_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub workspace_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub runtime_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub worker_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub flow_selector: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub flow_definition_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub flow_definition_revision: Option<u64>,
|
||||
}
|
||||
|
||||
/// Host-resolved source/evidence metadata for an individual staging claim.
|
||||
///
|
||||
/// This deliberately stores only bounded anchor metadata: stable ids, entry
|
||||
@@ -86,6 +120,9 @@ pub struct SourceEvidenceRef {
|
||||
/// Host-assigned evidence id within the referenced evidence set.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub evidence_id: Option<String>,
|
||||
/// Trusted typed origin snapshot for this logical evidence entry.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub origin: Option<EvidenceOrigin>,
|
||||
/// Extensible evidence kind tag.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub evidence_kind: Option<EvidenceKind>,
|
||||
|
||||
@@ -10,7 +10,10 @@ mod decision;
|
||||
mod request;
|
||||
mod summary;
|
||||
|
||||
pub use common::{EvidenceKind, Frontmatter, SourceEvidenceRef, SourceRef, split_frontmatter};
|
||||
pub use common::{
|
||||
EvidenceKind, EvidenceOrigin, EvidenceOriginKind, Frontmatter, SourceEvidenceRef, SourceRef,
|
||||
split_frontmatter,
|
||||
};
|
||||
pub use decision::{DecisionFrontmatter, DecisionStatus};
|
||||
pub use request::RequestFrontmatter;
|
||||
pub use summary::SummaryFrontmatter;
|
||||
|
||||
@@ -416,7 +416,7 @@ impl MergeRequestStore {
|
||||
let conflict:bool=t.query_row("SELECT EXISTS(SELECT 1 FROM merge_request_ticket_relations rel JOIN merge_requests mr ON mr.workspace_id=rel.workspace_id AND mr.merge_request_id=rel.merge_request_id WHERE rel.workspace_id=?1 AND rel.ticket_id=?2 AND mr.state='open')",params![i.auth.workspace_id,i.ticket_id],|r|r.get(0))?;
|
||||
if conflict {
|
||||
return Err(MergeRequestError::Conflict(
|
||||
"Ticket already has an open Merge Request".into(),
|
||||
"Ticket already has an open Merge Request; use ShowMergeRequest and advance the existing selector_from with a normal non-force push instead of opening a replacement Merge Request or adding a revision".into(),
|
||||
));
|
||||
}
|
||||
let now = i.now.to_rfc3339();
|
||||
@@ -575,12 +575,16 @@ impl MergeRequestStore {
|
||||
));
|
||||
};
|
||||
if subject != i.current_subject_ref {
|
||||
let reason = format!(
|
||||
"selector_from moved from requested subject {subject} to current subject {}; fresh review of the exact current source ref is required",
|
||||
i.current_subject_ref
|
||||
);
|
||||
let e = ReviewCancelledEvent {
|
||||
event_id: Uuid::now_v7().to_string(),
|
||||
sequence: next_seq(&t, &ws, &mr)?,
|
||||
request_event_id: req,
|
||||
subject_ref: subject,
|
||||
reason: "selector_from moved before submission".into(),
|
||||
reason,
|
||||
created_at: i.now,
|
||||
};
|
||||
insert_event(&t, &ws, &mr, "review_cancelled", &e, i.now, None)?;
|
||||
@@ -667,9 +671,27 @@ impl MergeRequestStore {
|
||||
}
|
||||
match (&i.current_subject_ref, &review) {
|
||||
(None, _) => b.push("selector_from could not be resolved".into()),
|
||||
(Some(_), None) => b.push("current source ref has no valid review".into()),
|
||||
(_, Some(r)) if r.decision == ReviewDecision::RequestChanges => {
|
||||
b.push("current source ref requests changes".into())
|
||||
(Some(subject_ref), None) => {
|
||||
let previous_review_subject = mr.thread.iter().rev().find_map(|event| match event {
|
||||
MergeRequestThreadEvent::ReviewRequested(value) => {
|
||||
Some(value.subject_ref.as_str())
|
||||
}
|
||||
MergeRequestThreadEvent::Review(value) => Some(value.subject_ref.as_str()),
|
||||
_ => None,
|
||||
});
|
||||
match previous_review_subject.filter(|previous| *previous != subject_ref) {
|
||||
Some(previous) => b.push(format!(
|
||||
"selector_from moved from reviewed/requested subject {previous} to current subject {subject_ref}; request a fresh review for this exact source ref (selector_to movement alone does not invalidate source approval)"
|
||||
)),
|
||||
None => b.push(format!(
|
||||
"current source ref {subject_ref} has no valid review; request a fresh review for this exact source ref"
|
||||
)),
|
||||
}
|
||||
}
|
||||
(Some(subject_ref), Some(r)) if r.decision == ReviewDecision::RequestChanges => {
|
||||
b.push(format!(
|
||||
"current source ref {subject_ref} requests changes; advance the existing selector_from with a normal non-force push, then request a fresh review for the exact new source ref"
|
||||
))
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
@@ -181,16 +181,85 @@ fn source_move_cancels_submission_and_old_approval_is_reusable_when_source_retur
|
||||
.is_err()
|
||||
);
|
||||
let mr = s.get("W", "T").unwrap();
|
||||
let cancellation = mr.thread.iter().find_map(|event| match event {
|
||||
MergeRequestThreadEvent::ReviewCancelled(value) => Some(value),
|
||||
_ => None,
|
||||
});
|
||||
assert!(
|
||||
mr.thread
|
||||
.iter()
|
||||
.any(|e| matches!(e, MergeRequestThreadEvent::ReviewCancelled(_)))
|
||||
cancellation
|
||||
.as_ref()
|
||||
.is_some_and(|value| value.reason.contains("selector_from moved")
|
||||
&& value.reason.contains("fresh review"))
|
||||
);
|
||||
assert_eq!(
|
||||
mr.effective_review("source-a").map(|r| &r.event_id),
|
||||
Some(&approved.event_id)
|
||||
);
|
||||
}
|
||||
#[test]
|
||||
fn same_selector_source_advancement_requires_fresh_review_and_preserves_target_only_approval() {
|
||||
let (_d, s) = fixture();
|
||||
open(&s);
|
||||
let first = approve(&s, "source-1", "one");
|
||||
|
||||
let stale = s
|
||||
.readiness(ReadinessCheck {
|
||||
ticket_id: "T".into(),
|
||||
current_subject_ref: Some("source-2".into()),
|
||||
auth: auth(),
|
||||
})
|
||||
.unwrap();
|
||||
assert!(!stale.ready);
|
||||
assert!(stale.review.is_none());
|
||||
assert!(stale.blockers.iter().any(|blocker| {
|
||||
blocker.contains("selector_from moved from reviewed/requested subject source-1")
|
||||
&& blocker.contains("current subject source-2")
|
||||
&& blocker.contains("fresh review")
|
||||
}));
|
||||
assert_eq!(
|
||||
s.get("W", "T")
|
||||
.unwrap()
|
||||
.effective_review("source-1")
|
||||
.map(|review| &review.event_id),
|
||||
Some(&first.event_id)
|
||||
);
|
||||
|
||||
let second = approve(&s, "source-2", "two");
|
||||
let ready = s
|
||||
.readiness(ReadinessCheck {
|
||||
ticket_id: "T".into(),
|
||||
current_subject_ref: Some("source-2".into()),
|
||||
auth: auth(),
|
||||
})
|
||||
.unwrap();
|
||||
assert!(ready.ready);
|
||||
assert_eq!(
|
||||
ready.review.as_ref().map(|review| &review.event_id),
|
||||
Some(&second.event_id)
|
||||
);
|
||||
|
||||
// The target can move from target-1 to target-2 without changing selector_from
|
||||
// or invalidating the exact-source approval. Completion consumes refreshed
|
||||
// integration evidence for the current target pair.
|
||||
let merged = s
|
||||
.complete(CompleteMergeRequest {
|
||||
operation_id: "target-moved".into(),
|
||||
ticket_id: "T".into(),
|
||||
current_subject_ref: "source-2".into(),
|
||||
target_ref_before: "target-2".into(),
|
||||
target_ref_after: "integrated-target-2".into(),
|
||||
approval_event_id: second.event_id,
|
||||
strategy: MergeStrategy::FastForward,
|
||||
resolution: ConflictResolution::None,
|
||||
auth: auth(),
|
||||
now: at(5),
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(merged.approved_source_ref, "source-2");
|
||||
assert_eq!(merged.target_ref_before, "target-2");
|
||||
assert_eq!(merged.target_ref_after, "integrated-target-2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn review_revocation_invalidates_readiness() {
|
||||
let (_d, s) = fixture();
|
||||
|
||||
+247
-57
@@ -281,11 +281,44 @@ impl Method {
|
||||
/// Presentation category for an Internal Worker exposed through its parent's
|
||||
/// protocol stream. Internal Workers never become independently addressable
|
||||
/// protocol subjects.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum InternalWorkerKind {
|
||||
SubWorker,
|
||||
Service { kind: String },
|
||||
}
|
||||
|
||||
/// Stable parent-owned lifecycle for one compaction run.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct CompactionLifecycle {
|
||||
pub schema_version: u32,
|
||||
pub compaction_id: String,
|
||||
pub revision: u64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub internal_worker: Option<InternalWorkerRef>,
|
||||
pub state: CompactionLifecycleState,
|
||||
/// Milliseconds since the Unix epoch.
|
||||
pub started_at_ms: u64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ended_at_ms: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub summary: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub error: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub new_segment_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum CompactionLifecycleState {
|
||||
Running,
|
||||
Done,
|
||||
Failed,
|
||||
Interrupted,
|
||||
}
|
||||
|
||||
/// Stable presentation identity for one parent-owned Internal Worker session.
|
||||
@@ -307,8 +340,7 @@ pub struct InternalWorkerRef {
|
||||
pub struct InternalWorkerSnapshot {
|
||||
pub worker: InternalWorkerRef,
|
||||
pub revision: u64,
|
||||
#[cfg_attr(feature = "typescript", ts(type = "Array<unknown>"))]
|
||||
pub entries: Vec<serde_json::Value>,
|
||||
pub session: SessionSnapshot,
|
||||
#[serde(default)]
|
||||
pub status: WorkerStatus,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
@@ -319,12 +351,126 @@ pub struct InternalWorkerSnapshot {
|
||||
pub internal_workers: Vec<InternalWorkerSnapshot>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ToolResultDisposition {
|
||||
#[default]
|
||||
Success,
|
||||
Error,
|
||||
Interrupted,
|
||||
Cancelled,
|
||||
OutcomeUnknown,
|
||||
}
|
||||
|
||||
/// Canonical, storage-independent projection of committed session history.
|
||||
///
|
||||
/// Worker protocols expose this DTO instead of append-log records. New
|
||||
/// storage variants can therefore be added without teaching every client how
|
||||
/// to replay the durable log format.
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct SessionSnapshot {
|
||||
pub entries: Vec<SessionSnapshotEntry>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum SessionEntryProvenance {
|
||||
HumanInput,
|
||||
WorkerInput,
|
||||
FlowInstruction,
|
||||
BackendInstruction,
|
||||
ModelOutput,
|
||||
ToolOutput,
|
||||
DerivedSummary,
|
||||
LegacyUnknown,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct SessionSnapshotEntry {
|
||||
/// Stable identity from durable history metadata, or a deterministic
|
||||
/// identity derived from the legacy segment and log position.
|
||||
pub entry_id: String,
|
||||
/// Timestamp copied from the durable log record that commits this entry.
|
||||
pub timestamp: u64,
|
||||
pub provenance: SessionEntryProvenance,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub derived_from: Vec<String>,
|
||||
#[serde(flatten)]
|
||||
pub data: SessionSnapshotEntryData,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum SessionSnapshotEntryData {
|
||||
UserInput {
|
||||
segments: Vec<Segment>,
|
||||
},
|
||||
Message {
|
||||
role: SessionMessageRole,
|
||||
content: Vec<SessionContentPart>,
|
||||
},
|
||||
ToolCall {
|
||||
call_id: String,
|
||||
name: String,
|
||||
arguments: String,
|
||||
},
|
||||
ToolResult {
|
||||
call_id: String,
|
||||
summary: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
content: Option<String>,
|
||||
is_error: bool,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
attachments: Vec<SessionToolAttachment>,
|
||||
},
|
||||
SystemItem {
|
||||
item_kind: String,
|
||||
content: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
#[cfg_attr(feature = "typescript", ts(type = "unknown"))]
|
||||
data: Option<serde_json::Value>,
|
||||
},
|
||||
RunError {
|
||||
message: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum SessionMessageRole {
|
||||
User,
|
||||
Assistant,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum SessionContentPart {
|
||||
Text { text: String },
|
||||
Refusal { refusal: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
pub struct SessionToolAttachment {
|
||||
pub media_type: String,
|
||||
/// Base64-encoded durable attachment body. Public snapshots preserve the
|
||||
/// committed multimodal value instead of replacing it with placeholder text.
|
||||
pub data_base64: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
|
||||
#[serde(tag = "event", content = "data", rename_all = "snake_case")]
|
||||
pub enum Event {
|
||||
/// A user input message was accepted, persisted as
|
||||
/// `LogEntry::UserInput`, and is about to start a new turn.
|
||||
/// `LogEntry::AnnotatedUserInput`, and is about to start a new turn.
|
||||
/// Broadcast to every subscribed client so TUI / GUI instances show
|
||||
/// the same user line that reconnect snapshots would replay from
|
||||
/// history; clients must not synthesize a separate pending/fake
|
||||
@@ -345,7 +491,7 @@ pub enum Event {
|
||||
/// of parsing free-text prefixes like `[Notification] …` or
|
||||
/// `[File: …]`.
|
||||
///
|
||||
/// One event per `LogEntry::SystemItem` commit. Disk-side and
|
||||
/// One event per `LogEntry::AnnotatedSystemItem` commit. Disk-side and
|
||||
/// wire-side are 1:1.
|
||||
SystemItem {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "unknown"))]
|
||||
@@ -468,6 +614,8 @@ pub enum Event {
|
||||
/// summary-only, or when the result was pruned.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
output: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
disposition: Option<ToolResultDisposition>,
|
||||
#[serde(default)]
|
||||
is_error: bool,
|
||||
},
|
||||
@@ -508,8 +656,7 @@ pub enum Event {
|
||||
/// role-specific entry events (`SegmentRotated` / `SystemItem`) —
|
||||
/// there is no generic "every committed entry" broadcast.
|
||||
Snapshot {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "Array<unknown>"))]
|
||||
entries: Vec<serde_json::Value>,
|
||||
session: SessionSnapshot,
|
||||
greeting: Greeting,
|
||||
#[serde(default)]
|
||||
status: WorkerStatus,
|
||||
@@ -542,14 +689,10 @@ pub enum Event {
|
||||
/// Server-side segment log rotated to a fresh `SegmentStart`.
|
||||
///
|
||||
/// Fires on compaction and on auto-fork when the store head drifts
|
||||
/// from the live writer's cached head. Clients drop their derived
|
||||
/// view and reseed from `entry.history` exactly the way they would
|
||||
/// from a connect-time `Snapshot`.
|
||||
///
|
||||
/// Payload is the JSON form of `session_store::LogEntry::SegmentStart`.
|
||||
/// A compaction/fork has replaced the authoritative segment. Clients drop
|
||||
/// their derived view and reseed from the canonical committed snapshot.
|
||||
SegmentRotated {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "unknown"))]
|
||||
entry: serde_json::Value,
|
||||
session: SessionSnapshot,
|
||||
},
|
||||
/// Current Worker controller status. Broadcast on every controller-level
|
||||
/// transition and included in `History` snapshots for late attach.
|
||||
@@ -576,11 +719,10 @@ pub enum Event {
|
||||
head_entries: usize,
|
||||
targets: Vec<RewindTarget>,
|
||||
},
|
||||
/// A rewind has truncated the authoritative session. `entries` is the
|
||||
/// retained session-log prefix clients should use to reseed display state.
|
||||
/// A rewind has truncated the authoritative session. `session` is the
|
||||
/// retained canonical snapshot clients should use to reseed display state.
|
||||
RewindApplied {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "Array<unknown>"))]
|
||||
entries: Vec<serde_json::Value>,
|
||||
session: SessionSnapshot,
|
||||
input: Vec<Segment>,
|
||||
summary: RewindSummary,
|
||||
},
|
||||
@@ -607,23 +749,18 @@ pub enum Event {
|
||||
/// This is not part of LLM history or prompt context; clients may display it
|
||||
/// briefly as operational status.
|
||||
MemoryWorker(MemoryWorkerEvent),
|
||||
/// Worker has started compacting the current session.
|
||||
///
|
||||
/// Fired immediately before a compaction run. Success is signalled by
|
||||
/// `CompactDone` (with the new `SegmentId`); failure by `CompactFailed`.
|
||||
/// Broadcast to all clients; not replayed to late subscribers.
|
||||
CompactStart,
|
||||
/// Compaction completed and the session was rotated.
|
||||
///
|
||||
/// `new_segment_id` is the UUID of the freshly created session that
|
||||
/// replaced the old history.
|
||||
CompactDone {
|
||||
#[cfg_attr(feature = "typescript", ts(type = "string"))]
|
||||
new_segment_id: uuid::Uuid,
|
||||
/// Worker has started compacting the current session, or bound the run to its
|
||||
/// observable Internal Worker. Revisions upsert one stable lifecycle item.
|
||||
CompactStart {
|
||||
lifecycle: CompactionLifecycle,
|
||||
},
|
||||
/// Compaction failed. The session is unchanged.
|
||||
/// Compaction completed and the session was rotated.
|
||||
CompactDone {
|
||||
lifecycle: CompactionLifecycle,
|
||||
},
|
||||
/// Compaction failed or was cancelled. The session is unchanged.
|
||||
CompactFailed {
|
||||
error: String,
|
||||
lifecycle: CompactionLifecycle,
|
||||
},
|
||||
Shutdown,
|
||||
}
|
||||
@@ -895,6 +1032,7 @@ pub enum WorkerStatus {
|
||||
Idle,
|
||||
Running,
|
||||
Paused,
|
||||
Stopped,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
@@ -1377,7 +1515,7 @@ mod tests {
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "completions");
|
||||
assert_eq!(parsed["data"]["kind"], "file");
|
||||
assert_eq!(parsed["data"]["entries"][0]["value"], "clear");
|
||||
assert_eq!(parsed["data"]["entries"][0]["value"], "src/main.rs");
|
||||
|
||||
// is_dir defaults to false on inbound payloads that omit it.
|
||||
let inbound =
|
||||
@@ -1397,7 +1535,17 @@ mod tests {
|
||||
#[test]
|
||||
fn event_snapshot_format() {
|
||||
let event = Event::Snapshot {
|
||||
entries: vec![serde_json::json!({"kind": "user_input", "ts": 1, "segments": []})],
|
||||
session: SessionSnapshot {
|
||||
entries: vec![SessionSnapshotEntry {
|
||||
entry_id: "entry-1".into(),
|
||||
timestamp: 1,
|
||||
provenance: SessionEntryProvenance::HumanInput,
|
||||
derived_from: Vec::new(),
|
||||
data: SessionSnapshotEntryData::UserInput {
|
||||
segments: Vec::new(),
|
||||
},
|
||||
}],
|
||||
},
|
||||
greeting: Greeting {
|
||||
worker_name: "test".into(),
|
||||
cwd: "/tmp".into(),
|
||||
@@ -1415,8 +1563,12 @@ mod tests {
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "snapshot");
|
||||
assert!(parsed["data"]["entries"].is_array());
|
||||
assert_eq!(parsed["data"]["entries"][0]["kind"], "user_input");
|
||||
assert!(parsed["data"]["session"]["entries"].is_array());
|
||||
assert_eq!(
|
||||
parsed["data"]["session"]["entries"][0]["kind"],
|
||||
"user_input"
|
||||
);
|
||||
assert_eq!(parsed["data"]["session"]["entries"][0]["timestamp"], 1);
|
||||
assert_eq!(parsed["data"]["greeting"]["worker_name"], "test");
|
||||
assert_eq!(parsed["data"]["greeting"]["tools"][0], "Read");
|
||||
assert_eq!(parsed["data"]["greeting"]["context_window"], 200_000);
|
||||
@@ -1426,7 +1578,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn event_snapshot_in_flight_roundtrip_and_default() {
|
||||
let inbound = r#"{"event":"snapshot","data":{"entries":[],"greeting":{"worker_name":"test","cwd":"/tmp","provider":"p","model":"m","scope_summary":"s","tools":[]},"status":"running"}}"#;
|
||||
let inbound = r#"{"event":"snapshot","data":{"session":{"entries":[]},"greeting":{"worker_name":"test","cwd":"/tmp","provider":"p","model":"m","scope_summary":"s","tools":[]},"status":"running"}}"#;
|
||||
let decoded: Event = serde_json::from_str(inbound).unwrap();
|
||||
match decoded {
|
||||
Event::Snapshot { in_flight, .. } => assert!(in_flight.is_empty()),
|
||||
@@ -1434,7 +1586,9 @@ mod tests {
|
||||
}
|
||||
|
||||
let event = Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
session: SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
greeting: Greeting {
|
||||
worker_name: "test".into(),
|
||||
cwd: "/tmp".into(),
|
||||
@@ -1500,15 +1654,17 @@ mod tests {
|
||||
#[test]
|
||||
fn event_segment_rotated_roundtrip() {
|
||||
let event = Event::SegmentRotated {
|
||||
entry: serde_json::json!({"kind": "segment_start", "ts": 1, "history": []}),
|
||||
session: SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "segment_rotated");
|
||||
assert_eq!(parsed["data"]["entry"]["kind"], "segment_start");
|
||||
assert!(parsed["data"]["session"]["entries"].is_array());
|
||||
let decoded: Event = serde_json::from_str(&json).unwrap();
|
||||
match decoded {
|
||||
Event::SegmentRotated { entry } => assert_eq!(entry["kind"], "segment_start"),
|
||||
Event::SegmentRotated { session } => assert!(session.entries.is_empty()),
|
||||
other => panic!("expected SegmentRotated, got {other:?}"),
|
||||
}
|
||||
}
|
||||
@@ -1584,8 +1740,8 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_snapshot_legacy_without_status_defaults_to_idle() {
|
||||
let json = r#"{"event":"snapshot","data":{"entries":[],"greeting":{"worker_name":"test","cwd":"/tmp","provider":"anthropic","model":"claude","scope_summary":"","tools":[]}}}"#;
|
||||
fn event_snapshot_without_status_defaults_to_idle() {
|
||||
let json = r#"{"event":"snapshot","data":{"session":{"entries":[]},"greeting":{"worker_name":"test","cwd":"/tmp","provider":"anthropic","model":"claude","scope_summary":"","tools":[]}}}"#;
|
||||
let decoded: Event = serde_json::from_str(json).unwrap();
|
||||
match decoded {
|
||||
Event::Snapshot {
|
||||
@@ -1732,45 +1888,74 @@ mod tests {
|
||||
assert_eq!(parsed["data"]["timestamp_ms"], 1_700_000_000_000i64);
|
||||
}
|
||||
|
||||
fn test_compaction_lifecycle(state: CompactionLifecycleState) -> CompactionLifecycle {
|
||||
CompactionLifecycle {
|
||||
schema_version: 2,
|
||||
compaction_id: "0192f0e8-4d84-7d6e-a000-000000000000".into(),
|
||||
revision: 1,
|
||||
internal_worker: None,
|
||||
state,
|
||||
started_at_ms: 1_700_000_000_000,
|
||||
ended_at_ms: None,
|
||||
summary: None,
|
||||
error: None,
|
||||
new_segment_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_compact_start_roundtrip() {
|
||||
let event = Event::CompactStart;
|
||||
let event = Event::CompactStart {
|
||||
lifecycle: test_compaction_lifecycle(CompactionLifecycleState::Running),
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
assert_eq!(json, r#"{"event":"compact_start"}"#);
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "compact_start");
|
||||
assert_eq!(parsed["data"]["lifecycle"]["state"], "running");
|
||||
let decoded: Event = serde_json::from_str(&json).unwrap();
|
||||
assert!(matches!(decoded, Event::CompactStart));
|
||||
assert!(matches!(decoded, Event::CompactStart { lifecycle } if lifecycle.revision == 1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_compact_done_roundtrip() {
|
||||
let id = uuid::Uuid::parse_str("0192f0e8-4d84-7d6e-a000-000000000001").unwrap();
|
||||
let event = Event::CompactDone { new_segment_id: id };
|
||||
let mut lifecycle = test_compaction_lifecycle(CompactionLifecycleState::Done);
|
||||
lifecycle.new_segment_id = Some(id.to_string());
|
||||
lifecycle.summary = Some("accepted summary".into());
|
||||
let event = Event::CompactDone { lifecycle };
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "compact_done");
|
||||
assert_eq!(
|
||||
parsed["data"]["new_segment_id"],
|
||||
parsed["data"]["lifecycle"]["new_segment_id"],
|
||||
"0192f0e8-4d84-7d6e-a000-000000000001"
|
||||
);
|
||||
let decoded: Event = serde_json::from_str(&json).unwrap();
|
||||
match decoded {
|
||||
Event::CompactDone { new_segment_id } => assert_eq!(new_segment_id, id),
|
||||
Event::CompactDone { lifecycle } => {
|
||||
assert_eq!(
|
||||
lifecycle.new_segment_id.as_deref(),
|
||||
Some(id.to_string().as_str())
|
||||
)
|
||||
}
|
||||
other => panic!("expected CompactDone, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_compact_failed_roundtrip() {
|
||||
let event = Event::CompactFailed {
|
||||
error: "provider 429".into(),
|
||||
};
|
||||
let mut lifecycle = test_compaction_lifecycle(CompactionLifecycleState::Failed);
|
||||
lifecycle.error = Some("provider 429".into());
|
||||
let event = Event::CompactFailed { lifecycle };
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["event"], "compact_failed");
|
||||
assert_eq!(parsed["data"]["error"], "provider 429");
|
||||
assert_eq!(parsed["data"]["lifecycle"]["error"], "provider 429");
|
||||
let decoded: Event = serde_json::from_str(&json).unwrap();
|
||||
match decoded {
|
||||
Event::CompactFailed { error } => assert_eq!(error, "provider 429"),
|
||||
Event::CompactFailed { lifecycle } => {
|
||||
assert_eq!(lifecycle.error.as_deref(), Some("provider 429"))
|
||||
}
|
||||
other => panic!("expected CompactFailed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
@@ -1781,6 +1966,7 @@ mod tests {
|
||||
id: "call_1".into(),
|
||||
summary: "Read 128 bytes".into(),
|
||||
output: Some("hello world".into()),
|
||||
disposition: Some(ToolResultDisposition::Success),
|
||||
is_error: false,
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
@@ -1797,11 +1983,13 @@ mod tests {
|
||||
id,
|
||||
summary,
|
||||
output,
|
||||
disposition,
|
||||
is_error,
|
||||
} => {
|
||||
assert_eq!(id, "call_1");
|
||||
assert_eq!(summary, "Read 128 bytes");
|
||||
assert_eq!(output.as_deref(), Some("hello world"));
|
||||
assert_eq!(disposition, Some(ToolResultDisposition::Success));
|
||||
assert!(!is_error);
|
||||
}
|
||||
other => panic!("expected ToolResult, got {other:?}"),
|
||||
@@ -1814,6 +2002,7 @@ mod tests {
|
||||
id: "call_2".into(),
|
||||
summary: "ok".into(),
|
||||
output: None,
|
||||
disposition: Some(ToolResultDisposition::Success),
|
||||
is_error: false,
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
@@ -1829,6 +2018,7 @@ mod tests {
|
||||
id: "call_3".into(),
|
||||
summary: "invalid argument".into(),
|
||||
output: None,
|
||||
disposition: Some(ToolResultDisposition::Error),
|
||||
is_error: true,
|
||||
};
|
||||
let json = serde_json::to_string(&event).unwrap();
|
||||
@@ -1962,11 +2152,11 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_snapshot_defaults_internal_workers_to_empty() {
|
||||
fn snapshot_defaults_internal_workers_to_empty() {
|
||||
let snapshot: Event = serde_json::from_value(serde_json::json!({
|
||||
"event": "snapshot",
|
||||
"data": {
|
||||
"entries": [],
|
||||
"session": { "entries": [] },
|
||||
"greeting": {
|
||||
"worker_name": "parent",
|
||||
"cwd": ".",
|
||||
|
||||
@@ -4,11 +4,13 @@ use ts_rs::{Config, TS};
|
||||
|
||||
use crate::{
|
||||
Alert, AlertLevel, AlertSource, CommandEvent, CommandSnapshot, CommandStatus, CommandStream,
|
||||
CommandStreamSlice, CompletionEntry, CompletionKind, ErrorCode, Event, Greeting, InFlightBlock,
|
||||
InFlightSnapshot, InFlightToolCallState, InternalWorkerKind, InternalWorkerRef,
|
||||
InternalWorkerSnapshot, InvokeKind, MemoryWorkerEvent, Method, Permission, RewindSummary,
|
||||
RewindTarget, RewindTargetId, RunResult, ScopeRule, Segment, TurnResult, WorkerEvent,
|
||||
WorkerStatus,
|
||||
CommandStreamSlice, CompactionLifecycle, CompactionLifecycleState, CompletionEntry,
|
||||
CompletionKind, ErrorCode, Event, Greeting, InFlightBlock, InFlightSnapshot,
|
||||
InFlightToolCallState, InternalWorkerKind, InternalWorkerRef, InternalWorkerSnapshot,
|
||||
InvokeKind, MemoryWorkerEvent, Method, Permission, RewindSummary, RewindTarget, RewindTargetId,
|
||||
RunResult, ScopeRule, Segment, SessionContentPart, SessionEntryProvenance, SessionMessageRole,
|
||||
SessionSnapshot, SessionSnapshotEntry, SessionSnapshotEntryData, SessionToolAttachment,
|
||||
ToolResultDisposition, TurnResult, WorkerEvent, WorkerStatus,
|
||||
subscription::{
|
||||
EventSubscriptionSelector, SubscriptionEvent, SubscriptionEventPayload, SubscriptionFrame,
|
||||
SubscriptionFramePayload, SubscriptionId, SubscriptionRejectionCode, SubscriptionRequest,
|
||||
@@ -45,6 +47,7 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<TurnResult>(&cfg, &mut output);
|
||||
push_decl::<InvokeKind>(&cfg, &mut output);
|
||||
push_decl::<RunResult>(&cfg, &mut output);
|
||||
push_decl::<ToolResultDisposition>(&cfg, &mut output);
|
||||
push_decl::<ErrorCode>(&cfg, &mut output);
|
||||
push_decl::<Permission>(&cfg, &mut output);
|
||||
push_decl::<InFlightToolCallState>(&cfg, &mut output);
|
||||
@@ -53,6 +56,8 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<CommandStreamSlice>(&cfg, &mut output);
|
||||
push_decl::<CommandSnapshot>(&cfg, &mut output);
|
||||
push_decl::<CommandEvent>(&cfg, &mut output);
|
||||
push_decl::<CompactionLifecycleState>(&cfg, &mut output);
|
||||
push_decl::<CompactionLifecycle>(&cfg, &mut output);
|
||||
push_decl::<ScopeRule>(&cfg, &mut output);
|
||||
push_decl::<CompletionEntry>(&cfg, &mut output);
|
||||
push_decl::<RewindTargetId>(&cfg, &mut output);
|
||||
@@ -60,6 +65,13 @@ pub fn generated_protocol_types() -> String {
|
||||
push_decl::<RewindSummary>(&cfg, &mut output);
|
||||
push_decl::<InFlightBlock>(&cfg, &mut output);
|
||||
push_decl::<InFlightSnapshot>(&cfg, &mut output);
|
||||
push_decl::<SessionEntryProvenance>(&cfg, &mut output);
|
||||
push_decl::<SessionMessageRole>(&cfg, &mut output);
|
||||
push_decl::<SessionContentPart>(&cfg, &mut output);
|
||||
push_decl::<SessionToolAttachment>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshotEntryData>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshotEntry>(&cfg, &mut output);
|
||||
push_decl::<SessionSnapshot>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerKind>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerRef>(&cfg, &mut output);
|
||||
push_decl::<InternalWorkerSnapshot>(&cfg, &mut output);
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
//! Serializable history entries with restore-authoritative logical identity and origin.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::LoggedItem;
|
||||
|
||||
/// Stable logical identity of one model-visible history entry.
|
||||
///
|
||||
/// This value is generated at the trusted Worker session boundary and copied
|
||||
/// unchanged across fork, rewind, compaction retention, restore, and reboot.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct LoggedSessionHistoryEntryId(pub String);
|
||||
|
||||
impl LoggedSessionHistoryEntryId {
|
||||
pub fn new() -> Self {
|
||||
Self(uuid::Uuid::now_v7().to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for LoggedSessionHistoryEntryId {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// Bounded subject snapshot. It is evidence, not a live authorization handle.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct LoggedWorkerSubject {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub workspace_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub runtime_id: Option<String>,
|
||||
pub worker_id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum LoggedSessionHistoryOrigin {
|
||||
HumanInput {
|
||||
account_id: String,
|
||||
},
|
||||
WorkerInput {
|
||||
actor: LoggedWorkerSubject,
|
||||
},
|
||||
FlowInstruction {
|
||||
selector: String,
|
||||
definition_id: String,
|
||||
definition_revision: u64,
|
||||
instance_id: String,
|
||||
state_id: String,
|
||||
},
|
||||
BackendInstruction {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
operation_id: Option<String>,
|
||||
},
|
||||
ModelOutput {
|
||||
worker: LoggedWorkerSubject,
|
||||
},
|
||||
ToolOutput {
|
||||
worker: LoggedWorkerSubject,
|
||||
},
|
||||
DerivedSummary,
|
||||
LegacyUnknown,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct LoggedHistoryDerivation {
|
||||
pub sources: Vec<LoggedSessionHistoryEntryId>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct LoggedSessionHistoryMetadata {
|
||||
pub entry_id: LoggedSessionHistoryEntryId,
|
||||
pub origin: LoggedSessionHistoryOrigin,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub derivation: Option<LoggedHistoryDerivation>,
|
||||
}
|
||||
|
||||
impl LoggedSessionHistoryMetadata {
|
||||
pub fn legacy_unknown() -> Self {
|
||||
Self {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Persisted item and metadata are one value so transforms cannot reorder or
|
||||
/// truncate one without the other.
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct LoggedHistoryEntry {
|
||||
pub item: LoggedItem,
|
||||
pub metadata: LoggedSessionHistoryMetadata,
|
||||
}
|
||||
|
||||
/// Typed system-item history record. The typed system event remains available
|
||||
/// to client replay while its model-visible projection carries the same stable
|
||||
/// metadata used by live history.
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub struct LoggedSystemHistoryEntry {
|
||||
pub item: crate::SystemItem,
|
||||
pub metadata: LoggedSessionHistoryMetadata,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::LoggedRole;
|
||||
use agen::llm_client::RequestConfig;
|
||||
|
||||
#[test]
|
||||
fn logged_history_entry_round_trip_preserves_id_origin_and_derivation() {
|
||||
let source_id = LoggedSessionHistoryEntryId::new();
|
||||
let entry = LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::User,
|
||||
content: vec![crate::LoggedContentPart::Text {
|
||||
text: "preference".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::HumanInput {
|
||||
account_id: "account-1".into(),
|
||||
},
|
||||
derivation: Some(LoggedHistoryDerivation {
|
||||
sources: vec![source_id.clone()],
|
||||
}),
|
||||
},
|
||||
};
|
||||
let encoded = serde_json::to_vec(&entry).unwrap();
|
||||
let decoded: LoggedHistoryEntry = serde_json::from_slice(&encoded).unwrap();
|
||||
assert_eq!(decoded, entry);
|
||||
assert_eq!(
|
||||
decoded.metadata.derivation.unwrap().sources,
|
||||
vec![source_id]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn annotated_segment_start_is_restore_visible_without_projecting_metadata() {
|
||||
let session_id = uuid::Uuid::now_v7();
|
||||
let history_entry = LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::Assistant,
|
||||
content: vec![crate::LoggedContentPart::Text {
|
||||
text: "answer".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata::legacy_unknown(),
|
||||
};
|
||||
let state = crate::collect_state(&[crate::LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![history_entry],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}]);
|
||||
assert_eq!(state.history[0].as_text(), Some("answer"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
//! Versioned decoder for Session schemas that predate canonical annotated history.
|
||||
//!
|
||||
//! These types are intentionally private to `session-store`. Current writers,
|
||||
//! replay, and public projections use [`crate::LogEntry`] exclusively; only the
|
||||
//! Worker Session schema migration is allowed to deserialize these shapes.
|
||||
|
||||
use agen::llm_client::types::RequestConfig;
|
||||
use protocol::Segment;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{
|
||||
LogEntry, LoggedHistoryEntry, LoggedItem, LoggedSessionHistoryEntryId,
|
||||
LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin, LoggedSystemHistoryEntry, SegmentId,
|
||||
SegmentOrigin, SessionExtension, SessionId, SystemItem,
|
||||
};
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
enum LegacyHistoryLogEntry {
|
||||
SegmentStart {
|
||||
ts: u64,
|
||||
session_id: SessionId,
|
||||
system_prompt: Option<String>,
|
||||
config: RequestConfig,
|
||||
history: Vec<LoggedItem>,
|
||||
#[serde(default)]
|
||||
forked_from: Option<SegmentOrigin>,
|
||||
#[serde(default)]
|
||||
compacted_from: Option<SegmentOrigin>,
|
||||
},
|
||||
UserInput {
|
||||
ts: u64,
|
||||
segments: Vec<Segment>,
|
||||
#[serde(default)]
|
||||
extensions: Vec<SessionExtension>,
|
||||
},
|
||||
AssistantItem {
|
||||
ts: u64,
|
||||
item: LoggedItem,
|
||||
},
|
||||
ToolResult {
|
||||
ts: u64,
|
||||
item: LoggedItem,
|
||||
},
|
||||
SystemItem {
|
||||
ts: u64,
|
||||
item: SystemItem,
|
||||
},
|
||||
}
|
||||
|
||||
/// Schema-v1 decoder. Non-history records already had their current shape, so
|
||||
/// they pass through `LogEntry`; legacy history records are converted below.
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum LegacySessionLogEntryV1 {
|
||||
History(LegacyHistoryLogEntry),
|
||||
Current(LogEntry),
|
||||
}
|
||||
|
||||
/// Schema v2 retained the v1 history shapes while adding non-history records.
|
||||
/// Keep a distinct type so supported source versions remain explicit rather
|
||||
/// than turning migration compatibility into the current `LogEntry` contract.
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum LegacySessionLogEntryV2 {
|
||||
History(LegacyHistoryLogEntry),
|
||||
Current(LogEntry),
|
||||
}
|
||||
|
||||
pub(crate) fn decode_entry(
|
||||
schema_version: u32,
|
||||
line: &str,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
line_index: usize,
|
||||
) -> Result<LogEntry, serde_json::Error> {
|
||||
let entry = match schema_version {
|
||||
1 => match serde_json::from_str::<LegacySessionLogEntryV1>(line)? {
|
||||
LegacySessionLogEntryV1::History(entry) => Entry::History(entry),
|
||||
LegacySessionLogEntryV1::Current(entry) => Entry::Current(entry),
|
||||
},
|
||||
2 => match serde_json::from_str::<LegacySessionLogEntryV2>(line)? {
|
||||
LegacySessionLogEntryV2::History(entry) => Entry::History(entry),
|
||||
LegacySessionLogEntryV2::Current(entry) => Entry::Current(entry),
|
||||
},
|
||||
_ => unreachable!("legacy decoder called for unsupported schema {schema_version}"),
|
||||
};
|
||||
Ok(match entry {
|
||||
Entry::History(entry) => {
|
||||
canonicalize_history_entry(session_id, segment_id, line_index, entry)
|
||||
}
|
||||
Entry::Current(entry) => entry,
|
||||
})
|
||||
}
|
||||
|
||||
enum Entry {
|
||||
History(LegacyHistoryLogEntry),
|
||||
Current(LogEntry),
|
||||
}
|
||||
|
||||
fn legacy_metadata(
|
||||
segment_id: SegmentId,
|
||||
line_index: usize,
|
||||
item_index: usize,
|
||||
) -> LoggedSessionHistoryMetadata {
|
||||
let mut identity = Vec::with_capacity(32);
|
||||
identity.extend_from_slice(segment_id.as_bytes());
|
||||
identity.extend_from_slice(&(line_index as u64).to_be_bytes());
|
||||
identity.extend_from_slice(&(item_index as u64).to_be_bytes());
|
||||
LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId(format!(
|
||||
"l-{}",
|
||||
base64::Engine::encode(&base64::engine::general_purpose::URL_SAFE_NO_PAD, identity)
|
||||
)),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn canonicalize_history_entry(
|
||||
_session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
line_index: usize,
|
||||
entry: LegacyHistoryLogEntry,
|
||||
) -> LogEntry {
|
||||
match entry {
|
||||
LegacyHistoryLogEntry::SegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
system_prompt,
|
||||
config,
|
||||
history,
|
||||
forked_from,
|
||||
compacted_from,
|
||||
} => LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
system_prompt,
|
||||
config,
|
||||
history: history
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(item_index, item)| LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, item_index),
|
||||
})
|
||||
.collect(),
|
||||
forked_from,
|
||||
compacted_from,
|
||||
},
|
||||
LegacyHistoryLogEntry::UserInput {
|
||||
ts,
|
||||
segments,
|
||||
extensions,
|
||||
} => LogEntry::AnnotatedUserInput {
|
||||
ts,
|
||||
history: vec![LoggedHistoryEntry {
|
||||
item: LoggedItem::from(agen::Item::user_message(Segment::flatten_to_text(
|
||||
&segments,
|
||||
))),
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
}],
|
||||
segments,
|
||||
extensions,
|
||||
},
|
||||
LegacyHistoryLogEntry::AssistantItem { ts, item } => LogEntry::AnnotatedAssistantItem {
|
||||
ts,
|
||||
entry: LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
},
|
||||
},
|
||||
LegacyHistoryLogEntry::ToolResult { ts, item } => LogEntry::AnnotatedToolResult {
|
||||
ts,
|
||||
entry: LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
},
|
||||
},
|
||||
LegacyHistoryLogEntry::SystemItem { ts, item } => LogEntry::AnnotatedSystemItem {
|
||||
ts,
|
||||
entry: LoggedSystemHistoryEntry {
|
||||
item,
|
||||
metadata: legacy_metadata(segment_id, line_index, 0),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -26,13 +26,16 @@
|
||||
//! let (session_id, segment_id) = create_segment(&store, SegmentStartState {
|
||||
//! system_prompt: None,
|
||||
//! config: &config,
|
||||
//! history: &[],
|
||||
//! history: Vec::new(),
|
||||
//! })?;
|
||||
//! ```
|
||||
|
||||
pub mod event_trace;
|
||||
pub mod fs_store;
|
||||
pub mod history;
|
||||
mod legacy_session_log;
|
||||
pub mod logged_item;
|
||||
pub mod public_snapshot;
|
||||
pub mod segment;
|
||||
pub mod segment_log;
|
||||
pub mod store;
|
||||
@@ -44,9 +47,14 @@ pub use agen::UsageRecord;
|
||||
pub use agen::llm_client::types::{ContentPart, Item, Role};
|
||||
pub use event_trace::{TraceEntry, TracePayload};
|
||||
pub use fs_store::FsStore;
|
||||
pub use history::{
|
||||
LoggedHistoryDerivation, LoggedHistoryEntry, LoggedSessionHistoryEntryId,
|
||||
LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin, LoggedSystemHistoryEntry,
|
||||
LoggedWorkerSubject,
|
||||
};
|
||||
pub use logged_item::{LoggedContentPart, LoggedItem, LoggedRole, from_logged, to_logged};
|
||||
pub use segment::{
|
||||
SegmentStartState, append_entry, append_system_item, classify_history_item,
|
||||
SegmentStartState, append_entry, append_system_item, classify_logged_history_entry,
|
||||
create_compacted_segment, create_segment, create_segment_with_ids, ensure_head_or_fork, fork,
|
||||
fork_at, restore, restore_by_segment, save_config_changed, save_delta, save_extension,
|
||||
save_run_completed, save_run_errored, save_turn_end, save_usage, save_user_input,
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
|
||||
use agen::{
|
||||
llm_client::types::{ContentPart, Item, Role},
|
||||
tool::{Attachment, ImageAttachment},
|
||||
tool::{Attachment, ImageAttachment, ToolResultDisposition},
|
||||
};
|
||||
use base64::{Engine as _, engine::general_purpose::STANDARD};
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Error as _};
|
||||
@@ -61,6 +61,8 @@ pub enum LoggedItem {
|
||||
content: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
attachments: Vec<LoggedAttachment>,
|
||||
#[serde(default, skip_serializing_if = "ToolResultDisposition::is_success")]
|
||||
disposition: ToolResultDisposition,
|
||||
#[serde(default, skip_serializing_if = "is_false")]
|
||||
is_error: bool,
|
||||
},
|
||||
@@ -128,6 +130,7 @@ impl From<&Item> for LoggedItem {
|
||||
summary,
|
||||
content,
|
||||
attachments,
|
||||
disposition,
|
||||
is_error,
|
||||
..
|
||||
} => Self::ToolResult {
|
||||
@@ -135,6 +138,7 @@ impl From<&Item> for LoggedItem {
|
||||
summary: summary.clone(),
|
||||
content: content.clone(),
|
||||
attachments: attachments.iter().map(LoggedAttachment::from).collect(),
|
||||
disposition: *disposition,
|
||||
is_error: *is_error,
|
||||
},
|
||||
Item::Reasoning {
|
||||
@@ -184,15 +188,24 @@ impl From<LoggedItem> for Item {
|
||||
summary,
|
||||
content,
|
||||
attachments,
|
||||
disposition,
|
||||
is_error,
|
||||
} => Item::ToolResult {
|
||||
id: None,
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
is_error,
|
||||
attachments: attachments.into_iter().map(Attachment::from).collect(),
|
||||
},
|
||||
} => {
|
||||
let disposition = if is_error && disposition.is_success() {
|
||||
ToolResultDisposition::Error
|
||||
} else {
|
||||
disposition
|
||||
};
|
||||
Item::ToolResult {
|
||||
id: None,
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
disposition,
|
||||
is_error,
|
||||
attachments: attachments.into_iter().map(Attachment::from).collect(),
|
||||
}
|
||||
}
|
||||
LoggedItem::Reasoning {
|
||||
text,
|
||||
summary,
|
||||
@@ -430,6 +443,42 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn outcome_unknown_tool_result_round_trips_as_terminal() {
|
||||
let original = Item::tool_result_item_with_disposition_and_attachments(
|
||||
"call_unknown",
|
||||
"outcome unknown",
|
||||
Some("bounded progress".to_string()),
|
||||
ToolResultDisposition::OutcomeUnknown,
|
||||
Vec::new(),
|
||||
);
|
||||
let logged: LoggedItem = (&original).into();
|
||||
let json = serde_json::to_string(&logged).unwrap();
|
||||
assert!(json.contains(r#""disposition":"outcome_unknown""#));
|
||||
match Item::from(serde_json::from_str::<LoggedItem>(&json).unwrap()) {
|
||||
Item::ToolResult {
|
||||
disposition,
|
||||
is_error,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(disposition, ToolResultDisposition::OutcomeUnknown);
|
||||
assert!(is_error);
|
||||
}
|
||||
other => panic!("unexpected variant: {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_error_tool_result_infers_error_disposition() {
|
||||
let legacy = r#"{"kind":"tool_result","call_id":"call_old","summary":"failed","content":null,"is_error":true}"#;
|
||||
match Item::from(serde_json::from_str::<LoggedItem>(legacy).unwrap()) {
|
||||
Item::ToolResult { disposition, .. } => {
|
||||
assert_eq!(disposition, ToolResultDisposition::Error)
|
||||
}
|
||||
other => panic!("unexpected variant: {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_result_persistence_round_trips_binary_attachments() {
|
||||
let original = Item::tool_result_item_with_attachments(
|
||||
|
||||
@@ -0,0 +1,484 @@
|
||||
use base64::{
|
||||
Engine as _,
|
||||
engine::general_purpose::{STANDARD as BASE64, URL_SAFE_NO_PAD},
|
||||
};
|
||||
use protocol::{
|
||||
Segment, SessionContentPart, SessionEntryProvenance, SessionMessageRole, SessionSnapshot,
|
||||
SessionSnapshotEntry, SessionSnapshotEntryData, SessionToolAttachment,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
LogEntry, LoggedContentPart, LoggedHistoryEntry, LoggedItem, LoggedRole,
|
||||
LoggedSessionHistoryOrigin, SessionId, SystemItem,
|
||||
};
|
||||
|
||||
/// Project a complete current-segment log. A valid segment starts with one
|
||||
/// canonical annotated SegmentStart record; malformed partial input uses the
|
||||
/// nil session only to keep the public failure projection deterministic.
|
||||
pub fn project_current_session_snapshot(log: &[LogEntry]) -> SessionSnapshot {
|
||||
let session_id = log.iter().find_map(|entry| match entry {
|
||||
LogEntry::AnnotatedSegmentStart { session_id, .. } => Some(*session_id),
|
||||
_ => None,
|
||||
});
|
||||
project_session_snapshot(session_id.unwrap_or_else(SessionId::nil), log)
|
||||
}
|
||||
|
||||
/// Project the current durable segment into the only public session-history
|
||||
/// representation. Append-log records remain an internal persistence format.
|
||||
pub fn project_session_snapshot(session_id: SessionId, log: &[LogEntry]) -> SessionSnapshot {
|
||||
let mut session_key = session_id;
|
||||
let mut entries = Vec::new();
|
||||
|
||||
for (log_index, record) in log.iter().enumerate() {
|
||||
match record {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
history,
|
||||
..
|
||||
} => {
|
||||
session_key = *session_id;
|
||||
entries.clear();
|
||||
extend_history(&mut entries, history, None, *ts);
|
||||
}
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts,
|
||||
segments,
|
||||
history,
|
||||
..
|
||||
} => extend_history(&mut entries, history, Some(segments), *ts),
|
||||
LogEntry::AnnotatedAssistantItem { ts, entry }
|
||||
| LogEntry::AnnotatedToolResult { ts, entry } => {
|
||||
if let Some(data) = project_item(&entry.item) {
|
||||
entries.push(history_entry(entry, *ts, data));
|
||||
}
|
||||
}
|
||||
LogEntry::AnnotatedSystemItem { ts, entry } => entries.push(system_entry(
|
||||
&entry.item,
|
||||
entry.metadata.entry_id.0.clone(),
|
||||
*ts,
|
||||
provenance(&entry.metadata.origin),
|
||||
derivation_ids(entry),
|
||||
)),
|
||||
LogEntry::RunErrored { ts, message, .. } => entries.push(legacy_entry(
|
||||
&session_key,
|
||||
log_index,
|
||||
0,
|
||||
*ts,
|
||||
SessionSnapshotEntryData::RunError {
|
||||
message: message.clone(),
|
||||
},
|
||||
)),
|
||||
// Run checkpoints, configuration, usage, and extension state are
|
||||
// controller/storage authority rather than committed conversation.
|
||||
LogEntry::Invoke { .. }
|
||||
| LogEntry::TurnEnd { .. }
|
||||
| LogEntry::RunCompleted { .. }
|
||||
| LogEntry::ActiveRunCheckpoint { .. }
|
||||
| LogEntry::PausedTurnAbandoned { .. }
|
||||
| LogEntry::ConfigChanged { .. }
|
||||
| LogEntry::LlmUsage { .. }
|
||||
| LogEntry::Extension { .. } => {}
|
||||
}
|
||||
}
|
||||
|
||||
SessionSnapshot { entries }
|
||||
}
|
||||
|
||||
fn extend_history(
|
||||
output: &mut Vec<SessionSnapshotEntry>,
|
||||
history: &[LoggedHistoryEntry],
|
||||
input_segments: Option<&Vec<Segment>>,
|
||||
timestamp: u64,
|
||||
) {
|
||||
let mut attached_segments = false;
|
||||
for entry in history {
|
||||
let data = if !attached_segments
|
||||
&& input_segments.is_some()
|
||||
&& matches!(
|
||||
&entry.item,
|
||||
LoggedItem::Message {
|
||||
role: LoggedRole::User,
|
||||
..
|
||||
}
|
||||
) {
|
||||
attached_segments = true;
|
||||
SessionSnapshotEntryData::UserInput {
|
||||
segments: input_segments.cloned().unwrap_or_default(),
|
||||
}
|
||||
} else {
|
||||
let Some(data) = project_item(&entry.item) else {
|
||||
continue;
|
||||
};
|
||||
data
|
||||
};
|
||||
output.push(history_entry(entry, timestamp, data));
|
||||
}
|
||||
}
|
||||
|
||||
fn history_entry(
|
||||
entry: &LoggedHistoryEntry,
|
||||
timestamp: u64,
|
||||
data: SessionSnapshotEntryData,
|
||||
) -> SessionSnapshotEntry {
|
||||
SessionSnapshotEntry {
|
||||
entry_id: entry.metadata.entry_id.0.clone(),
|
||||
timestamp,
|
||||
provenance: provenance(&entry.metadata.origin),
|
||||
derived_from: entry
|
||||
.metadata
|
||||
.derivation
|
||||
.as_ref()
|
||||
.map(|derivation| {
|
||||
derivation
|
||||
.sources
|
||||
.iter()
|
||||
.map(|source| source.0.clone())
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
data,
|
||||
}
|
||||
}
|
||||
|
||||
fn derivation_ids(entry: &crate::LoggedSystemHistoryEntry) -> Vec<String> {
|
||||
entry
|
||||
.metadata
|
||||
.derivation
|
||||
.as_ref()
|
||||
.map(|derivation| {
|
||||
derivation
|
||||
.sources
|
||||
.iter()
|
||||
.map(|source| source.0.clone())
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn legacy_entry(
|
||||
session_key: &SessionId,
|
||||
log_index: usize,
|
||||
item_index: usize,
|
||||
timestamp: u64,
|
||||
data: SessionSnapshotEntryData,
|
||||
) -> SessionSnapshotEntry {
|
||||
SessionSnapshotEntry {
|
||||
entry_id: legacy_entry_id(session_key, log_index, item_index),
|
||||
timestamp,
|
||||
provenance: SessionEntryProvenance::LegacyUnknown,
|
||||
derived_from: Vec::new(),
|
||||
data,
|
||||
}
|
||||
}
|
||||
|
||||
fn legacy_entry_id(session_key: &SessionId, log_index: usize, item_index: usize) -> String {
|
||||
let mut identity = Vec::with_capacity(32);
|
||||
identity.extend_from_slice(session_key.as_bytes());
|
||||
identity.extend_from_slice(&(log_index as u64).to_be_bytes());
|
||||
identity.extend_from_slice(&(item_index as u64).to_be_bytes());
|
||||
format!("l-{}", URL_SAFE_NO_PAD.encode(identity))
|
||||
}
|
||||
|
||||
fn provenance(origin: &LoggedSessionHistoryOrigin) -> SessionEntryProvenance {
|
||||
match origin {
|
||||
LoggedSessionHistoryOrigin::HumanInput { .. } => SessionEntryProvenance::HumanInput,
|
||||
LoggedSessionHistoryOrigin::WorkerInput { .. } => SessionEntryProvenance::WorkerInput,
|
||||
LoggedSessionHistoryOrigin::FlowInstruction { .. } => {
|
||||
SessionEntryProvenance::FlowInstruction
|
||||
}
|
||||
LoggedSessionHistoryOrigin::BackendInstruction { .. } => {
|
||||
SessionEntryProvenance::BackendInstruction
|
||||
}
|
||||
LoggedSessionHistoryOrigin::ModelOutput { .. } => SessionEntryProvenance::ModelOutput,
|
||||
LoggedSessionHistoryOrigin::ToolOutput { .. } => SessionEntryProvenance::ToolOutput,
|
||||
LoggedSessionHistoryOrigin::DerivedSummary => SessionEntryProvenance::DerivedSummary,
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown => SessionEntryProvenance::LegacyUnknown,
|
||||
}
|
||||
}
|
||||
|
||||
fn project_item(item: &LoggedItem) -> Option<SessionSnapshotEntryData> {
|
||||
match item {
|
||||
LoggedItem::Message { role, content } => {
|
||||
let role = match role {
|
||||
LoggedRole::User => SessionMessageRole::User,
|
||||
LoggedRole::Assistant => SessionMessageRole::Assistant,
|
||||
// System prompts and instruction history never cross the public
|
||||
// snapshot boundary. Typed SystemItems have separate records.
|
||||
LoggedRole::System => return None,
|
||||
};
|
||||
Some(SessionSnapshotEntryData::Message {
|
||||
role,
|
||||
content: content
|
||||
.iter()
|
||||
.map(|part| match part {
|
||||
LoggedContentPart::Text { text } => {
|
||||
SessionContentPart::Text { text: text.clone() }
|
||||
}
|
||||
LoggedContentPart::Refusal { refusal } => SessionContentPart::Refusal {
|
||||
refusal: refusal.clone(),
|
||||
},
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
LoggedItem::ToolCall {
|
||||
call_id,
|
||||
name,
|
||||
arguments,
|
||||
} => Some(SessionSnapshotEntryData::ToolCall {
|
||||
call_id: call_id.clone(),
|
||||
name: name.clone(),
|
||||
arguments: arguments.clone(),
|
||||
}),
|
||||
LoggedItem::ToolResult {
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
is_error,
|
||||
attachments,
|
||||
..
|
||||
} => Some(SessionSnapshotEntryData::ToolResult {
|
||||
call_id: call_id.clone(),
|
||||
summary: summary.clone(),
|
||||
content: content.clone(),
|
||||
is_error: *is_error,
|
||||
attachments: attachments
|
||||
.iter()
|
||||
.map(|attachment| match attachment {
|
||||
crate::logged_item::LoggedAttachment::Image { mime_type, data } => {
|
||||
SessionToolAttachment {
|
||||
media_type: mime_type.clone(),
|
||||
data_base64: BASE64.encode(data),
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
}),
|
||||
// Hidden model reasoning is never observable.
|
||||
LoggedItem::Reasoning { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn system_entry(
|
||||
item: &SystemItem,
|
||||
entry_id: String,
|
||||
timestamp: u64,
|
||||
provenance: SessionEntryProvenance,
|
||||
derived_from: Vec<String>,
|
||||
) -> SessionSnapshotEntry {
|
||||
let mut data = serde_json::to_value(item).ok();
|
||||
if let Some(serde_json::Value::Object(object)) = data.as_mut() {
|
||||
object.remove("prompt_provenance");
|
||||
}
|
||||
let item_kind = data
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("kind"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("system_item")
|
||||
.to_owned();
|
||||
SessionSnapshotEntry {
|
||||
entry_id,
|
||||
timestamp,
|
||||
provenance,
|
||||
derived_from,
|
||||
data: SessionSnapshotEntryData::SystemItem {
|
||||
item_kind,
|
||||
content: item.history_text(),
|
||||
data,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use agen::llm_client::RequestConfig;
|
||||
|
||||
use super::*;
|
||||
use crate::{
|
||||
LoggedHistoryDerivation, LoggedSessionHistoryEntryId, LoggedSessionHistoryMetadata,
|
||||
LoggedWorkerSubject,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn current_projection_is_stable_and_hides_reasoning_and_system_prompts() {
|
||||
let session_id = crate::new_session_id();
|
||||
let log = vec![LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![
|
||||
LoggedItem::Message {
|
||||
role: LoggedRole::System,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "secret prompt".into(),
|
||||
}],
|
||||
},
|
||||
LoggedItem::Reasoning {
|
||||
text: "secret reasoning".into(),
|
||||
summary: Vec::new(),
|
||||
encrypted_content: None,
|
||||
signature: None,
|
||||
},
|
||||
LoggedItem::Message {
|
||||
role: LoggedRole::Assistant,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "visible".into(),
|
||||
}],
|
||||
},
|
||||
]
|
||||
.into_iter()
|
||||
.map(|item| LoggedHistoryEntry {
|
||||
item,
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
})
|
||||
.collect(),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}];
|
||||
|
||||
let first = project_session_snapshot(session_id, &log);
|
||||
let second = project_session_snapshot(session_id, &log);
|
||||
assert_eq!(first, second);
|
||||
assert_eq!(first.entries.len(), 1);
|
||||
assert_eq!(first.entries[0].timestamp, 1);
|
||||
assert_eq!(
|
||||
first.entries[0].provenance,
|
||||
SessionEntryProvenance::LegacyUnknown
|
||||
);
|
||||
let json = serde_json::to_string(&first).unwrap();
|
||||
assert!(!json.contains("secret prompt"));
|
||||
assert!(!json.contains("secret reasoning"));
|
||||
assert!(json.contains("visible"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn annotated_user_input_attaches_segments_to_first_user_role_entry_for_any_origin() {
|
||||
let session_id = crate::new_session_id();
|
||||
let segments = vec![Segment::Text {
|
||||
content: "normal submit".into(),
|
||||
}];
|
||||
|
||||
for origin in [
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
LoggedSessionHistoryOrigin::FlowInstruction {
|
||||
selector: "builtin:coder-review".into(),
|
||||
definition_id: "flow-definition".into(),
|
||||
definition_revision: 7,
|
||||
instance_id: "flow-instance".into(),
|
||||
state_id: "implement".into(),
|
||||
},
|
||||
] {
|
||||
let user_entry_id = LoggedSessionHistoryEntryId::new();
|
||||
let source_entry_id = LoggedSessionHistoryEntryId::new();
|
||||
let log = vec![
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: Vec::new(),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2,
|
||||
segments: segments.clone(),
|
||||
history: vec![
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::System,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "flow instruction".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::FlowInstruction {
|
||||
selector: "builtin:coder-review".into(),
|
||||
definition_id: "flow-definition".into(),
|
||||
definition_revision: 7,
|
||||
instance_id: "flow-instance".into(),
|
||||
state_id: "implement".into(),
|
||||
},
|
||||
derivation: None,
|
||||
},
|
||||
},
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::User,
|
||||
content: vec![LoggedContentPart::Text {
|
||||
text: "normal submit".into(),
|
||||
}],
|
||||
},
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: user_entry_id.clone(),
|
||||
origin: origin.clone(),
|
||||
derivation: Some(LoggedHistoryDerivation {
|
||||
sources: vec![source_entry_id.clone()],
|
||||
}),
|
||||
},
|
||||
},
|
||||
],
|
||||
extensions: Vec::new(),
|
||||
},
|
||||
];
|
||||
|
||||
let snapshot = project_current_session_snapshot(&log);
|
||||
assert_eq!(snapshot.entries.len(), 1);
|
||||
assert_eq!(snapshot.entries[0].entry_id, user_entry_id.0);
|
||||
assert_eq!(snapshot.entries[0].provenance, provenance(&origin));
|
||||
assert_eq!(snapshot.entries[0].derived_from, vec![source_entry_id.0]);
|
||||
assert_eq!(
|
||||
snapshot.entries[0].data,
|
||||
SessionSnapshotEntryData::UserInput {
|
||||
segments: segments.clone(),
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn annotated_projection_preserves_identity_and_provenance() {
|
||||
let session_id = crate::new_session_id();
|
||||
let metadata = LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::ModelOutput {
|
||||
worker: LoggedWorkerSubject {
|
||||
workspace_id: None,
|
||||
runtime_id: None,
|
||||
worker_id: "worker".into(),
|
||||
},
|
||||
},
|
||||
derivation: None,
|
||||
};
|
||||
let expected_id = metadata.entry_id.0.clone();
|
||||
let log = vec![LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![LoggedHistoryEntry {
|
||||
item: LoggedItem::Message {
|
||||
role: LoggedRole::Assistant,
|
||||
content: vec![LoggedContentPart::Text { text: "ok".into() }],
|
||||
},
|
||||
metadata,
|
||||
}],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}];
|
||||
|
||||
let snapshot = project_session_snapshot(session_id, &log);
|
||||
assert_eq!(snapshot.entries[0].entry_id, expected_id);
|
||||
assert_eq!(
|
||||
snapshot.entries[0].provenance,
|
||||
SessionEntryProvenance::ModelOutput
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -4,11 +4,9 @@
|
||||
//! The caller (typically Worker) holds the Engine directly and calls these
|
||||
//! functions after state-mutating operations.
|
||||
|
||||
use crate::logged_item::{LoggedItem, to_logged};
|
||||
use crate::segment_log::{self, LogEntry, SegmentOrigin};
|
||||
use crate::store::{Store, StoreError};
|
||||
use crate::system_item::SystemItem;
|
||||
use crate::{SegmentId, SessionId};
|
||||
use crate::{LoggedHistoryEntry, LoggedSystemHistoryEntry, SegmentId, SessionId};
|
||||
use agen::EngineResult;
|
||||
use agen::llm_client::RequestConfig;
|
||||
use agen::llm_client::types::Item;
|
||||
@@ -18,7 +16,7 @@ use protocol::Segment;
|
||||
pub struct SegmentStartState<'a> {
|
||||
pub system_prompt: Option<&'a str>,
|
||||
pub config: &'a RequestConfig,
|
||||
pub history: &'a [Item],
|
||||
pub history: Vec<LoggedHistoryEntry>,
|
||||
}
|
||||
|
||||
/// Create a new session + initial segment, writing the initial
|
||||
@@ -44,12 +42,12 @@ pub fn create_segment_with_ids(
|
||||
segment_id: SegmentId,
|
||||
state: SegmentStartState<'_>,
|
||||
) -> Result<(), StoreError> {
|
||||
let entry = LogEntry::SegmentStart {
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
history: state.history.to_vec(),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
};
|
||||
@@ -70,12 +68,12 @@ pub fn create_compacted_segment(
|
||||
source_turn_count: usize,
|
||||
) -> Result<SegmentId, StoreError> {
|
||||
let segment_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id: source_session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
history: state.history.to_vec(),
|
||||
forked_from: None,
|
||||
compacted_from: Some(SegmentOrigin {
|
||||
segment_id: source_segment_id,
|
||||
@@ -154,12 +152,12 @@ pub fn ensure_head_or_fork(
|
||||
}
|
||||
let source_segment_id = *segment_id;
|
||||
let fork_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
history: state.history.to_vec(),
|
||||
forked_from: Some(SegmentOrigin {
|
||||
segment_id: source_segment_id,
|
||||
at_turn_index,
|
||||
@@ -183,8 +181,9 @@ pub fn save_user_input(
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
segments: Vec<Segment>,
|
||||
history: Vec<LoggedHistoryEntry>,
|
||||
) -> Result<(), StoreError> {
|
||||
save_user_input_with_extensions(store, session_id, segment_id, segments, Vec::new())
|
||||
save_user_input_with_extensions(store, session_id, segment_id, segments, history, Vec::new())
|
||||
}
|
||||
|
||||
/// Atomically persist one typed user submission and Runtime-owned session
|
||||
@@ -194,15 +193,17 @@ pub fn save_user_input_with_extensions(
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
segments: Vec<Segment>,
|
||||
history: Vec<LoggedHistoryEntry>,
|
||||
extensions: Vec<segment_log::SessionExtension>,
|
||||
) -> Result<(), StoreError> {
|
||||
append_entry(
|
||||
store,
|
||||
session_id,
|
||||
segment_id,
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: segment_log::now_millis(),
|
||||
segments,
|
||||
history,
|
||||
extensions,
|
||||
},
|
||||
)
|
||||
@@ -220,64 +221,57 @@ pub fn save_delta(
|
||||
store: &impl Store,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
new_items: &[Item],
|
||||
new_items: &[LoggedHistoryEntry],
|
||||
) -> Result<(), StoreError> {
|
||||
if new_items.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let ts = segment_log::now_millis();
|
||||
for item in new_items {
|
||||
for entry in new_items {
|
||||
let item = Item::from(entry.item.clone());
|
||||
if item.is_user_message() {
|
||||
// Already persisted by save_user_input at submit time.
|
||||
continue;
|
||||
}
|
||||
let entry = classify_history_item(item, ts);
|
||||
let entry = classify_logged_history_entry(entry.clone(), ts);
|
||||
append_entry(store, session_id, segment_id, entry)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Map one history item to its singular `LogEntry` form. Used by the
|
||||
/// fallback `save_delta` path and the controller's worker-callback
|
||||
/// classifier so write classification lives in one place.
|
||||
pub fn classify_history_item(item: &Item, ts: u64) -> LogEntry {
|
||||
/// Map one annotated history entry to its singular `LogEntry` form. Used by
|
||||
/// the fallback `save_delta` path and the controller's worker-callback
|
||||
/// classifier so write classification lives in one place without discarding
|
||||
/// identity or provenance.
|
||||
/// Map one already-annotated history entry to its singular canonical record
|
||||
/// without changing its identity or provenance.
|
||||
pub fn classify_logged_history_entry(entry: LoggedHistoryEntry, ts: u64) -> LogEntry {
|
||||
let item = Item::from(entry.item.clone());
|
||||
if item.is_tool_result() {
|
||||
LogEntry::ToolResult {
|
||||
ts,
|
||||
item: LoggedItem::from(item),
|
||||
}
|
||||
} else if item.is_assistant_message() || item.is_tool_call() || item.is_reasoning() {
|
||||
LogEntry::AssistantItem {
|
||||
ts,
|
||||
item: LoggedItem::from(item),
|
||||
}
|
||||
LogEntry::AnnotatedToolResult { ts, entry }
|
||||
} else {
|
||||
// Defensive: anything else (future Item kinds) routes through
|
||||
// AssistantItem rather than getting silently dropped.
|
||||
LogEntry::AssistantItem {
|
||||
ts,
|
||||
item: LoggedItem::from(item),
|
||||
}
|
||||
// Assistant messages, tool calls, reasoning, and future non-user
|
||||
// items all use the assistant-side canonical record.
|
||||
LogEntry::AnnotatedAssistantItem { ts, entry }
|
||||
}
|
||||
}
|
||||
|
||||
/// Append a single typed system item as `LogEntry::SystemItem`. Helper
|
||||
/// for the Worker-side interceptor commit path; mirrors the per-item
|
||||
/// commit shape used for assistant / tool result entries.
|
||||
/// Append one typed system item and its history metadata as a canonical
|
||||
/// `LogEntry::AnnotatedSystemItem`.
|
||||
pub fn append_system_item(
|
||||
store: &impl Store,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
item: SystemItem,
|
||||
entry: LoggedSystemHistoryEntry,
|
||||
) -> Result<(), StoreError> {
|
||||
append_entry(
|
||||
store,
|
||||
session_id,
|
||||
segment_id,
|
||||
LogEntry::SystemItem {
|
||||
LogEntry::AnnotatedSystemItem {
|
||||
ts: segment_log::now_millis(),
|
||||
item,
|
||||
entry,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -307,6 +301,7 @@ pub fn save_run_completed(
|
||||
segment_id: SegmentId,
|
||||
result: EngineResult,
|
||||
interrupted: bool,
|
||||
active_run_turn_count: Option<usize>,
|
||||
) -> Result<(), StoreError> {
|
||||
append_entry(
|
||||
store,
|
||||
@@ -316,6 +311,7 @@ pub fn save_run_completed(
|
||||
ts: segment_log::now_millis(),
|
||||
interrupted,
|
||||
result,
|
||||
active_run_turn_count,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -428,12 +424,12 @@ pub fn fork(
|
||||
) -> Result<(SessionId, SegmentId), StoreError> {
|
||||
let session_id = crate::new_session_id();
|
||||
let fork_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id,
|
||||
system_prompt: state.system_prompt.map(String::from),
|
||||
config: state.config.clone(),
|
||||
history: to_logged(state.history),
|
||||
history: state.history.to_vec(),
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
};
|
||||
@@ -468,7 +464,7 @@ pub fn fork_at(
|
||||
// segment), before any turn completes.
|
||||
entries
|
||||
.iter()
|
||||
.position(|e| !matches!(e, LogEntry::SegmentStart { .. }))
|
||||
.position(|e| !matches!(e, LogEntry::AnnotatedSegmentStart { .. }))
|
||||
.unwrap_or(entries.len())
|
||||
} else {
|
||||
entries
|
||||
@@ -480,12 +476,12 @@ pub fn fork_at(
|
||||
let state = segment_log::collect_state(&entries[..cut]);
|
||||
|
||||
let fork_id = crate::new_segment_id();
|
||||
let entry = LogEntry::SegmentStart {
|
||||
let entry = LogEntry::AnnotatedSegmentStart {
|
||||
ts: segment_log::now_millis(),
|
||||
session_id: source_session_id,
|
||||
system_prompt: state.system_prompt,
|
||||
config: state.config,
|
||||
history: to_logged(&state.history),
|
||||
history: state.annotated_history,
|
||||
forked_from: Some(SegmentOrigin {
|
||||
segment_id: source_id,
|
||||
at_turn_index,
|
||||
|
||||
@@ -14,8 +14,8 @@ use agen::{EngineResult, UsageRecord};
|
||||
use protocol::{InvokeKind, Segment};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::history::{LoggedHistoryEntry, LoggedSystemHistoryEntry};
|
||||
use crate::logged_item::LoggedItem;
|
||||
use crate::system_item::SystemItem;
|
||||
|
||||
/// A single segment log entry, serialized as one JSONL line.
|
||||
///
|
||||
@@ -49,23 +49,16 @@ impl SessionExtension {
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum LogEntry {
|
||||
/// Segment start. Always the first entry in a segment log.
|
||||
/// For forked segments, `history` contains the seed state from the parent.
|
||||
SegmentStart {
|
||||
/// Canonical segment seed. Retained entries keep their stable logical
|
||||
/// identity and origin across fork/compaction/restore.
|
||||
AnnotatedSegmentStart {
|
||||
ts: u64,
|
||||
/// Session this segment belongs to. Compaction / fork inherits
|
||||
/// the source segment's session_id; only fresh "new conversation"
|
||||
/// segments mint a new session_id.
|
||||
session_id: crate::SessionId,
|
||||
system_prompt: Option<String>,
|
||||
config: RequestConfig,
|
||||
history: Vec<LoggedItem>,
|
||||
/// Origin: forked from a sibling segment at a specific turn boundary.
|
||||
/// The referenced segment is guaranteed to share `session_id`.
|
||||
history: Vec<LoggedHistoryEntry>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
forked_from: Option<SegmentOrigin>,
|
||||
/// Origin: compacted from a sibling segment at a specific turn boundary.
|
||||
/// The referenced segment is guaranteed to share `session_id`.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
compacted_from: Option<SegmentOrigin>,
|
||||
},
|
||||
@@ -90,46 +83,43 @@ pub enum LogEntry {
|
||||
/// restore conservatively instead of re-running a dangling tool call.
|
||||
Invoke { ts: u64, trigger: InvokeKind },
|
||||
|
||||
/// User input accepted at submit time. Carries the original typed
|
||||
/// `Vec<Segment>` so clients can re-render typed atoms (paste chips,
|
||||
/// file refs) on segment restore.
|
||||
/// Replay flattens these into a `Item::user_message` for the worker
|
||||
/// history; the worker layer never sees segments directly.
|
||||
UserInput {
|
||||
/// Canonical user submission with its exact model-visible entries. Typed
|
||||
/// Flow instructions and caller-attributed input remain separate entries.
|
||||
AnnotatedUserInput {
|
||||
ts: u64,
|
||||
segments: Vec<Segment>,
|
||||
/// Typed durable state committed atomically with this input record.
|
||||
/// Runtime-owned Flow invocation uses this to avoid a Backend-instance
|
||||
/// commit that can get ahead of Worker history.
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
extensions: Vec<SessionExtension>,
|
||||
history: Vec<LoggedHistoryEntry>,
|
||||
},
|
||||
|
||||
/// One assistant-side item appended to history — assistant message,
|
||||
/// reasoning, or tool call. Singular: one entry per history item so
|
||||
/// the wire-side `Event::*` lane and on-disk LogEntry stay 1:1.
|
||||
AssistantItem { ts: u64, item: LoggedItem },
|
||||
/// Canonical model output and metadata committed as one journal record.
|
||||
AnnotatedAssistantItem { ts: u64, entry: LoggedHistoryEntry },
|
||||
|
||||
/// One tool-execution result appended to history.
|
||||
ToolResult { ts: u64, item: LoggedItem },
|
||||
/// Canonical tool output and metadata committed as one journal record.
|
||||
AnnotatedToolResult { ts: u64, entry: LoggedHistoryEntry },
|
||||
|
||||
/// One typed agent-injected system item: notification, child-Worker
|
||||
/// lifecycle event, `@<path>` / `/<slug>` resolution payload. Each
|
||||
/// `SystemItem` carries kind metadata that the LLM
|
||||
/// itself never sees (the LLM gets `Item::system_message` with the
|
||||
/// item's denormalised `body`), but live clients and replay paths
|
||||
/// dispatch on `kind` for typed rendering.
|
||||
SystemItem { ts: u64, item: SystemItem },
|
||||
/// Canonical typed system event and model-visible metadata committed
|
||||
/// together.
|
||||
AnnotatedSystemItem {
|
||||
ts: u64,
|
||||
entry: LoggedSystemHistoryEntry,
|
||||
},
|
||||
|
||||
/// Turn boundary. Records the turn count after increment.
|
||||
TurnEnd { ts: u64, turn_count: usize },
|
||||
|
||||
/// `run()` / `resume()` が `EngineResult` で正常終了した。
|
||||
/// Audit-only metadata: replay は `interrupted` のみ反映する。
|
||||
/// Replay restores both interruption state and any resumable logical-run
|
||||
/// turn budget.
|
||||
RunCompleted {
|
||||
ts: u64,
|
||||
interrupted: bool,
|
||||
result: EngineResult,
|
||||
/// AgentTurns consumed by a paused/yielded logical run. Terminal
|
||||
/// outcomes persist `None`.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
active_run_turn_count: Option<usize>,
|
||||
},
|
||||
|
||||
/// `run()` / `resume()` が `EngineError` で終了した。
|
||||
@@ -141,6 +131,15 @@ pub enum LogEntry {
|
||||
message: String,
|
||||
},
|
||||
|
||||
/// Restores an active logical-run budget at a segment boundary, notably
|
||||
/// after compaction replaced the segment that held the original Invoke and
|
||||
/// RunCompleted entries.
|
||||
ActiveRunCheckpoint {
|
||||
ts: u64,
|
||||
active_turn_count: usize,
|
||||
total_turn_count: usize,
|
||||
},
|
||||
|
||||
/// A paused interrupted turn was explicitly abandoned without calling
|
||||
/// `run()` or `resume()` again. Replay clears the interrupted marker so
|
||||
/// the restored Worker is idle and future user input starts a normal new turn.
|
||||
@@ -208,7 +207,13 @@ pub struct RestoredState {
|
||||
pub system_prompt: Option<String>,
|
||||
pub config: RequestConfig,
|
||||
pub history: Vec<Item>,
|
||||
/// Canonical persisted history with stable identity and provenance. This is
|
||||
/// the authority for rewrites, forks, and annotated restore; `history` is
|
||||
/// retained as the model-facing item projection.
|
||||
pub annotated_history: Vec<LoggedHistoryEntry>,
|
||||
pub turn_count: usize,
|
||||
/// AgentTurns consumed by the active paused/yielded logical run.
|
||||
pub active_run_turn_count: Option<usize>,
|
||||
pub last_run_interrupted: bool,
|
||||
/// Number of entries replayed. `0` means the segment log was empty.
|
||||
/// Writers track their own append count via the same counter so
|
||||
@@ -222,7 +227,7 @@ pub struct RestoredState {
|
||||
/// session-store は domain を不透明扱いし、各ドメインが自前で fold する。
|
||||
pub extensions: Vec<(String, serde_json::Value)>,
|
||||
/// User submissions in original typed form, in submit order.
|
||||
/// One entry per `LogEntry::UserInput`; the K-th entry corresponds to
|
||||
/// One entry per `LogEntry::AnnotatedUserInput`; the K-th entry corresponds to
|
||||
/// the K-th `Item::user_message` derived during replay (modulo
|
||||
/// pre-compaction history seeded via `SegmentStart.history`, whose
|
||||
/// original segments are not preserved). Used by clients to re-render
|
||||
@@ -237,7 +242,9 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: Vec::new(),
|
||||
annotated_history: Vec::new(),
|
||||
turn_count: 0,
|
||||
active_run_turn_count: None,
|
||||
last_run_interrupted: false,
|
||||
entries_count: 0,
|
||||
usage_history: Vec::new(),
|
||||
@@ -249,7 +256,7 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
state.entries_count += 1;
|
||||
|
||||
match entry {
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
session_id,
|
||||
system_prompt,
|
||||
config,
|
||||
@@ -259,20 +266,29 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
state.session_id = Some(*session_id);
|
||||
state.system_prompt = system_prompt.clone();
|
||||
state.config = config.clone();
|
||||
state.history = history.iter().cloned().map(Item::from).collect();
|
||||
state.annotated_history = history.clone();
|
||||
state.history = history
|
||||
.iter()
|
||||
.cloned()
|
||||
.map(|entry| Item::from(entry.item))
|
||||
.collect();
|
||||
}
|
||||
LogEntry::Invoke { .. } => {
|
||||
// A terminal run record below clears or refines this. If the
|
||||
// log ends first, restore must treat the turn as interrupted.
|
||||
state.last_run_interrupted = true;
|
||||
state.active_run_turn_count = Some(0);
|
||||
}
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
segments,
|
||||
extensions,
|
||||
history,
|
||||
..
|
||||
} => {
|
||||
let text = Segment::flatten_to_text(segments);
|
||||
state.history.push(Item::user_message(text));
|
||||
state.annotated_history.extend(history.iter().cloned());
|
||||
state
|
||||
.history
|
||||
.extend(history.iter().cloned().map(|entry| Item::from(entry.item)));
|
||||
state.user_segments.push(segments.clone());
|
||||
state.extensions.extend(
|
||||
extensions
|
||||
@@ -280,26 +296,57 @@ pub fn collect_state(entries: &[LogEntry]) -> RestoredState {
|
||||
.map(|extension| (extension.domain.clone(), extension.payload.clone())),
|
||||
);
|
||||
}
|
||||
LogEntry::AssistantItem { item, .. } => {
|
||||
state.history.push(Item::from(item.clone()));
|
||||
LogEntry::AnnotatedAssistantItem { entry, .. }
|
||||
| LogEntry::AnnotatedToolResult { entry, .. } => {
|
||||
state.annotated_history.push(entry.clone());
|
||||
state.history.push(Item::from(entry.item.clone()));
|
||||
}
|
||||
LogEntry::ToolResult { item, .. } => {
|
||||
state.history.push(Item::from(item.clone()));
|
||||
}
|
||||
LogEntry::SystemItem { item, .. } => {
|
||||
state.history.push(item.to_history_item());
|
||||
LogEntry::AnnotatedSystemItem { entry, .. } => {
|
||||
state.annotated_history.push(LoggedHistoryEntry {
|
||||
item: LoggedItem::from(entry.item.to_history_item()),
|
||||
metadata: entry.metadata.clone(),
|
||||
});
|
||||
state.history.push(entry.item.to_history_item());
|
||||
}
|
||||
LogEntry::TurnEnd { turn_count, .. } => {
|
||||
if let Some(active_turn_count) = &mut state.active_run_turn_count {
|
||||
*active_turn_count += turn_count.saturating_sub(state.turn_count);
|
||||
}
|
||||
state.turn_count = *turn_count;
|
||||
}
|
||||
LogEntry::RunCompleted { interrupted, .. } => {
|
||||
LogEntry::RunCompleted {
|
||||
interrupted,
|
||||
result,
|
||||
active_run_turn_count,
|
||||
..
|
||||
} => {
|
||||
state.last_run_interrupted = *interrupted;
|
||||
if *interrupted && matches!(result, EngineResult::Paused | EngineResult::Yielded) {
|
||||
// Legacy entries omit the explicit field; retain the
|
||||
// Invoke/TurnEnd-derived count in that case.
|
||||
if let Some(turn_count) = active_run_turn_count {
|
||||
state.active_run_turn_count = Some(*turn_count);
|
||||
}
|
||||
} else {
|
||||
state.active_run_turn_count = None;
|
||||
}
|
||||
}
|
||||
LogEntry::RunErrored { interrupted, .. } => {
|
||||
state.last_run_interrupted = *interrupted;
|
||||
state.active_run_turn_count = None;
|
||||
}
|
||||
LogEntry::ActiveRunCheckpoint {
|
||||
active_turn_count,
|
||||
total_turn_count,
|
||||
..
|
||||
} => {
|
||||
state.active_run_turn_count = Some(*active_turn_count);
|
||||
state.turn_count = *total_turn_count;
|
||||
state.last_run_interrupted = true;
|
||||
}
|
||||
LogEntry::PausedTurnAbandoned { .. } => {
|
||||
state.last_run_interrupted = false;
|
||||
state.active_run_turn_count = None;
|
||||
}
|
||||
LogEntry::ConfigChanged { config, .. } => {
|
||||
state.config = config.clone();
|
||||
@@ -342,6 +389,20 @@ pub fn now_millis() -> u64 {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::{
|
||||
LoggedSessionHistoryEntryId, LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin,
|
||||
};
|
||||
|
||||
fn annotated(item: Item) -> LoggedHistoryEntry {
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::from(item),
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replay_empty() {
|
||||
@@ -353,12 +414,12 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn replay_segment_start_sets_initial_state() {
|
||||
let state = collect_state(&[LogEntry::SegmentStart {
|
||||
let state = collect_state(&[LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: Some("You are helpful.".into()),
|
||||
config: RequestConfig::default().with_max_tokens(1024),
|
||||
history: vec![Item::user_message("seed").into()],
|
||||
history: vec![annotated(Item::user_message("seed"))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}]);
|
||||
@@ -371,7 +432,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_full_turn() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -380,14 +441,15 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("Hello")],
|
||||
history: vec![annotated(Item::user_message("Hello"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::assistant_message("Hi!").into(),
|
||||
entry: annotated(Item::assistant_message("Hi!")),
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 3100,
|
||||
@@ -397,6 +459,7 @@ mod tests {
|
||||
ts: 3200,
|
||||
interrupted: false,
|
||||
result: EngineResult::Finished,
|
||||
active_run_turn_count: None,
|
||||
},
|
||||
]);
|
||||
assert_eq!(state.history.len(), 2);
|
||||
@@ -407,7 +470,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_incomplete_invoke_is_interrupted() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -420,14 +483,15 @@ mod tests {
|
||||
ts: 2000,
|
||||
trigger: InvokeKind::UserSend,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2001,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("run a tool")],
|
||||
history: vec![annotated(Item::user_message("run a tool"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::tool_call("call_1", "side_effect", "{}").into(),
|
||||
entry: annotated(Item::tool_call("call_1", "side_effect", "{}")),
|
||||
},
|
||||
]);
|
||||
|
||||
@@ -437,7 +501,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_with_tool_calls() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -446,22 +510,27 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("Check weather")],
|
||||
history: vec![annotated(Item::user_message("Check weather"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::tool_call("call_1", "get_weather", r#"{"city":"Tokyo"}"#).into(),
|
||||
entry: annotated(Item::tool_call(
|
||||
"call_1",
|
||||
"get_weather",
|
||||
r#"{"city":"Tokyo"}"#,
|
||||
)),
|
||||
},
|
||||
LogEntry::ToolResult {
|
||||
LogEntry::AnnotatedToolResult {
|
||||
ts: 3500,
|
||||
item: Item::tool_result("call_1", "Sunny, 25C").into(),
|
||||
entry: annotated(Item::tool_result("call_1", "Sunny, 25C")),
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 4000,
|
||||
item: Item::assistant_message("It's sunny in Tokyo!").into(),
|
||||
entry: annotated(Item::assistant_message("It's sunny in Tokyo!")),
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 4100,
|
||||
@@ -475,9 +544,9 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn replay_restores_durable_tool_image_detail() {
|
||||
let entry = LogEntry::ToolResult {
|
||||
let entry = LogEntry::AnnotatedToolResult {
|
||||
ts: 3500,
|
||||
item: Item::tool_result_item_with_attachments(
|
||||
entry: annotated(Item::tool_result_item_with_attachments(
|
||||
"call_image",
|
||||
"attached",
|
||||
None,
|
||||
@@ -485,8 +554,7 @@ mod tests {
|
||||
vec![agen::tool::Attachment::Image(
|
||||
agen::tool::ImageAttachment::new("image/png", b"durable-image".to_vec()),
|
||||
)],
|
||||
)
|
||||
.into(),
|
||||
)),
|
||||
};
|
||||
let persisted = serde_json::to_string(&entry).unwrap();
|
||||
let restored_entry: LogEntry = serde_json::from_str(&persisted).unwrap();
|
||||
@@ -506,7 +574,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_config_changed() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -526,7 +594,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_llm_usage_appends_to_usage_history() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -535,10 +603,11 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("hi")],
|
||||
history: vec![annotated(Item::user_message("hi"))],
|
||||
},
|
||||
LogEntry::LlmUsage {
|
||||
ts: 2100,
|
||||
@@ -548,9 +617,9 @@ mod tests {
|
||||
cache_write_tokens: 0,
|
||||
output_tokens: 10,
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 2200,
|
||||
item: Item::assistant_message("yo").into(),
|
||||
entry: annotated(Item::assistant_message("yo")),
|
||||
},
|
||||
LogEntry::LlmUsage {
|
||||
ts: 3100,
|
||||
@@ -574,7 +643,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_without_llm_usage_keeps_usage_history_empty() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -583,10 +652,11 @@ mod tests {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("hi")],
|
||||
history: vec![annotated(Item::user_message("hi"))],
|
||||
},
|
||||
]);
|
||||
assert!(state.usage_history.is_empty());
|
||||
@@ -647,7 +717,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_invoke_marker_only_mutates_interrupted_state() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -660,10 +730,11 @@ mod tests {
|
||||
ts: 100,
|
||||
trigger: InvokeKind::UserSend,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 101,
|
||||
extensions: vec![],
|
||||
segments: vec![Segment::text("hi")],
|
||||
history: vec![annotated(Item::user_message("hi"))],
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 200,
|
||||
@@ -682,7 +753,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_paused_turn_abandoned_clears_interrupted_marker() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -695,10 +766,93 @@ mod tests {
|
||||
ts: 100,
|
||||
interrupted: true,
|
||||
result: EngineResult::Paused,
|
||||
active_run_turn_count: Some(1),
|
||||
},
|
||||
LogEntry::PausedTurnAbandoned { ts: 200 },
|
||||
]);
|
||||
assert!(!state.last_run_interrupted);
|
||||
assert_eq!(state.active_run_turn_count, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replay_restores_active_run_budget_across_compaction_checkpoint() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::ActiveRunCheckpoint {
|
||||
ts: 100,
|
||||
active_turn_count: 3,
|
||||
total_turn_count: 9,
|
||||
},
|
||||
]);
|
||||
|
||||
assert_eq!(state.turn_count, 9);
|
||||
assert_eq!(state.active_run_turn_count, Some(3));
|
||||
assert!(state.last_run_interrupted);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_interrupted_run_derives_budget_from_invoke_and_turn_end() {
|
||||
let entry: LogEntry = serde_json::from_value(serde_json::json!({
|
||||
"kind": "run_completed",
|
||||
"ts": 300,
|
||||
"interrupted": true,
|
||||
"result": "paused"
|
||||
}))
|
||||
.expect("legacy run-completed entry");
|
||||
let state = collect_state(&[
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 0,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::Invoke {
|
||||
ts: 100,
|
||||
trigger: InvokeKind::UserSend,
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 200,
|
||||
turn_count: 2,
|
||||
},
|
||||
entry,
|
||||
]);
|
||||
|
||||
assert_eq!(state.active_run_turn_count, Some(2));
|
||||
assert!(state.last_run_interrupted);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_resumable_interruption_clears_the_active_run_budget() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::Invoke {
|
||||
ts: 100,
|
||||
trigger: InvokeKind::UserSend,
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 200,
|
||||
turn_count: 2,
|
||||
},
|
||||
LogEntry::RunCompleted {
|
||||
ts: 300,
|
||||
interrupted: true,
|
||||
result: EngineResult::LimitReached,
|
||||
active_run_turn_count: None,
|
||||
},
|
||||
]);
|
||||
|
||||
assert!(state.last_run_interrupted);
|
||||
assert_eq!(state.active_run_turn_count, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -717,7 +871,7 @@ mod tests {
|
||||
#[test]
|
||||
fn replay_extension_collects_domain_payload_pairs() {
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -776,9 +930,12 @@ mod tests {
|
||||
#[test]
|
||||
fn user_input_extensions_restore_with_the_same_committed_input() {
|
||||
let segments = vec![Segment::text("Flow instructions"), Segment::text("Ticket")];
|
||||
let entry = LogEntry::UserInput {
|
||||
let entry = LogEntry::AnnotatedUserInput {
|
||||
ts: 9999,
|
||||
segments: segments.clone(),
|
||||
history: vec![annotated(Item::user_message(Segment::flatten_to_text(
|
||||
&segments,
|
||||
)))],
|
||||
extensions: vec![SessionExtension::new(
|
||||
"flow.runtime.v1",
|
||||
serde_json::json!({ "state": "implement", "revision": 0 }),
|
||||
@@ -793,7 +950,7 @@ mod tests {
|
||||
assert_eq!(state.extensions[0].1["state"], "implement");
|
||||
}
|
||||
|
||||
/// Mixed segments survive a JSON round-trip through `LogEntry::UserInput`,
|
||||
/// Mixed segments survive a JSON round-trip through `LogEntry::AnnotatedUserInput`,
|
||||
/// and `collect_state` derives `Item::user_message` from the flattened
|
||||
/// text while preserving the original segments separately. This covers
|
||||
/// the segments → flatten → Item replay path from the ticket.
|
||||
@@ -813,16 +970,19 @@ mod tests {
|
||||
path: "src/main.rs".into(),
|
||||
},
|
||||
];
|
||||
let entry = LogEntry::UserInput {
|
||||
let entry = LogEntry::AnnotatedUserInput {
|
||||
ts: 4242,
|
||||
extensions: vec![],
|
||||
segments: segments.clone(),
|
||||
history: vec![annotated(Item::user_message(Segment::flatten_to_text(
|
||||
&segments,
|
||||
)))],
|
||||
};
|
||||
// JSON round-trip preserves the variant byte-for-byte.
|
||||
let json = serde_json::to_string(&entry).unwrap();
|
||||
let parsed: LogEntry = serde_json::from_str(&json).unwrap();
|
||||
let state = collect_state(&[
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
//! `kind` instead of parsing text prefixes like `[Notification] …` or
|
||||
//! `[File: …]`.
|
||||
//!
|
||||
//! Persisted as the payload of [`crate::LogEntry::SystemItem`] (one
|
||||
//! Persisted as the payload of [`crate::LogEntry::AnnotatedSystemItem`] (one
|
||||
//! entry per item), and broadcast live as the payload of
|
||||
//! `Event::SystemItem` on the wire.
|
||||
//!
|
||||
|
||||
@@ -20,7 +20,9 @@ use std::path::{Path, PathBuf};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::SystemTime;
|
||||
|
||||
const SESSION_SCHEMA_VERSION: u32 = 1;
|
||||
const SESSION_SCHEMA_VERSION: u32 = 3;
|
||||
const PREVIOUS_SESSION_SCHEMA_VERSION: u32 = 2;
|
||||
const LEGACY_SESSION_SCHEMA_VERSION: u32 = 1;
|
||||
const SESSION_FILE: &str = "session.json";
|
||||
const SEGMENTS_DIR: &str = "segments";
|
||||
|
||||
@@ -44,15 +46,28 @@ impl WorkerSessionStore {
|
||||
fs::create_dir_all(root.join(SEGMENTS_DIR))?;
|
||||
let session_id = match fs::read(root.join(SESSION_FILE)) {
|
||||
Ok(bytes) => {
|
||||
let manifest: SessionManifest = serde_json::from_slice(&bytes)?;
|
||||
if manifest.schema_version != SESSION_SCHEMA_VERSION {
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"unsupported Worker Session schema version {}, expected {}",
|
||||
manifest.schema_version, SESSION_SCHEMA_VERSION
|
||||
),
|
||||
});
|
||||
let mut manifest: SessionManifest = serde_json::from_slice(&bytes)?;
|
||||
match manifest.schema_version {
|
||||
SESSION_SCHEMA_VERSION => {
|
||||
validate_canonical_segment_logs(&root)?;
|
||||
}
|
||||
PREVIOUS_SESSION_SCHEMA_VERSION | LEGACY_SESSION_SCHEMA_VERSION => {
|
||||
migrate_segment_logs_to_v3(
|
||||
&root,
|
||||
manifest.session_id,
|
||||
manifest.schema_version,
|
||||
)?;
|
||||
manifest.schema_version = SESSION_SCHEMA_VERSION;
|
||||
atomic_write_json(&root.join(SESSION_FILE), &manifest)?;
|
||||
}
|
||||
version => {
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"unsupported Worker Session schema version {version}, expected {SESSION_SCHEMA_VERSION}"
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
Some(manifest.session_id)
|
||||
}
|
||||
@@ -136,6 +151,41 @@ impl WorkerSessionStore {
|
||||
.join(format!("{segment_id}.trace.jsonl"))
|
||||
}
|
||||
|
||||
fn append_log_entry(&self, path: &Path, entry: &LogEntry) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
.lock()
|
||||
.map_err(|_| std::io::Error::other("Worker Session append lock was poisoned"))?;
|
||||
let mut file = OpenOptions::new()
|
||||
.create(true)
|
||||
.read(true)
|
||||
.write(true)
|
||||
.append(true)
|
||||
.open(path)?;
|
||||
let committed_len = truncate_uncommitted_tail(&mut file)?;
|
||||
file.seek(SeekFrom::Start(0))?;
|
||||
let mut existing = Vec::new();
|
||||
file.read_to_end(&mut existing)?;
|
||||
parse_jsonl::<LogEntry>(&existing)?;
|
||||
let line = serde_json::to_string(entry)?;
|
||||
let mut record = Vec::with_capacity(line.len() + 1);
|
||||
record.extend_from_slice(line.as_bytes());
|
||||
record.push(b'\n');
|
||||
if let Err(write_error) = file.write_all(&record) {
|
||||
return match file.set_len(committed_len) {
|
||||
Ok(()) => Err(write_error.into()),
|
||||
Err(rollback_error) => Err(std::io::Error::new(
|
||||
rollback_error.kind(),
|
||||
format!(
|
||||
"session append failed ({write_error}) and rollback failed: {rollback_error}"
|
||||
),
|
||||
)
|
||||
.into()),
|
||||
};
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn append_line(&self, path: &Path, line: &str) -> Result<(), StoreError> {
|
||||
let _guard = self
|
||||
.append_lock
|
||||
@@ -175,7 +225,7 @@ impl Store for WorkerSessionStore {
|
||||
entry: &LogEntry,
|
||||
) -> Result<(), StoreError> {
|
||||
self.ensure_session(session_id, true)?;
|
||||
self.append_line(&self.log_path(segment_id), &serde_json::to_string(entry)?)
|
||||
self.append_log_entry(&self.log_path(segment_id), entry)
|
||||
}
|
||||
|
||||
fn read_all(
|
||||
@@ -278,6 +328,138 @@ impl Store for WorkerSessionStore {
|
||||
}
|
||||
}
|
||||
|
||||
fn segment_log_paths(root: &Path) -> Result<Vec<(SegmentId, PathBuf)>, StoreError> {
|
||||
let segments = root.join(SEGMENTS_DIR);
|
||||
if !segments.exists() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut paths = Vec::new();
|
||||
for entry in fs::read_dir(&segments)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
let metadata = fs::symlink_metadata(&path)?;
|
||||
let Some(name) = path.file_name().and_then(|name| name.to_str()) else {
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!("non-UTF-8 Worker Session segment path: {}", path.display()),
|
||||
});
|
||||
};
|
||||
if name.ends_with(".trace.jsonl") || name.starts_with('.') {
|
||||
continue;
|
||||
}
|
||||
if !name.ends_with(".jsonl") {
|
||||
continue;
|
||||
}
|
||||
if !metadata.file_type().is_file() {
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"Worker Session segment is not a regular file: {}",
|
||||
path.display()
|
||||
),
|
||||
});
|
||||
}
|
||||
let segment_id =
|
||||
name.trim_end_matches(".jsonl")
|
||||
.parse()
|
||||
.map_err(|_| StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!("invalid Worker Session segment name: {name}"),
|
||||
})?;
|
||||
paths.push((segment_id, path));
|
||||
}
|
||||
paths.sort_by_key(|(segment_id, _)| *segment_id);
|
||||
Ok(paths)
|
||||
}
|
||||
|
||||
fn migrate_segment_logs_to_v3(
|
||||
root: &Path,
|
||||
session_id: SessionId,
|
||||
source_schema_version: u32,
|
||||
) -> Result<(), StoreError> {
|
||||
struct MigrationPlan {
|
||||
path: PathBuf,
|
||||
source: Vec<u8>,
|
||||
output: Vec<u8>,
|
||||
}
|
||||
|
||||
// Phase 1 is strictly read-only. Every segment must parse and canonicalize
|
||||
// successfully before the first authoritative byte is replaced.
|
||||
let mut plans = Vec::new();
|
||||
for (segment_id, path) in segment_log_paths(root)? {
|
||||
let source = fs::read(&path)?;
|
||||
let canonical = parse_legacy_jsonl(source_schema_version, session_id, segment_id, &source)
|
||||
.map_err(|error| StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"cannot migrate Worker Session log {}: {error}",
|
||||
path.display()
|
||||
),
|
||||
})?;
|
||||
let mut output = Vec::new();
|
||||
for entry in canonical {
|
||||
serde_json::to_writer(&mut output, &entry)?;
|
||||
output.push(b'\n');
|
||||
}
|
||||
plans.push(MigrationPlan {
|
||||
path,
|
||||
source,
|
||||
output,
|
||||
});
|
||||
}
|
||||
|
||||
// Fence the complete preflight snapshot before starting phase 2. Session
|
||||
// open is the exclusive restore boundary; this additionally fails closed
|
||||
// if an unexpected writer raced the preflight.
|
||||
for plan in &plans {
|
||||
if fs::read(&plan.path)? != plan.source {
|
||||
return Err(StoreError::Corrupt {
|
||||
line: 0,
|
||||
message: format!(
|
||||
"Worker Session segment changed during migration: {}",
|
||||
plan.path.display()
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
for plan in plans {
|
||||
atomic_write_bytes(&plan.path, &plan.output)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_canonical_segment_logs(root: &Path) -> Result<(), StoreError> {
|
||||
for (_, path) in segment_log_paths(root)? {
|
||||
let _: Vec<LogEntry> = parse_jsonl(&fs::read(&path)?)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn parse_legacy_jsonl(
|
||||
schema_version: u32,
|
||||
session_id: SessionId,
|
||||
segment_id: SegmentId,
|
||||
bytes: &[u8],
|
||||
) -> Result<Vec<LogEntry>, serde_json::Error> {
|
||||
let text = std::str::from_utf8(bytes).map_err(|error| {
|
||||
serde_json::Error::io(std::io::Error::new(std::io::ErrorKind::InvalidData, error))
|
||||
})?;
|
||||
text.lines()
|
||||
.enumerate()
|
||||
.filter(|(_, line)| !line.trim().is_empty())
|
||||
.map(|(line_index, line)| {
|
||||
crate::legacy_session_log::decode_entry(
|
||||
schema_version,
|
||||
line,
|
||||
session_id,
|
||||
segment_id,
|
||||
line_index,
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn atomic_write_json<T: Serialize>(path: &Path, value: &T) -> Result<(), StoreError> {
|
||||
let mut bytes = serde_json::to_vec_pretty(value)?;
|
||||
bytes.push(b'\n');
|
||||
@@ -379,7 +561,21 @@ fn truncate_uncommitted_tail(file: &mut File) -> std::io::Result<u64> {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::{Store, new_segment_id, new_session_id};
|
||||
use crate::{
|
||||
LoggedHistoryEntry, LoggedItem, LoggedSessionHistoryEntryId, LoggedSessionHistoryMetadata,
|
||||
LoggedSessionHistoryOrigin, Store, new_segment_id, new_session_id,
|
||||
};
|
||||
|
||||
fn annotated(item: agen::Item) -> LoggedHistoryEntry {
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::from(item),
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn canonical_layout_and_single_session_invariant() {
|
||||
@@ -405,6 +601,328 @@ mod tests {
|
||||
assert_eq!(store.list_sessions().unwrap(), vec![session_id]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v1_logs_are_rewritten_and_promoted_to_v3() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
WorkerSessionStore::new(root.path())
|
||||
.unwrap()
|
||||
.create_segment(session_id, segment_id, &[])
|
||||
.unwrap();
|
||||
let manifest_path = root.path().join(SESSION_FILE);
|
||||
let mut manifest: SessionManifest =
|
||||
serde_json::from_slice(&fs::read(&manifest_path).unwrap()).unwrap();
|
||||
manifest.schema_version = LEGACY_SESSION_SCHEMA_VERSION;
|
||||
atomic_write_json(&manifest_path, &manifest).unwrap();
|
||||
|
||||
let reopened = WorkerSessionStore::new(root.path()).unwrap();
|
||||
assert_eq!(reopened.session_id().unwrap(), Some(session_id));
|
||||
let migrated: SessionManifest =
|
||||
serde_json::from_slice(&fs::read(&manifest_path).unwrap()).unwrap();
|
||||
assert_eq!(migrated.schema_version, SESSION_SCHEMA_VERSION);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v1_migration_rejects_corrupt_log_before_v3_manifest_update() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let manifest = SessionManifest {
|
||||
schema_version: LEGACY_SESSION_SCHEMA_VERSION,
|
||||
session_id,
|
||||
};
|
||||
atomic_write_json(&root.path().join(SESSION_FILE), &manifest).unwrap();
|
||||
fs::create_dir_all(root.path().join(SEGMENTS_DIR)).unwrap();
|
||||
fs::write(
|
||||
root.path().join(SEGMENTS_DIR).join("broken.jsonl"),
|
||||
"{not-json}\n",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let error = match WorkerSessionStore::new(root.path()) {
|
||||
Ok(_) => panic!("corrupt legacy Session log must reject migration"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(matches!(error, StoreError::Corrupt { .. }));
|
||||
let persisted: SessionManifest =
|
||||
serde_json::from_slice(&fs::read(root.path().join(SESSION_FILE)).unwrap()).unwrap();
|
||||
assert_eq!(persisted.schema_version, LEGACY_SESSION_SCHEMA_VERSION);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v2_migration_rewrites_legacy_records_with_stable_unknown_provenance() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
fs::create_dir_all(root.path().join(SEGMENTS_DIR)).unwrap();
|
||||
atomic_write_json(
|
||||
&root.path().join(SESSION_FILE),
|
||||
&SessionManifest {
|
||||
schema_version: PREVIOUS_SESSION_SCHEMA_VERSION,
|
||||
session_id,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let source = vec![
|
||||
serde_json::json!({
|
||||
"kind": "segment_start",
|
||||
"ts": 1,
|
||||
"session_id": session_id,
|
||||
"system_prompt": null,
|
||||
"config": agen::llm_client::RequestConfig::default(),
|
||||
"history": [LoggedItem::from(agen::Item::assistant_message("prior"))],
|
||||
"forked_from": null,
|
||||
"compacted_from": null
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "user_input",
|
||||
"ts": 2,
|
||||
"segments": [{ "kind": "text", "content": "hello" }],
|
||||
"extensions": []
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 3,
|
||||
"item": LoggedItem::from(agen::Item::assistant_message("reply"))
|
||||
}),
|
||||
];
|
||||
let path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{segment_id}.jsonl"));
|
||||
let mut bytes = Vec::new();
|
||||
for entry in source {
|
||||
serde_json::to_writer(&mut bytes, &entry).unwrap();
|
||||
bytes.push(b'\n');
|
||||
}
|
||||
fs::write(&path, bytes).unwrap();
|
||||
|
||||
let store = WorkerSessionStore::new(root.path()).unwrap();
|
||||
let first = store.read_all(session_id, segment_id).unwrap();
|
||||
assert!(matches!(first[0], LogEntry::AnnotatedSegmentStart { .. }));
|
||||
assert!(matches!(first[1], LogEntry::AnnotatedUserInput { .. }));
|
||||
assert!(matches!(first[2], LogEntry::AnnotatedAssistantItem { .. }));
|
||||
let first_bytes = fs::read(&path).unwrap();
|
||||
drop(store);
|
||||
|
||||
let reopened = WorkerSessionStore::new(root.path()).unwrap();
|
||||
assert_eq!(fs::read(&path).unwrap(), first_bytes);
|
||||
let snapshot = crate::public_snapshot::project_current_session_snapshot(
|
||||
&reopened.read_all(session_id, segment_id).unwrap(),
|
||||
);
|
||||
assert_eq!(snapshot.entries.len(), 3);
|
||||
assert_eq!(
|
||||
snapshot
|
||||
.entries
|
||||
.iter()
|
||||
.map(|entry| entry.timestamp)
|
||||
.collect::<Vec<_>>(),
|
||||
vec![1, 2, 3]
|
||||
);
|
||||
assert!(snapshot.entries.iter().all(|entry| {
|
||||
entry.provenance == protocol::SessionEntryProvenance::LegacyUnknown
|
||||
&& entry.entry_id.len() <= 64
|
||||
}));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v2_preflight_keeps_earlier_segments_unchanged_when_later_is_corrupt() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let valid_segment = uuid::Uuid::from_u128(1);
|
||||
let corrupt_segment = uuid::Uuid::from_u128(2);
|
||||
fs::create_dir_all(root.path().join(SEGMENTS_DIR)).unwrap();
|
||||
atomic_write_json(
|
||||
&root.path().join(SESSION_FILE),
|
||||
&SessionManifest {
|
||||
schema_version: PREVIOUS_SESSION_SCHEMA_VERSION,
|
||||
session_id,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let manifest_before = fs::read(root.path().join(SESSION_FILE)).unwrap();
|
||||
|
||||
let valid_path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{valid_segment}.jsonl"));
|
||||
let valid_entry = serde_json::json!({
|
||||
"kind": "segment_start",
|
||||
"ts": 1,
|
||||
"session_id": session_id,
|
||||
"system_prompt": null,
|
||||
"config": agen::llm_client::RequestConfig::default(),
|
||||
"history": [LoggedItem::from(agen::Item::assistant_message("prior"))],
|
||||
"forked_from": null,
|
||||
"compacted_from": null
|
||||
});
|
||||
let mut valid_bytes = serde_json::to_vec(&valid_entry).unwrap();
|
||||
valid_bytes.push(b'\n');
|
||||
fs::write(&valid_path, &valid_bytes).unwrap();
|
||||
let corrupt_path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{corrupt_segment}.jsonl"));
|
||||
fs::write(&corrupt_path, b"{not-json}\n").unwrap();
|
||||
let corrupt_before = fs::read(&corrupt_path).unwrap();
|
||||
|
||||
let error = match WorkerSessionStore::new(root.path()) {
|
||||
Ok(_) => panic!("later corrupt segment must fail migration preflight"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(matches!(error, StoreError::Corrupt { .. }));
|
||||
assert_eq!(fs::read(&valid_path).unwrap(), valid_bytes);
|
||||
assert_eq!(fs::read(&corrupt_path).unwrap(), corrupt_before);
|
||||
assert_eq!(
|
||||
fs::read(root.path().join(SESSION_FILE)).unwrap(),
|
||||
manifest_before
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn current_jsonl_requires_annotations_across_append_rewrite_and_reopen() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
let store = WorkerSessionStore::new(root.path()).unwrap();
|
||||
store
|
||||
.create_segment(
|
||||
session_id,
|
||||
segment_id,
|
||||
&[LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: agen::llm_client::RequestConfig::default(),
|
||||
history: vec![annotated(agen::Item::user_message("seed"))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}],
|
||||
)
|
||||
.unwrap();
|
||||
store
|
||||
.append(
|
||||
session_id,
|
||||
segment_id,
|
||||
&LogEntry::AnnotatedAssistantItem {
|
||||
ts: 2,
|
||||
entry: annotated(agen::Item::assistant_message("reply")),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let before_rewrite = store.read_all(session_id, segment_id).unwrap();
|
||||
store
|
||||
.create_segment(session_id, segment_id, &before_rewrite)
|
||||
.unwrap();
|
||||
drop(store);
|
||||
|
||||
let reopened = WorkerSessionStore::new(root.path()).unwrap();
|
||||
let restored = reopened.read_all(session_id, segment_id).unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(&restored).unwrap(),
|
||||
serde_json::to_value(&before_rewrite).unwrap()
|
||||
);
|
||||
for entry in &restored {
|
||||
match entry {
|
||||
LogEntry::AnnotatedSegmentStart { history, .. } => assert!(history.iter().all(
|
||||
|entry| !entry.metadata.entry_id.0.is_empty()
|
||||
&& matches!(
|
||||
entry.metadata.origin,
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown
|
||||
)
|
||||
)),
|
||||
LogEntry::AnnotatedAssistantItem { entry, .. } => {
|
||||
assert!(!entry.metadata.entry_id.0.is_empty());
|
||||
assert!(matches!(
|
||||
entry.metadata.origin,
|
||||
LoggedSessionHistoryOrigin::LegacyUnknown
|
||||
));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
let log = fs::read_to_string(reopened.log_path(segment_id)).unwrap();
|
||||
for line in log.lines() {
|
||||
let value: serde_json::Value = serde_json::from_str(line).unwrap();
|
||||
let kind = value["kind"].as_str().unwrap();
|
||||
assert!(
|
||||
!matches!(
|
||||
kind,
|
||||
"segment_start"
|
||||
| "user_input"
|
||||
| "assistant_item"
|
||||
| "tool_result"
|
||||
| "system_item"
|
||||
),
|
||||
"current-schema JSONL contains legacy history record: {kind}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn schema_v3_rejects_legacy_records_and_new_writes_are_canonical() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let session_id = new_session_id();
|
||||
let segment_id = new_segment_id();
|
||||
let store = WorkerSessionStore::new(root.path()).unwrap();
|
||||
store
|
||||
.create_segment(
|
||||
session_id,
|
||||
segment_id,
|
||||
&[LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
config: agen::llm_client::RequestConfig::default(),
|
||||
history: vec![annotated(agen::Item::assistant_message("seed"))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
}],
|
||||
)
|
||||
.unwrap();
|
||||
store
|
||||
.append(
|
||||
session_id,
|
||||
segment_id,
|
||||
&LogEntry::AnnotatedUserInput {
|
||||
ts: 2,
|
||||
segments: vec![protocol::Segment::Text {
|
||||
content: "new".into(),
|
||||
}],
|
||||
history: vec![annotated(agen::Item::user_message("new"))],
|
||||
extensions: Vec::new(),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
let entries = store.read_all(session_id, segment_id).unwrap();
|
||||
assert!(matches!(entries[0], LogEntry::AnnotatedSegmentStart { .. }));
|
||||
assert!(matches!(entries[1], LogEntry::AnnotatedUserInput { .. }));
|
||||
drop(store);
|
||||
|
||||
let path = root
|
||||
.path()
|
||||
.join(SEGMENTS_DIR)
|
||||
.join(format!("{segment_id}.jsonl"));
|
||||
let mut file = OpenOptions::new().append(true).open(path).unwrap();
|
||||
serde_json::to_writer(
|
||||
&mut file,
|
||||
&serde_json::json!({
|
||||
"kind": "system_item",
|
||||
"ts": 3,
|
||||
"item": { "kind": "legacy_ignored", "slug": "legacy" }
|
||||
}),
|
||||
)
|
||||
.unwrap();
|
||||
file.write_all(b"\n").unwrap();
|
||||
let error = match WorkerSessionStore::new(root.path()) {
|
||||
Ok(_) => panic!("schema v3 must reject a legacy history record"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(matches!(error, StoreError::Corrupt { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reopen_preserves_session_and_segment_ids() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
|
||||
@@ -1,12 +1,25 @@
|
||||
use agen::EngineResult;
|
||||
use agen::llm_client::types::{Item, RequestConfig};
|
||||
use session_store::{
|
||||
FsStore, LogEntry, Store, TraceEntry, collect_state, new_segment_id, new_session_id,
|
||||
FsStore, LogEntry, LoggedHistoryEntry, LoggedItem, LoggedSessionHistoryEntryId,
|
||||
LoggedSessionHistoryMetadata, LoggedSessionHistoryOrigin, Store, TraceEntry, collect_state,
|
||||
new_segment_id, new_session_id,
|
||||
};
|
||||
use std::io::Write;
|
||||
|
||||
fn annotated(item: Item) -> LoggedHistoryEntry {
|
||||
LoggedHistoryEntry {
|
||||
item: LoggedItem::from(item),
|
||||
metadata: LoggedSessionHistoryMetadata {
|
||||
entry_id: LoggedSessionHistoryEntryId::new(),
|
||||
origin: LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn nil_session_start(ts: u64, session_id: uuid::Uuid) -> LogEntry {
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts,
|
||||
session_id,
|
||||
system_prompt: None,
|
||||
@@ -25,7 +38,7 @@ fn round_trip_write_and_read() {
|
||||
let segid = new_segment_id();
|
||||
|
||||
let entries = vec![
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: sid,
|
||||
system_prompt: Some("You are helpful.".into()),
|
||||
@@ -34,14 +47,15 @@ fn round_trip_write_and_read() {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("Hello")],
|
||||
history: vec![annotated(Item::user_message("Hello"))],
|
||||
},
|
||||
LogEntry::AssistantItem {
|
||||
LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3000,
|
||||
item: Item::assistant_message("Hi there!").into(),
|
||||
entry: annotated(Item::assistant_message("Hi there!")),
|
||||
},
|
||||
LogEntry::TurnEnd {
|
||||
ts: 3100,
|
||||
@@ -51,6 +65,7 @@ fn round_trip_write_and_read() {
|
||||
ts: 3200,
|
||||
interrupted: false,
|
||||
result: EngineResult::Finished,
|
||||
active_run_turn_count: None,
|
||||
},
|
||||
];
|
||||
|
||||
@@ -78,14 +93,14 @@ fn create_segment_writes_all_entries() {
|
||||
let sid = new_session_id();
|
||||
let segid = new_segment_id();
|
||||
|
||||
let entries = [LogEntry::SegmentStart {
|
||||
let entries = [LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: sid,
|
||||
system_prompt: None,
|
||||
config: RequestConfig::default(),
|
||||
history: vec![
|
||||
Item::user_message("seed").into(),
|
||||
Item::assistant_message("ok").into(),
|
||||
annotated(Item::user_message("seed")),
|
||||
annotated(Item::assistant_message("ok")),
|
||||
],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
@@ -204,7 +219,7 @@ fn read_entry_count_matches_append_tally() {
|
||||
let segid = new_segment_id();
|
||||
|
||||
let entries = [
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1000,
|
||||
session_id: sid,
|
||||
system_prompt: None,
|
||||
@@ -213,10 +228,11 @@ fn read_entry_count_matches_append_tally() {
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
},
|
||||
LogEntry::UserInput {
|
||||
LogEntry::AnnotatedUserInput {
|
||||
ts: 2000,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("Hello")],
|
||||
history: vec![annotated(Item::user_message("Hello"))],
|
||||
},
|
||||
];
|
||||
|
||||
@@ -253,10 +269,11 @@ fn unterminated_utf8_tail_is_ignored_and_replaced_on_append() {
|
||||
assert_eq!(store.read_all(sid, segid).unwrap().len(), 1);
|
||||
assert_eq!(store.read_entry_count(sid, segid).unwrap(), 1);
|
||||
|
||||
let next = LogEntry::UserInput {
|
||||
let next = LogEntry::AnnotatedUserInput {
|
||||
ts: 2,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("recovered")],
|
||||
history: vec![annotated(Item::user_message("recovered"))],
|
||||
};
|
||||
store.append(sid, segid, &next).unwrap();
|
||||
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
mod common;
|
||||
|
||||
use std::ops::{Deref, DerefMut};
|
||||
use std::sync::Arc;
|
||||
|
||||
use agen::Engine;
|
||||
use agen::interceptor::{Interceptor, TurnEndAction};
|
||||
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
||||
use agen::llm_client::types::{Item, RequestConfig};
|
||||
use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput};
|
||||
use agen::{Engine, History};
|
||||
use async_trait::async_trait;
|
||||
use common::MockLlmClient;
|
||||
use session_store::{FsStore, LogEntry, SegmentStartState, Store, collect_state};
|
||||
@@ -15,6 +16,21 @@ use session_store::{FsStore, LogEntry, SegmentStartState, Store, collect_state};
|
||||
// Helpers
|
||||
// =============================================================================
|
||||
|
||||
fn annotated(items: &[Item]) -> Vec<session_store::LoggedHistoryEntry> {
|
||||
items
|
||||
.iter()
|
||||
.cloned()
|
||||
.map(|item| session_store::LoggedHistoryEntry {
|
||||
item: session_store::LoggedItem::from(item),
|
||||
metadata: session_store::LoggedSessionHistoryMetadata {
|
||||
entry_id: session_store::LoggedSessionHistoryEntryId::new(),
|
||||
origin: session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn simple_text_events() -> Vec<Event> {
|
||||
vec![
|
||||
Event::text_block_start(0),
|
||||
@@ -94,15 +110,47 @@ fn make_store() -> (tempfile::TempDir, FsStore) {
|
||||
(dir, store)
|
||||
}
|
||||
|
||||
struct TestWorker {
|
||||
engine: Engine<MockLlmClient>,
|
||||
history: History,
|
||||
}
|
||||
|
||||
impl TestWorker {
|
||||
fn new(engine: Engine<MockLlmClient>) -> Self {
|
||||
Self {
|
||||
engine,
|
||||
history: History::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn history(&self) -> Vec<Item> {
|
||||
self.history.items_cloned()
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for TestWorker {
|
||||
type Target = Engine<MockLlmClient>;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.engine
|
||||
}
|
||||
}
|
||||
|
||||
impl DerefMut for TestWorker {
|
||||
fn deref_mut(&mut self) -> &mut Self::Target {
|
||||
&mut self.engine
|
||||
}
|
||||
}
|
||||
|
||||
/// Run a worker turn and persist via session-store functions.
|
||||
/// Takes ownership of the worker (needed for lock/unlock) and returns it.
|
||||
async fn run_and_persist(
|
||||
worker: Engine<MockLlmClient>,
|
||||
mut worker: TestWorker,
|
||||
store: &FsStore,
|
||||
session_id: session_store::SessionId,
|
||||
segment_id: session_store::SegmentId,
|
||||
input: &str,
|
||||
) -> (Engine<MockLlmClient>, agen::EngineResult) {
|
||||
) -> (TestWorker, agen::EngineRunExit) {
|
||||
// Mirror Worker's run-entry contract: log the user input as segments
|
||||
// before the worker pushes its flattened user_message; save_delta
|
||||
// skips the resulting user_message item to avoid double-write.
|
||||
@@ -111,44 +159,65 @@ async fn run_and_persist(
|
||||
session_id,
|
||||
segment_id,
|
||||
vec![protocol::Segment::text(input)],
|
||||
annotated(&[Item::user_message(input)]),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let history_before = worker.history().len();
|
||||
let history_before = worker.history.len();
|
||||
|
||||
let mut locked = worker.lock();
|
||||
let result = locked.run(input).await;
|
||||
let worker = locked.unlock();
|
||||
let mut locked = worker.engine.lock(&worker.history);
|
||||
let result = locked.run(&mut worker.history, input).await;
|
||||
worker.engine = locked.unlock();
|
||||
|
||||
let new_items = &worker.history()[history_before..];
|
||||
session_store::save_delta(store, session_id, segment_id, new_items).unwrap();
|
||||
let projected = worker.history();
|
||||
let new_items = annotated(&projected[history_before..]);
|
||||
session_store::save_delta(store, session_id, segment_id, &new_items).unwrap();
|
||||
session_store::save_turn_end(store, session_id, segment_id, worker.turn_count()).unwrap();
|
||||
|
||||
match &result {
|
||||
Ok(r) => {
|
||||
agen::EngineRunExit::Finished
|
||||
| agen::EngineRunExit::Paused
|
||||
| agen::EngineRunExit::Yielded => {
|
||||
let (legacy_result, interrupted) = match &result {
|
||||
agen::EngineRunExit::Finished => (agen::EngineResult::Finished, false),
|
||||
agen::EngineRunExit::Paused => (agen::EngineResult::Paused, true),
|
||||
agen::EngineRunExit::Yielded => (agen::EngineResult::Yielded, true),
|
||||
agen::EngineRunExit::Interrupted(_) => unreachable!(),
|
||||
};
|
||||
session_store::save_run_completed(
|
||||
store,
|
||||
session_id,
|
||||
segment_id,
|
||||
r.clone(),
|
||||
worker.last_run_interrupted(),
|
||||
legacy_result,
|
||||
interrupted,
|
||||
worker.active_run_turn_count(),
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
Err(e) => {
|
||||
agen::EngineRunExit::Interrupted(agen::StopReason::LimitReached) => {
|
||||
session_store::save_run_completed(
|
||||
store,
|
||||
session_id,
|
||||
segment_id,
|
||||
agen::EngineResult::LimitReached,
|
||||
false,
|
||||
worker.active_run_turn_count(),
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
agen::EngineRunExit::Interrupted(reason) => {
|
||||
session_store::save_run_errored(
|
||||
store,
|
||||
session_id,
|
||||
segment_id,
|
||||
e.to_string(),
|
||||
worker.last_run_interrupted(),
|
||||
format!("{reason:?}"),
|
||||
true,
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
let r = result.unwrap();
|
||||
(worker, r)
|
||||
(worker, result)
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -159,14 +228,14 @@ async fn run_and_persist(
|
||||
async fn session_run_logs_entries() {
|
||||
let (_dir, store) = make_store();
|
||||
let client = MockLlmClient::new(simple_text_events());
|
||||
let worker = Engine::new(client);
|
||||
let worker = TestWorker::new(Engine::new(client));
|
||||
|
||||
let (sid, segid) = session_store::create_segment(
|
||||
&store,
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -184,7 +253,10 @@ async fn session_run_logs_entries() {
|
||||
);
|
||||
|
||||
// First entry is SegmentStart
|
||||
assert!(matches!(&entries[0], LogEntry::SegmentStart { .. }));
|
||||
assert!(matches!(
|
||||
&entries[0],
|
||||
LogEntry::AnnotatedSegmentStart { .. }
|
||||
));
|
||||
|
||||
// Has a RunCompleted with Finished
|
||||
let has_finished = entries.iter().any(|e| {
|
||||
@@ -203,7 +275,7 @@ async fn session_run_logs_entries() {
|
||||
async fn session_restore_round_trip() {
|
||||
let (_dir, store) = make_store();
|
||||
let client = MockLlmClient::new(simple_text_events());
|
||||
let mut worker = Engine::new(client);
|
||||
let mut worker = TestWorker::new(Engine::new(client));
|
||||
worker.set_system_prompt("You are helpful.");
|
||||
|
||||
let (sid, segid) = session_store::create_segment(
|
||||
@@ -211,7 +283,7 @@ async fn session_restore_round_trip() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -242,7 +314,7 @@ async fn session_restore_round_trip() {
|
||||
async fn session_run_with_tool_call() {
|
||||
let (_dir, store) = make_store();
|
||||
let client = MockLlmClient::with_responses(tool_call_events());
|
||||
let mut worker = Engine::new(client);
|
||||
let mut worker = TestWorker::new(Engine::new(client));
|
||||
worker.register_tool(weather_tool_definition());
|
||||
|
||||
let (sid, segid) = session_store::create_segment(
|
||||
@@ -250,7 +322,7 @@ async fn session_run_with_tool_call() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -261,12 +333,12 @@ async fn session_run_with_tool_call() {
|
||||
|
||||
let has_tool_results = entries
|
||||
.iter()
|
||||
.any(|e| matches!(e, LogEntry::ToolResult { .. }));
|
||||
.any(|e| matches!(e, LogEntry::AnnotatedToolResult { .. }));
|
||||
assert!(has_tool_results, "should have ToolResult entry");
|
||||
|
||||
let has_assistant = entries
|
||||
.iter()
|
||||
.any(|e| matches!(e, LogEntry::AssistantItem { .. }));
|
||||
.any(|e| matches!(e, LogEntry::AnnotatedAssistantItem { .. }));
|
||||
assert!(has_assistant, "should have AssistantItem entry");
|
||||
}
|
||||
|
||||
@@ -276,7 +348,7 @@ async fn session_resume_after_pause() {
|
||||
|
||||
// First run: tool call with pause policy → Paused
|
||||
let client = MockLlmClient::with_responses(tool_call_events());
|
||||
let mut worker = Engine::new(client);
|
||||
let mut worker = TestWorker::new(Engine::new(client));
|
||||
worker.register_tool(weather_tool_definition());
|
||||
worker.set_interceptor(PausePolicy);
|
||||
|
||||
@@ -285,13 +357,13 @@ async fn session_resume_after_pause() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let (_worker, result) = run_and_persist(worker, &store, sid, segid, "Weather?").await;
|
||||
assert!(matches!(result, agen::EngineResult::Paused));
|
||||
assert!(matches!(result, agen::EngineRunExit::Paused));
|
||||
|
||||
// Check RunCompleted is Paused
|
||||
let entries = store.read_all(sid, segid).unwrap();
|
||||
@@ -309,13 +381,14 @@ async fn session_resume_after_pause() {
|
||||
// Restore state and verify
|
||||
let state = session_store::restore(&store, sid, segid).unwrap();
|
||||
assert!(state.last_run_interrupted);
|
||||
assert_eq!(state.active_run_turn_count, Some(2));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn session_fork_creates_new_session() {
|
||||
let (_dir, store) = make_store();
|
||||
let client = MockLlmClient::new(simple_text_events());
|
||||
let mut worker = Engine::new(client);
|
||||
let mut worker = TestWorker::new(Engine::new(client));
|
||||
worker.set_system_prompt("System prompt");
|
||||
|
||||
let (sid, segid) = session_store::create_segment(
|
||||
@@ -323,7 +396,7 @@ async fn session_fork_creates_new_session() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -336,7 +409,7 @@ async fn session_fork_creates_new_session() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -345,7 +418,10 @@ async fn session_fork_creates_new_session() {
|
||||
// Fork should have a SegmentStart with the current history
|
||||
let fork_entries = store.read_all(fork_sid, fork_segid).unwrap();
|
||||
assert_eq!(fork_entries.len(), 1);
|
||||
assert!(matches!(&fork_entries[0], LogEntry::SegmentStart { .. }));
|
||||
assert!(matches!(
|
||||
&fork_entries[0],
|
||||
LogEntry::AnnotatedSegmentStart { .. }
|
||||
));
|
||||
|
||||
let fork_state = collect_state(&fork_entries);
|
||||
assert_eq!(fork_state.session_id, Some(fork_sid));
|
||||
@@ -357,14 +433,14 @@ async fn session_fork_creates_new_session() {
|
||||
async fn session_fork_at_truncates_within_session() {
|
||||
let (_dir, store) = make_store();
|
||||
let client = MockLlmClient::new(simple_text_events());
|
||||
let worker = Engine::new(client);
|
||||
let worker = TestWorker::new(Engine::new(client));
|
||||
|
||||
let (sid, segid) = session_store::create_segment(
|
||||
&store,
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -391,6 +467,23 @@ async fn session_fork_at_truncates_within_session() {
|
||||
.expect("source segment has the matching TurnEnd");
|
||||
let source_state_at_fork = collect_state(&all_entries[..=turn_end_pos]);
|
||||
assert_eq!(fork_state.history.len(), source_state_at_fork.history.len());
|
||||
assert_eq!(
|
||||
fork_state.annotated_history, source_state_at_fork.annotated_history,
|
||||
"fork_at must preserve every retained history entry identity and provenance",
|
||||
);
|
||||
assert!(fork_state.annotated_history.iter().all(|entry| {
|
||||
!entry.metadata.entry_id.0.is_empty()
|
||||
&& matches!(
|
||||
entry.metadata.origin,
|
||||
session_store::LoggedSessionHistoryOrigin::LegacyUnknown
|
||||
| session_store::LoggedSessionHistoryOrigin::HumanInput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::WorkerInput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::BackendInstruction { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::ModelOutput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::ToolOutput { .. }
|
||||
| session_store::LoggedSessionHistoryOrigin::DerivedSummary
|
||||
)
|
||||
}));
|
||||
|
||||
// list_segments should show both source and fork in the same Session.
|
||||
let segs = store.list_segments(sid).unwrap();
|
||||
@@ -402,14 +495,14 @@ async fn session_fork_at_truncates_within_session() {
|
||||
async fn session_config_changed_logged() {
|
||||
let (_dir, store) = make_store();
|
||||
let client = MockLlmClient::new(vec![]);
|
||||
let mut worker = Engine::new(client);
|
||||
let mut worker = TestWorker::new(Engine::new(client));
|
||||
|
||||
let (sid, segid) = session_store::create_segment(
|
||||
&store,
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -435,14 +528,14 @@ async fn session_auto_forks_on_conflict() {
|
||||
|
||||
// Create a segment
|
||||
let client_a = MockLlmClient::new(simple_text_events());
|
||||
let worker_a = Engine::new(client_a);
|
||||
let worker_a = TestWorker::new(Engine::new(client_a));
|
||||
|
||||
let (sid, original_segid) = session_store::create_segment(
|
||||
&store,
|
||||
SegmentStartState {
|
||||
system_prompt: worker_a.get_system_prompt(),
|
||||
config: worker_a.request_config(),
|
||||
history: worker_a.history(),
|
||||
history: annotated(&worker_a.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -451,12 +544,14 @@ async fn session_auto_forks_on_conflict() {
|
||||
let mut entries_written: usize = 1;
|
||||
|
||||
// Simulate another Worker writing to the same segment behind our back.
|
||||
let extra_entry = LogEntry::UserInput {
|
||||
ts: 9999,
|
||||
extensions: vec![],
|
||||
segments: vec![protocol::Segment::text("Interloper")],
|
||||
};
|
||||
store.append(sid, original_segid, &extra_entry).unwrap();
|
||||
session_store::save_user_input(
|
||||
&store,
|
||||
sid,
|
||||
original_segid,
|
||||
vec![protocol::Segment::text("Interloper")],
|
||||
annotated(&[Item::user_message("Interloper")]),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
// Now the on-disk count exceeds our tally — ensure_head_or_fork should auto-fork.
|
||||
session_store::ensure_head_or_fork(
|
||||
@@ -468,7 +563,7 @@ async fn session_auto_forks_on_conflict() {
|
||||
SegmentStartState {
|
||||
system_prompt: worker_a.get_system_prompt(),
|
||||
config: worker_a.request_config(),
|
||||
history: worker_a.history(),
|
||||
history: annotated(&worker_a.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -489,7 +584,7 @@ async fn session_auto_forks_on_conflict() {
|
||||
// The new segment records its lineage forward via forked_from; the
|
||||
// source segment is left immutable (no terminal marker written back).
|
||||
match &fork_entries[0] {
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
forked_from: Some(origin),
|
||||
..
|
||||
} => {
|
||||
@@ -509,7 +604,7 @@ async fn session_auto_forks_on_conflict() {
|
||||
);
|
||||
let has_interloper = original_entries
|
||||
.iter()
|
||||
.any(|e| matches!(e, LogEntry::UserInput { .. }));
|
||||
.any(|e| matches!(e, LogEntry::AnnotatedUserInput { .. }));
|
||||
assert!(has_interloper);
|
||||
}
|
||||
|
||||
@@ -520,14 +615,14 @@ async fn session_auto_forks_on_conflict() {
|
||||
async fn nested_past_fork_leaves_ancestors_immutable() {
|
||||
let (_dir, store) = make_store();
|
||||
let client = MockLlmClient::new(simple_text_events());
|
||||
let worker = Engine::new(client);
|
||||
let worker = TestWorker::new(Engine::new(client));
|
||||
|
||||
let (sid, root_segid) = session_store::create_segment(
|
||||
&store,
|
||||
SegmentStartState {
|
||||
system_prompt: worker.get_system_prompt(),
|
||||
config: worker.request_config(),
|
||||
history: worker.history(),
|
||||
history: annotated(&worker.history()),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
@@ -564,7 +659,7 @@ async fn nested_past_fork_leaves_ancestors_immutable() {
|
||||
|
||||
// fork2's lineage points at fork1, not the root.
|
||||
match &store.read_all(sid, fork2).unwrap()[0] {
|
||||
LogEntry::SegmentStart {
|
||||
LogEntry::AnnotatedSegmentStart {
|
||||
forked_from: Some(origin),
|
||||
..
|
||||
} => assert_eq!(origin.segment_id, fork1),
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
[package]
|
||||
name = "standalone"
|
||||
description = "In-process standalone Worker host"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
agen.workspace = true
|
||||
fs4.workspace = true
|
||||
manifest.workspace = true
|
||||
protocol.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
session-store.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["rt", "sync", "time"] }
|
||||
uuid = { workspace = true, features = ["v7"] }
|
||||
worker.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
async-trait.workspace = true
|
||||
futures.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread", "time"] }
|
||||
@@ -0,0 +1,410 @@
|
||||
use std::path::PathBuf;
|
||||
use std::time::Duration;
|
||||
|
||||
use agen::llm_client::client::LlmClient;
|
||||
use protocol::{Event, Method};
|
||||
use session_store::{
|
||||
CombinedStore, FsStore, FsWorkerStore, WorkerActiveSegmentRef, WorkerMetadataStore,
|
||||
};
|
||||
use thiserror::Error;
|
||||
use tokio::sync::broadcast;
|
||||
use worker::bootstrap::{WorkerBootstrap, WorkerBootstrapError, WorkerBootstrapLayout};
|
||||
use worker::controller::WorkerControllerTransport;
|
||||
use worker::{BootstrappedWorker, WorkerError, WorkerFilesystemAuthority, WorkerWorkspaceContext};
|
||||
|
||||
use crate::launch::ResolvedStandaloneLaunch;
|
||||
use crate::store::{
|
||||
StaleLeasePolicy, StandaloneSessionId, StandaloneSessionLease, StandaloneSessionRecord,
|
||||
StandaloneSessionStore, StandaloneShutdownReason, StandaloneStoreError,
|
||||
};
|
||||
|
||||
const DEFAULT_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
type StandaloneBackingStore = CombinedStore<FsStore, FsWorkerStore>;
|
||||
|
||||
/// One client-owned top-level Worker and its standalone session authority.
|
||||
///
|
||||
/// The host deliberately exposes the existing typed Worker protocol rather than owning an
|
||||
/// HTTP/WebSocket server or creating Runtime/Workspace/Ticket/Workdir domain records.
|
||||
pub struct StandaloneHost {
|
||||
handle: worker::WorkerHandle,
|
||||
shutdown: Option<worker::controller::ShutdownReceiver>,
|
||||
shutdown_timeout: Duration,
|
||||
store: StandaloneSessionStore,
|
||||
worker_store: FsWorkerStore,
|
||||
record: StandaloneSessionRecord,
|
||||
lease: Option<StandaloneSessionLease>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneStartupError {
|
||||
#[error("the standalone state store could not be opened or validated")]
|
||||
StateStore,
|
||||
#[error("the standalone session is already active")]
|
||||
SessionActive,
|
||||
#[error("the standalone session lease cannot be observed safely; recovery is rejected")]
|
||||
LeaseLivenessUnknown,
|
||||
#[error("the standalone session working directory is unavailable or changed")]
|
||||
WorkingDirectoryUnavailable,
|
||||
#[error("the resolved Worker configuration or persisted history is invalid")]
|
||||
WorkerConfiguration,
|
||||
#[error("the configured model provider is unavailable")]
|
||||
ModelProvider,
|
||||
#[error("the fixed standalone feature composition could not be installed")]
|
||||
FeatureComposition,
|
||||
#[error("the in-process Worker controller could not start")]
|
||||
Controller,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneRequestError {
|
||||
#[error("the standalone Worker is no longer accepting requests")]
|
||||
WorkerUnavailable,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneShutdownError {
|
||||
#[error("the standalone Worker did not stop before the shutdown deadline")]
|
||||
DeadlineExceeded,
|
||||
#[error("the standalone Worker shutdown confirmation was lost")]
|
||||
ConfirmationLost,
|
||||
#[error("the standalone session final state could not be committed")]
|
||||
StateStore,
|
||||
}
|
||||
|
||||
impl StandaloneHost {
|
||||
pub async fn start(launch: ResolvedStandaloneLaunch) -> Result<Self, StandaloneStartupError> {
|
||||
Self::start_with_optional_model_client(launch, None).await
|
||||
}
|
||||
|
||||
pub async fn start_with_model_client<C>(
|
||||
launch: ResolvedStandaloneLaunch,
|
||||
model_client: C,
|
||||
) -> Result<Self, StandaloneStartupError>
|
||||
where
|
||||
C: LlmClient + 'static,
|
||||
{
|
||||
Self::start_with_optional_model_client(launch, Some(Box::new(model_client))).await
|
||||
}
|
||||
|
||||
async fn start_with_optional_model_client(
|
||||
mut launch: ResolvedStandaloneLaunch,
|
||||
model_client: Option<Box<dyn LlmClient>>,
|
||||
) -> Result<Self, StandaloneStartupError> {
|
||||
let store = StandaloneSessionStore::open(&launch.state_dir)
|
||||
.map_err(classify_store_startup_error)?;
|
||||
let allocation = store
|
||||
.allocate(&launch.cwd, StaleLeasePolicy::Reject)
|
||||
.map_err(classify_store_startup_error)?;
|
||||
let id = allocation.id();
|
||||
|
||||
// The standalone session ID is the local identity. A unique internal Worker name avoids
|
||||
// process-global allocation collisions without creating a Runtime/Workspace Worker ID.
|
||||
launch.profile.manifest.worker.name = format!("standalone-{id}");
|
||||
let manifest = launch.profile.manifest.clone();
|
||||
let worker_name = manifest.worker.name.clone();
|
||||
let (backing_store, worker_store) = match backing_store(&store, id) {
|
||||
Ok(stores) => stores,
|
||||
Err(error) => {
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let filesystem_authority =
|
||||
WorkerFilesystemAuthority::local(launch.cwd.clone(), launch.cwd.clone());
|
||||
let workspace_context = WorkerWorkspaceContext::local_filesystem(None);
|
||||
let runtime_base = store.runtime_dir(id);
|
||||
|
||||
let mut bootstrap = WorkerBootstrap::new(
|
||||
manifest.clone(),
|
||||
backing_store,
|
||||
launch.prompt_catalog,
|
||||
workspace_context,
|
||||
filesystem_authority,
|
||||
WorkerBootstrapLayout::Direct { runtime_base },
|
||||
WorkerControllerTransport::InProcess,
|
||||
);
|
||||
if let Some(model_client) = model_client {
|
||||
bootstrap = bootstrap.with_model_client(model_client);
|
||||
}
|
||||
let started = match bootstrap.start().await {
|
||||
Ok(started) => started,
|
||||
Err(error) => {
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(classify_startup_error(error));
|
||||
}
|
||||
};
|
||||
let active = match active_pointer(&worker_store, &worker_name) {
|
||||
Ok(active) => active,
|
||||
Err(error) => {
|
||||
stop_started_worker(started).await;
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let record =
|
||||
match store.commit_created(&allocation, manifest, active.session_id, active.segment_id)
|
||||
{
|
||||
Ok(record) => record,
|
||||
Err(_) => {
|
||||
stop_started_worker(started).await;
|
||||
let _ = store.abandon_allocation(allocation);
|
||||
return Err(StandaloneStartupError::StateStore);
|
||||
}
|
||||
};
|
||||
Ok(Self::from_started(
|
||||
started,
|
||||
store,
|
||||
worker_store,
|
||||
record,
|
||||
allocation.into_lease(),
|
||||
))
|
||||
}
|
||||
|
||||
pub async fn restore(
|
||||
state_dir: PathBuf,
|
||||
session_id: StandaloneSessionId,
|
||||
) -> Result<Self, StandaloneStartupError> {
|
||||
Self::restore_with_optional_model_client(state_dir, session_id, None).await
|
||||
}
|
||||
|
||||
pub async fn restore_with_model_client<C>(
|
||||
state_dir: PathBuf,
|
||||
session_id: StandaloneSessionId,
|
||||
model_client: C,
|
||||
) -> Result<Self, StandaloneStartupError>
|
||||
where
|
||||
C: LlmClient + 'static,
|
||||
{
|
||||
Self::restore_with_optional_model_client(
|
||||
state_dir,
|
||||
session_id,
|
||||
Some(Box::new(model_client)),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn restore_with_optional_model_client(
|
||||
state_dir: PathBuf,
|
||||
session_id: StandaloneSessionId,
|
||||
model_client: Option<Box<dyn LlmClient>>,
|
||||
) -> Result<Self, StandaloneStartupError> {
|
||||
let store =
|
||||
StandaloneSessionStore::open(state_dir).map_err(classify_store_startup_error)?;
|
||||
let record = store
|
||||
.load(session_id)
|
||||
.map_err(classify_store_startup_error)?;
|
||||
record.cwd.verify().map_err(classify_store_startup_error)?;
|
||||
let lease = store
|
||||
.acquire_lease(session_id, StaleLeasePolicy::Recover)
|
||||
.map_err(classify_store_startup_error)?;
|
||||
let (backing_store, worker_store) = backing_store(&store, session_id)?;
|
||||
let worker_name = record.worker_name.clone();
|
||||
let manifest = record.manifest.clone();
|
||||
let filesystem_authority = WorkerFilesystemAuthority::local(
|
||||
record.cwd.canonical_path.clone(),
|
||||
record.cwd.canonical_path.clone(),
|
||||
);
|
||||
let workspace_context = WorkerWorkspaceContext::local_filesystem(None);
|
||||
let runtime_base = store.runtime_dir(session_id);
|
||||
|
||||
let mut bootstrap = WorkerBootstrap::new(
|
||||
manifest,
|
||||
backing_store,
|
||||
worker::PromptCatalogSource::builtins_only(),
|
||||
workspace_context,
|
||||
filesystem_authority,
|
||||
WorkerBootstrapLayout::Direct { runtime_base },
|
||||
WorkerControllerTransport::InProcess,
|
||||
);
|
||||
if let Some(model_client) = model_client {
|
||||
bootstrap = bootstrap.with_model_client(model_client);
|
||||
}
|
||||
let prepared = bootstrap
|
||||
.prepare_restored(&worker_name)
|
||||
.await
|
||||
.map_err(classify_startup_error)?;
|
||||
let started = prepared.start().await.map_err(classify_startup_error)?;
|
||||
let active = match active_pointer(&worker_store, &worker_name) {
|
||||
Ok(active) => active,
|
||||
Err(error) => {
|
||||
stop_started_worker(started).await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let record =
|
||||
match store.update_active_pointer(&record, active.session_id, active.segment_id) {
|
||||
Ok(record) => record,
|
||||
Err(_) => {
|
||||
stop_started_worker(started).await;
|
||||
lease.retain();
|
||||
return Err(StandaloneStartupError::StateStore);
|
||||
}
|
||||
};
|
||||
Ok(Self::from_started(
|
||||
started,
|
||||
store,
|
||||
worker_store,
|
||||
record,
|
||||
lease,
|
||||
))
|
||||
}
|
||||
|
||||
fn from_started(
|
||||
started: BootstrappedWorker,
|
||||
store: StandaloneSessionStore,
|
||||
worker_store: FsWorkerStore,
|
||||
record: StandaloneSessionRecord,
|
||||
lease: StandaloneSessionLease,
|
||||
) -> Self {
|
||||
Self {
|
||||
handle: started.handle,
|
||||
shutdown: Some(started.shutdown),
|
||||
shutdown_timeout: DEFAULT_SHUTDOWN_TIMEOUT,
|
||||
store,
|
||||
worker_store,
|
||||
record,
|
||||
lease: Some(lease),
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn session_id(&self) -> StandaloneSessionId {
|
||||
self.record.session_id
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn record(&self) -> &StandaloneSessionRecord {
|
||||
&self.record
|
||||
}
|
||||
|
||||
pub async fn send(&self, method: Method) -> Result<(), StandaloneRequestError> {
|
||||
self.handle
|
||||
.send(method)
|
||||
.await
|
||||
.map_err(|_| StandaloneRequestError::WorkerUnavailable)
|
||||
}
|
||||
|
||||
pub fn subscribe(&self) -> broadcast::Receiver<Event> {
|
||||
self.handle.subscribe()
|
||||
}
|
||||
|
||||
pub fn snapshot(&self) -> Event {
|
||||
self.handle.snapshot_event()
|
||||
}
|
||||
|
||||
pub fn with_shutdown_timeout(mut self, shutdown_timeout: Duration) -> Self {
|
||||
self.shutdown_timeout = shutdown_timeout;
|
||||
self
|
||||
}
|
||||
|
||||
pub async fn shutdown(mut self) -> Result<(), StandaloneShutdownError> {
|
||||
let _ = self.handle.send(Method::Shutdown).await;
|
||||
let Some(shutdown) = self.shutdown.take() else {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::ConfirmationLost);
|
||||
};
|
||||
match tokio::time::timeout(self.shutdown_timeout, shutdown).await {
|
||||
Ok(Ok(())) => {}
|
||||
Ok(Err(_)) => {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::ConfirmationLost);
|
||||
}
|
||||
Err(_) => {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::DeadlineExceeded);
|
||||
}
|
||||
}
|
||||
let active = match active_pointer(&self.worker_store, &self.record.worker_name) {
|
||||
Ok(active) => active,
|
||||
Err(_) => {
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::StateStore);
|
||||
}
|
||||
};
|
||||
if self
|
||||
.store
|
||||
.mark_stopped(
|
||||
&self.record,
|
||||
active.session_id,
|
||||
active.segment_id,
|
||||
StandaloneShutdownReason::UserExit,
|
||||
)
|
||||
.is_err()
|
||||
{
|
||||
self.retain_lease();
|
||||
return Err(StandaloneShutdownError::StateStore);
|
||||
}
|
||||
if let Some(lease) = self.lease.take() {
|
||||
lease
|
||||
.release()
|
||||
.map_err(|_| StandaloneShutdownError::StateStore)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn retain_lease(&mut self) {
|
||||
if let Some(lease) = self.lease.take() {
|
||||
lease.retain();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn backing_store(
|
||||
store: &StandaloneSessionStore,
|
||||
id: StandaloneSessionId,
|
||||
) -> Result<(StandaloneBackingStore, FsWorkerStore), StandaloneStartupError> {
|
||||
let session_store =
|
||||
FsStore::new(store.session_log_dir(id)).map_err(|_| StandaloneStartupError::StateStore)?;
|
||||
let worker_store = FsWorkerStore::new(store.worker_metadata_dir(id))
|
||||
.map_err(|_| StandaloneStartupError::StateStore)?;
|
||||
Ok((
|
||||
CombinedStore::new(session_store, worker_store.clone()),
|
||||
worker_store,
|
||||
))
|
||||
}
|
||||
|
||||
fn active_pointer(
|
||||
worker_store: &FsWorkerStore,
|
||||
worker_name: &str,
|
||||
) -> Result<WorkerActiveSegmentRef, StandaloneStartupError> {
|
||||
worker_store
|
||||
.read_by_name(worker_name)
|
||||
.map_err(|_| StandaloneStartupError::StateStore)?
|
||||
.and_then(|metadata| metadata.active)
|
||||
.ok_or(StandaloneStartupError::StateStore)
|
||||
}
|
||||
|
||||
async fn stop_started_worker(started: BootstrappedWorker) {
|
||||
let _ = started.handle.send(Method::Shutdown).await;
|
||||
let _ = tokio::time::timeout(Duration::from_secs(2), started.shutdown).await;
|
||||
}
|
||||
|
||||
fn classify_store_startup_error(error: StandaloneStoreError) -> StandaloneStartupError {
|
||||
match error {
|
||||
StandaloneStoreError::SessionLeased(_) => StandaloneStartupError::SessionActive,
|
||||
StandaloneStoreError::LeaseLivenessUnknown(_) => {
|
||||
StandaloneStartupError::LeaseLivenessUnknown
|
||||
}
|
||||
StandaloneStoreError::CwdUnavailable(_)
|
||||
| StandaloneStoreError::CwdNotDirectory
|
||||
| StandaloneStoreError::CwdIdentityMismatch => {
|
||||
StandaloneStartupError::WorkingDirectoryUnavailable
|
||||
}
|
||||
_ => StandaloneStartupError::StateStore,
|
||||
}
|
||||
}
|
||||
|
||||
fn classify_startup_error(error: WorkerBootstrapError) -> StandaloneStartupError {
|
||||
match error {
|
||||
WorkerBootstrapError::Worker(WorkerError::Provider(_)) => {
|
||||
StandaloneStartupError::ModelProvider
|
||||
}
|
||||
WorkerBootstrapError::Worker(_) => StandaloneStartupError::WorkerConfiguration,
|
||||
WorkerBootstrapError::Controller { source, .. }
|
||||
if source.kind() == std::io::ErrorKind::Other =>
|
||||
{
|
||||
StandaloneStartupError::FeatureComposition
|
||||
}
|
||||
WorkerBootstrapError::Controller { .. } => StandaloneStartupError::Controller,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use manifest::{
|
||||
ProfileExecutionTarget, ProfileResolveOptions, ProfileResolver, ProfileSelector,
|
||||
ResolvedProfile,
|
||||
};
|
||||
use thiserror::Error;
|
||||
use worker::PromptCatalogSource;
|
||||
|
||||
/// Process launch input resolved before any Worker/session side effect occurs.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StandaloneLaunchConfig {
|
||||
pub cwd: PathBuf,
|
||||
pub state_dir: PathBuf,
|
||||
pub profile: ProfileSelector,
|
||||
pub worker_name: String,
|
||||
}
|
||||
|
||||
pub struct ResolvedStandaloneLaunch {
|
||||
pub cwd: PathBuf,
|
||||
pub state_dir: PathBuf,
|
||||
pub profile: ResolvedProfile,
|
||||
pub prompt_catalog: PromptCatalogSource,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
|
||||
pub enum StandaloneLaunchError {
|
||||
#[error("the standalone working directory is unavailable")]
|
||||
WorkingDirectoryUnavailable,
|
||||
#[error("path-based profiles are not standalone launch authority")]
|
||||
PathProfileUnsupported,
|
||||
#[error("the standalone profile could not be resolved")]
|
||||
ProfileResolutionFailed,
|
||||
}
|
||||
|
||||
impl StandaloneLaunchConfig {
|
||||
pub fn new(
|
||||
cwd: impl Into<PathBuf>,
|
||||
state_dir: impl Into<PathBuf>,
|
||||
profile: ProfileSelector,
|
||||
worker_name: impl Into<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
cwd: cwd.into(),
|
||||
state_dir: state_dir.into(),
|
||||
profile,
|
||||
worker_name: worker_name.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve only built-in/XDG profile authority and bind standalone scope
|
||||
/// to the canonical process cwd. Repository-local profile discovery is
|
||||
/// deliberately not part of this path.
|
||||
pub fn resolve(self) -> Result<ResolvedStandaloneLaunch, StandaloneLaunchError> {
|
||||
if matches!(self.profile, ProfileSelector::Path { .. }) {
|
||||
return Err(StandaloneLaunchError::PathProfileUnsupported);
|
||||
}
|
||||
let cwd = canonical_directory(&self.cwd)?;
|
||||
let profile = ProfileResolver::new()
|
||||
.with_workspace_base(&cwd)
|
||||
.resolve_for_target(
|
||||
&self.profile,
|
||||
ProfileResolveOptions {
|
||||
worker_name: Some(self.worker_name),
|
||||
},
|
||||
ProfileExecutionTarget::Standalone,
|
||||
)
|
||||
.map_err(|_| StandaloneLaunchError::ProfileResolutionFailed)?;
|
||||
|
||||
Ok(ResolvedStandaloneLaunch {
|
||||
cwd,
|
||||
state_dir: self.state_dir,
|
||||
profile,
|
||||
prompt_catalog: PromptCatalogSource::builtins_only(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn canonical_directory(path: &Path) -> Result<PathBuf, StandaloneLaunchError> {
|
||||
let path = std::fs::canonicalize(path)
|
||||
.map_err(|_| StandaloneLaunchError::WorkingDirectoryUnavailable)?;
|
||||
if !path.is_dir() {
|
||||
return Err(StandaloneLaunchError::WorkingDirectoryUnavailable);
|
||||
}
|
||||
Ok(path)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
//! In-process standalone host for one top-level Yoi Worker.
|
||||
//!
|
||||
//! The crate composes existing `worker`, `manifest`, `session-store`, and
|
||||
//! `workdir` contracts. It intentionally owns no TUI, Runtime, Workspace
|
||||
//! Server, HTTP, WebSocket, subprocess Worker, or alternative execution path.
|
||||
|
||||
pub mod host;
|
||||
pub mod launch;
|
||||
pub mod store;
|
||||
|
||||
pub use host::{
|
||||
StandaloneHost, StandaloneRequestError, StandaloneShutdownError, StandaloneStartupError,
|
||||
};
|
||||
pub use launch::{ResolvedStandaloneLaunch, StandaloneLaunchConfig, StandaloneLaunchError};
|
||||
pub use store::{
|
||||
StaleLeasePolicy, StandaloneCwdIdentity, StandaloneListScope, StandaloneSessionId,
|
||||
StandaloneSessionRecord, StandaloneSessionStatus, StandaloneSessionStore,
|
||||
StandaloneShutdownReason, StandaloneStoreError,
|
||||
};
|
||||
@@ -0,0 +1,778 @@
|
||||
use std::fmt;
|
||||
use std::fs::{self, File, OpenOptions};
|
||||
use std::io::{self, Write};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::str::FromStr;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use fs4::fs_std::FileExt;
|
||||
use manifest::WorkerManifest;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use session_store::{SegmentId, SessionId};
|
||||
use thiserror::Error;
|
||||
use uuid::Uuid;
|
||||
|
||||
const RECORD_FILE: &str = "record.json";
|
||||
const COMMIT_MARKER: &str = "commit.pending";
|
||||
const LEASE_FILE: &str = "lease.json";
|
||||
const LEASE_LOCK_FILE: &str = "lease.lock";
|
||||
const SESSION_DIR: &str = "session";
|
||||
const WORKER_DIR: &str = "worker";
|
||||
const SCHEMA_VERSION: u32 = 1;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct StandaloneSessionId(Uuid);
|
||||
|
||||
impl StandaloneSessionId {
|
||||
#[must_use]
|
||||
pub fn new() -> Self {
|
||||
Self(Uuid::now_v7())
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn short(self) -> String {
|
||||
let simple = self.0.simple().to_string();
|
||||
simple[simple.len() - 12..].to_string()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for StandaloneSessionId {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for StandaloneSessionId {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
self.0.fmt(formatter)
|
||||
}
|
||||
}
|
||||
|
||||
impl FromStr for StandaloneSessionId {
|
||||
type Err = uuid::Error;
|
||||
|
||||
fn from_str(value: &str) -> Result<Self, Self::Err> {
|
||||
Uuid::parse_str(value).map(Self)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct StandaloneCwdIdentity {
|
||||
pub canonical_path: PathBuf,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub device: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub inode: Option<u64>,
|
||||
}
|
||||
|
||||
impl StandaloneCwdIdentity {
|
||||
pub fn capture(path: impl AsRef<Path>) -> Result<Self, StandaloneStoreError> {
|
||||
let canonical_path =
|
||||
fs::canonicalize(path).map_err(StandaloneStoreError::CwdUnavailable)?;
|
||||
let metadata =
|
||||
fs::metadata(&canonical_path).map_err(StandaloneStoreError::CwdUnavailable)?;
|
||||
if !metadata.is_dir() {
|
||||
return Err(StandaloneStoreError::CwdNotDirectory);
|
||||
}
|
||||
#[cfg(unix)]
|
||||
let (device, inode) = {
|
||||
use std::os::unix::fs::MetadataExt;
|
||||
(Some(metadata.dev()), Some(metadata.ino()))
|
||||
};
|
||||
#[cfg(not(unix))]
|
||||
let (device, inode) = (None, None);
|
||||
Ok(Self {
|
||||
canonical_path,
|
||||
device,
|
||||
inode,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn verify(&self) -> Result<PathBuf, StandaloneStoreError> {
|
||||
let current = Self::capture(&self.canonical_path)?;
|
||||
if current != *self {
|
||||
return Err(StandaloneStoreError::CwdIdentityMismatch);
|
||||
}
|
||||
Ok(current.canonical_path)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum StandaloneSessionStatus {
|
||||
Active,
|
||||
Stopped,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum StandaloneShutdownReason {
|
||||
UserExit,
|
||||
StartupFailed,
|
||||
ControllerError,
|
||||
ProcessInterrupted,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct StandaloneSessionRecord {
|
||||
pub schema_version: u32,
|
||||
pub revision: u64,
|
||||
pub session_id: StandaloneSessionId,
|
||||
pub worker_name: String,
|
||||
pub cwd: StandaloneCwdIdentity,
|
||||
pub manifest: WorkerManifest,
|
||||
pub active_session_id: SessionId,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub active_segment_id: Option<SegmentId>,
|
||||
pub status: StandaloneSessionStatus,
|
||||
pub created_at_unix_ms: u64,
|
||||
pub updated_at_unix_ms: u64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub shutdown_reason: Option<StandaloneShutdownReason>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum StandaloneListScope {
|
||||
CurrentCwd,
|
||||
All,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum StaleLeasePolicy {
|
||||
Reject,
|
||||
Recover,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StandaloneSessionStore {
|
||||
root: PathBuf,
|
||||
}
|
||||
|
||||
impl StandaloneSessionStore {
|
||||
pub fn open(root: impl Into<PathBuf>) -> Result<Self, StandaloneStoreError> {
|
||||
let root = root.into();
|
||||
fs::create_dir_all(&root).map_err(StandaloneStoreError::Io)?;
|
||||
if !fs::metadata(&root)
|
||||
.map_err(StandaloneStoreError::Io)?
|
||||
.is_dir()
|
||||
{
|
||||
return Err(StandaloneStoreError::NotDirectory);
|
||||
}
|
||||
Ok(Self { root })
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn root(&self) -> &Path {
|
||||
&self.root
|
||||
}
|
||||
|
||||
pub fn allocate(
|
||||
&self,
|
||||
cwd: impl AsRef<Path>,
|
||||
policy: StaleLeasePolicy,
|
||||
) -> Result<StandaloneSessionAllocation, StandaloneStoreError> {
|
||||
let id = StandaloneSessionId::new();
|
||||
let cwd = StandaloneCwdIdentity::capture(cwd)?;
|
||||
let dir = self.session_dir(id);
|
||||
fs::create_dir(&dir).map_err(StandaloneStoreError::Io)?;
|
||||
fs::create_dir(dir.join(SESSION_DIR)).map_err(StandaloneStoreError::Io)?;
|
||||
fs::create_dir(dir.join(WORKER_DIR)).map_err(StandaloneStoreError::Io)?;
|
||||
let lease = self.acquire_lease(id, policy)?;
|
||||
Ok(StandaloneSessionAllocation { id, cwd, lease })
|
||||
}
|
||||
|
||||
pub fn commit_created(
|
||||
&self,
|
||||
allocation: &StandaloneSessionAllocation,
|
||||
manifest: WorkerManifest,
|
||||
active_session_id: SessionId,
|
||||
active_segment_id: Option<SegmentId>,
|
||||
) -> Result<StandaloneSessionRecord, StandaloneStoreError> {
|
||||
let now = now_unix_ms()?;
|
||||
let record = StandaloneSessionRecord {
|
||||
schema_version: SCHEMA_VERSION,
|
||||
revision: 1,
|
||||
session_id: allocation.id,
|
||||
worker_name: manifest.worker.name.clone(),
|
||||
cwd: allocation.cwd.clone(),
|
||||
manifest,
|
||||
active_session_id,
|
||||
active_segment_id,
|
||||
status: StandaloneSessionStatus::Active,
|
||||
created_at_unix_ms: now,
|
||||
updated_at_unix_ms: now,
|
||||
shutdown_reason: None,
|
||||
};
|
||||
self.commit_record(None, &record)?;
|
||||
Ok(record)
|
||||
}
|
||||
|
||||
pub fn load(
|
||||
&self,
|
||||
id: StandaloneSessionId,
|
||||
) -> Result<StandaloneSessionRecord, StandaloneStoreError> {
|
||||
let dir = self.session_dir(id);
|
||||
if dir.join(COMMIT_MARKER).exists() {
|
||||
return Err(StandaloneStoreError::IncompleteCommit(id));
|
||||
}
|
||||
let bytes = fs::read(dir.join(RECORD_FILE)).map_err(|error| {
|
||||
if error.kind() == io::ErrorKind::NotFound {
|
||||
StandaloneStoreError::SessionNotFound(id)
|
||||
} else {
|
||||
StandaloneStoreError::Io(error)
|
||||
}
|
||||
})?;
|
||||
let record: StandaloneSessionRecord = serde_json::from_slice(&bytes)
|
||||
.map_err(|source| StandaloneStoreError::CorruptRecord { id, source })?;
|
||||
if record.schema_version > SCHEMA_VERSION {
|
||||
return Err(StandaloneStoreError::NewerSchema {
|
||||
id,
|
||||
found: record.schema_version,
|
||||
supported: SCHEMA_VERSION,
|
||||
});
|
||||
}
|
||||
if record.schema_version != SCHEMA_VERSION || record.session_id != id {
|
||||
return Err(StandaloneStoreError::InvalidRecord(id));
|
||||
}
|
||||
Ok(record)
|
||||
}
|
||||
|
||||
pub fn list(
|
||||
&self,
|
||||
cwd: impl AsRef<Path>,
|
||||
scope: StandaloneListScope,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StandaloneSessionRecord>, StandaloneStoreError> {
|
||||
let current_cwd = (scope == StandaloneListScope::CurrentCwd)
|
||||
.then(|| StandaloneCwdIdentity::capture(cwd))
|
||||
.transpose()?;
|
||||
let mut records = Vec::new();
|
||||
for entry in fs::read_dir(&self.root).map_err(StandaloneStoreError::Io)? {
|
||||
let entry = entry.map_err(StandaloneStoreError::Io)?;
|
||||
if !entry
|
||||
.file_type()
|
||||
.map_err(StandaloneStoreError::Io)?
|
||||
.is_dir()
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let Ok(id) = entry.file_name().to_string_lossy().parse() else {
|
||||
continue;
|
||||
};
|
||||
let record = self.load(id)?;
|
||||
if current_cwd.as_ref().is_none_or(|cwd| &record.cwd == cwd) {
|
||||
records.push(record);
|
||||
}
|
||||
}
|
||||
records.sort_by(|left, right| {
|
||||
right
|
||||
.updated_at_unix_ms
|
||||
.cmp(&left.updated_at_unix_ms)
|
||||
.then_with(|| {
|
||||
right
|
||||
.session_id
|
||||
.to_string()
|
||||
.cmp(&left.session_id.to_string())
|
||||
})
|
||||
});
|
||||
records.truncate(limit);
|
||||
Ok(records)
|
||||
}
|
||||
|
||||
pub fn acquire_lease(
|
||||
&self,
|
||||
id: StandaloneSessionId,
|
||||
policy: StaleLeasePolicy,
|
||||
) -> Result<StandaloneSessionLease, StandaloneStoreError> {
|
||||
let dir = self.session_dir(id);
|
||||
let path = dir.join(LEASE_FILE);
|
||||
let _guard = LeaseMutationGuard::acquire(&dir)?;
|
||||
let lease = LeaseRecord::current()?;
|
||||
loop {
|
||||
match OpenOptions::new().write(true).create_new(true).open(&path) {
|
||||
Ok(mut file) => {
|
||||
serde_json::to_writer(&mut file, &lease).map_err(StandaloneStoreError::Json)?;
|
||||
file.write_all(b"\n").map_err(StandaloneStoreError::Io)?;
|
||||
file.sync_all().map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
return Ok(StandaloneSessionLease {
|
||||
path,
|
||||
lease_id: lease.lease_id,
|
||||
released: false,
|
||||
});
|
||||
}
|
||||
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
|
||||
let existing = read_lease(&path, id)?;
|
||||
match existing.liveness() {
|
||||
LeaseLiveness::Live => {
|
||||
return Err(StandaloneStoreError::SessionLeased(id));
|
||||
}
|
||||
LeaseLiveness::Unknown => {
|
||||
return Err(StandaloneStoreError::LeaseLivenessUnknown(id));
|
||||
}
|
||||
LeaseLiveness::Stale => {}
|
||||
}
|
||||
if policy == StaleLeasePolicy::Reject {
|
||||
return Err(StandaloneStoreError::StaleLease(id));
|
||||
}
|
||||
fs::remove_file(&path).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
}
|
||||
Err(error) => return Err(StandaloneStoreError::Io(error)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn update_active_pointer(
|
||||
&self,
|
||||
record: &StandaloneSessionRecord,
|
||||
active_session_id: SessionId,
|
||||
active_segment_id: Option<SegmentId>,
|
||||
) -> Result<StandaloneSessionRecord, StandaloneStoreError> {
|
||||
let mut next = record.clone();
|
||||
next.revision = next.revision.saturating_add(1);
|
||||
next.updated_at_unix_ms = now_unix_ms()?;
|
||||
next.active_session_id = active_session_id;
|
||||
next.active_segment_id = active_segment_id;
|
||||
next.status = StandaloneSessionStatus::Active;
|
||||
next.shutdown_reason = None;
|
||||
self.commit_record(Some(record.revision), &next)?;
|
||||
Ok(next)
|
||||
}
|
||||
|
||||
pub fn mark_stopped(
|
||||
&self,
|
||||
record: &StandaloneSessionRecord,
|
||||
active_session_id: SessionId,
|
||||
active_segment_id: Option<SegmentId>,
|
||||
reason: StandaloneShutdownReason,
|
||||
) -> Result<StandaloneSessionRecord, StandaloneStoreError> {
|
||||
let mut next = record.clone();
|
||||
next.revision = next.revision.saturating_add(1);
|
||||
next.updated_at_unix_ms = now_unix_ms()?;
|
||||
next.active_session_id = active_session_id;
|
||||
next.active_segment_id = active_segment_id;
|
||||
next.status = StandaloneSessionStatus::Stopped;
|
||||
next.shutdown_reason = Some(reason);
|
||||
self.commit_record(Some(record.revision), &next)?;
|
||||
Ok(next)
|
||||
}
|
||||
|
||||
pub fn delete(&self, id: StandaloneSessionId) -> Result<(), StandaloneStoreError> {
|
||||
let record = self.load(id)?;
|
||||
if record.status != StandaloneSessionStatus::Stopped {
|
||||
return Err(StandaloneStoreError::DeleteActive(id));
|
||||
}
|
||||
let session_dir = self.session_dir(id);
|
||||
let _guard = LeaseMutationGuard::acquire(&session_dir)?;
|
||||
let lease_path = session_dir.join(LEASE_FILE);
|
||||
if lease_path.exists() {
|
||||
let lease = read_lease(&lease_path, id)?;
|
||||
return Err(match lease.liveness() {
|
||||
LeaseLiveness::Live => StandaloneStoreError::SessionLeased(id),
|
||||
LeaseLiveness::Stale => StandaloneStoreError::StaleLease(id),
|
||||
LeaseLiveness::Unknown => StandaloneStoreError::LeaseLivenessUnknown(id),
|
||||
});
|
||||
}
|
||||
fs::remove_dir_all(self.session_dir(id)).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&self.root)
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn session_log_dir(&self, id: StandaloneSessionId) -> PathBuf {
|
||||
self.session_dir(id).join(SESSION_DIR)
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn worker_metadata_dir(&self, id: StandaloneSessionId) -> PathBuf {
|
||||
self.session_dir(id).join(WORKER_DIR)
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub(crate) fn runtime_dir(&self, id: StandaloneSessionId) -> PathBuf {
|
||||
self.session_dir(id).join("runtime")
|
||||
}
|
||||
|
||||
pub(crate) fn abandon_allocation(
|
||||
&self,
|
||||
allocation: StandaloneSessionAllocation,
|
||||
) -> Result<(), StandaloneStoreError> {
|
||||
let id = allocation.id;
|
||||
allocation.lease.release()?;
|
||||
fs::remove_dir_all(self.session_dir(id)).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&self.root)
|
||||
}
|
||||
|
||||
fn commit_record(
|
||||
&self,
|
||||
expected_revision: Option<u64>,
|
||||
next: &StandaloneSessionRecord,
|
||||
) -> Result<(), StandaloneStoreError> {
|
||||
let dir = self.session_dir(next.session_id);
|
||||
let marker = dir.join(COMMIT_MARKER);
|
||||
let mut marker_file = OpenOptions::new()
|
||||
.write(true)
|
||||
.create_new(true)
|
||||
.open(&marker)
|
||||
.map_err(|error| {
|
||||
if error.kind() == io::ErrorKind::AlreadyExists {
|
||||
StandaloneStoreError::IncompleteCommit(next.session_id)
|
||||
} else {
|
||||
StandaloneStoreError::Io(error)
|
||||
}
|
||||
})?;
|
||||
writeln!(marker_file, "{}", next.revision).map_err(StandaloneStoreError::Io)?;
|
||||
marker_file.sync_all().map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
|
||||
if let Some(expected) = expected_revision {
|
||||
let current = self.load_record_while_committing(next.session_id)?;
|
||||
if current.revision != expected {
|
||||
let _ = fs::remove_file(&marker);
|
||||
return Err(StandaloneStoreError::RevisionConflict {
|
||||
id: next.session_id,
|
||||
expected,
|
||||
found: current.revision,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let temporary = dir.join(format!("record.{}.tmp", Uuid::now_v7()));
|
||||
let result = (|| {
|
||||
let mut file = OpenOptions::new()
|
||||
.write(true)
|
||||
.create_new(true)
|
||||
.open(&temporary)
|
||||
.map_err(StandaloneStoreError::Io)?;
|
||||
serde_json::to_writer_pretty(&mut file, next).map_err(StandaloneStoreError::Json)?;
|
||||
file.write_all(b"\n").map_err(StandaloneStoreError::Io)?;
|
||||
file.sync_all().map_err(StandaloneStoreError::Io)?;
|
||||
fs::rename(&temporary, dir.join(RECORD_FILE)).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)?;
|
||||
fs::remove_file(&marker).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(&dir)
|
||||
})();
|
||||
if result.is_err() {
|
||||
let _ = fs::remove_file(&temporary);
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn load_record_while_committing(
|
||||
&self,
|
||||
id: StandaloneSessionId,
|
||||
) -> Result<StandaloneSessionRecord, StandaloneStoreError> {
|
||||
let bytes =
|
||||
fs::read(self.session_dir(id).join(RECORD_FILE)).map_err(StandaloneStoreError::Io)?;
|
||||
serde_json::from_slice(&bytes)
|
||||
.map_err(|source| StandaloneStoreError::CorruptRecord { id, source })
|
||||
}
|
||||
|
||||
fn session_dir(&self, id: StandaloneSessionId) -> PathBuf {
|
||||
self.root.join(id.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct StandaloneSessionAllocation {
|
||||
id: StandaloneSessionId,
|
||||
cwd: StandaloneCwdIdentity,
|
||||
lease: StandaloneSessionLease,
|
||||
}
|
||||
|
||||
impl StandaloneSessionAllocation {
|
||||
#[must_use]
|
||||
pub fn id(&self) -> StandaloneSessionId {
|
||||
self.id
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn cwd(&self) -> &StandaloneCwdIdentity {
|
||||
&self.cwd
|
||||
}
|
||||
|
||||
pub fn into_lease(self) -> StandaloneSessionLease {
|
||||
self.lease
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct StandaloneSessionLease {
|
||||
path: PathBuf,
|
||||
lease_id: Uuid,
|
||||
released: bool,
|
||||
}
|
||||
|
||||
impl StandaloneSessionLease {
|
||||
pub fn release(mut self) -> Result<(), StandaloneStoreError> {
|
||||
self.release_inner()
|
||||
}
|
||||
|
||||
pub(crate) fn retain(mut self) {
|
||||
self.released = true;
|
||||
}
|
||||
|
||||
fn release_inner(&mut self) -> Result<(), StandaloneStoreError> {
|
||||
if self.released {
|
||||
return Ok(());
|
||||
}
|
||||
if self.path.exists() {
|
||||
let parent = self.path.parent().expect("lease parent");
|
||||
let _guard = LeaseMutationGuard::acquire(parent)?;
|
||||
let bytes = fs::read(&self.path).map_err(StandaloneStoreError::Io)?;
|
||||
let current: LeaseRecord =
|
||||
serde_json::from_slice(&bytes).map_err(StandaloneStoreError::Json)?;
|
||||
if current.lease_id != self.lease_id {
|
||||
return Err(StandaloneStoreError::LeaseOwnershipLost);
|
||||
}
|
||||
fs::remove_file(&self.path).map_err(StandaloneStoreError::Io)?;
|
||||
sync_directory(self.path.parent().expect("lease parent"))?;
|
||||
}
|
||||
self.released = true;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for StandaloneSessionLease {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.release_inner();
|
||||
}
|
||||
}
|
||||
|
||||
struct LeaseMutationGuard {
|
||||
file: File,
|
||||
}
|
||||
|
||||
impl LeaseMutationGuard {
|
||||
fn acquire(dir: &Path) -> Result<Self, StandaloneStoreError> {
|
||||
let file = OpenOptions::new()
|
||||
.read(true)
|
||||
.write(true)
|
||||
.create(true)
|
||||
.truncate(false)
|
||||
.open(dir.join(LEASE_LOCK_FILE))
|
||||
.map_err(StandaloneStoreError::Io)?;
|
||||
file.lock_exclusive().map_err(StandaloneStoreError::Io)?;
|
||||
Ok(Self { file })
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for LeaseMutationGuard {
|
||||
fn drop(&mut self) {
|
||||
let _ = FileExt::unlock(&self.file);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
struct LeaseRecord {
|
||||
lease_id: Uuid,
|
||||
pid: u32,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
process_start_marker: Option<u64>,
|
||||
acquired_at_unix_ms: u64,
|
||||
}
|
||||
|
||||
impl LeaseRecord {
|
||||
fn current() -> Result<Self, StandaloneStoreError> {
|
||||
Ok(Self {
|
||||
lease_id: Uuid::now_v7(),
|
||||
pid: std::process::id(),
|
||||
process_start_marker: match observe_process(std::process::id()) {
|
||||
ProcessObservation::Running { start_marker } => Some(start_marker),
|
||||
ProcessObservation::Missing | ProcessObservation::Unobservable => None,
|
||||
},
|
||||
acquired_at_unix_ms: now_unix_ms()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn liveness(&self) -> LeaseLiveness {
|
||||
classify_lease_liveness(self.process_start_marker, observe_process(self.pid))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum LeaseLiveness {
|
||||
Live,
|
||||
Stale,
|
||||
Unknown,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum ProcessObservation {
|
||||
Running { start_marker: u64 },
|
||||
Missing,
|
||||
Unobservable,
|
||||
}
|
||||
|
||||
fn classify_lease_liveness(
|
||||
recorded_start_marker: Option<u64>,
|
||||
observation: ProcessObservation,
|
||||
) -> LeaseLiveness {
|
||||
match (recorded_start_marker, observation) {
|
||||
(Some(recorded), ProcessObservation::Running { start_marker })
|
||||
if recorded == start_marker =>
|
||||
{
|
||||
LeaseLiveness::Live
|
||||
}
|
||||
(Some(_), ProcessObservation::Running { .. }) | (_, ProcessObservation::Missing) => {
|
||||
LeaseLiveness::Stale
|
||||
}
|
||||
(None, ProcessObservation::Running { .. }) | (_, ProcessObservation::Unobservable) => {
|
||||
LeaseLiveness::Unknown
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn read_lease(path: &Path, id: StandaloneSessionId) -> Result<LeaseRecord, StandaloneStoreError> {
|
||||
let bytes = fs::read(path).map_err(StandaloneStoreError::Io)?;
|
||||
serde_json::from_slice(&bytes)
|
||||
.map_err(|source| StandaloneStoreError::CorruptLease { id, source })
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn observe_process(pid: u32) -> ProcessObservation {
|
||||
let stat = match fs::read_to_string(format!("/proc/{pid}/stat")) {
|
||||
Ok(stat) => stat,
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => {
|
||||
return if pid != std::process::id() && linux_proc_is_observable() {
|
||||
ProcessObservation::Missing
|
||||
} else {
|
||||
ProcessObservation::Unobservable
|
||||
};
|
||||
}
|
||||
Err(_) => return ProcessObservation::Unobservable,
|
||||
};
|
||||
parse_linux_process_start_marker(&stat)
|
||||
.map(|start_marker| ProcessObservation::Running { start_marker })
|
||||
.unwrap_or(ProcessObservation::Unobservable)
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn linux_proc_is_observable() -> bool {
|
||||
fs::read_to_string("/proc/self/stat")
|
||||
.ok()
|
||||
.and_then(|stat| parse_linux_process_start_marker(&stat))
|
||||
.is_some()
|
||||
}
|
||||
|
||||
#[cfg(target_os = "linux")]
|
||||
fn parse_linux_process_start_marker(stat: &str) -> Option<u64> {
|
||||
let (_, tail) = stat.rsplit_once(") ")?;
|
||||
tail.split_whitespace().nth(19)?.parse().ok()
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "linux"))]
|
||||
fn observe_process(pid: u32) -> ProcessObservation {
|
||||
if pid == std::process::id() {
|
||||
ProcessObservation::Running { start_marker: 0 }
|
||||
} else {
|
||||
ProcessObservation::Unobservable
|
||||
}
|
||||
}
|
||||
|
||||
fn now_unix_ms() -> Result<u64, StandaloneStoreError> {
|
||||
let duration = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map_err(|_| StandaloneStoreError::Clock)?;
|
||||
u64::try_from(duration.as_millis()).map_err(|_| StandaloneStoreError::Clock)
|
||||
}
|
||||
|
||||
fn sync_directory(path: &Path) -> Result<(), StandaloneStoreError> {
|
||||
File::open(path)
|
||||
.and_then(|file| file.sync_all())
|
||||
.map_err(StandaloneStoreError::Io)
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum StandaloneStoreError {
|
||||
#[error("standalone state path is not a directory")]
|
||||
NotDirectory,
|
||||
#[error("standalone cwd is unavailable")]
|
||||
CwdUnavailable(#[source] io::Error),
|
||||
#[error("standalone cwd is not a directory")]
|
||||
CwdNotDirectory,
|
||||
#[error("standalone cwd identity no longer matches the persisted session")]
|
||||
CwdIdentityMismatch,
|
||||
#[error("standalone session {0} was not found")]
|
||||
SessionNotFound(StandaloneSessionId),
|
||||
#[error("standalone session {0} has an incomplete metadata commit")]
|
||||
IncompleteCommit(StandaloneSessionId),
|
||||
#[error("standalone session {0} has invalid metadata")]
|
||||
InvalidRecord(StandaloneSessionId),
|
||||
#[error("standalone session {id} metadata is corrupt")]
|
||||
CorruptRecord {
|
||||
id: StandaloneSessionId,
|
||||
#[source]
|
||||
source: serde_json::Error,
|
||||
},
|
||||
#[error("standalone session {id} lease is corrupt")]
|
||||
CorruptLease {
|
||||
id: StandaloneSessionId,
|
||||
#[source]
|
||||
source: serde_json::Error,
|
||||
},
|
||||
#[error("standalone session {id} uses schema {found}, newer than supported schema {supported}")]
|
||||
NewerSchema {
|
||||
id: StandaloneSessionId,
|
||||
found: u32,
|
||||
supported: u32,
|
||||
},
|
||||
#[error("standalone session {0} is already active")]
|
||||
SessionLeased(StandaloneSessionId),
|
||||
#[error("standalone session {0} lease liveness cannot be proven; recovery is rejected")]
|
||||
LeaseLivenessUnknown(StandaloneSessionId),
|
||||
#[error("standalone session {0} has a stale lease; explicit recovery is required")]
|
||||
StaleLease(StandaloneSessionId),
|
||||
#[error("standalone session lease ownership changed")]
|
||||
LeaseOwnershipLost,
|
||||
#[error("standalone session {0} must be stopped before deletion")]
|
||||
DeleteActive(StandaloneSessionId),
|
||||
#[error(
|
||||
"standalone session {id} metadata revision changed (expected {expected}, found {found})"
|
||||
)]
|
||||
RevisionConflict {
|
||||
id: StandaloneSessionId,
|
||||
expected: u64,
|
||||
found: u64,
|
||||
},
|
||||
#[error("system clock is before the Unix epoch or out of range")]
|
||||
Clock,
|
||||
#[error("standalone metadata serialization failed")]
|
||||
Json(#[source] serde_json::Error),
|
||||
#[error("standalone state I/O failed")]
|
||||
Io(#[source] io::Error),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{LeaseLiveness, ProcessObservation, classify_lease_liveness};
|
||||
|
||||
#[test]
|
||||
fn lease_liveness_requires_positive_live_or_stale_evidence() {
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Running { start_marker: 41 }),
|
||||
LeaseLiveness::Live
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Running { start_marker: 42 }),
|
||||
LeaseLiveness::Stale
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Missing),
|
||||
LeaseLiveness::Stale
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(None, ProcessObservation::Running { start_marker: 41 }),
|
||||
LeaseLiveness::Unknown
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(Some(41), ProcessObservation::Unobservable),
|
||||
LeaseLiveness::Unknown
|
||||
);
|
||||
assert_eq!(
|
||||
classify_lease_liveness(None, ProcessObservation::Unobservable),
|
||||
LeaseLiveness::Unknown
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,516 @@
|
||||
use std::collections::VecDeque;
|
||||
use std::pin::Pin;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
use agen::llm_client::client::LlmClient;
|
||||
use agen::llm_client::error::ClientError;
|
||||
use agen::llm_client::event::{Event as LlmEvent, StopReason};
|
||||
use agen::llm_client::types::Request;
|
||||
use async_trait::async_trait;
|
||||
use futures::{Stream, stream};
|
||||
use protocol::{Event, Method};
|
||||
use standalone::{
|
||||
StaleLeasePolicy, StandaloneHost, StandaloneLaunchConfig, StandaloneListScope,
|
||||
StandaloneSessionStatus, StandaloneSessionStore, StandaloneStartupError, StandaloneStoreError,
|
||||
};
|
||||
use uuid::Uuid;
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ScriptedClient {
|
||||
responses: Arc<Mutex<VecDeque<Vec<LlmEvent>>>>,
|
||||
requests: Arc<Mutex<Vec<Request>>>,
|
||||
}
|
||||
|
||||
impl ScriptedClient {
|
||||
fn new(responses: Vec<Vec<LlmEvent>>) -> Self {
|
||||
Self {
|
||||
responses: Arc::new(Mutex::new(responses.into())),
|
||||
requests: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn requests(&self) -> Vec<Request> {
|
||||
self.requests.lock().expect("requests lock").clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmClient for ScriptedClient {
|
||||
async fn stream(
|
||||
&self,
|
||||
request: Request,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = Result<LlmEvent, ClientError>> + Send>>, ClientError>
|
||||
{
|
||||
self.requests.lock().expect("requests lock").push(request);
|
||||
let response = self
|
||||
.responses
|
||||
.lock()
|
||||
.expect("responses lock")
|
||||
.pop_front()
|
||||
.expect("scripted response");
|
||||
Ok(Box::pin(stream::iter(response.into_iter().map(Ok))))
|
||||
}
|
||||
|
||||
fn clone_boxed(&self) -> Box<dyn LlmClient> {
|
||||
Box::new(self.clone())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn in_process_host_runs_text_and_read_tool_then_shuts_down() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
std::fs::write(temp.path().join("probe.txt"), "standalone tool evidence\n")
|
||||
.expect("write probe");
|
||||
let worker_name = format!("standalone-{}", Uuid::now_v7());
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
temp.path().join("state"),
|
||||
manifest::ProfileSelector::Default,
|
||||
&worker_name,
|
||||
)
|
||||
.resolve()
|
||||
.expect("resolve standalone profile");
|
||||
|
||||
let client = ScriptedClient::new(vec![
|
||||
vec![
|
||||
LlmEvent::tool_use_start(0, "read-1", "Read"),
|
||||
LlmEvent::tool_input_delta(0, r#"{"file_path":"probe.txt"}"#),
|
||||
LlmEvent::tool_use_stop(0),
|
||||
],
|
||||
vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "standalone response"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
],
|
||||
]);
|
||||
let inspection = client.clone();
|
||||
let host = StandaloneHost::start_with_model_client(launch, client)
|
||||
.await
|
||||
.expect("start in-process host");
|
||||
let mut events = host.subscribe();
|
||||
|
||||
host.send(Method::run_text("read the probe"))
|
||||
.await
|
||||
.expect("submit input");
|
||||
|
||||
tokio::time::timeout(Duration::from_secs(30), async {
|
||||
let mut saw_text = false;
|
||||
let mut saw_tool_result = false;
|
||||
loop {
|
||||
match events.recv().await.expect("worker event") {
|
||||
Event::TextDelta { text } if text.contains("standalone response") => {
|
||||
saw_text = true;
|
||||
}
|
||||
Event::ToolResult { .. } => {
|
||||
saw_tool_result = true;
|
||||
}
|
||||
Event::RunEnd { .. } => {
|
||||
assert!(saw_text, "stream must expose the model text delta");
|
||||
assert!(saw_tool_result, "stream must expose the tool result");
|
||||
break;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("run completed");
|
||||
|
||||
let requests = inspection.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
let tool_names = requests[0]
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert!(tool_names.contains(&"Read"));
|
||||
assert!(tool_names.contains(&"TaskCreate"));
|
||||
assert!(tool_names.contains(&"SubWorkerSpawn"));
|
||||
assert!(format!("{:?}", requests[1].items).contains("standalone tool evidence"));
|
||||
assert!(
|
||||
!temp
|
||||
.path()
|
||||
.join("state/runtime")
|
||||
.join(&worker_name)
|
||||
.join("worker.sock")
|
||||
.exists()
|
||||
);
|
||||
|
||||
host.shutdown().await.expect("graceful shutdown");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn state_store_failure_is_redacted_and_starts_no_controller() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let state_path = temp.path().join("state-file-with-secret-name");
|
||||
std::fs::write(&state_path, "not a directory").expect("write blocking file");
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_path,
|
||||
manifest::ProfileSelector::Default,
|
||||
format!("standalone-failure-{}", Uuid::now_v7()),
|
||||
)
|
||||
.resolve()
|
||||
.expect("resolve launch");
|
||||
let client = ScriptedClient::new(Vec::new());
|
||||
|
||||
let error = StandaloneHost::start_with_model_client(launch, client)
|
||||
.await
|
||||
.err()
|
||||
.expect("state store startup rejected");
|
||||
assert_eq!(error, standalone::StandaloneStartupError::StateStore);
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"the standalone state store could not be opened or validated"
|
||||
);
|
||||
assert!(!error.to_string().contains("secret-name"));
|
||||
assert!(
|
||||
!temp
|
||||
.path()
|
||||
.join("state-file-with-secret-name/runtime")
|
||||
.exists()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standalone_crate_has_no_tui_runtime_or_workspace_server_dependency() {
|
||||
let manifest = include_str!("../Cargo.toml");
|
||||
let dependencies = manifest
|
||||
.split("[dependencies]")
|
||||
.nth(1)
|
||||
.expect("dependencies section")
|
||||
.split("[dev-dependencies]")
|
||||
.next()
|
||||
.expect("dependency body");
|
||||
for forbidden in ["tui", "worker-runtime", "yoi-workspace-server"] {
|
||||
assert!(
|
||||
!dependencies.lines().any(|line| {
|
||||
line.split_once('=')
|
||||
.is_some_and(|(name, _)| name.trim() == forbidden)
|
||||
}),
|
||||
"standalone must not depend on {forbidden}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn launch_rejects_path_profile_before_worker_startup() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let error = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
temp.path().join("state"),
|
||||
manifest::ProfileSelector::Path {
|
||||
path: temp.path().join("profile.dcdl"),
|
||||
},
|
||||
"standalone-path-profile",
|
||||
)
|
||||
.resolve()
|
||||
.err()
|
||||
.expect("path profile rejected");
|
||||
assert_eq!(
|
||||
error,
|
||||
standalone::StandaloneLaunchError::PathProfileUnsupported
|
||||
);
|
||||
}
|
||||
|
||||
type TestResult = Result<(), Box<dyn std::error::Error>>;
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_preserves_history_tasks_notifications_and_cwd_scope() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let cwd = temp.path().join("project");
|
||||
let state_dir = temp.path().join("client").join("standalone-sessions");
|
||||
std::fs::create_dir_all(&cwd)?;
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
&cwd,
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"display-name-is-not-session-identity",
|
||||
)
|
||||
.resolve()?;
|
||||
let first_client = ScriptedClient::new(vec![
|
||||
vec![
|
||||
LlmEvent::tool_use_start(0, "task-1", "TaskCreate"),
|
||||
LlmEvent::tool_input_delta(
|
||||
0,
|
||||
r#"{"subject":"persisted task","description":"survives restore"}"#,
|
||||
),
|
||||
LlmEvent::tool_use_stop(0),
|
||||
],
|
||||
vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "first answer"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
],
|
||||
vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "notification acknowledged"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
],
|
||||
]);
|
||||
let host = StandaloneHost::start_with_model_client(launch, first_client).await?;
|
||||
let session_id = host.session_id();
|
||||
let mut events = host.subscribe();
|
||||
host.send(Method::run_text("first request")).await?;
|
||||
wait_for_run_end(&mut events).await?;
|
||||
host.send(Method::Notify {
|
||||
message: "persisted notification".to_string(),
|
||||
auto_run: true,
|
||||
})
|
||||
.await?;
|
||||
wait_for_run_end(&mut events).await?;
|
||||
host.shutdown().await?;
|
||||
|
||||
let store = StandaloneSessionStore::open(&state_dir)?;
|
||||
let current = store.list(&cwd, StandaloneListScope::CurrentCwd, 100)?;
|
||||
assert_eq!(current.len(), 1);
|
||||
assert_eq!(current[0].session_id, session_id);
|
||||
assert_eq!(current[0].status, StandaloneSessionStatus::Stopped);
|
||||
let other_cwd = temp.path().join("other");
|
||||
std::fs::create_dir(&other_cwd)?;
|
||||
assert!(
|
||||
store
|
||||
.list(&other_cwd, StandaloneListScope::CurrentCwd, 100)?
|
||||
.is_empty()
|
||||
);
|
||||
assert_eq!(
|
||||
store.list(&other_cwd, StandaloneListScope::All, 100)?.len(),
|
||||
1
|
||||
);
|
||||
|
||||
let second_client = ScriptedClient::new(vec![vec![
|
||||
LlmEvent::text_block_start(0),
|
||||
LlmEvent::text_delta(0, "second answer"),
|
||||
LlmEvent::text_block_stop(0, Some(StopReason::EndTurn)),
|
||||
]]);
|
||||
let second_inspection = second_client.clone();
|
||||
let host =
|
||||
StandaloneHost::restore_with_model_client(state_dir.clone(), session_id, second_client)
|
||||
.await?;
|
||||
let snapshot = format!("{:?}", host.snapshot());
|
||||
assert!(snapshot.contains("first request"), "{snapshot}");
|
||||
assert!(snapshot.contains("first answer"), "{snapshot}");
|
||||
assert!(snapshot.contains("persisted task"), "{snapshot}");
|
||||
assert!(snapshot.contains("persisted notification"), "{snapshot}");
|
||||
|
||||
let mut events = host.subscribe();
|
||||
host.send(Method::run_text("continue after restore"))
|
||||
.await?;
|
||||
wait_for_run_end(&mut events).await?;
|
||||
let request = second_inspection
|
||||
.requests()
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("restored run request");
|
||||
let projected = format!("{:?}", request.items);
|
||||
assert!(projected.contains("first answer"), "{projected}");
|
||||
assert!(projected.contains("persisted notification"), "{projected}");
|
||||
assert!(projected.contains("persisted task"), "{projected}");
|
||||
host.shutdown().await?;
|
||||
|
||||
store.delete(session_id)?;
|
||||
assert!(cwd.exists(), "deleting session state must not mutate cwd");
|
||||
assert!(matches!(
|
||||
store.load(session_id),
|
||||
Err(StandaloneStoreError::SessionNotFound(_))
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_rejects_concurrent_lease_and_missing_cwd() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let cwd = temp.path().join("project");
|
||||
let moved = temp.path().join("moved-project");
|
||||
let state_dir = temp.path().join("state");
|
||||
std::fs::create_dir(&cwd)?;
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
&cwd,
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-lease-test",
|
||||
)
|
||||
.resolve()?;
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let session_id = host.session_id();
|
||||
let store = StandaloneSessionStore::open(&state_dir)?;
|
||||
assert!(matches!(
|
||||
store.acquire_lease(session_id, StaleLeasePolicy::Recover),
|
||||
Err(StandaloneStoreError::SessionLeased(id)) if id == session_id
|
||||
));
|
||||
let restore = StandaloneHost::restore_with_model_client(
|
||||
state_dir.clone(),
|
||||
session_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(
|
||||
restore,
|
||||
Err(StandaloneStartupError::SessionActive)
|
||||
));
|
||||
host.shutdown().await?;
|
||||
|
||||
std::fs::rename(&cwd, &moved)?;
|
||||
let restore = StandaloneHost::restore_with_model_client(
|
||||
state_dir,
|
||||
session_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(
|
||||
restore,
|
||||
Err(StandaloneStartupError::WorkingDirectoryUnavailable)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_recovers_only_a_proven_stale_lease() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let state_dir = temp.path().join("state");
|
||||
let mut launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-stale-lease-test",
|
||||
)
|
||||
.resolve()?;
|
||||
launch.profile.manifest.profile = Some(manifest::ProfileManifestSnapshot {
|
||||
source: manifest::ProfileSource::Registry {
|
||||
source: manifest::ProfileRegistrySource::User,
|
||||
name: "user-standalone".to_string(),
|
||||
path: None,
|
||||
provenance: Some("user-config-revision-7".to_string()),
|
||||
},
|
||||
profile: Some(manifest::ProfileMetadata {
|
||||
name: Some("User standalone".to_string()),
|
||||
description: None,
|
||||
format: None,
|
||||
}),
|
||||
});
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let session_id = host.session_id();
|
||||
host.shutdown().await?;
|
||||
let store = StandaloneSessionStore::open(&state_dir)?;
|
||||
assert!(matches!(
|
||||
store.load(session_id)?.manifest.profile,
|
||||
Some(manifest::ProfileManifestSnapshot {
|
||||
source: manifest::ProfileSource::Registry {
|
||||
source: manifest::ProfileRegistrySource::User,
|
||||
..
|
||||
},
|
||||
..
|
||||
})
|
||||
));
|
||||
let session_dir = state_dir.join(session_id.to_string());
|
||||
std::fs::write(
|
||||
session_dir.join("lease.json"),
|
||||
serde_json::to_vec(&serde_json::json!({
|
||||
"lease_id": uuid::Uuid::now_v7(),
|
||||
"pid": u32::MAX,
|
||||
"process_start_marker": 1,
|
||||
"acquired_at_unix_ms": 1
|
||||
}))?,
|
||||
)?;
|
||||
|
||||
let host = StandaloneHost::restore_with_model_client(
|
||||
state_dir,
|
||||
session_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await?;
|
||||
host.shutdown().await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_restore_rejects_lease_with_missing_start_marker() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let state_dir = temp.path().join("state");
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-unknown-lease-test",
|
||||
)
|
||||
.resolve()?;
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let session_id = host.session_id();
|
||||
host.shutdown().await?;
|
||||
let session_dir = state_dir.join(session_id.to_string());
|
||||
std::fs::write(
|
||||
session_dir.join("lease.json"),
|
||||
serde_json::to_vec(&serde_json::json!({
|
||||
"lease_id": uuid::Uuid::now_v7(),
|
||||
"pid": std::process::id(),
|
||||
"acquired_at_unix_ms": 1
|
||||
}))?,
|
||||
)?;
|
||||
|
||||
let store = StandaloneSessionStore::open(&state_dir)?;
|
||||
assert!(matches!(
|
||||
store.acquire_lease(session_id, StaleLeasePolicy::Recover),
|
||||
Err(StandaloneStoreError::LeaseLivenessUnknown(id)) if id == session_id
|
||||
));
|
||||
let restore = StandaloneHost::restore_with_model_client(
|
||||
state_dir,
|
||||
session_id,
|
||||
ScriptedClient::new(Vec::new()),
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(
|
||||
restore,
|
||||
Err(StandaloneStartupError::LeaseLivenessUnknown)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standalone_metadata_fails_closed_on_incomplete_or_newer_records() -> TestResult {
|
||||
let temp = tempfile::tempdir()?;
|
||||
let state_dir = temp.path().join("state");
|
||||
let launch = StandaloneLaunchConfig::new(
|
||||
temp.path(),
|
||||
&state_dir,
|
||||
manifest::ProfileSelector::Default,
|
||||
"standalone-schema-test",
|
||||
)
|
||||
.resolve()?;
|
||||
let host =
|
||||
StandaloneHost::start_with_model_client(launch, ScriptedClient::new(Vec::new())).await?;
|
||||
let session_id = host.session_id();
|
||||
host.shutdown().await?;
|
||||
let store = StandaloneSessionStore::open(&state_dir)?;
|
||||
let session_dir = state_dir.join(session_id.to_string());
|
||||
std::fs::write(session_dir.join("commit.pending"), b"interrupted\n")?;
|
||||
assert!(matches!(
|
||||
store.load(session_id),
|
||||
Err(StandaloneStoreError::IncompleteCommit(id)) if id == session_id
|
||||
));
|
||||
std::fs::remove_file(session_dir.join("commit.pending"))?;
|
||||
let record_path = session_dir.join("record.json");
|
||||
let mut record: serde_json::Value = serde_json::from_slice(&std::fs::read(&record_path)?)?;
|
||||
record["schema_version"] = serde_json::json!(u32::MAX);
|
||||
std::fs::write(&record_path, serde_json::to_vec_pretty(&record)?)?;
|
||||
assert!(matches!(
|
||||
store.load(session_id),
|
||||
Err(StandaloneStoreError::NewerSchema { id, .. }) if id == session_id
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn wait_for_run_end(events: &mut tokio::sync::broadcast::Receiver<Event>) -> TestResult {
|
||||
tokio::time::timeout(Duration::from_secs(10), async {
|
||||
loop {
|
||||
if matches!(events.recv().await, Ok(Event::RunEnd { .. })) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
})
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
+1137
-232
File diff suppressed because it is too large
Load Diff
+169
-21
@@ -142,8 +142,8 @@ const INTAKE_READY_DESCRIPTION: &str = "Record a bounded intake summary and mark
|
||||
The backend applies the same target validation and lock as TicketMarkReady and commits the summary, \
|
||||
state_changed event, effective target, and planning -> ready transition atomically.";
|
||||
const QUEUE_DESCRIPTION: &str = "Queue a ready Ticket for Orchestrator routing through the typed \
|
||||
Ticket backend. The backend performs the gated ready -> queued transition, records queued_by/queued_at, \
|
||||
and rejects unresolved blocking relations.";
|
||||
Ticket backend. The backend rejects transitive planning dependencies and cycles, atomically queues the \
|
||||
requested Ticket plus every transitive ready dependency, and leaves queued or in-progress dependencies unchanged.";
|
||||
const WORKFLOW_STATE_DESCRIPTION: &str = "Transition Ticket `state` through the typed \
|
||||
Ticket backend with a bounded `state_changed` event. Treat `queued -> inprogress` \
|
||||
as the implementation acceptance step: implementation side effects should happen only after that \
|
||||
@@ -316,7 +316,11 @@ impl TicketBackend for TicketToolBackend {
|
||||
self.backend.mark_ready(id, request)
|
||||
}
|
||||
|
||||
fn queue_ready(&self, id: TicketIdOrSlug, queued_by: &str) -> TicketResult<()> {
|
||||
fn queue_ready(
|
||||
&self,
|
||||
id: TicketIdOrSlug,
|
||||
queued_by: &str,
|
||||
) -> TicketResult<crate::TicketQueueOutcome> {
|
||||
self.backend.queue_ready(id, queued_by)
|
||||
}
|
||||
|
||||
@@ -406,7 +410,7 @@ struct TicketCreateParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketEditItemParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
/// Optional replacement title.
|
||||
#[serde(default)]
|
||||
@@ -535,7 +539,7 @@ impl QueryTicketParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct ShowTicketParams {
|
||||
/// Ticket id. Exactly one of `id` or `query` must be provided.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility. Exactly one of `id` or `query` must be provided.
|
||||
#[serde(default)]
|
||||
id: Option<String>,
|
||||
/// Exact ticket id query. Exactly one of `id` or `query` must be provided.
|
||||
@@ -554,7 +558,7 @@ struct ShowTicketParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketThreadEventParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
/// Markdown event body.
|
||||
body: String,
|
||||
@@ -562,7 +566,7 @@ struct TicketThreadEventParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketMarkReadyParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
/// Optional reason attached to the state_changed event.
|
||||
#[serde(default)]
|
||||
@@ -571,7 +575,7 @@ struct TicketMarkReadyParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketIntakeReadyParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
/// Concise bounded intake summary appended before the ready transition.
|
||||
intake_summary: String,
|
||||
@@ -582,13 +586,13 @@ struct TicketIntakeReadyParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketQueueParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketWorkflowStateParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
/// Expected current state. The backend rejects stale transitions.
|
||||
from: TicketWorkflowStateParam,
|
||||
@@ -602,7 +606,7 @@ struct TicketWorkflowStateParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketCloseParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
/// Markdown resolution written to resolution.md and thread.md.
|
||||
resolution: String,
|
||||
@@ -610,7 +614,7 @@ struct TicketCloseParams {
|
||||
|
||||
#[derive(Debug, Deserialize, schemars::JsonSchema)]
|
||||
struct TicketDependencyCheckParams {
|
||||
/// Ticket id.
|
||||
/// Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
ticket: String,
|
||||
}
|
||||
|
||||
@@ -642,7 +646,7 @@ struct TicketRelationRecordParams {
|
||||
ticket: String,
|
||||
/// Forward relation kind: depends_on, blocks, related, supersedes, or duplicate_of.
|
||||
kind: TicketRelationKindParam,
|
||||
/// Target canonical Ticket id. Title/slug words are not accepted as relation authority.
|
||||
/// Target Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
target: String,
|
||||
/// Optional bounded rationale/note.
|
||||
#[serde(default)]
|
||||
@@ -655,7 +659,7 @@ struct TicketRelationRemoveParams {
|
||||
ticket: String,
|
||||
/// Forward relation kind to remove.
|
||||
kind: TicketRelationKindParam,
|
||||
/// Target canonical Ticket id.
|
||||
/// Target Ticket reference. Prefer `T-*`; canonical internal ids remain accepted for compatibility.
|
||||
target: String,
|
||||
}
|
||||
|
||||
@@ -1219,12 +1223,29 @@ impl Tool for TicketQueueTool {
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let params: TicketQueueParams = parse_input("TicketQueue", input_json)?;
|
||||
let queued_by = default_author();
|
||||
self.backend
|
||||
let mut outcome = self
|
||||
.backend
|
||||
.queue_ready(TicketIdOrSlug::Query(params.ticket.clone()), &queued_by)
|
||||
.map_err(|error| backend_error("TicketQueue", error))?;
|
||||
outcome.requested_ticket =
|
||||
model_ticket_reference(&self.backend, &outcome.requested_ticket, "TicketQueue")?;
|
||||
outcome.queued_tickets = outcome
|
||||
.queued_tickets
|
||||
.into_iter()
|
||||
.map(|ticket| model_ticket_reference(&self.backend, &ticket, "TicketQueue"))
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
Ok(json_output(
|
||||
format!("Queued ticket {} for Orchestrator", params.ticket),
|
||||
json!({ "ticket": params.ticket, "state": "queued", "queued_by": queued_by, "ok": true }),
|
||||
format!(
|
||||
"Queued {} ticket(s) for Orchestrator",
|
||||
outcome.queued_tickets.len()
|
||||
),
|
||||
json!({
|
||||
"ticket": outcome.requested_ticket,
|
||||
"queued_tickets": outcome.queued_tickets,
|
||||
"state": "queued",
|
||||
"queued_by": queued_by,
|
||||
"ok": true
|
||||
}),
|
||||
))
|
||||
}
|
||||
}
|
||||
@@ -1250,15 +1271,17 @@ impl Tool for TicketWorkflowStateTool {
|
||||
self.backend
|
||||
.set_workflow_state(TicketIdOrSlug::Query(params.ticket.clone()), change)
|
||||
.map_err(|error| backend_error("TicketWorkflowState", error))?;
|
||||
let ticket_ref =
|
||||
model_ticket_reference(&self.backend, ¶ms.ticket, "TicketWorkflowState")?;
|
||||
Ok(json_output(
|
||||
format!(
|
||||
"Transitioned ticket {} state {} -> {}",
|
||||
params.ticket,
|
||||
ticket_ref,
|
||||
from.as_str(),
|
||||
to.as_str()
|
||||
),
|
||||
json!({
|
||||
"ticket": params.ticket,
|
||||
"ticket": ticket_ref,
|
||||
"from": from.as_str(),
|
||||
"to": to.as_str(),
|
||||
"state": to.as_str(),
|
||||
@@ -1282,9 +1305,10 @@ impl Tool for TicketCloseTool {
|
||||
MarkdownText::new(params.resolution),
|
||||
)
|
||||
.map_err(|error| backend_error("TicketClose", error))?;
|
||||
let ticket_ref = model_ticket_reference(&self.backend, ¶ms.ticket, "TicketClose")?;
|
||||
Ok(json_output(
|
||||
format!("Closed ticket {}", params.ticket),
|
||||
json!({ "ticket": params.ticket, "state": "closed", "ok": true }),
|
||||
format!("Closed ticket {ticket_ref}"),
|
||||
json!({ "ticket": ticket_ref, "state": "closed", "ok": true }),
|
||||
))
|
||||
}
|
||||
}
|
||||
@@ -1511,6 +1535,29 @@ impl Tool for TicketDependencyCheckTool {
|
||||
}
|
||||
}
|
||||
|
||||
fn model_ticket_reference(
|
||||
backend: &TicketToolBackend,
|
||||
reference: &str,
|
||||
tool_name: &str,
|
||||
) -> Result<String, ToolError> {
|
||||
let ticket = backend
|
||||
.show(TicketIdOrSlug::Id(reference.to_string()))
|
||||
.map_err(|error| backend_error(tool_name, error))?;
|
||||
match ticket.meta.resource_key {
|
||||
Some(resource_key) if is_canonical_ticket_resource_key(&resource_key) => Ok(resource_key),
|
||||
Some(_) => Err(ToolError::ExecutionFailed(format!(
|
||||
"{tool_name} failed: required Ticket key is unavailable"
|
||||
))),
|
||||
None => Ok(ticket.meta.id),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_canonical_ticket_resource_key(resource_key: &str) -> bool {
|
||||
resource_key.strip_prefix("T-").is_some_and(|sequence| {
|
||||
!sequence.is_empty() && sequence.bytes().all(|byte| byte.is_ascii_digit())
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_input<T: for<'de> Deserialize<'de>>(tool: &str, input_json: &str) -> Result<T, ToolError> {
|
||||
serde_json::from_str(input_json)
|
||||
.map_err(|error| ToolError::InvalidArgument(format!("invalid {tool} input: {error}")))
|
||||
@@ -1908,6 +1955,12 @@ mod tests {
|
||||
.with_target_authority(Arc::new(TestTargetAuthority))
|
||||
}
|
||||
|
||||
fn sqlite_backend(temp: &TempDir) -> crate::SqliteTicketBackend {
|
||||
crate::SqliteTicketBackend::open(temp.path().join("tickets.db"), "workspace")
|
||||
.unwrap()
|
||||
.with_target_authority(Arc::new(TestTargetAuthority))
|
||||
}
|
||||
|
||||
fn tool(definition: ToolDefinition) -> Arc<dyn Tool> {
|
||||
let (_, tool) = definition();
|
||||
tool
|
||||
@@ -2535,6 +2588,101 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn queue_workflow_and_close_project_internal_inputs_to_ticket_keys() {
|
||||
let temp = TempDir::new().unwrap();
|
||||
let inner = sqlite_backend(&temp);
|
||||
let mut dependency_input = NewTicket::new("Dependency");
|
||||
dependency_input.repository_id = Some("main".to_string());
|
||||
let dependency = inner.create(dependency_input).unwrap();
|
||||
let mut target_input = NewTicket::new("Target");
|
||||
target_input.repository_id = Some("main".to_string());
|
||||
let target = inner.create(target_input).unwrap();
|
||||
inner
|
||||
.add_ticket_relation(
|
||||
TicketIdOrSlug::Id(target.id.clone()),
|
||||
NewTicketRelation {
|
||||
kind: TicketRelationKind::DependsOn,
|
||||
target: dependency.id.clone(),
|
||||
note: None,
|
||||
author: None,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
for id in [&dependency.id, &target.id] {
|
||||
inner
|
||||
.mark_ready(
|
||||
TicketIdOrSlug::Id(id.clone()),
|
||||
TicketMarkReady {
|
||||
operation_key: format!("ready-{id}"),
|
||||
reason: None,
|
||||
author: None,
|
||||
intake_summary: None,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
let target_key = target.resource_key.clone().unwrap();
|
||||
let dependency_key = dependency.resource_key.clone().unwrap();
|
||||
let backend = inner;
|
||||
let queue = tool_by_name(TicketToolBackend::new(backend.clone()), "TicketQueue");
|
||||
let workflow = tool_by_name(
|
||||
TicketToolBackend::new(backend.clone()),
|
||||
"TicketWorkflowState",
|
||||
);
|
||||
let close = tool_by_name(TicketToolBackend::new(backend), "TicketClose");
|
||||
|
||||
let queued = queue
|
||||
.execute(
|
||||
&json!({"ticket": target.id.clone()}).to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(queued.summary.contains("2 ticket(s)"));
|
||||
let queued_content = queued.content.unwrap();
|
||||
assert!(queued_content.contains(&target_key));
|
||||
assert!(queued_content.contains(&dependency_key));
|
||||
assert!(!queued_content.contains(&target.id));
|
||||
assert!(!queued_content.contains(&dependency.id));
|
||||
|
||||
for (from, to) in [("queued", "inprogress"), ("inprogress", "done")] {
|
||||
let transitioned = workflow
|
||||
.execute(
|
||||
&json!({
|
||||
"ticket": target.id.clone(),
|
||||
"from": from,
|
||||
"to": to,
|
||||
"reason": "test_transition",
|
||||
"body": "transitioned",
|
||||
"author": "tester"
|
||||
})
|
||||
.to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(transitioned.summary.contains(&target_key));
|
||||
assert!(!transitioned.summary.contains(&target.id));
|
||||
let content = transitioned.content.unwrap();
|
||||
assert!(content.contains(&target_key));
|
||||
assert!(!content.contains(&target.id));
|
||||
}
|
||||
|
||||
let closed = close
|
||||
.execute(
|
||||
&json!({"ticket": target.id.clone(), "resolution": "Done"}).to_string(),
|
||||
Default::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(closed.summary.contains(&target_key));
|
||||
assert!(!closed.summary.contains(&target.id));
|
||||
let content = closed.content.unwrap();
|
||||
assert!(content.contains(&target_key));
|
||||
assert!(!content.contains(&target.id));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ticket_workflow_tools_mark_ready_and_transition_state() {
|
||||
let temp = TempDir::new().unwrap();
|
||||
|
||||
+159
-14
@@ -1,5 +1,6 @@
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput};
|
||||
use async_trait::async_trait;
|
||||
@@ -20,21 +21,65 @@ struct BashParams {
|
||||
|
||||
pub(crate) struct BashTool {
|
||||
session: WorkdirSessionHandle,
|
||||
state: Arc<Mutex<BashExecutionState>>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ActiveCommand {
|
||||
call_id: String,
|
||||
execution_nonce: u64,
|
||||
handle: CommandHandle,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct BashExecutionState {
|
||||
active: HashMap<String, ActiveCommand>,
|
||||
cancellation_requested: HashSet<String>,
|
||||
legacy_cancellation_requested: HashSet<String>,
|
||||
next_execution_nonce: u64,
|
||||
}
|
||||
|
||||
struct CommandGuard {
|
||||
session: WorkdirSessionHandle,
|
||||
state: Arc<Mutex<BashExecutionState>>,
|
||||
execution_id: String,
|
||||
execution_nonce: u64,
|
||||
handle: Option<CommandHandle>,
|
||||
}
|
||||
|
||||
impl Drop for CommandGuard {
|
||||
fn drop(&mut self) {
|
||||
if let Some(handle) = self.handle.take() {
|
||||
let workdir = self.session.clone();
|
||||
tokio::spawn(async move {
|
||||
let _ = workdir.cancel_command(handle).await;
|
||||
});
|
||||
}
|
||||
let Some(handle) = self.handle.take() else {
|
||||
return;
|
||||
};
|
||||
let workdir = self.session.clone();
|
||||
let state = Arc::clone(&self.state);
|
||||
let execution_id = self.execution_id.clone();
|
||||
let execution_nonce = self.execution_nonce;
|
||||
// A dropped provider future is not terminal confirmation. Keep the live
|
||||
// execution registered until cleanup has both requested cancellation and
|
||||
// observed terminal command output, so cancellation/session teardown
|
||||
// cannot race with an apparently empty registry.
|
||||
tokio::spawn(async move {
|
||||
let _ = workdir.cancel_command(handle.clone()).await;
|
||||
let _ = workdir
|
||||
.command_output(CommandOutputRequest {
|
||||
handle,
|
||||
cursor: 0,
|
||||
limit: INLINE_BYTE_BUDGET,
|
||||
wait: true,
|
||||
})
|
||||
.await;
|
||||
let mut state = state.lock().unwrap();
|
||||
if state
|
||||
.active
|
||||
.get(&execution_id)
|
||||
.is_some_and(|active| active.execution_nonce == execution_nonce)
|
||||
{
|
||||
state.active.remove(&execution_id);
|
||||
state.cancellation_requested.remove(&execution_id);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,20 +97,50 @@ impl Tool for BashTool {
|
||||
.unwrap_or(DEFAULT_TIMEOUT_SECS)
|
||||
.clamp(1, MAX_TIMEOUT_SECS);
|
||||
let cmd_summary = truncate_for_summary(¶ms.command);
|
||||
let execution_id = ctx.execution_id();
|
||||
let call_id = ctx.call_id;
|
||||
let execution_nonce = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
state.next_execution_nonce = state.next_execution_nonce.wrapping_add(1);
|
||||
state.next_execution_nonce
|
||||
};
|
||||
let mut guard = CommandGuard {
|
||||
session: self.session.clone(),
|
||||
state: self.state.clone(),
|
||||
execution_id: execution_id.clone(),
|
||||
execution_nonce,
|
||||
handle: None,
|
||||
};
|
||||
let handle = self
|
||||
.session
|
||||
.start_command(CommandRequest {
|
||||
command: params.command,
|
||||
timeout_secs,
|
||||
output_limit: INLINE_BYTE_BUDGET,
|
||||
tool_call_id: Some(ctx.call_id),
|
||||
tool_call_id: Some(call_id.clone()),
|
||||
})
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
let mut guard = CommandGuard {
|
||||
session: self.session.clone(),
|
||||
handle: Some(handle.clone()),
|
||||
let cancel_after_start = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
state.active.insert(
|
||||
execution_id.clone(),
|
||||
ActiveCommand {
|
||||
call_id: call_id.clone(),
|
||||
execution_nonce,
|
||||
handle: handle.clone(),
|
||||
},
|
||||
);
|
||||
state.cancellation_requested.contains(&execution_id)
|
||||
|| state.legacy_cancellation_requested.contains(&call_id)
|
||||
};
|
||||
guard.handle = Some(handle.clone());
|
||||
if cancel_after_start {
|
||||
self.session
|
||||
.cancel_command(handle.clone())
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
}
|
||||
let output = self
|
||||
.session
|
||||
.command_output(CommandOutputRequest {
|
||||
@@ -76,9 +151,27 @@ impl Tool for BashTool {
|
||||
})
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
let cancellation_requested = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
let owns_registration = state
|
||||
.active
|
||||
.get(&execution_id)
|
||||
.is_some_and(|active| active.execution_nonce == execution_nonce);
|
||||
let exact = if owns_registration {
|
||||
state.active.remove(&execution_id);
|
||||
state.cancellation_requested.remove(&execution_id)
|
||||
} else {
|
||||
false
|
||||
};
|
||||
let legacy = state.legacy_cancellation_requested.remove(&call_id);
|
||||
exact || legacy
|
||||
};
|
||||
guard.handle = None;
|
||||
|
||||
let summary = if output.timed_out {
|
||||
let timed_out = output.timed_out;
|
||||
let summary = if cancellation_requested {
|
||||
format!("$ {cmd_summary} (cancelled)")
|
||||
} else if output.timed_out {
|
||||
format!("$ {cmd_summary} (timed out after {timeout_secs}s)")
|
||||
} else {
|
||||
match output.exit_code {
|
||||
@@ -97,11 +190,62 @@ impl Tool for BashTool {
|
||||
} else {
|
||||
Some(output.content)
|
||||
};
|
||||
Ok(ToolOutput {
|
||||
let output = ToolOutput {
|
||||
summary,
|
||||
content,
|
||||
attachments: Vec::new(),
|
||||
})
|
||||
};
|
||||
if cancellation_requested {
|
||||
Err(ToolError::Cancelled(output))
|
||||
} else if timed_out {
|
||||
Err(ToolError::Interrupted(output))
|
||||
} else {
|
||||
Ok(output)
|
||||
}
|
||||
}
|
||||
|
||||
async fn cancel(&self, call_id: &str) -> Result<(), ToolError> {
|
||||
let handles = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
state
|
||||
.legacy_cancellation_requested
|
||||
.insert(call_id.to_string());
|
||||
state
|
||||
.active
|
||||
.values()
|
||||
.filter(|active| active.call_id == call_id)
|
||||
.map(|active| active.handle.clone())
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
for handle in handles {
|
||||
self.session
|
||||
.cancel_command(handle)
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn cancel_execution(
|
||||
&self,
|
||||
ctx: &agen::tool::ToolExecutionContext,
|
||||
) -> Result<(), ToolError> {
|
||||
let execution_id = ctx.execution_id();
|
||||
let handle = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
state.cancellation_requested.insert(execution_id.clone());
|
||||
state
|
||||
.active
|
||||
.get(&execution_id)
|
||||
.map(|active| active.handle.clone())
|
||||
};
|
||||
if let Some(handle) = handle {
|
||||
self.session
|
||||
.cancel_command(handle)
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -123,6 +267,7 @@ pub fn bash_tool(session: WorkdirSessionHandle, _output_dir: PathBuf) -> ToolDef
|
||||
.input_schema(serde_json::to_value(schema).expect("Bash schema serialization"));
|
||||
let tool: Arc<dyn Tool> = Arc::new(BashTool {
|
||||
session: session.clone(),
|
||||
state: Arc::new(Mutex::new(BashExecutionState::default())),
|
||||
});
|
||||
(meta, tool)
|
||||
})
|
||||
|
||||
@@ -22,7 +22,7 @@ enum OutputMode {
|
||||
#[derive(Debug, Deserialize, JsonSchema)]
|
||||
struct GrepParams {
|
||||
pattern: String,
|
||||
/// Logical Workdir-relative path to search. Defaults to the Workdir root.
|
||||
/// Logical Workdir-relative file or directory to search. Defaults to the Workdir root.
|
||||
#[serde(default)]
|
||||
path: Option<String>,
|
||||
#[serde(default)]
|
||||
@@ -129,7 +129,7 @@ pub fn grep_tool(session: WorkdirSessionHandle) -> ToolDefinition {
|
||||
Arc::new(move || {
|
||||
let schema = schemars::schema_for!(GrepParams);
|
||||
let meta = ToolMeta::new("Grep")
|
||||
.description("Search Workdir file contents with a regex. Glob/Grep traversal executes inside the WorkdirSession provider. Results are bounded and Workdir-relative.")
|
||||
.description("Search a Workdir file or directory with a regex. Content results group lines by file; `>` marks matching lines and unmarked lines are context. Directory traversal executes inside the WorkdirSession provider. Results are bounded and Workdir-relative.")
|
||||
.input_schema(serde_json::to_value(schema).expect("Grep schema serialization"));
|
||||
let tool: Arc<dyn Tool> = Arc::new(GrepTool {
|
||||
session: session.clone(),
|
||||
|
||||
@@ -131,10 +131,11 @@ async fn symlink_to_outside_scope_is_rejected_for_write() {
|
||||
assert!(
|
||||
msg.contains("outside allowed read scope")
|
||||
|| msg.contains("outside allowed write scope")
|
||||
|| msg.contains("outside allowed scope")
|
||||
|| msg.contains("has not been read"),
|
||||
"symlink escape not rejected: {msg}"
|
||||
);
|
||||
if !msg.contains("has not been read") {
|
||||
if msg.contains("outside allowed read scope") || msg.contains("outside allowed write scope") {
|
||||
assert!(
|
||||
msg.contains("add the symlink target"),
|
||||
"symlink escape diagnostic should include remediation: {msg}"
|
||||
@@ -233,12 +234,16 @@ async fn absolute_path_is_rejected() {
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(format!("{err}").contains("invalid Workdir path"));
|
||||
let msg = format!("{err}");
|
||||
assert!(
|
||||
msg.contains("invalid logical filesystem path"),
|
||||
"absolute path was not rejected as invalid: {msg}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn directory_target_is_rejected_for_read() {
|
||||
let (dir, _spill, reg) = setup();
|
||||
let (_dir, _spill, reg) = setup();
|
||||
let read = reg.get("Read");
|
||||
let err = read
|
||||
.execute(&json!({ "file_path": "." }).to_string(), Default::default())
|
||||
|
||||
@@ -7,7 +7,10 @@
|
||||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use agen::tool::{Tool, ToolDefinition, ToolMeta};
|
||||
use agen::tool::{
|
||||
Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolExecutionHandle,
|
||||
ToolExecutionTerminal, ToolMeta,
|
||||
};
|
||||
use manifest::{Permission, Scope, ScopeConfig, ScopeRule};
|
||||
use serde_json::json;
|
||||
use tempfile::TempDir;
|
||||
@@ -191,7 +194,7 @@ async fn write_then_grep_finds_content() {
|
||||
|
||||
#[tokio::test]
|
||||
async fn glob_finds_written_files() {
|
||||
let (dir, _spill, reg) = setup();
|
||||
let (_dir, _spill, reg) = setup();
|
||||
let write = reg.get("Write");
|
||||
let glob = reg.get("Glob");
|
||||
|
||||
@@ -229,7 +232,10 @@ async fn absolute_path_is_rejected() {
|
||||
.await;
|
||||
// Absolute paths are rejected at the logical WorkdirSession boundary.
|
||||
let msg = format!("{err}");
|
||||
assert!(msg.contains("invalid Workdir path"), "unexpected: {msg}");
|
||||
assert!(
|
||||
msg.contains("invalid logical filesystem path"),
|
||||
"unexpected: {msg}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -340,7 +346,7 @@ async fn tracker_recent_files_tracks_read_write_edit() {
|
||||
));
|
||||
|
||||
let a = dir.path().join("a.txt");
|
||||
let b = dir.path().join("b.txt");
|
||||
let _b = dir.path().join("b.txt");
|
||||
std::fs::write(&a, "one\n").unwrap();
|
||||
|
||||
// Read `a` — should appear in recency.
|
||||
@@ -398,5 +404,84 @@ async fn bash_provider_output_does_not_expose_internal_paths() {
|
||||
assert_eq!(std::fs::read_dir(spill.path()).unwrap().count(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bash_cancellation_returns_bounded_progress_as_terminal_output() {
|
||||
let (dir, _spill, reg) = setup();
|
||||
let marker = dir.path().join("must-not-run-after-cancel");
|
||||
let command = format!(
|
||||
"printf 'before\\n'; printf 'err-before\\n' >&2; sleep 1; touch {}; printf 'after\\n'",
|
||||
marker.display()
|
||||
);
|
||||
let input = serde_json::to_string(&json!({ "command": command })).unwrap();
|
||||
let context = ToolExecutionContext::new("call-heavy", "attempt-heavy", 0);
|
||||
let bash = reg.get("Bash");
|
||||
let executing = bash.clone();
|
||||
let execution_context = context.clone();
|
||||
let execution = tokio::spawn(async move { executing.execute(&input, execution_context).await });
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
bash.cancel_execution(&context)
|
||||
.await
|
||||
.expect("signal exact execution cancellation");
|
||||
let error = tokio::time::timeout(std::time::Duration::from_secs(2), execution)
|
||||
.await
|
||||
.expect("cancelled Bash should terminate inside the Engine grace budget")
|
||||
.expect("Bash task join");
|
||||
|
||||
let ToolError::Cancelled(output) = error.expect_err("cancelled command is non-success") else {
|
||||
panic!("expected typed cancellation result");
|
||||
};
|
||||
let content = output.content.expect("bounded progress output");
|
||||
assert!(
|
||||
content.contains("before"),
|
||||
"missing pre-cancel stdout: {content}"
|
||||
);
|
||||
assert!(
|
||||
content.contains("err-before"),
|
||||
"missing pre-cancel stderr: {content}"
|
||||
);
|
||||
assert!(
|
||||
!content.contains("after"),
|
||||
"post-cancel output leaked: {content}"
|
||||
);
|
||||
assert!(content.len() <= 16 * 1024, "output must remain bounded");
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(1_100)).await;
|
||||
assert!(
|
||||
!marker.exists(),
|
||||
"the cancelled command continued executing after terminal confirmation"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bash_force_close_cleanup_stops_command_and_keeps_session_reusable() {
|
||||
let (dir, _spill, reg) = setup();
|
||||
let marker = dir.path().join("must-not-survive-force-close");
|
||||
let command = format!("sleep 1; touch {}", marker.display());
|
||||
let input = serde_json::to_string(&json!({ "command": command })).unwrap();
|
||||
let bash = reg.get("Bash");
|
||||
let context = ToolExecutionContext::new("call-force", "attempt-force", 0);
|
||||
let (handle, terminal) = ToolExecutionHandle::start(bash.clone(), input, context);
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
handle.force_close();
|
||||
assert!(matches!(
|
||||
terminal.await,
|
||||
ToolExecutionTerminal::OutcomeUnknown
|
||||
));
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(1_100)).await;
|
||||
assert!(
|
||||
!marker.exists(),
|
||||
"CommandGuard cleanup allowed a force-closed command to continue"
|
||||
);
|
||||
|
||||
let output = bash
|
||||
.execute(r#"{"command":"printf 'reused'"}"#, Default::default())
|
||||
.await
|
||||
.expect("workdir session remains reusable after cleanup");
|
||||
assert_eq!(output.content.as_deref(), Some("reused"));
|
||||
}
|
||||
|
||||
// Sanity: unused Path import guard
|
||||
const _: fn() -> &'static Path = || Path::new("/");
|
||||
|
||||
@@ -10,11 +10,13 @@ e2e-test = []
|
||||
|
||||
[dependencies]
|
||||
client = { workspace = true }
|
||||
standalone = { workspace = true }
|
||||
thiserror.workspace = true
|
||||
protocol = { workspace = true }
|
||||
ratatui = { version = "0.30.0", features = ["scrolling-regions"] }
|
||||
base64 = "0.22.1"
|
||||
crossterm = "0.28"
|
||||
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "io-util", "sync", "time", "process"] }
|
||||
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "sync", "time"] }
|
||||
serde_json = { workspace = true }
|
||||
unicode-width = "0.2.2"
|
||||
uuid = { workspace = true }
|
||||
@@ -22,10 +24,8 @@ toml = { workspace = true }
|
||||
manifest = { workspace = true }
|
||||
secrets = { workspace = true }
|
||||
session-store = { workspace = true }
|
||||
fs4 = { workspace = true }
|
||||
ticket = { workspace = true }
|
||||
serde = { workspace = true, features = ["derive"] }
|
||||
worker = { path = "../worker" }
|
||||
pulldown-cmark = { version = "0.13.3", default-features = false }
|
||||
agen.workspace = true
|
||||
|
||||
|
||||
+225
-208
@@ -765,7 +765,7 @@ impl App {
|
||||
|
||||
fn method_for_run(&mut self, segments: Vec<Segment>) -> Method {
|
||||
// TurnHeader / UserMessage blocks are pushed only after the Worker
|
||||
// emits `Event::UserMessage` from a committed `LogEntry::UserInput`.
|
||||
// emits `Event::UserMessage` from a committed `LogEntry::AnnotatedUserInput`.
|
||||
// Locally we only clear the input buffer and forward the method,
|
||||
// while remembering enough local state to undo the visible submit if
|
||||
// the accepted run produced no assistant output and was rolled back.
|
||||
@@ -1098,10 +1098,9 @@ impl App {
|
||||
self.blocks.push(Block::UserMessage { segments });
|
||||
self.assistant_streaming = false;
|
||||
}
|
||||
Event::SegmentRotated { entry } => {
|
||||
Event::SegmentRotated { session } => {
|
||||
let retained_run_errors = self.run_error_messages.clone();
|
||||
self.reset_for_rotation();
|
||||
self.apply_log_entry_raw(&entry);
|
||||
self.restore_session(&session, self.greeting.clone());
|
||||
for message in retained_run_errors {
|
||||
self.blocks.push(Block::Alert {
|
||||
level: AlertLevel::Error,
|
||||
@@ -1244,6 +1243,7 @@ impl App {
|
||||
id,
|
||||
summary,
|
||||
output,
|
||||
disposition: _,
|
||||
is_error,
|
||||
} => {
|
||||
self.latest_llm_wait_event = None;
|
||||
@@ -1342,13 +1342,20 @@ impl App {
|
||||
}
|
||||
}
|
||||
}
|
||||
Event::CompactStart => {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Streaming {
|
||||
started_at: Instant::now(),
|
||||
}));
|
||||
Event::CompactStart { .. } => {
|
||||
if self.last_streaming_compact_mut().is_none() {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Streaming {
|
||||
started_at: Instant::now(),
|
||||
}));
|
||||
}
|
||||
}
|
||||
Event::CompactDone { new_segment_id } => {
|
||||
Event::CompactDone { lifecycle } => {
|
||||
self.session_context_tokens = 0;
|
||||
let new_segment_id = lifecycle
|
||||
.new_segment_id
|
||||
.as_deref()
|
||||
.and_then(|value| uuid::Uuid::parse_str(value).ok())
|
||||
.unwrap_or_default();
|
||||
if let Some(evt) = self.last_streaming_compact_mut() {
|
||||
let elapsed_secs = match evt {
|
||||
CompactEvent::Streaming { started_at } => {
|
||||
@@ -1367,7 +1374,10 @@ impl App {
|
||||
}));
|
||||
}
|
||||
}
|
||||
Event::CompactFailed { error } => {
|
||||
Event::CompactFailed { lifecycle } => {
|
||||
let error = lifecycle
|
||||
.error
|
||||
.unwrap_or_else(|| "compaction failed".to_string());
|
||||
if let Some(evt) = self.last_streaming_compact_mut() {
|
||||
let elapsed_secs = match evt {
|
||||
CompactEvent::Streaming { started_at } => {
|
||||
@@ -1397,14 +1407,14 @@ impl App {
|
||||
self.latest_memory_worker_event = Some(event.message);
|
||||
}
|
||||
Event::Snapshot {
|
||||
entries,
|
||||
session,
|
||||
greeting,
|
||||
status,
|
||||
in_flight,
|
||||
internal_workers,
|
||||
} => {
|
||||
self.rewind_refresh_fence = false;
|
||||
self.restore_snapshot(&entries, greeting, in_flight);
|
||||
self.restore_snapshot(&session, greeting, in_flight);
|
||||
self.replace_internal_worker_snapshots(internal_workers);
|
||||
self.set_worker_status(status);
|
||||
}
|
||||
@@ -1444,11 +1454,11 @@ impl App {
|
||||
}
|
||||
}
|
||||
Event::RewindApplied {
|
||||
entries,
|
||||
session,
|
||||
input,
|
||||
summary,
|
||||
} => {
|
||||
self.restore_rewind_snapshot(&entries);
|
||||
self.restore_rewind_snapshot(&session);
|
||||
self.rewind_refresh_fence = true;
|
||||
let restored_composer = if self.input.is_empty() {
|
||||
self.input.replace_with_segments(&input);
|
||||
@@ -2162,7 +2172,7 @@ impl App {
|
||||
) -> InternalWorkerView {
|
||||
let mut app = App::new(snapshot.worker.name.clone());
|
||||
app.mode = mode;
|
||||
app.restore_entries(&snapshot.entries, None);
|
||||
app.restore_session(&snapshot.session, None);
|
||||
app.apply_in_flight_snapshot(snapshot.in_flight);
|
||||
app.set_worker_status(snapshot.status);
|
||||
if let Some(error) = snapshot.error {
|
||||
@@ -2243,14 +2253,14 @@ impl App {
|
||||
|
||||
fn restore_snapshot(
|
||||
&mut self,
|
||||
entries: &[serde_json::Value],
|
||||
session: &protocol::SessionSnapshot,
|
||||
greeting: protocol::Greeting,
|
||||
in_flight: InFlightSnapshot,
|
||||
) {
|
||||
self.greeting = Some(greeting.clone());
|
||||
self.context_window = greeting.context_window;
|
||||
self.session_context_tokens = greeting.context_tokens;
|
||||
self.restore_entries(entries, Some(greeting));
|
||||
self.restore_session(session, Some(greeting));
|
||||
self.apply_in_flight_snapshot(in_flight);
|
||||
}
|
||||
|
||||
@@ -2259,7 +2269,7 @@ impl App {
|
||||
/// session tail; always clear/replay from it even if this TUI instance has
|
||||
/// somehow lost connect-time greeting metadata. Skipping the restore in
|
||||
/// that case would leave old post-target output visible after success.
|
||||
fn restore_rewind_snapshot(&mut self, entries: &[serde_json::Value]) {
|
||||
fn restore_rewind_snapshot(&mut self, session: &protocol::SessionSnapshot) {
|
||||
let greeting = self.greeting.clone().or_else(|| {
|
||||
self.blocks.iter().find_map(|b| match b {
|
||||
Block::Greeting(g) => Some(g.clone()),
|
||||
@@ -2272,7 +2282,7 @@ impl App {
|
||||
self.session_context_tokens = greeting.context_tokens;
|
||||
}
|
||||
let missing_greeting = greeting.is_none();
|
||||
self.restore_entries(entries, greeting);
|
||||
self.restore_session(session, greeting);
|
||||
if missing_greeting {
|
||||
self.blocks.push(Block::Alert {
|
||||
level: AlertLevel::Warn,
|
||||
@@ -2282,9 +2292,9 @@ impl App {
|
||||
}
|
||||
}
|
||||
|
||||
fn restore_entries(
|
||||
fn restore_session(
|
||||
&mut self,
|
||||
entries: &[serde_json::Value],
|
||||
session: &protocol::SessionSnapshot,
|
||||
greeting: Option<protocol::Greeting>,
|
||||
) {
|
||||
self.run_error_messages.clear();
|
||||
@@ -2298,137 +2308,90 @@ impl App {
|
||||
}
|
||||
self.assistant_streaming = false;
|
||||
|
||||
for entry in entries {
|
||||
self.apply_log_entry_raw(entry);
|
||||
for entry in &session.entries {
|
||||
use protocol::{SessionContentPart, SessionMessageRole, SessionSnapshotEntryData};
|
||||
match &entry.data {
|
||||
SessionSnapshotEntryData::UserInput { segments } => {
|
||||
self.turn_index += 1;
|
||||
self.blocks.push(Block::TurnHeader {
|
||||
turn: self.turn_index,
|
||||
});
|
||||
if !segments.is_empty() {
|
||||
self.blocks.push(Block::UserMessage {
|
||||
segments: segments.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
SessionSnapshotEntryData::Message { role, content } => {
|
||||
let role = match role {
|
||||
SessionMessageRole::User => agen::Role::User,
|
||||
SessionMessageRole::Assistant => agen::Role::Assistant,
|
||||
};
|
||||
let item = agen::Item::Message {
|
||||
id: None,
|
||||
role,
|
||||
content: content
|
||||
.iter()
|
||||
.map(|part| match part {
|
||||
SessionContentPart::Text { text } => {
|
||||
agen::ContentPart::Text { text: text.clone() }
|
||||
}
|
||||
SessionContentPart::Refusal { refusal } => {
|
||||
agen::ContentPart::Refusal {
|
||||
refusal: refusal.clone(),
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect(),
|
||||
status: None,
|
||||
};
|
||||
let value = serde_json::to_value(item).expect("Item is Serialize");
|
||||
self.push_history_item(&value);
|
||||
}
|
||||
SessionSnapshotEntryData::ToolCall {
|
||||
call_id,
|
||||
name,
|
||||
arguments,
|
||||
} => {
|
||||
let item =
|
||||
agen::Item::tool_call(call_id.clone(), name.clone(), arguments.clone());
|
||||
let value = serde_json::to_value(item).expect("Item is Serialize");
|
||||
self.push_history_item(&value);
|
||||
}
|
||||
SessionSnapshotEntryData::ToolResult {
|
||||
call_id,
|
||||
summary,
|
||||
content,
|
||||
is_error,
|
||||
..
|
||||
} => {
|
||||
let item = agen::Item::tool_result_item(
|
||||
call_id.clone(),
|
||||
summary.clone(),
|
||||
content.clone(),
|
||||
*is_error,
|
||||
);
|
||||
let value = serde_json::to_value(item).expect("Item is Serialize");
|
||||
self.push_history_item(&value);
|
||||
}
|
||||
SessionSnapshotEntryData::SystemItem { data, .. } => {
|
||||
if let Some(data) = data {
|
||||
self.apply_system_item(data);
|
||||
}
|
||||
}
|
||||
SessionSnapshotEntryData::RunError { message } => {
|
||||
self.push_run_error(message.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.mark_orphan_tool_calls_incomplete_pass();
|
||||
}
|
||||
|
||||
/// Drop the derived view in preparation for replaying a new
|
||||
/// `SegmentStart` (compaction / fork). Greeting is preserved
|
||||
/// because the Worker identity hasn't changed.
|
||||
fn reset_for_rotation(&mut self) {
|
||||
let greeting = self.blocks.iter().find_map(|b| match b {
|
||||
Block::Greeting(g) => Some(g.clone()),
|
||||
_ => None,
|
||||
});
|
||||
self.turn_index = 0;
|
||||
self.blocks.clear();
|
||||
self.cache = FileCache::new();
|
||||
self.task_store = TaskStore::new();
|
||||
self.task_pane_scroll = 0;
|
||||
if let Some(g) = greeting {
|
||||
self.greeting = Some(g.clone());
|
||||
self.blocks.push(Block::Greeting(g));
|
||||
}
|
||||
}
|
||||
|
||||
/// Walk a single `LogEntry` JSON value and translate it into blocks
|
||||
/// the live event path would have produced. Shared between
|
||||
/// `restore_snapshot` (replay path) and `apply_log_entry` (live
|
||||
/// path).
|
||||
fn apply_log_entry_raw(&mut self, value: &serde_json::Value) {
|
||||
let Ok(entry) = serde_json::from_value::<session_store::LogEntry>(value.clone()) else {
|
||||
return;
|
||||
};
|
||||
match entry {
|
||||
session_store::LogEntry::SegmentStart { history, .. } => {
|
||||
for logged in history {
|
||||
let item: agen::Item = logged.into();
|
||||
let item_value = serde_json::to_value(&item).expect("Item is Serialize");
|
||||
self.push_history_item(&item_value);
|
||||
}
|
||||
}
|
||||
session_store::LogEntry::UserInput { segments, .. } => {
|
||||
self.turn_index += 1;
|
||||
self.blocks.push(Block::TurnHeader {
|
||||
turn: self.turn_index,
|
||||
});
|
||||
if !segments.is_empty() {
|
||||
self.blocks.push(Block::UserMessage { segments });
|
||||
}
|
||||
}
|
||||
session_store::LogEntry::AssistantItem { item, .. }
|
||||
| session_store::LogEntry::ToolResult { item, .. } => {
|
||||
let it: agen::Item = item.into();
|
||||
let item_value = serde_json::to_value(&it).expect("Item is Serialize");
|
||||
self.push_history_item(&item_value);
|
||||
}
|
||||
session_store::LogEntry::SystemItem { item, .. } => {
|
||||
let value = serde_json::to_value(&item).expect("SystemItem is Serialize");
|
||||
self.apply_system_item(&value);
|
||||
}
|
||||
session_store::LogEntry::Extension {
|
||||
domain, payload, ..
|
||||
} if domain == "yoi.compaction" => {
|
||||
self.apply_compaction_extension(&payload);
|
||||
}
|
||||
session_store::LogEntry::RunErrored { message, .. } => {
|
||||
self.push_run_error(message);
|
||||
}
|
||||
// Non-history-bearing variants don't affect the block view.
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
/// Dispatch one `SystemItem` JSON value into the appropriate block.
|
||||
///
|
||||
/// Kind-based routing replaces the old free-text `[Notification]` /
|
||||
/// `[File: …]` parsing path: each kind maps directly to a typed
|
||||
/// block (`Block::Notify`, `Block::WorkerEvent`, …).
|
||||
fn apply_compaction_extension(&mut self, payload: &serde_json::Value) {
|
||||
if payload.get("kind").and_then(|value| value.as_str()) != Some("compaction_block") {
|
||||
return;
|
||||
}
|
||||
match payload.get("state").and_then(|value| value.as_str()) {
|
||||
Some("running") => {
|
||||
if self.last_streaming_compact_mut().is_none() {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Streaming {
|
||||
started_at: Instant::now(),
|
||||
}));
|
||||
}
|
||||
}
|
||||
Some("done") => {
|
||||
let new_segment_id = payload
|
||||
.get("new_segment_id")
|
||||
.and_then(|value| value.as_str())
|
||||
.and_then(|value| value.parse::<uuid::Uuid>().ok())
|
||||
.unwrap_or_else(uuid::Uuid::nil);
|
||||
if let Some(evt) = self.last_streaming_compact_mut() {
|
||||
*evt = CompactEvent::Done {
|
||||
new_segment_id,
|
||||
elapsed_secs: None,
|
||||
};
|
||||
} else {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Done {
|
||||
new_segment_id,
|
||||
elapsed_secs: None,
|
||||
}));
|
||||
}
|
||||
}
|
||||
Some("failed") => {
|
||||
let error = payload
|
||||
.get("error")
|
||||
.and_then(|value| value.as_str())
|
||||
.unwrap_or("compact failed")
|
||||
.to_string();
|
||||
if let Some(evt) = self.last_streaming_compact_mut() {
|
||||
*evt = CompactEvent::Failed {
|
||||
error,
|
||||
elapsed_secs: None,
|
||||
};
|
||||
} else {
|
||||
self.blocks.push(Block::Compact(CompactEvent::Failed {
|
||||
error,
|
||||
elapsed_secs: None,
|
||||
}));
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_system_item(&mut self, value: &serde_json::Value) {
|
||||
let Ok(item) = serde_json::from_value::<session_store::SystemItem>(value.clone()) else {
|
||||
// Unknown / forward-compat shape: fall back to rendering the
|
||||
@@ -2486,7 +2449,7 @@ fn event_is_stale_after_rewind(event: &Event) -> bool {
|
||||
event,
|
||||
Event::Alert(_)
|
||||
| Event::MemoryWorker(_)
|
||||
| Event::CompactStart
|
||||
| Event::CompactStart { .. }
|
||||
| Event::CompactDone { .. }
|
||||
| Event::CompactFailed { .. }
|
||||
| Event::SegmentRotated { .. }
|
||||
@@ -2531,6 +2494,15 @@ fn fmt_millis(ms: u64) -> String {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn public_session(values: Vec<serde_json::Value>) -> protocol::SessionSnapshot {
|
||||
let entries = values
|
||||
.into_iter()
|
||||
.map(|value| serde_json::from_value(value).expect("LogEntry deserializes"))
|
||||
.collect::<Vec<session_store::LogEntry>>();
|
||||
session_store::public_snapshot::project_current_session_snapshot(&entries)
|
||||
}
|
||||
|
||||
fn message_text(item: &serde_json::Value) -> String {
|
||||
item["content"]
|
||||
.as_array()
|
||||
@@ -2674,7 +2646,7 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::text("selected rewind input")],
|
||||
summary: summary(3),
|
||||
});
|
||||
@@ -2693,7 +2665,7 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::text("rewound input")],
|
||||
summary: summary(1),
|
||||
});
|
||||
@@ -2736,7 +2708,7 @@ mod rewind_refresh_tests {
|
||||
});
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::text("rewound input")],
|
||||
summary: summary(2),
|
||||
});
|
||||
@@ -2965,6 +2937,17 @@ mod composer_history_persistence_tests {
|
||||
mod completion_flow_tests {
|
||||
use super::*;
|
||||
|
||||
fn annotated(item: agen::Item) -> session_store::LoggedHistoryEntry {
|
||||
session_store::LoggedHistoryEntry {
|
||||
item: session_store::LoggedItem::from(item),
|
||||
metadata: session_store::LoggedSessionHistoryMetadata {
|
||||
entry_id: session_store::LoggedSessionHistoryEntryId::new(),
|
||||
origin: session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
|
||||
derivation: None,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typing_at_creates_completion_state_and_emits_query() {
|
||||
let mut app = App::new("test".into());
|
||||
@@ -3267,7 +3250,7 @@ mod completion_flow_tests {
|
||||
#[test]
|
||||
fn committed_user_message_survives_fresh_segment_rotation() {
|
||||
let mut app = App::new("test".into());
|
||||
let start = session_store::LogEntry::SegmentStart {
|
||||
let start = session_store::LogEntry::AnnotatedSegmentStart {
|
||||
ts: session_store::segment_log::now_millis(),
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -3278,7 +3261,9 @@ mod completion_flow_tests {
|
||||
};
|
||||
|
||||
app.handle_worker_event(Event::SegmentRotated {
|
||||
entry: serde_json::to_value(start).expect("LogEntry is Serialize"),
|
||||
session: public_session(vec![
|
||||
serde_json::to_value(start).expect("LogEntry is Serialize"),
|
||||
]),
|
||||
});
|
||||
app.handle_worker_event(Event::UserMessage {
|
||||
segments: vec![Segment::text("first persisted message")],
|
||||
@@ -3522,23 +3507,23 @@ mod completion_flow_tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_renders_system_message_block_from_session_start() {
|
||||
fn snapshot_excludes_system_prompt_history_from_public_blocks() {
|
||||
let mut app = App::new("test".into());
|
||||
let session_start = session_store::LogEntry::SegmentStart {
|
||||
let session_start = session_store::LogEntry::AnnotatedSegmentStart {
|
||||
ts: 1,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
config: Default::default(),
|
||||
history: vec![session_store::LoggedItem::from(
|
||||
&agen::Item::system_message("[File: src/main.rs]\nfn main() {}"),
|
||||
)],
|
||||
history: vec![annotated(agen::Item::system_message(
|
||||
"[File: src/main.rs]\nfn main() {}",
|
||||
))],
|
||||
forked_from: None,
|
||||
compacted_from: None,
|
||||
};
|
||||
let session_start_value = serde_json::to_value(&session_start).unwrap();
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: vec![session_start_value],
|
||||
session: public_session(vec![session_start_value]),
|
||||
status: WorkerStatus::Running,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -3546,10 +3531,8 @@ mod completion_flow_tests {
|
||||
|
||||
assert!(matches!(app.worker_status, WorkerStatus::Running));
|
||||
assert!(app.running);
|
||||
assert!(matches!(
|
||||
app.blocks.get(1),
|
||||
Some(Block::SystemMessage { text }) if text == "[File: src/main.rs]\nfn main() {}"
|
||||
));
|
||||
assert_eq!(app.blocks.len(), 1);
|
||||
assert!(matches!(app.blocks.first(), Some(Block::Greeting(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -3584,7 +3567,7 @@ mod completion_flow_tests {
|
||||
};
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: vec![serde_json::to_value(run_errored).unwrap()],
|
||||
session: public_session(vec![serde_json::to_value(run_errored).unwrap()]),
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -3612,7 +3595,7 @@ mod completion_flow_tests {
|
||||
code: ErrorCode::ProviderError,
|
||||
message: "provider unavailable".into(),
|
||||
});
|
||||
let segment_start = session_store::LogEntry::SegmentStart {
|
||||
let segment_start = session_store::LogEntry::AnnotatedSegmentStart {
|
||||
ts: 5,
|
||||
session_id: uuid::Uuid::nil(),
|
||||
system_prompt: None,
|
||||
@@ -3622,7 +3605,7 @@ mod completion_flow_tests {
|
||||
compacted_from: None,
|
||||
};
|
||||
app.handle_worker_event(Event::SegmentRotated {
|
||||
entry: serde_json::to_value(segment_start).unwrap(),
|
||||
session: public_session(vec![serde_json::to_value(segment_start).unwrap()]),
|
||||
});
|
||||
|
||||
let errors = app
|
||||
@@ -3645,7 +3628,9 @@ mod completion_flow_tests {
|
||||
let mut app = App::new("test".into());
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Running,
|
||||
in_flight: InFlightSnapshot {
|
||||
blocks: vec![
|
||||
@@ -3751,7 +3736,9 @@ mod completion_flow_tests {
|
||||
},
|
||||
revision,
|
||||
status: WorkerStatus::Idle,
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
in_flight: protocol::InFlightSnapshot::default(),
|
||||
error: None,
|
||||
internal_workers: Vec::new(),
|
||||
@@ -3966,7 +3953,9 @@ mod completion_flow_tests {
|
||||
assert_eq!(app.selected_worker_view().worker_name, "parent");
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -4015,7 +4004,9 @@ mod completion_flow_tests {
|
||||
});
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: vec![InternalWorkerSnapshot {
|
||||
@@ -4026,7 +4017,9 @@ mod completion_flow_tests {
|
||||
kind: protocol::InternalWorkerKind::SubWorker,
|
||||
},
|
||||
revision: 4,
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Running,
|
||||
error: None,
|
||||
in_flight: Default::default(),
|
||||
@@ -4076,13 +4069,34 @@ mod completion_flow_tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn test_compaction_lifecycle(
|
||||
state: protocol::CompactionLifecycleState,
|
||||
) -> protocol::CompactionLifecycle {
|
||||
protocol::CompactionLifecycle {
|
||||
schema_version: 2,
|
||||
compaction_id: "compaction-test".into(),
|
||||
revision: 1,
|
||||
internal_worker: None,
|
||||
state,
|
||||
started_at_ms: 1,
|
||||
ended_at_ms: None,
|
||||
summary: None,
|
||||
error: None,
|
||||
new_segment_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compact_done_replaces_live_block() {
|
||||
let mut app = App::new("test".into());
|
||||
let id = uuid::Uuid::parse_str("12345678-1234-5678-1234-567812345678").unwrap();
|
||||
|
||||
app.handle_worker_event(Event::CompactStart);
|
||||
app.handle_worker_event(Event::CompactDone { new_segment_id: id });
|
||||
app.handle_worker_event(Event::CompactStart {
|
||||
lifecycle: test_compaction_lifecycle(protocol::CompactionLifecycleState::Running),
|
||||
});
|
||||
let mut lifecycle = test_compaction_lifecycle(protocol::CompactionLifecycleState::Done);
|
||||
lifecycle.new_segment_id = Some(id.to_string());
|
||||
app.handle_worker_event(Event::CompactDone { lifecycle });
|
||||
|
||||
assert_eq!(compact_block_count(&app), 1);
|
||||
assert!(matches!(
|
||||
@@ -4098,10 +4112,12 @@ mod completion_flow_tests {
|
||||
fn compact_failed_replaces_live_block() {
|
||||
let mut app = App::new("test".into());
|
||||
|
||||
app.handle_worker_event(Event::CompactStart);
|
||||
app.handle_worker_event(Event::CompactFailed {
|
||||
error: "provider 429".into(),
|
||||
app.handle_worker_event(Event::CompactStart {
|
||||
lifecycle: test_compaction_lifecycle(protocol::CompactionLifecycleState::Running),
|
||||
});
|
||||
let mut lifecycle = test_compaction_lifecycle(protocol::CompactionLifecycleState::Failed);
|
||||
lifecycle.error = Some("provider 429".into());
|
||||
app.handle_worker_event(Event::CompactFailed { lifecycle });
|
||||
|
||||
assert_eq!(compact_block_count(&app), 1);
|
||||
assert!(matches!(
|
||||
@@ -4117,7 +4133,9 @@ mod completion_flow_tests {
|
||||
fn shutdown_marks_live_compact_incomplete() {
|
||||
let mut app = App::new("test".into());
|
||||
|
||||
app.handle_worker_event(Event::CompactStart);
|
||||
app.handle_worker_event(Event::CompactStart {
|
||||
lifecycle: test_compaction_lifecycle(protocol::CompactionLifecycleState::Running),
|
||||
});
|
||||
app.handle_worker_event(Event::Shutdown);
|
||||
|
||||
assert!(app.quit);
|
||||
@@ -4157,7 +4175,9 @@ mod completion_flow_tests {
|
||||
greeting.context_tokens = 45_000;
|
||||
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
greeting,
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
@@ -4208,9 +4228,9 @@ mod completion_flow_tests {
|
||||
let mut app = App::new("test".into());
|
||||
app.session_context_tokens = 42_000;
|
||||
|
||||
app.handle_worker_event(Event::CompactDone {
|
||||
new_segment_id: uuid::Uuid::nil(),
|
||||
});
|
||||
let mut lifecycle = test_compaction_lifecycle(protocol::CompactionLifecycleState::Done);
|
||||
lifecycle.new_segment_id = Some(uuid::Uuid::nil().to_string());
|
||||
app.handle_worker_event(Event::CompactDone { lifecycle });
|
||||
|
||||
assert_eq!(app.session_context_tokens, 0);
|
||||
}
|
||||
@@ -4327,40 +4347,37 @@ mod completion_flow_tests {
|
||||
});
|
||||
|
||||
let assistant_item_entries = vec![
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 1,
|
||||
"item": {
|
||||
"kind": "tool_call",
|
||||
"call_id": "c1",
|
||||
"name": "TaskCreate",
|
||||
"arguments": r#"{"subject":"a","description":"A"}"#,
|
||||
},
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 2,
|
||||
"item": {
|
||||
"kind": "tool_call",
|
||||
"call_id": "c2",
|
||||
"name": "TaskCreate",
|
||||
"arguments": r#"{"subject":"b","description":"B"}"#,
|
||||
},
|
||||
}),
|
||||
serde_json::json!({
|
||||
"kind": "assistant_item",
|
||||
"ts": 3,
|
||||
"item": {
|
||||
"kind": "tool_call",
|
||||
"call_id": "u1",
|
||||
"name": "TaskUpdate",
|
||||
"arguments": r#"{"taskid":2,"status":"inprogress"}"#,
|
||||
},
|
||||
}),
|
||||
serde_json::to_value(session_store::LogEntry::AnnotatedAssistantItem {
|
||||
ts: 1,
|
||||
entry: annotated(agen::Item::tool_call(
|
||||
"c1",
|
||||
"TaskCreate",
|
||||
r#"{"subject":"a","description":"A"}"#,
|
||||
)),
|
||||
})
|
||||
.unwrap(),
|
||||
serde_json::to_value(session_store::LogEntry::AnnotatedAssistantItem {
|
||||
ts: 2,
|
||||
entry: annotated(agen::Item::tool_call(
|
||||
"c2",
|
||||
"TaskCreate",
|
||||
r#"{"subject":"b","description":"B"}"#,
|
||||
)),
|
||||
})
|
||||
.unwrap(),
|
||||
serde_json::to_value(session_store::LogEntry::AnnotatedAssistantItem {
|
||||
ts: 3,
|
||||
entry: annotated(agen::Item::tool_call(
|
||||
"u1",
|
||||
"TaskUpdate",
|
||||
r#"{"taskid":2,"status":"inprogress"}"#,
|
||||
)),
|
||||
})
|
||||
.unwrap(),
|
||||
];
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: assistant_item_entries,
|
||||
session: public_session(assistant_item_entries),
|
||||
status: WorkerStatus::Running,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
|
||||
+164
-329
@@ -1,8 +1,7 @@
|
||||
use std::error::Error;
|
||||
use std::fmt;
|
||||
use std::future::Future;
|
||||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, Ordering},
|
||||
@@ -21,26 +20,18 @@ use protocol::{Event, Method, WorkerStatus};
|
||||
use protocol::{Greeting, RewindSummary, RewindTarget, RewindTargetId, Segment};
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use session_store::SegmentId;
|
||||
use tokio::sync::mpsc;
|
||||
use standalone::{StandaloneHost, StandaloneLaunchConfig};
|
||||
use tokio::sync::{broadcast, mpsc};
|
||||
|
||||
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD};
|
||||
use client::{BackendRuntimeClient, BackendRuntimeTarget, WorkerClient, WorkerRuntimeCommand};
|
||||
use client::{BackendRuntimeClient, BackendRuntimeTarget, StandaloneSessionResumeIntent};
|
||||
|
||||
use crate::app::{ActionbarNoticeLevel, ActionbarNoticeSource, App};
|
||||
use crate::composer_keys::{ComposerEditAction, composer_edit_action};
|
||||
use crate::picker::PickerOutcome;
|
||||
use crate::spawn::{SpawnOutcome, SpawnReady};
|
||||
use crate::{picker, spawn, ui};
|
||||
use crate::ui;
|
||||
|
||||
pub(crate) type ConsoleTerminal = Terminal<CrosstermBackend<io::Stdout>>;
|
||||
|
||||
/// Narrow request bridge used when the workspace Dashboard opens a Worker Console.
|
||||
pub(crate) struct DashboardConsoleOpenRequest {
|
||||
pub(crate) worker_name: String,
|
||||
pub(crate) socket_override: Option<PathBuf>,
|
||||
}
|
||||
|
||||
/// Enable SGR coordinates plus normal mouse tracking. This captures clicks,
|
||||
/// releases, and wheel events without drag-capture modes (`?1002h`/`?1003h`)
|
||||
/// so terminal-native drag selection remains available during startup.
|
||||
@@ -128,75 +119,161 @@ fn copy_selection_to_terminal(app: &mut App) -> bool {
|
||||
copy_selection_to_writer(app, &mut stdout)
|
||||
}
|
||||
|
||||
fn resolve_socket(worker_name: &str, override_path: Option<PathBuf>) -> PathBuf {
|
||||
if let Some(p) = override_path {
|
||||
return p;
|
||||
}
|
||||
manifest::paths::worker_socket_path(worker_name).unwrap_or_else(|| {
|
||||
PathBuf::from("/tmp")
|
||||
.join("yoi")
|
||||
.join(worker_name)
|
||||
.join("sock")
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn run_worker_name(
|
||||
worker_name: String,
|
||||
socket_override: Option<PathBuf>,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
#[cfg(feature = "e2e-test")]
|
||||
if std::env::var_os("YOI_TUI_TEST_REWIND_FIXTURE").is_some() {
|
||||
let mut terminal = enter_fullscreen()?;
|
||||
terminal.clear()?;
|
||||
let result = run_e2e_rewind_fixture(&mut terminal, worker_name).await;
|
||||
let _ = leave_fullscreen(&mut terminal);
|
||||
return result;
|
||||
}
|
||||
|
||||
if let Some(client) = try_connect_live_pod(&worker_name, socket_override.clone()).await {
|
||||
let mut terminal = enter_fullscreen()?;
|
||||
run_connected_pod(&mut terminal, worker_name, client, runtime_command.clone()).await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let ready = match spawn::run_worker_name(worker_name, runtime_command.clone()).await? {
|
||||
SpawnOutcome::Ready(r) => r,
|
||||
SpawnOutcome::Cancelled => return Ok(()),
|
||||
};
|
||||
let mut terminal = enter_fullscreen()?;
|
||||
terminal.clear()?;
|
||||
let result = run_ready_pod(&mut terminal, ready, runtime_command).await;
|
||||
let _ = leave_fullscreen(&mut terminal);
|
||||
result
|
||||
}
|
||||
|
||||
enum ConsoleConnection {
|
||||
LegacySocket(WorkerClient),
|
||||
BackendRuntime(BackendRuntimeClient),
|
||||
Standalone {
|
||||
host: Option<StandaloneHost>,
|
||||
events: broadcast::Receiver<Event>,
|
||||
initial_snapshot: Option<Event>,
|
||||
},
|
||||
}
|
||||
|
||||
impl ConsoleConnection {
|
||||
fn standalone(host: StandaloneHost) -> Self {
|
||||
let events = host.subscribe();
|
||||
let initial_snapshot = Some(host.snapshot());
|
||||
Self::Standalone {
|
||||
host: Some(host),
|
||||
events,
|
||||
initial_snapshot,
|
||||
}
|
||||
}
|
||||
|
||||
fn try_next_event(&mut self) -> Option<Event> {
|
||||
match self {
|
||||
Self::LegacySocket(client) => client.try_next_event(),
|
||||
Self::BackendRuntime(client) => client.try_next_event(),
|
||||
Self::Standalone {
|
||||
events,
|
||||
initial_snapshot,
|
||||
..
|
||||
} => initial_snapshot.take().or_else(|| events.try_recv().ok()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn next_event(&mut self) -> Option<Event> {
|
||||
match self {
|
||||
Self::LegacySocket(client) => client.next_event().await,
|
||||
Self::BackendRuntime(client) => client.next_event().await,
|
||||
Self::Standalone { host, events, .. } => loop {
|
||||
match events.recv().await {
|
||||
Ok(event) => break Some(event),
|
||||
Err(broadcast::error::RecvError::Lagged(_)) => {
|
||||
let Some(host) = host.as_ref() else {
|
||||
break None;
|
||||
};
|
||||
break Some(host.snapshot());
|
||||
}
|
||||
Err(broadcast::error::RecvError::Closed) => break None,
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
async fn send(&mut self, method: &Method) -> Result<(), Box<dyn std::error::Error>> {
|
||||
match self {
|
||||
Self::LegacySocket(client) => Ok(client.send(method).await?),
|
||||
Self::BackendRuntime(client) => Ok(client.send(method).await?),
|
||||
Self::Standalone { host, .. } => {
|
||||
let host = host.as_ref().ok_or_else(|| {
|
||||
io::Error::new(
|
||||
io::ErrorKind::BrokenPipe,
|
||||
"Standalone Worker has already shut down",
|
||||
)
|
||||
})?;
|
||||
Ok(host.send(method.clone()).await?)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn shutdown(&mut self) -> Result<(), Box<dyn std::error::Error>> {
|
||||
if let Self::Standalone { host, .. } = self
|
||||
&& let Some(host) = host.take()
|
||||
{
|
||||
host.shutdown().await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn run_standalone(
|
||||
workspace_root: PathBuf,
|
||||
state_dir: PathBuf,
|
||||
worker_name: Option<String>,
|
||||
profile: Option<String>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let worker_name = worker_name.unwrap_or_else(|| "local".to_string());
|
||||
let profile = profile.map_or(manifest::ProfileSelector::Default, |profile| {
|
||||
manifest::ProfileSelector::parse_cli(&profile)
|
||||
});
|
||||
let history_root = workspace_root.clone();
|
||||
let launch = StandaloneLaunchConfig {
|
||||
state_dir,
|
||||
cwd: workspace_root,
|
||||
profile,
|
||||
worker_name: worker_name.clone(),
|
||||
}
|
||||
.resolve()
|
||||
.map_err(|error| {
|
||||
io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!("Standalone launch configuration failed: {error}"),
|
||||
)
|
||||
})?;
|
||||
let host = StandaloneHost::start(launch)
|
||||
.await
|
||||
.map_err(|error| io::Error::other(format!("Standalone Worker startup failed: {error}")))?;
|
||||
run_standalone_host(host, worker_name, history_root).await
|
||||
}
|
||||
|
||||
pub(crate) async fn run_standalone_restore(
|
||||
intent: StandaloneSessionResumeIntent,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let session_id = intent.session_id.parse().map_err(|error| {
|
||||
io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
format!("Invalid standalone session ID: {error}"),
|
||||
)
|
||||
})?;
|
||||
let host = StandaloneHost::restore(intent.state_dir, session_id)
|
||||
.await
|
||||
.map_err(|error| io::Error::other(format!("Standalone restore failed: {error}")))?;
|
||||
let worker_label = format!("standalone-{}", session_id.short());
|
||||
let history_root = host.record().cwd.canonical_path.clone();
|
||||
run_standalone_host(host, worker_label, history_root).await
|
||||
}
|
||||
|
||||
fn standalone_console_app(worker_label: String, history_root: &Path) -> App {
|
||||
let mut app = App::new_with_persistent_input_history(worker_label, history_root);
|
||||
app.connected = true;
|
||||
app
|
||||
}
|
||||
|
||||
async fn run_standalone_host(
|
||||
host: StandaloneHost,
|
||||
worker_label: String,
|
||||
history_root: PathBuf,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let mut connection = ConsoleConnection::standalone(host);
|
||||
|
||||
let mut terminal = match enter_fullscreen() {
|
||||
Ok(terminal) => terminal,
|
||||
Err(error) => {
|
||||
let _ = connection.shutdown().await;
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let mut app = standalone_console_app(worker_label, &history_root);
|
||||
let run_result = run_loop(&mut terminal, &mut app, &mut connection).await;
|
||||
let shutdown_result = connection
|
||||
.shutdown()
|
||||
.await
|
||||
.map_err(|error| io::Error::other(format!("Standalone Worker shutdown failed: {error}")));
|
||||
let leave_result = leave_fullscreen(&mut terminal);
|
||||
|
||||
if let Err(error) = run_result {
|
||||
return Err(error);
|
||||
}
|
||||
shutdown_result?;
|
||||
leave_result?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) async fn run_backend_runtime(
|
||||
@@ -208,201 +285,12 @@ pub(crate) async fn run_backend_runtime(
|
||||
let workspace_root = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."));
|
||||
let mut app = App::new_with_persistent_input_history(worker_label, &workspace_root);
|
||||
app.connected = true;
|
||||
let result = run_loop(
|
||||
&mut terminal,
|
||||
&mut app,
|
||||
ConsoleConnection::BackendRuntime(client),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
let mut connection = ConsoleConnection::BackendRuntime(client);
|
||||
let result = run_loop(&mut terminal, &mut app, &mut connection).await;
|
||||
let _ = leave_fullscreen(&mut terminal);
|
||||
result
|
||||
}
|
||||
|
||||
async fn run_connected_pod(
|
||||
terminal: &mut ConsoleTerminal,
|
||||
worker_name: String,
|
||||
client: WorkerClient,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let workspace_root = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."));
|
||||
let mut app = App::new_with_persistent_input_history(worker_name, &workspace_root);
|
||||
app.connected = true;
|
||||
run_loop(
|
||||
terminal,
|
||||
&mut app,
|
||||
ConsoleConnection::LegacySocket(client),
|
||||
Some(runtime_command),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn open_from_dashboard(
|
||||
terminal: &mut ConsoleTerminal,
|
||||
request: DashboardConsoleOpenRequest,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let DashboardConsoleOpenRequest {
|
||||
worker_name,
|
||||
socket_override,
|
||||
} = request;
|
||||
|
||||
if let Some(client) = try_connect_live_pod(&worker_name, socket_override).await {
|
||||
return run_connected_pod(terminal, worker_name, client, runtime_command.clone()).await;
|
||||
}
|
||||
|
||||
let ready =
|
||||
spawn_worker_name_from_fullscreen(terminal, &worker_name, runtime_command.clone()).await?;
|
||||
run_ready_pod(terminal, ready, runtime_command).await
|
||||
}
|
||||
|
||||
async fn spawn_worker_name_from_fullscreen(
|
||||
terminal: &mut ConsoleTerminal,
|
||||
worker_name: &str,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<SpawnReady, Box<dyn std::error::Error>> {
|
||||
leave_fullscreen(terminal)?;
|
||||
let outcome = spawn::run_worker_name(worker_name.to_string(), runtime_command).await;
|
||||
enter_fullscreen_existing(terminal)?;
|
||||
terminal.clear()?;
|
||||
|
||||
match outcome? {
|
||||
SpawnOutcome::Ready(ready) => Ok(ready),
|
||||
SpawnOutcome::Cancelled => Err(Box::new(NestedOpenCancelled)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn try_connect_live_pod(
|
||||
worker_name: &str,
|
||||
socket_override: Option<PathBuf>,
|
||||
) -> Option<WorkerClient> {
|
||||
let preferred_socket = resolve_socket(worker_name, socket_override.clone());
|
||||
connect_live_pod(worker_name, preferred_socket, socket_override.is_none())
|
||||
.await
|
||||
.map(|(_, client)| client)
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct NestedOpenCancelled;
|
||||
|
||||
impl std::fmt::Display for NestedOpenCancelled {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.write_str("Worker open was cancelled")
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for NestedOpenCancelled {}
|
||||
|
||||
async fn run_ready_pod(
|
||||
terminal: &mut ConsoleTerminal,
|
||||
ready: SpawnReady,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let SpawnReady {
|
||||
worker_name,
|
||||
socket_path,
|
||||
} = ready;
|
||||
run(terminal, worker_name, &socket_path, runtime_command).await
|
||||
}
|
||||
|
||||
async fn connect_live_pod(
|
||||
worker_name: &str,
|
||||
preferred_socket: PathBuf,
|
||||
allow_registry_fallback: bool,
|
||||
) -> Option<(PathBuf, WorkerClient)> {
|
||||
if let Ok(client) = WorkerClient::connect(&preferred_socket).await {
|
||||
return Some((preferred_socket, client));
|
||||
}
|
||||
|
||||
if !allow_registry_fallback {
|
||||
return None;
|
||||
}
|
||||
let registry_socket = picker::live_socket_for_worker(worker_name)?;
|
||||
if registry_socket == preferred_socket {
|
||||
return None;
|
||||
}
|
||||
WorkerClient::connect(®istry_socket)
|
||||
.await
|
||||
.ok()
|
||||
.map(|client| (registry_socket, client))
|
||||
}
|
||||
|
||||
pub(crate) async fn run_resume(
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
workspace_root: PathBuf,
|
||||
all: bool,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
run_worker_picker(runtime_command, workspace_root, all, true).await
|
||||
}
|
||||
|
||||
pub(crate) async fn run_worker_picker(
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
workspace_root: PathBuf,
|
||||
all: bool,
|
||||
include_stopped: bool,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
// Pick a Worker in its own inline viewport, dropping the viewport before
|
||||
// attaching/restoring so each phase gets fresh vertical room.
|
||||
let picker_options = if all {
|
||||
picker::PickerOptions::all()
|
||||
} else {
|
||||
picker::PickerOptions::workspace(workspace_root)
|
||||
}
|
||||
.with_stopped(include_stopped);
|
||||
let (worker_name, socket_override) = match picker::run(picker_options).await? {
|
||||
PickerOutcome::Picked {
|
||||
worker_name,
|
||||
socket_override,
|
||||
} => (worker_name, socket_override),
|
||||
PickerOutcome::Cancelled => return Ok(()),
|
||||
};
|
||||
run_worker_name(worker_name, socket_override, runtime_command).await
|
||||
}
|
||||
|
||||
pub(crate) fn is_recoverable_dashboard_open_error(error: &(dyn Error + 'static)) -> bool {
|
||||
error.is::<spawn::SpawnError>() || error.is::<NestedOpenCancelled>()
|
||||
}
|
||||
|
||||
pub(crate) async fn run_spawn(
|
||||
resume_from: Option<SegmentId>,
|
||||
worker_name: Option<String>,
|
||||
profile: Option<String>,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
#[cfg(feature = "e2e-test")]
|
||||
if std::env::var_os("YOI_TUI_TEST_REWIND_FIXTURE").is_some() {
|
||||
let mut terminal = enter_fullscreen()?;
|
||||
terminal.clear()?;
|
||||
let fixture_worker_name = worker_name.unwrap_or_else(|| "e2e-rewind".to_string());
|
||||
let result = run_e2e_rewind_fixture(&mut terminal, fixture_worker_name).await;
|
||||
let _ = leave_fullscreen(&mut terminal);
|
||||
return result;
|
||||
}
|
||||
|
||||
let ready = match spawn::run(resume_from, worker_name, profile, runtime_command.clone()).await?
|
||||
{
|
||||
SpawnOutcome::Ready(r) => r,
|
||||
SpawnOutcome::Cancelled => return Ok(()),
|
||||
};
|
||||
|
||||
let SpawnReady {
|
||||
worker_name,
|
||||
socket_path,
|
||||
} = ready;
|
||||
|
||||
let mut terminal = enter_fullscreen()?;
|
||||
let result = run(&mut terminal, worker_name, &socket_path, runtime_command).await;
|
||||
|
||||
// Leave alt-screen explicitly before `main`'s terminal restore path.
|
||||
let _ = execute!(
|
||||
terminal.backend_mut(),
|
||||
DisableMouseCapture,
|
||||
LeaveAlternateScreen
|
||||
);
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
fn enter_fullscreen() -> Result<ConsoleTerminal, Box<dyn std::error::Error>> {
|
||||
let mut stdout = io::stdout();
|
||||
// Enable button-event tracking so the transcript can own drag selection;
|
||||
@@ -421,19 +309,6 @@ pub(crate) fn enter_dashboard_fullscreen() -> Result<ConsoleTerminal, Box<dyn st
|
||||
Ok(Terminal::new(backend)?)
|
||||
}
|
||||
|
||||
fn enter_fullscreen_existing(
|
||||
terminal: &mut ConsoleTerminal,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
// Re-enable the same least-intrusive wheel mouse mode after returning from
|
||||
// nested inline screens.
|
||||
execute!(
|
||||
terminal.backend_mut(),
|
||||
EnterAlternateScreen,
|
||||
EnableSinglePodMouseCapture
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn leave_fullscreen(terminal: &mut ConsoleTerminal) -> io::Result<()> {
|
||||
execute!(
|
||||
terminal.backend_mut(),
|
||||
@@ -446,40 +321,6 @@ pub(crate) fn leave_dashboard_fullscreen(terminal: &mut ConsoleTerminal) -> io::
|
||||
leave_fullscreen(terminal)
|
||||
}
|
||||
|
||||
async fn run(
|
||||
terminal: &mut ConsoleTerminal,
|
||||
worker_name: String,
|
||||
socket_path: &std::path::Path,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let workspace_root = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."));
|
||||
let mut app = App::new_with_persistent_input_history(worker_name, &workspace_root);
|
||||
|
||||
match WorkerClient::connect(socket_path).await {
|
||||
Ok(client) => {
|
||||
app.connected = true;
|
||||
// The Worker sends `Event::Snapshot` automatically on connect;
|
||||
// no explicit method call is required to fetch history.
|
||||
run_loop(
|
||||
terminal,
|
||||
&mut app,
|
||||
ConsoleConnection::LegacySocket(client),
|
||||
Some(runtime_command),
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
Err(e) => {
|
||||
app.push_error(format!(
|
||||
"Failed to connect to {}: {e}",
|
||||
socket_path.display()
|
||||
));
|
||||
terminal.draw(|f| ui::draw(f, &mut app))?;
|
||||
run_disconnected(&mut app)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
type TerminalEventResult = io::Result<TermEvent>;
|
||||
|
||||
const TERMINAL_POLL_INTERVAL: Duration = Duration::from_millis(50);
|
||||
@@ -547,7 +388,9 @@ async fn run_e2e_rewind_fixture(
|
||||
let mut app = App::new_with_persistent_input_history(worker_name.clone(), &workspace_root);
|
||||
app.connected = true;
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
status: WorkerStatus::Idle,
|
||||
greeting: Greeting {
|
||||
worker_name: worker_name.clone(),
|
||||
@@ -673,7 +516,9 @@ async fn run_e2e_rewind_fixture(
|
||||
if let Some(submitted_at) = pending_apply {
|
||||
if submitted_at.elapsed() >= apply_delay {
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
input: vec![Segment::text("rewind-live-refresh")],
|
||||
summary: RewindSummary {
|
||||
truncated_to_entries: 1,
|
||||
@@ -745,14 +590,13 @@ async fn drain_terminal_events(
|
||||
app: &mut App,
|
||||
client: &mut ConsoleConnection,
|
||||
term_rx: &mut mpsc::UnboundedReceiver<TerminalEventResult>,
|
||||
runtime_command: Option<&WorkerRuntimeCommand>,
|
||||
) -> Result<bool, Box<dyn std::error::Error>> {
|
||||
let mut handled = false;
|
||||
for _ in 0..TERMINAL_EVENT_DRAIN_LIMIT {
|
||||
match term_rx.try_recv() {
|
||||
Ok(event) => {
|
||||
handled = true;
|
||||
handle_terminal_event(app, client, event?, runtime_command).await?;
|
||||
handle_terminal_event(app, client, event?).await?;
|
||||
if app.quit {
|
||||
break;
|
||||
}
|
||||
@@ -791,8 +635,7 @@ async fn drain_worker_events(
|
||||
async fn run_loop(
|
||||
terminal: &mut Terminal<CrosstermBackend<io::Stdout>>,
|
||||
app: &mut App,
|
||||
mut client: ConsoleConnection,
|
||||
runtime_command: Option<WorkerRuntimeCommand>,
|
||||
client: &mut ConsoleConnection,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let (_terminal_reader, mut term_rx) = TerminalEventReader::spawn()?;
|
||||
|
||||
@@ -803,12 +646,11 @@ async fn run_loop(
|
||||
break;
|
||||
}
|
||||
|
||||
let handled_term_event =
|
||||
drain_terminal_events(app, &mut client, &mut term_rx, runtime_command.as_ref()).await?;
|
||||
let handled_term_event = drain_terminal_events(app, client, &mut term_rx).await?;
|
||||
if app.quit {
|
||||
break;
|
||||
}
|
||||
let handled_worker_event = drain_worker_events(app, &mut client).await?;
|
||||
let handled_worker_event = drain_worker_events(app, client).await?;
|
||||
if handled_term_event || handled_worker_event {
|
||||
terminal.draw(|f| ui::draw(f, app))?;
|
||||
continue;
|
||||
@@ -816,8 +658,7 @@ async fn run_loop(
|
||||
|
||||
match next_loop_input(&mut term_rx, app.connected, client.next_event()).await {
|
||||
LoopInput::Terminal(term_event) => {
|
||||
handle_terminal_event(app, &mut client, term_event?, runtime_command.as_ref())
|
||||
.await?;
|
||||
handle_terminal_event(app, client, term_event?).await?;
|
||||
}
|
||||
LoopInput::Worker(event) => match event {
|
||||
Some(ev) => {
|
||||
@@ -843,7 +684,6 @@ async fn handle_terminal_event(
|
||||
app: &mut App,
|
||||
client: &mut ConsoleConnection,
|
||||
event: TermEvent,
|
||||
_runtime_command: Option<&WorkerRuntimeCommand>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
match event {
|
||||
TermEvent::Key(key) => {
|
||||
@@ -865,19 +705,6 @@ async fn handle_terminal_event(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn run_disconnected(_app: &mut App) -> Result<(), Box<dyn std::error::Error>> {
|
||||
loop {
|
||||
if event::poll(std::time::Duration::from_millis(100))?
|
||||
&& let TermEvent::Key(key) = event::read()?
|
||||
&& let KeyCode::Char('c') = key.code
|
||||
&& key.modifiers.contains(KeyModifiers::CONTROL)
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Lines per wheel notch. Faster than Shift+↑/↓ (which is 1 line) so
|
||||
/// hand-rolling through long histories isn't tedious, but slow enough
|
||||
/// that a single notch doesn't blow past the section the user is
|
||||
@@ -1016,7 +843,7 @@ fn handle_key(app: &mut App, key: KeyEvent) -> Option<Method> {
|
||||
app.clear_queued_inputs();
|
||||
Some(Method::Cancel)
|
||||
}
|
||||
WorkerStatus::Idle => Some(Method::Shutdown),
|
||||
WorkerStatus::Idle | WorkerStatus::Stopped => Some(Method::Shutdown),
|
||||
}),
|
||||
KeyCode::Char('d') if ctrl => {
|
||||
app.quit = true;
|
||||
@@ -1304,6 +1131,14 @@ mod tests {
|
||||
use crate::text_selection::{HistoryViewport, SelectionRow};
|
||||
use protocol::{Event, RewindTarget, RewindTargetId, Segment};
|
||||
|
||||
#[test]
|
||||
fn standalone_console_starts_with_in_process_connection_ready() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let app = standalone_console_app("standalone".to_string(), temp.path());
|
||||
|
||||
assert!(app.connected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_worker_mouse_capture_avoids_drag_and_all_motion_modes() {
|
||||
let mut ansi = String::new();
|
||||
@@ -2023,13 +1858,13 @@ mod tests {
|
||||
let mut app = App::new("agent".to_string());
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
});
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::Text {
|
||||
content: "retry this".into(),
|
||||
}],
|
||||
@@ -2050,7 +1885,7 @@ mod tests {
|
||||
let mut app = App::new("agent".to_string());
|
||||
app.handle_worker_event(Event::Snapshot {
|
||||
greeting: test_greeting(),
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
status: WorkerStatus::Idle,
|
||||
in_flight: Default::default(),
|
||||
internal_workers: Vec::new(),
|
||||
@@ -2058,7 +1893,7 @@ mod tests {
|
||||
type_keys(&mut app, "draft");
|
||||
|
||||
app.handle_worker_event(Event::RewindApplied {
|
||||
entries: vec![],
|
||||
session: protocol::SessionSnapshot { entries: vec![] },
|
||||
input: vec![Segment::Text {
|
||||
content: "retry this".into(),
|
||||
}],
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,784 +0,0 @@
|
||||
use super::*;
|
||||
|
||||
pub(super) fn draw(frame: &mut Frame<'_>, app: &mut DashboardApp) {
|
||||
let area = frame.area();
|
||||
let input_content_width = area.width.saturating_sub(2).max(1);
|
||||
let mut input_render = app.input.render(input_content_width);
|
||||
let input_height = input_area_height(&input_render, area.height);
|
||||
app.input
|
||||
.apply_cursor_viewport(&mut input_render, input_height);
|
||||
let layout = dashboard_layout(area, input_height);
|
||||
|
||||
draw_title(frame, app, layout.title);
|
||||
draw_list(frame, app, layout.list);
|
||||
draw_separator(frame, layout.boundary);
|
||||
draw_target_status(frame, app, layout.target_status);
|
||||
draw_input(frame, &input_render, layout.input);
|
||||
draw_actionbar(frame, app, layout.actionbar);
|
||||
if app.panel_diagnostic_open {
|
||||
render_panel_diagnostic(frame, app, area);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_diagnostic_area(area: Rect) -> Rect {
|
||||
let width = if area.width <= 20 {
|
||||
area.width
|
||||
} else {
|
||||
area.width.saturating_sub(4).min(100).max(20)
|
||||
};
|
||||
let height = if area.height <= 8 {
|
||||
area.height
|
||||
} else {
|
||||
area.height.saturating_sub(4).min(24).max(8)
|
||||
};
|
||||
let x = area.x + area.width.saturating_sub(width) / 2;
|
||||
let y = area.y + area.height.saturating_sub(height) / 2;
|
||||
Rect::new(x, y, width, height)
|
||||
}
|
||||
|
||||
pub(super) fn render_panel_diagnostic(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
let Some(diagnostic) = app.panel_diagnostic.as_ref() else {
|
||||
return;
|
||||
};
|
||||
let popup_area = panel_diagnostic_area(area);
|
||||
let title = format!(" {} ", diagnostic.title);
|
||||
let text = format!("{}\n\nF2/Esc: close", diagnostic.details);
|
||||
let paragraph = Paragraph::new(text)
|
||||
.block(Block::default().title(title).borders(Borders::ALL))
|
||||
.wrap(Wrap { trim: false });
|
||||
frame.render_widget(Clear, popup_area);
|
||||
frame.render_widget(paragraph, popup_area);
|
||||
}
|
||||
|
||||
pub(super) fn input_area_height(render: &crate::input::InputRender, terminal_height: u16) -> u16 {
|
||||
let needed = render.lines.len().max(1) as u16;
|
||||
let cap = (terminal_height / 3).max(1).min(10);
|
||||
needed.clamp(1, cap)
|
||||
}
|
||||
|
||||
pub(super) fn draw_title(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
frame.render_widget(Paragraph::new(title_line(app)), area);
|
||||
}
|
||||
|
||||
pub(super) fn title_line(app: &DashboardApp) -> Line<'static> {
|
||||
let mut spans = vec![Span::styled(
|
||||
"workspace dashboard",
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
)];
|
||||
if let Some(companion) = &app.panel.header.companion {
|
||||
spans.push(Span::styled(
|
||||
" · companion ",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
spans.push(Span::styled(
|
||||
companion.status.label(),
|
||||
companion_status_style(companion.status),
|
||||
));
|
||||
if let Some(detail) = companion.detail.as_deref() {
|
||||
spans.push(Span::styled(
|
||||
format!(" ({detail})"),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(orchestrator) = &app.panel.header.orchestrator {
|
||||
spans.push(Span::styled(
|
||||
" · orchestrator ",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
));
|
||||
spans.push(Span::styled(
|
||||
orchestrator.status.label(),
|
||||
orchestrator_status_style(orchestrator.status),
|
||||
));
|
||||
}
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn companion_status_style(status: CompanionPanelStatus) -> Style {
|
||||
match status {
|
||||
CompanionPanelStatus::Live
|
||||
| CompanionPanelStatus::Restored
|
||||
| CompanionPanelStatus::Spawned => Style::default().fg(Color::Green),
|
||||
CompanionPanelStatus::Stopped | CompanionPanelStatus::Missing => {
|
||||
Style::default().fg(Color::Yellow)
|
||||
}
|
||||
CompanionPanelStatus::Unavailable => Style::default().fg(Color::Red),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn orchestrator_status_style(status: OrchestratorPanelStatus) -> Style {
|
||||
match status {
|
||||
OrchestratorPanelStatus::Live
|
||||
| OrchestratorPanelStatus::Restored
|
||||
| OrchestratorPanelStatus::Spawned => Style::default().fg(Color::Green),
|
||||
OrchestratorPanelStatus::Stopped | OrchestratorPanelStatus::Missing => {
|
||||
Style::default().fg(Color::Yellow)
|
||||
}
|
||||
OrchestratorPanelStatus::Unavailable => Style::default().fg(Color::Red),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn draw_list(frame: &mut Frame<'_>, app: &mut DashboardApp, area: Rect) {
|
||||
if area.width == 0 || area.height == 0 {
|
||||
app.row_hit_boxes.clear();
|
||||
return;
|
||||
}
|
||||
let rows = list_rows(app, area.width, area.height);
|
||||
app.set_row_hit_boxes(&rows, area);
|
||||
let lines = rows.into_iter().map(|row| row.line).collect::<Vec<_>>();
|
||||
Paragraph::new(lines).render(area, frame.buffer_mut());
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(super) struct PanelListRow {
|
||||
pub(super) line: Line<'static>,
|
||||
pub(super) key: Option<PanelRowKey>,
|
||||
}
|
||||
|
||||
impl PanelListRow {
|
||||
fn inert(line: Line<'static>) -> Self {
|
||||
Self { line, key: None }
|
||||
}
|
||||
|
||||
fn selectable(line: Line<'static>, key: PanelRowKey) -> Self {
|
||||
Self {
|
||||
line,
|
||||
key: Some(key),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(super) fn list_lines(app: &DashboardApp, width: u16, height: u16) -> Vec<Line<'static>> {
|
||||
list_rows(app, width, height)
|
||||
.into_iter()
|
||||
.map(|row| row.line)
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn list_rows(app: &DashboardApp, width: u16, height: u16) -> Vec<PanelListRow> {
|
||||
let sections = sectioned_entries(&app.list);
|
||||
let selected = app.selected_row.as_ref();
|
||||
let diagnostic_rows = panel_diagnostic_lines(&app.panel, width)
|
||||
.into_iter()
|
||||
.map(PanelListRow::inert)
|
||||
.collect::<Vec<_>>();
|
||||
let action_rows = panel_action_rows(&app.panel, selected, width);
|
||||
let live_rows = sections
|
||||
.iter()
|
||||
.filter(|section| section.kind != DashboardSectionKind::Closed)
|
||||
.flat_map(|section| section_rows(&app.list, section, selected, width))
|
||||
.collect::<Vec<_>>();
|
||||
let closed_rows = sections
|
||||
.iter()
|
||||
.find(|section| section.kind == DashboardSectionKind::Closed)
|
||||
.map(|section| section_rows(&app.list, section, selected, width))
|
||||
.unwrap_or_default();
|
||||
|
||||
let available = height as usize;
|
||||
let diagnostic_len = diagnostic_rows.len().min(available);
|
||||
let remaining_after_diagnostics = available.saturating_sub(diagnostic_len);
|
||||
let action_len = action_rows.len().min(remaining_after_diagnostics);
|
||||
let remaining_after_actions = remaining_after_diagnostics.saturating_sub(action_len);
|
||||
let closed_len = closed_rows.len().min(remaining_after_actions);
|
||||
let live_len = live_rows
|
||||
.len()
|
||||
.min(remaining_after_actions.saturating_sub(closed_len));
|
||||
let spacer_len = available.saturating_sub(diagnostic_len + action_len + live_len + closed_len);
|
||||
|
||||
let mut rows = Vec::with_capacity(available);
|
||||
rows.extend(diagnostic_rows.into_iter().take(diagnostic_len));
|
||||
rows.extend(action_rows.into_iter().take(action_len));
|
||||
rows.extend(live_rows.into_iter().take(live_len));
|
||||
rows.extend(
|
||||
std::iter::repeat_with(|| PanelListRow::inert(Line::from(Span::raw("")))).take(spacer_len),
|
||||
);
|
||||
rows.extend(closed_rows.into_iter().take(closed_len));
|
||||
rows
|
||||
}
|
||||
|
||||
pub(super) fn row_hit_boxes(rows: &[PanelListRow], area: Rect) -> Vec<PanelRowHitBox> {
|
||||
if area.width == 0 || area.height == 0 {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let mut hit_boxes: Vec<PanelRowHitBox> = Vec::new();
|
||||
for (offset, row) in rows.iter().enumerate() {
|
||||
let Some(key) = row.key.clone() else {
|
||||
continue;
|
||||
};
|
||||
let Some(y) = area.y.checked_add(offset as u16) else {
|
||||
continue;
|
||||
};
|
||||
if y >= area.y.saturating_add(area.height) {
|
||||
continue;
|
||||
}
|
||||
if let Some(last) = hit_boxes.last_mut() {
|
||||
if last.key == key
|
||||
&& last.rect.x == area.x
|
||||
&& last.rect.width == area.width
|
||||
&& last.rect.y.saturating_add(last.rect.height) == y
|
||||
{
|
||||
last.rect.height = last.rect.height.saturating_add(1);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
hit_boxes.push(PanelRowHitBox {
|
||||
rect: Rect::new(area.x, y, area.width, 1),
|
||||
key,
|
||||
});
|
||||
}
|
||||
hit_boxes
|
||||
}
|
||||
|
||||
pub(super) fn panel_diagnostic_lines(
|
||||
panel: &WorkspacePanelViewModel,
|
||||
width: u16,
|
||||
) -> Vec<Line<'static>> {
|
||||
panel
|
||||
.header
|
||||
.diagnostics
|
||||
.iter()
|
||||
.map(|diagnostic| {
|
||||
Line::from(vec![
|
||||
Span::styled("⚠ ", Style::default().fg(Color::Yellow)),
|
||||
Span::styled(
|
||||
truncate_with_ellipsis(diagnostic, width.saturating_sub(2) as usize),
|
||||
Style::default().fg(Color::Yellow),
|
||||
),
|
||||
])
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn panel_action_rows(
|
||||
panel: &WorkspacePanelViewModel,
|
||||
selected: Option<&PanelRowKey>,
|
||||
width: u16,
|
||||
) -> Vec<PanelListRow> {
|
||||
let rows = panel
|
||||
.rows
|
||||
.iter()
|
||||
.filter(|row| row.is_ticket_section_row())
|
||||
.collect::<Vec<_>>();
|
||||
if rows.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
let mut lines = Vec::with_capacity((rows.len() * 2) + 1);
|
||||
lines.push(PanelListRow::inert(panel_action_header_line(
|
||||
rows.len(),
|
||||
width,
|
||||
)));
|
||||
for row in rows {
|
||||
for line in panel_row_lines(row, selected == Some(&row.key), width) {
|
||||
lines.push(PanelListRow::selectable(line, row.key.clone()));
|
||||
}
|
||||
}
|
||||
lines
|
||||
}
|
||||
|
||||
pub(super) fn panel_action_header_line(total: usize, width: u16) -> Line<'static> {
|
||||
let detail = if total == 1 {
|
||||
" 1 row".to_string()
|
||||
} else {
|
||||
format!(" {total} rows")
|
||||
};
|
||||
let text = truncate_with_ellipsis(&format!("--tickets{detail}---"), width as usize);
|
||||
Line::from(Span::styled(
|
||||
text,
|
||||
Style::default()
|
||||
.fg(Color::DarkGray)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) const TICKET_STATE_COLUMN_WIDTH: usize = 10;
|
||||
pub(super) const POD_STATUS_COLUMN_WIDTH: usize = 18;
|
||||
|
||||
pub(super) fn panel_row_lines(row: &PanelRow, selected: bool, width: u16) -> Vec<Line<'static>> {
|
||||
if row.kind == PanelRowKind::TicketIntakeWorker {
|
||||
vec![panel_intake_child_line(row, selected, width)]
|
||||
} else {
|
||||
vec![
|
||||
panel_row_title_line(row, selected, width),
|
||||
panel_row_detail_line(row, selected, width),
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_row_title_line(row: &PanelRow, selected: bool, width: u16) -> Line<'static> {
|
||||
let title_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Magenta)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Magenta)
|
||||
};
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_ticket_primary_marker_span(&mut spans, selected, &mut remaining);
|
||||
push_column_span(
|
||||
&mut spans,
|
||||
&row.status,
|
||||
TICKET_STATE_COLUMN_WIDTH,
|
||||
panel_priority_style(row.priority),
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(&mut spans, row.title.as_str(), title_style, &mut remaining);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn panel_intake_child_line(row: &PanelRow, selected: bool, width: u16) -> Line<'static> {
|
||||
let title_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Cyan)
|
||||
};
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_intake_child_marker_span(&mut spans, selected, &mut remaining);
|
||||
push_column_span(
|
||||
&mut spans,
|
||||
&row.status,
|
||||
TICKET_STATE_COLUMN_WIDTH,
|
||||
intake_status_style(&row.status),
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(&mut spans, row.title.as_str(), title_style, &mut remaining);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn panel_row_detail_line(row: &PanelRow, selected: bool, width: u16) -> Line<'static> {
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_ticket_detail_marker_span(&mut spans, selected, &mut remaining);
|
||||
push_bounded_span(
|
||||
&mut spans,
|
||||
"meta ",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(
|
||||
&mut spans,
|
||||
&panel_ticket_detail(row),
|
||||
ticket_detail_style(row),
|
||||
&mut remaining,
|
||||
);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn push_ticket_primary_marker_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
selected: bool,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
let (marker, style) = if selected {
|
||||
(
|
||||
"▶ ",
|
||||
Style::default()
|
||||
.fg(Color::Magenta)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
)
|
||||
} else {
|
||||
(" ", Style::default().fg(Color::DarkGray))
|
||||
};
|
||||
push_bounded_span(spans, marker, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn push_ticket_detail_marker_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
selected: bool,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
let (marker, style) = if selected {
|
||||
(
|
||||
"│ ",
|
||||
Style::default()
|
||||
.fg(Color::Magenta)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
)
|
||||
} else {
|
||||
(" ", Style::default().fg(Color::DarkGray))
|
||||
};
|
||||
push_bounded_span(spans, marker, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn push_intake_child_marker_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
selected: bool,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
let (marker, style) = if selected {
|
||||
(
|
||||
" ▶ ",
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
)
|
||||
} else {
|
||||
(" └ ", Style::default().fg(Color::DarkGray))
|
||||
};
|
||||
push_bounded_span(spans, marker, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_detail(row: &PanelRow) -> String {
|
||||
if row.kind == PanelRowKind::InvalidTicket {
|
||||
let mut parts = vec![panel_ticket_reference(row), "Gate: unavailable".to_string()];
|
||||
if let Some(reason) = panel_ticket_reason(row) {
|
||||
parts.push(format!("Reason: {reason}"));
|
||||
}
|
||||
return parts.join(" · ");
|
||||
}
|
||||
|
||||
if row.kind == PanelRowKind::TicketIntakeWorker {
|
||||
let mut parts = row
|
||||
.subtitle
|
||||
.as_ref()
|
||||
.map(|subtitle| vec![subtitle.clone()])
|
||||
.unwrap_or_else(|| vec![panel_ticket_reference(row)]);
|
||||
if let Some(action) = row.next_action {
|
||||
parts.push(format!("Action: {}", action.label()));
|
||||
}
|
||||
if let Some(reason) = panel_ticket_reason(row) {
|
||||
parts.push(format!("Reason: {reason}"));
|
||||
}
|
||||
return parts.join(" · ");
|
||||
}
|
||||
|
||||
let mut parts = vec![panel_ticket_reference(row)];
|
||||
if let Some(overlay_detail) = panel_ticket_overlay_detail(row) {
|
||||
parts.push(overlay_detail);
|
||||
}
|
||||
if let Some(blocked_reason) = row
|
||||
.ticket
|
||||
.as_ref()
|
||||
.and_then(|ticket| ticket.blocked_reason.as_deref())
|
||||
{
|
||||
parts.push(format!("Gate: waiting for {blocked_reason}"));
|
||||
} else {
|
||||
parts.push("Gate: clear".to_string());
|
||||
}
|
||||
if let Some(action) = row.next_action {
|
||||
parts.push(format!(
|
||||
"Action: {}",
|
||||
panel_ticket_action_label(row, action)
|
||||
));
|
||||
}
|
||||
if let Some(reason) = panel_ticket_reason(row) {
|
||||
parts.push(format!("Reason: {reason}"));
|
||||
}
|
||||
parts.join(" · ")
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_action_label(row: &PanelRow, action: NextUserAction) -> &'static str {
|
||||
if action == NextUserAction::Wait
|
||||
&& row
|
||||
.ticket
|
||||
.as_ref()
|
||||
.and_then(|ticket| ticket.blocked_reason.as_ref())
|
||||
.is_some()
|
||||
{
|
||||
"queue disabled"
|
||||
} else {
|
||||
action.label()
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_overlay_detail(row: &PanelRow) -> Option<String> {
|
||||
let ticket = row.ticket.as_ref()?;
|
||||
let overlay = ticket.orchestration_overlay.as_ref()?;
|
||||
let mut detail = format!(
|
||||
"Overlay: local {} · {} {}",
|
||||
ticket.workflow_state.as_str(),
|
||||
overlay.source,
|
||||
overlay.workflow_state.as_str()
|
||||
);
|
||||
if matches!(
|
||||
overlay.workflow_state,
|
||||
TicketWorkflowState::Done | TicketWorkflowState::Closed
|
||||
) {
|
||||
detail.push_str(" · merge pending");
|
||||
}
|
||||
Some(detail)
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_reason(row: &PanelRow) -> Option<&str> {
|
||||
row.disabled_reason
|
||||
.as_deref()
|
||||
.or_else(|| row.key_hint.as_deref())
|
||||
}
|
||||
|
||||
pub(super) fn ticket_detail_style(row: &PanelRow) -> Style {
|
||||
if row.kind == PanelRowKind::InvalidTicket {
|
||||
return Style::default().fg(Color::Yellow);
|
||||
}
|
||||
if row
|
||||
.ticket
|
||||
.as_ref()
|
||||
.and_then(|ticket| ticket.blocked_reason.as_ref())
|
||||
.is_some()
|
||||
{
|
||||
Style::default().fg(Color::Yellow)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn panel_ticket_reference(row: &PanelRow) -> String {
|
||||
row.ticket
|
||||
.as_ref()
|
||||
.map(|ticket| {
|
||||
ticket
|
||||
.resource_key
|
||||
.clone()
|
||||
.unwrap_or_else(|| "resource key unavailable".to_string())
|
||||
})
|
||||
.unwrap_or_else(|| match &row.key {
|
||||
PanelRowKey::Ticket(id) | PanelRowKey::InvalidTicket(id) => id.clone(),
|
||||
PanelRowKey::TicketIntakeWorker { ticket_id, .. } => ticket_id.clone(),
|
||||
PanelRowKey::Worker(name) => name.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn push_column_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
value: &str,
|
||||
column_width: usize,
|
||||
style: Style,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
if *remaining == 0 {
|
||||
return;
|
||||
}
|
||||
let mut content = padded_cell(value, column_width);
|
||||
content.push(' ');
|
||||
push_bounded_span(spans, &content, style, remaining);
|
||||
}
|
||||
|
||||
pub(super) fn push_bounded_span(
|
||||
spans: &mut Vec<Span<'static>>,
|
||||
value: &str,
|
||||
style: Style,
|
||||
remaining: &mut usize,
|
||||
) {
|
||||
if *remaining == 0 || value.is_empty() {
|
||||
return;
|
||||
}
|
||||
let content = truncate_with_ellipsis(value, *remaining);
|
||||
*remaining = remaining.saturating_sub(content.width());
|
||||
spans.push(Span::styled(content, style));
|
||||
}
|
||||
|
||||
pub(super) fn padded_cell(value: &str, width: usize) -> String {
|
||||
let mut cell = truncate_with_ellipsis(value, width);
|
||||
let padding = width.saturating_sub(cell.width());
|
||||
cell.extend(std::iter::repeat_n(' ', padding));
|
||||
cell
|
||||
}
|
||||
|
||||
pub(super) fn panel_priority_style(priority: ActionPriority) -> Style {
|
||||
match priority {
|
||||
ActionPriority::ReadyForQueue => Style::default().fg(Color::Green),
|
||||
ActionPriority::ActiveWork => Style::default().fg(Color::Cyan),
|
||||
ActionPriority::Background => Style::default().fg(Color::DarkGray),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn intake_status_style(status: &str) -> Style {
|
||||
match status {
|
||||
"live" => Style::default().fg(Color::Green),
|
||||
"restorable" => Style::default().fg(Color::Yellow),
|
||||
"stale" => Style::default().fg(Color::DarkGray),
|
||||
_ => Style::default().fg(Color::Cyan),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn section_rows(
|
||||
list: &WorkerList,
|
||||
section: &DashboardSection,
|
||||
selected: Option<&PanelRowKey>,
|
||||
width: u16,
|
||||
) -> Vec<PanelListRow> {
|
||||
let visible = visible_section_indices(section);
|
||||
if visible.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let mut rows = Vec::with_capacity(visible.len() + 1);
|
||||
rows.push(PanelListRow::inert(section_header_line(
|
||||
section.kind,
|
||||
section.entries.len(),
|
||||
section.hidden_count(),
|
||||
width,
|
||||
)));
|
||||
for index in visible {
|
||||
if let Some(entry) = list.entries.get(index) {
|
||||
let key = PanelRowKey::Worker(entry.name.clone());
|
||||
let selected = selected == Some(&key);
|
||||
rows.push(PanelListRow::selectable(
|
||||
row_line(entry, selected, width),
|
||||
key,
|
||||
));
|
||||
}
|
||||
}
|
||||
rows
|
||||
}
|
||||
|
||||
pub(super) fn row_line(entry: &WorkerListEntry, selected: bool, width: u16) -> Line<'static> {
|
||||
let marker = if selected { "▶ " } else { " " };
|
||||
let name_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Cyan)
|
||||
};
|
||||
let (status, status_style) = row_status_label(entry);
|
||||
let mut spans = Vec::new();
|
||||
let mut remaining = width as usize;
|
||||
|
||||
push_bounded_span(
|
||||
&mut spans,
|
||||
marker,
|
||||
if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
},
|
||||
&mut remaining,
|
||||
);
|
||||
push_column_span(
|
||||
&mut spans,
|
||||
status,
|
||||
POD_STATUS_COLUMN_WIDTH,
|
||||
status_style,
|
||||
&mut remaining,
|
||||
);
|
||||
push_bounded_span(&mut spans, entry.name.as_str(), name_style, &mut remaining);
|
||||
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
pub(super) fn draw_separator(frame: &mut Frame<'_>, area: Rect) {
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
"─".repeat(area.width as usize),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
))),
|
||||
area,
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn draw_target_status(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
frame.render_widget(Paragraph::new(target_status_line(app)), area);
|
||||
}
|
||||
|
||||
pub(super) fn target_status_line(_app: &DashboardApp) -> Line<'static> {
|
||||
Line::from(Span::raw(""))
|
||||
}
|
||||
|
||||
pub(super) fn draw_input(frame: &mut Frame<'_>, render: &crate::input::InputRender, area: Rect) {
|
||||
let mut lines: Vec<Line<'static>> = Vec::with_capacity(render.lines.len());
|
||||
for (i, src) in render.lines.iter().enumerate() {
|
||||
let absolute_row = render.viewport_start_row as usize + i;
|
||||
let prefix = if absolute_row == 0 { "> " } else { " " };
|
||||
let mut spans = vec![Span::styled(prefix, Style::default().fg(Color::DarkGray))];
|
||||
spans.extend(src.spans.iter().cloned());
|
||||
lines.push(Line::from(spans));
|
||||
}
|
||||
frame.render_widget(Paragraph::new(lines), area);
|
||||
|
||||
let cursor_x = area.x + 2 + render.cursor_col;
|
||||
let cursor_y = area.y + render.cursor_row;
|
||||
if cursor_y < area.y + area.height {
|
||||
frame.set_cursor_position(Position::new(cursor_x, cursor_y));
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn actionbar_left_text(app: &DashboardApp) -> String {
|
||||
if app.sending && app.composer_target() == ComposerTarget::TicketIntake {
|
||||
"launching Ticket Intake…".to_string()
|
||||
} else if app.sending {
|
||||
"working…".to_string()
|
||||
} else if app.refreshing {
|
||||
match app.notice.as_deref() {
|
||||
Some(notice) if notice.contains("Refreshing") || notice.contains("refreshing") => {
|
||||
notice.to_string()
|
||||
}
|
||||
Some(notice) => format!("{notice} Refreshing workspace…"),
|
||||
None => "Refreshing workspace…".to_string(),
|
||||
}
|
||||
} else if let Some(notice) = app.notice.as_deref() {
|
||||
notice.to_string()
|
||||
} else {
|
||||
String::new()
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn actionbar_right_text(app: &DashboardApp) -> &'static str {
|
||||
if app.panel_diagnostic_open {
|
||||
"F2/Esc close details"
|
||||
} else if app.panel_diagnostic.is_some() {
|
||||
"F2 details"
|
||||
} else {
|
||||
""
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn draw_actionbar(frame: &mut Frame<'_>, app: &DashboardApp, area: Rect) {
|
||||
let left = actionbar_left_text(app);
|
||||
let right = actionbar_right_text(app);
|
||||
let left_width = area
|
||||
.width
|
||||
.saturating_sub(right.width() as u16)
|
||||
.saturating_sub(2) as usize;
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
truncate_with_ellipsis(&left, left_width),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
))),
|
||||
area,
|
||||
);
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
right,
|
||||
Style::default().fg(Color::DarkGray),
|
||||
)))
|
||||
.alignment(ratatui::layout::Alignment::Right),
|
||||
area,
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn truncate_with_ellipsis(s: &str, max_width: usize) -> String {
|
||||
if max_width == 0 {
|
||||
return String::new();
|
||||
}
|
||||
if s.width() <= max_width {
|
||||
return s.to_string();
|
||||
}
|
||||
if max_width == 1 {
|
||||
return "…".to_string();
|
||||
}
|
||||
let mut out = String::new();
|
||||
let mut width = 0usize;
|
||||
for c in s.chars() {
|
||||
let cw = unicode_width::UnicodeWidthChar::width(c).unwrap_or(0);
|
||||
if width + cw > max_width - 1 {
|
||||
break;
|
||||
}
|
||||
out.push(c);
|
||||
width += cw;
|
||||
}
|
||||
out.push('…');
|
||||
out
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
+70
-94
@@ -8,24 +8,19 @@ mod command;
|
||||
mod composer_history;
|
||||
mod composer_keys;
|
||||
mod console;
|
||||
mod dashboard;
|
||||
#[cfg(feature = "e2e-test")]
|
||||
mod e2e_observer;
|
||||
mod input;
|
||||
pub mod keys;
|
||||
mod markdown;
|
||||
mod picker;
|
||||
mod role_session_registry;
|
||||
mod scroll;
|
||||
pub mod setup_model;
|
||||
mod spawn;
|
||||
mod standalone_picker;
|
||||
mod task;
|
||||
mod text_selection;
|
||||
mod tool;
|
||||
mod ui;
|
||||
mod view_mode;
|
||||
mod worker_list;
|
||||
mod workspace_panel;
|
||||
|
||||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
@@ -34,7 +29,6 @@ use std::process::ExitCode;
|
||||
use crossterm::event::{DisableBracketedPaste, DisableMouseCapture, EnableBracketedPaste};
|
||||
use crossterm::execute;
|
||||
use crossterm::terminal::{LeaveAlternateScreen, disable_raw_mode, enable_raw_mode};
|
||||
use session_store::SegmentId;
|
||||
|
||||
use client::{Target, WorkerConnectionSelector, WorkerListRequest};
|
||||
|
||||
@@ -47,42 +41,69 @@ pub struct LaunchOptions {
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum LaunchMode {
|
||||
/// Start one client-owned in-process Standalone Worker.
|
||||
Spawn {
|
||||
worker_name: Option<String>,
|
||||
profile: Option<String>,
|
||||
},
|
||||
/// `yoi --worker <name>`: attach to a live Worker by name if possible;
|
||||
/// otherwise launch the Worker runtime command with `--worker <name>` so it
|
||||
/// resumes from name-keyed state or creates a fresh same-name Worker.
|
||||
WorkerName {
|
||||
worker_name: String,
|
||||
socket_override: Option<PathBuf>,
|
||||
},
|
||||
/// `yoi workers` / `yoi --backend <url>`: list workers through the selected
|
||||
/// connection target, then attach to the selected Worker.
|
||||
/// Restore one client-owned standalone session. The current cwd is the default scope;
|
||||
/// `include_all` opts into all standalone sessions under the same client data root.
|
||||
StandaloneResume { include_all: bool },
|
||||
/// List Backend Workers and attach to the selected Worker.
|
||||
Workers {
|
||||
runtime_id: Option<String>,
|
||||
include_stopped: bool,
|
||||
all: bool,
|
||||
},
|
||||
/// `yoi --backend <url> --runtime-id <id> --worker-id <id>`: open one Worker
|
||||
/// through the selected connection target.
|
||||
/// Open one Backend Worker through the selected connection target.
|
||||
OpenWorker {
|
||||
runtime_id: String,
|
||||
worker_id: String,
|
||||
},
|
||||
/// `yoi resume`: open the Worker picker, then attach to the selected live Worker
|
||||
/// or restore the selected stopped Worker by name. Without `--all`, the picker
|
||||
/// is scoped to the current runtime workspace.
|
||||
Resume { all: bool },
|
||||
/// `yoi --session <UUID>`: skip the picker, go straight to the
|
||||
/// resume name dialog with `id` baked in.
|
||||
ResumeWithSession {
|
||||
id: SegmentId,
|
||||
worker_name: Option<String>,
|
||||
},
|
||||
/// `yoi panel`: open the workspace Dashboard from the current workspace.
|
||||
Panel { include_stopped: bool },
|
||||
/// Open the Backend Workspace dashboard.
|
||||
Panel,
|
||||
}
|
||||
|
||||
struct TerminalModeGuard {
|
||||
active: bool,
|
||||
}
|
||||
|
||||
impl TerminalModeGuard {
|
||||
fn new() -> Self {
|
||||
Self { active: true }
|
||||
}
|
||||
|
||||
fn restore(&mut self) -> io::Result<()> {
|
||||
if !self.active {
|
||||
return Ok(());
|
||||
}
|
||||
self.active = false;
|
||||
let mut stdout = io::stdout();
|
||||
execute!(
|
||||
stdout,
|
||||
DisableMouseCapture,
|
||||
LeaveAlternateScreen,
|
||||
DisableBracketedPaste,
|
||||
crossterm::cursor::Show
|
||||
)?;
|
||||
disable_raw_mode()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TerminalModeGuard {
|
||||
fn drop(&mut self) {
|
||||
if self.active {
|
||||
let mut stdout = io::stdout();
|
||||
let _ = execute!(
|
||||
stdout,
|
||||
DisableMouseCapture,
|
||||
LeaveAlternateScreen,
|
||||
DisableBracketedPaste,
|
||||
crossterm::cursor::Show
|
||||
);
|
||||
let _ = disable_raw_mode();
|
||||
self.active = false;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
@@ -109,6 +130,7 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
eprintln!("yoi: {e}");
|
||||
return ExitCode::FAILURE;
|
||||
}
|
||||
let mut terminal_mode = TerminalModeGuard::new();
|
||||
|
||||
let result = match mode {
|
||||
LaunchMode::Spawn {
|
||||
@@ -116,49 +138,34 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
profile,
|
||||
} => match target.spawn_worker() {
|
||||
Ok(spawn) => {
|
||||
console::run_spawn(None, worker_name, profile, spawn.runtime_command).await
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::WorkerName {
|
||||
worker_name,
|
||||
socket_override,
|
||||
} => match target.worker_by_name() {
|
||||
Ok(worker_by_name) => {
|
||||
console::run_worker_name(
|
||||
console::run_standalone(
|
||||
workspace_root.clone(),
|
||||
spawn.state_dir,
|
||||
worker_name,
|
||||
socket_override,
|
||||
worker_by_name.runtime_command,
|
||||
profile,
|
||||
)
|
||||
.await
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::StandaloneResume { include_all } => {
|
||||
match standalone_picker::pick(target.as_ref(), include_all) {
|
||||
Ok(Some(intent)) => console::run_standalone_restore(intent).await,
|
||||
Ok(None) => Ok(()),
|
||||
Err(error) => Err(Box::new(error) as Box<dyn std::error::Error>),
|
||||
}
|
||||
}
|
||||
LaunchMode::Workers {
|
||||
runtime_id,
|
||||
include_stopped,
|
||||
all,
|
||||
} => match target.list_workers(if include_stopped {
|
||||
WorkerListRequest::with_stopped(runtime_id)
|
||||
} else {
|
||||
WorkerListRequest::new(runtime_id)
|
||||
}) {
|
||||
Ok(worker_list) => {
|
||||
if let Some(target) = worker_list.backend_target {
|
||||
backend_worker_picker::run(target, worker_list.include_stopped).await
|
||||
} else if let Some(runtime_command) = worker_list.local_runtime_command {
|
||||
console::run_worker_picker(
|
||||
runtime_command,
|
||||
workspace_root.clone(),
|
||||
all,
|
||||
worker_list.include_stopped,
|
||||
)
|
||||
backend_worker_picker::run(worker_list.backend_target, worker_list.include_stopped)
|
||||
.await
|
||||
} else {
|
||||
Err(Box::new(io::Error::other(
|
||||
"worker list target did not include a local or backend source",
|
||||
)) as Box<dyn std::error::Error>)
|
||||
}
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
@@ -169,28 +176,12 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
Ok(connection) => console::run_backend_runtime(connection.target).await,
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::Resume { all } => match target.resume_worker() {
|
||||
Ok(resume) => {
|
||||
console::run_resume(resume.runtime_command, workspace_root.clone(), all).await
|
||||
LaunchMode::Panel => match target.dashboard() {
|
||||
Ok(dashboard) => {
|
||||
backend_dashboard::launch(dashboard.base_url, dashboard.workspace_id).await
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::ResumeWithSession { id, worker_name } => match target.spawn_worker() {
|
||||
Ok(spawn) => {
|
||||
console::run_spawn(Some(id), worker_name, None, spawn.runtime_command).await
|
||||
}
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
LaunchMode::Panel { include_stopped } => match target.dashboard() {
|
||||
Ok(client::Dashboard::Local { runtime_command }) => {
|
||||
dashboard::launch(runtime_command, include_stopped).await
|
||||
}
|
||||
Ok(client::Dashboard::Backend {
|
||||
base_url,
|
||||
workspace_id,
|
||||
}) => backend_dashboard::launch(base_url, workspace_id).await,
|
||||
Err(e) => Err(Box::new(e) as Box<dyn std::error::Error>),
|
||||
},
|
||||
};
|
||||
|
||||
// Always restore the terminal first so any pending eprintln below
|
||||
@@ -198,15 +189,7 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
// alternate-screen buffer.
|
||||
#[cfg(feature = "e2e-test")]
|
||||
e2e_observer::emit("tui", "terminal_cleanup_started", serde_json::json!({}));
|
||||
let mut stdout = io::stdout();
|
||||
let _ = execute!(
|
||||
stdout,
|
||||
DisableMouseCapture,
|
||||
LeaveAlternateScreen,
|
||||
DisableBracketedPaste
|
||||
);
|
||||
let _ = disable_raw_mode();
|
||||
let _ = execute!(stdout, crossterm::cursor::Show);
|
||||
let _ = terminal_mode.restore();
|
||||
#[cfg(feature = "e2e-test")]
|
||||
e2e_observer::emit("tui", "terminal_cleanup_finished", serde_json::json!({}));
|
||||
|
||||
@@ -217,14 +200,7 @@ pub async fn launch(options: LaunchOptions) -> ExitCode {
|
||||
ExitCode::SUCCESS
|
||||
}
|
||||
Err(e) => {
|
||||
// SpawnError has already been painted into the inline
|
||||
// viewport's final frame, so it's already visible in the
|
||||
// user's scrollback — printing it again would be a noisy
|
||||
// duplicate. Other errors (worker-name failures, terminal setup
|
||||
// hiccups, etc.) need surfacing here.
|
||||
if e.downcast_ref::<spawn::SpawnError>().is_none() {
|
||||
eprintln!("yoi: {e}");
|
||||
}
|
||||
eprintln!("yoi: {e}");
|
||||
#[cfg(feature = "e2e-test")]
|
||||
e2e_observer::emit("tui", "exit", serde_json::json!({ "status": "failure" }));
|
||||
ExitCode::FAILURE
|
||||
|
||||
@@ -1,525 +0,0 @@
|
||||
//! Inline-viewport "pick a Worker to attach or restore" UX.
|
||||
//!
|
||||
//! Reads live Worker allocations from the runtime registry and stopped Worker state
|
||||
//! from the session-store worker metadata name-keyed metadata. Picking a live row attaches to
|
||||
//! its socket; picking a stopped row restores via the Worker runtime command.
|
||||
|
||||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
use std::time::Duration;
|
||||
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{Frame, TerminalOptions, Viewport};
|
||||
use session_store::FsStore;
|
||||
use session_store::FsWorkerStore;
|
||||
|
||||
use crate::worker_list::{
|
||||
LiveWorkerInfo, StoredMetadataState, StoredWorkerInfo, WorkerList, WorkerListEntry,
|
||||
WorkerVisibilitySource, live_socket_for_worker as worker_list_live_socket_for_worker,
|
||||
read_reachable_live_worker_infos, read_stored_worker_infos,
|
||||
};
|
||||
|
||||
const MAX_ROWS: usize = 10;
|
||||
const VIEWPORT_LINES: u16 = MAX_ROWS as u16 + 4;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum PickerError {
|
||||
Io(io::Error),
|
||||
Store(session_store::StoreError),
|
||||
NoWorkers { all: bool },
|
||||
}
|
||||
|
||||
impl std::fmt::Display for PickerError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(e) => write!(f, "io error: {e}"),
|
||||
Self::Store(e) => write!(f, "session store error: {e}"),
|
||||
Self::NoWorkers { all: true } => write!(
|
||||
f,
|
||||
"no workers found — start a fresh Worker with `yoi` and try again"
|
||||
),
|
||||
Self::NoWorkers { all: false } => write!(
|
||||
f,
|
||||
"no workers found in this workspace — use `yoi resume --all` to list all host/data-dir Workers"
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for PickerError {}
|
||||
|
||||
impl From<io::Error> for PickerError {
|
||||
fn from(e: io::Error) -> Self {
|
||||
Self::Io(e)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<session_store::StoreError> for PickerError {
|
||||
fn from(e: session_store::StoreError) -> Self {
|
||||
Self::Store(e)
|
||||
}
|
||||
}
|
||||
|
||||
pub enum PickerOutcome {
|
||||
/// User picked a Worker. `socket_override` is set for live rows when the
|
||||
/// runtime registry knows the exact socket path; stopped rows leave it
|
||||
/// empty so the caller restores by spawning the Worker runtime command.
|
||||
Picked {
|
||||
worker_name: String,
|
||||
socket_override: Option<PathBuf>,
|
||||
},
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct PickerOptions {
|
||||
scope: PickerScope,
|
||||
include_stopped: bool,
|
||||
}
|
||||
|
||||
impl PickerOptions {
|
||||
pub(crate) fn workspace(workspace_root: PathBuf) -> Self {
|
||||
Self {
|
||||
scope: PickerScope::Workspace(workspace_root),
|
||||
include_stopped: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn all() -> Self {
|
||||
Self {
|
||||
scope: PickerScope::All,
|
||||
include_stopped: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn with_stopped(mut self, include_stopped: bool) -> Self {
|
||||
self.include_stopped = include_stopped;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
enum PickerScope {
|
||||
Workspace(PathBuf),
|
||||
All,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum WorkerRowState {
|
||||
Live,
|
||||
Stopped,
|
||||
Corrupt,
|
||||
}
|
||||
|
||||
impl WorkerRowState {
|
||||
fn label(self) -> &'static str {
|
||||
match self {
|
||||
Self::Live => "live",
|
||||
Self::Stopped => "stopped",
|
||||
Self::Corrupt => "corrupt",
|
||||
}
|
||||
}
|
||||
|
||||
fn style(self) -> Style {
|
||||
match self {
|
||||
Self::Live => Style::default()
|
||||
.fg(Color::Green)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
Self::Stopped => Style::default().fg(Color::Yellow),
|
||||
Self::Corrupt => Style::default().fg(Color::Red).add_modifier(Modifier::BOLD),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn list_for_options(
|
||||
options: &PickerOptions,
|
||||
stored_workers: Vec<StoredWorkerInfo>,
|
||||
live_workers: Vec<LiveWorkerInfo>,
|
||||
) -> WorkerList {
|
||||
let stored_workers = if options.include_stopped {
|
||||
stored_workers
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
match &options.scope {
|
||||
PickerScope::Workspace(workspace_root) => WorkerList::from_workspace_sources(
|
||||
WorkerVisibilitySource::ResumePicker,
|
||||
stored_workers,
|
||||
live_workers,
|
||||
None,
|
||||
MAX_ROWS,
|
||||
workspace_root,
|
||||
),
|
||||
PickerScope::All => WorkerList::from_sources(
|
||||
WorkerVisibilitySource::ResumePicker,
|
||||
stored_workers,
|
||||
live_workers,
|
||||
None,
|
||||
MAX_ROWS,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn run(options: PickerOptions) -> Result<PickerOutcome, PickerError> {
|
||||
let store_dir = default_store_dir()?;
|
||||
let store = FsStore::new(&store_dir)?;
|
||||
let worker_metadata_store =
|
||||
FsWorkerStore::new(default_worker_metadata_dir()?).map_err(io::Error::other)?;
|
||||
let stored_workers = read_stored_worker_infos(&store, &worker_metadata_store)?;
|
||||
let live_workers = read_reachable_live_worker_infos(&store)
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let mut list = list_for_options(&options, stored_workers, live_workers);
|
||||
if list.entries.is_empty() {
|
||||
return Err(PickerError::NoWorkers {
|
||||
all: matches!(options.scope, PickerScope::All),
|
||||
});
|
||||
}
|
||||
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
loop {
|
||||
terminal.draw(|f| draw(f, &list))?;
|
||||
match poll_event()? {
|
||||
None => continue,
|
||||
Some(Action::Up) => {
|
||||
let selected = list.selected_index().saturating_sub(1);
|
||||
list.select_index(selected);
|
||||
}
|
||||
Some(Action::Down) => {
|
||||
let selected = list.selected_index();
|
||||
if selected + 1 < list.entries.len() {
|
||||
list.select_index(selected + 1);
|
||||
}
|
||||
}
|
||||
Some(Action::Submit) => {
|
||||
close_viewport(&mut terminal)?;
|
||||
let entry = list.selected_entry().expect("non-empty worker list");
|
||||
return Ok(PickerOutcome::Picked {
|
||||
worker_name: entry.name.clone(),
|
||||
socket_override: entry.attach_socket_path().map(PathBuf::from),
|
||||
});
|
||||
}
|
||||
Some(Action::Cancel) => {
|
||||
close_viewport(&mut terminal)?;
|
||||
return Ok(PickerOutcome::Cancelled);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Park the cursor at the very bottom of the picker's inline viewport and emit
|
||||
/// one newline before dropping the terminal. This keeps any next inline viewport
|
||||
/// from drawing over the lower picker rows.
|
||||
fn close_viewport(terminal: &mut Terminal<CrosstermBackend<io::Stdout>>) -> io::Result<()> {
|
||||
let area = terminal.get_frame().area();
|
||||
let last_row = area.bottom().saturating_sub(1);
|
||||
terminal.set_cursor_position((0, last_row))?;
|
||||
use std::io::Write;
|
||||
let mut out = io::stdout();
|
||||
out.write_all(b"\r\n")?;
|
||||
out.flush()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn default_store_dir() -> Result<PathBuf, PickerError> {
|
||||
manifest::paths::sessions_dir().ok_or_else(|| {
|
||||
PickerError::Io(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
"could not resolve sessions directory \
|
||||
(set YOI_DATA_DIR, YOI_HOME, XDG_DATA_HOME, or HOME)",
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn default_worker_metadata_dir() -> Result<PathBuf, PickerError> {
|
||||
manifest::paths::data_dir()
|
||||
.map(|dir| dir.join("workers"))
|
||||
.ok_or_else(|| {
|
||||
PickerError::Io(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
"could not resolve worker state directory \
|
||||
(set YOI_DATA_DIR, YOI_HOME, XDG_DATA_HOME, or HOME)",
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn live_socket_for_worker(worker_name: &str) -> Option<PathBuf> {
|
||||
worker_list_live_socket_for_worker(worker_name)
|
||||
}
|
||||
|
||||
fn make_inline_terminal() -> io::Result<Terminal<CrosstermBackend<io::Stdout>>> {
|
||||
let backend = CrosstermBackend::new(io::stdout());
|
||||
Terminal::with_options(
|
||||
backend,
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(VIEWPORT_LINES),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
enum Action {
|
||||
Up,
|
||||
Down,
|
||||
Submit,
|
||||
Cancel,
|
||||
}
|
||||
|
||||
fn poll_event() -> io::Result<Option<Action>> {
|
||||
if !event::poll(Duration::from_millis(100))? {
|
||||
return Ok(None);
|
||||
}
|
||||
match event::read()? {
|
||||
TermEvent::Key(k) if k.kind != KeyEventKind::Release => {
|
||||
let ctrl = k.modifiers.contains(KeyModifiers::CONTROL);
|
||||
Ok(match k.code {
|
||||
KeyCode::Up => Some(Action::Up),
|
||||
KeyCode::Down => Some(Action::Down),
|
||||
KeyCode::Char('k') if !ctrl => Some(Action::Up),
|
||||
KeyCode::Char('j') if !ctrl => Some(Action::Down),
|
||||
KeyCode::Enter => Some(Action::Submit),
|
||||
KeyCode::Esc => Some(Action::Cancel),
|
||||
KeyCode::Char('c') if ctrl => Some(Action::Cancel),
|
||||
_ => None,
|
||||
})
|
||||
}
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn draw(f: &mut Frame<'_>, list: &WorkerList) {
|
||||
let area = f.area();
|
||||
let mut constraints: Vec<Constraint> = Vec::with_capacity(list.entries.len() + 3);
|
||||
constraints.push(Constraint::Length(1)); // title
|
||||
for _ in &list.entries {
|
||||
constraints.push(Constraint::Length(1));
|
||||
}
|
||||
constraints.push(Constraint::Length(1)); // hint
|
||||
constraints.push(Constraint::Length(1)); // spacer
|
||||
let layout = Layout::vertical(constraints).split(area);
|
||||
|
||||
f.render_widget(
|
||||
Paragraph::new(Line::from(vec![Span::styled(
|
||||
picker_title(),
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
)])),
|
||||
layout[0],
|
||||
);
|
||||
|
||||
let selected = list.selected_index();
|
||||
for (i, entry) in list.entries.iter().enumerate() {
|
||||
f.render_widget(
|
||||
Paragraph::new(row_line(entry, i == selected)),
|
||||
layout[i + 1],
|
||||
);
|
||||
}
|
||||
|
||||
f.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("[↑/↓]", Style::default().fg(Color::DarkGray)),
|
||||
Span::raw(" select "),
|
||||
Span::styled("[enter]", Style::default().fg(Color::Green)),
|
||||
Span::raw(" open/restore "),
|
||||
Span::styled("[esc]", Style::default().fg(Color::Yellow)),
|
||||
Span::raw(" cancel"),
|
||||
])),
|
||||
layout[list.entries.len() + 1],
|
||||
);
|
||||
}
|
||||
|
||||
fn picker_title() -> &'static str {
|
||||
"resume worker pick a worker"
|
||||
}
|
||||
|
||||
fn row_line(entry: &WorkerListEntry, selected: bool) -> Line<'_> {
|
||||
let marker = if selected { "▶ " } else { " " };
|
||||
let name_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::Cyan)
|
||||
};
|
||||
let preview_style = if selected {
|
||||
Style::default().fg(Color::White)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
};
|
||||
let state = row_state(entry);
|
||||
let _visibility = entry.visibility;
|
||||
let _source_kinds = &entry.source_kinds;
|
||||
|
||||
let mut spans = vec![
|
||||
Span::raw(marker),
|
||||
Span::styled(entry.name.as_str(), name_style),
|
||||
Span::raw(" "),
|
||||
Span::styled(format!("[{}]", state.label()), state.style()),
|
||||
Span::raw(" "),
|
||||
Span::styled(
|
||||
format_updated_at(entry.summary.updated_at),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
),
|
||||
Span::raw(" "),
|
||||
Span::styled(debug_ids(entry), Style::default().fg(Color::DarkGray)),
|
||||
];
|
||||
if let Some(preview) = entry.summary.preview.as_ref() {
|
||||
spans.push(Span::raw(" "));
|
||||
spans.push(Span::styled(preview.as_str(), preview_style));
|
||||
}
|
||||
Line::from(spans)
|
||||
}
|
||||
|
||||
fn row_state(entry: &WorkerListEntry) -> WorkerRowState {
|
||||
if entry.live.as_ref().is_some_and(|live| live.reachable) {
|
||||
return WorkerRowState::Live;
|
||||
}
|
||||
if entry
|
||||
.stored
|
||||
.as_ref()
|
||||
.is_some_and(|stored| matches!(stored.metadata_state, StoredMetadataState::Corrupt(_)))
|
||||
{
|
||||
return WorkerRowState::Corrupt;
|
||||
}
|
||||
WorkerRowState::Stopped
|
||||
}
|
||||
|
||||
fn format_updated_at(updated_at: u64) -> String {
|
||||
if updated_at == 0 {
|
||||
"updated: —".to_string()
|
||||
} else {
|
||||
format!("updated: {updated_at}")
|
||||
}
|
||||
}
|
||||
|
||||
fn debug_ids(entry: &WorkerListEntry) -> String {
|
||||
let session = entry
|
||||
.summary
|
||||
.active_session_id
|
||||
.map(short_id)
|
||||
.unwrap_or_else(|| "--------".to_string());
|
||||
let segment = entry
|
||||
.summary
|
||||
.active_segment_id
|
||||
.map(short_id)
|
||||
.unwrap_or_else(|| "--------".to_string());
|
||||
format!("s:{session} g:{segment}")
|
||||
}
|
||||
|
||||
fn short_id<T: ToString>(id: T) -> String {
|
||||
id.to_string().chars().take(8).collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn picker_title_names_pods_not_sessions() {
|
||||
assert_eq!(picker_title(), "resume worker pick a worker");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_no_pods_message_mentions_all_for_workspace_scope() {
|
||||
let message = PickerError::NoWorkers { all: false }.to_string();
|
||||
assert!(message.contains("no workers found in this workspace"));
|
||||
assert!(message.contains("yoi resume --all"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_no_pods_message_keeps_fresh_pod_hint_for_all_scope() {
|
||||
let message = PickerError::NoWorkers { all: true }.to_string();
|
||||
assert!(message.contains("start a fresh Worker with `yoi`"));
|
||||
assert!(!message.contains("yoi resume --all"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_workspace_options_filter_by_workspace_metadata() {
|
||||
let list = list_for_options(
|
||||
&PickerOptions::workspace(PathBuf::from("/workspace/current")),
|
||||
vec![
|
||||
stored_pod("current", Some("/workspace/current"), 3),
|
||||
stored_pod("other", Some("/workspace/other"), 2),
|
||||
stored_pod("legacy", None, 1),
|
||||
],
|
||||
vec![],
|
||||
);
|
||||
|
||||
let names: Vec<_> = list
|
||||
.entries
|
||||
.iter()
|
||||
.map(|entry| entry.name.as_str())
|
||||
.collect();
|
||||
assert_eq!(names, vec!["current"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_all_options_include_host_wide_and_legacy_pods() {
|
||||
let list = list_for_options(
|
||||
&PickerOptions::all(),
|
||||
vec![
|
||||
stored_pod("current", Some("/workspace/current"), 3),
|
||||
stored_pod("other", Some("/workspace/other"), 2),
|
||||
stored_pod("legacy", None, 1),
|
||||
],
|
||||
vec![],
|
||||
);
|
||||
|
||||
let names: Vec<_> = list
|
||||
.entries
|
||||
.iter()
|
||||
.map(|entry| entry.name.as_str())
|
||||
.collect();
|
||||
assert_eq!(names, vec!["current", "other", "legacy"]);
|
||||
}
|
||||
|
||||
fn stored_pod(name: &str, workspace_root: Option<&str>, updated_at: u64) -> StoredWorkerInfo {
|
||||
StoredWorkerInfo {
|
||||
worker_name: name.to_string(),
|
||||
metadata_state: StoredMetadataState::Present,
|
||||
active_session_id: None,
|
||||
active_segment_id: None,
|
||||
updated_at,
|
||||
workspace_root: workspace_root.map(PathBuf::from),
|
||||
preview: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn picker_row_shows_live_pending_preview_and_runtime_segment_id() {
|
||||
let segment_id = session_store::new_segment_id();
|
||||
let entry = WorkerList::from_sources(
|
||||
WorkerVisibilitySource::ResumePicker,
|
||||
vec![],
|
||||
vec![crate::worker_list::LiveWorkerInfo {
|
||||
worker_name: "pending".to_string(),
|
||||
socket_path: PathBuf::from("/tmp/pending.sock"),
|
||||
status: Some(protocol::WorkerStatus::Idle),
|
||||
reachable: true,
|
||||
segment_id: Some(segment_id),
|
||||
summary: crate::worker_list::WorkerEntrySummary::default(),
|
||||
}],
|
||||
None,
|
||||
10,
|
||||
)
|
||||
.entries
|
||||
.into_iter()
|
||||
.next()
|
||||
.unwrap();
|
||||
|
||||
let text = row_line(&entry, false)
|
||||
.spans
|
||||
.iter()
|
||||
.map(|span| span.content.as_ref())
|
||||
.collect::<String>();
|
||||
|
||||
assert!(text.contains("[live]"));
|
||||
assert!(text.contains("[live, pending segment]"));
|
||||
assert!(text.contains(&format!("g:{}", short_id(segment_id))));
|
||||
}
|
||||
}
|
||||
@@ -1,556 +0,0 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::fs::{self, OpenOptions};
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::thread;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
const REGISTRY_VERSION: u32 = 1;
|
||||
const REGISTRY_FILE: &str = "role-sessions.json";
|
||||
const REGISTRY_LOCK_FILE: &str = "role-sessions.lock";
|
||||
const CLAIMS_DIR: &str = "ticket-claims";
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct PanelRegistryStore {
|
||||
root: PathBuf,
|
||||
workspace_root: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub(crate) struct RoleSessionRegistry {
|
||||
pub version: u32,
|
||||
pub workspace_root: String,
|
||||
pub sessions: BTreeMap<String, RoleSessionRecord>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub(crate) struct RoleSessionRecord {
|
||||
pub role: String,
|
||||
pub worker_name: String,
|
||||
pub origin: RoleSessionOrigin,
|
||||
pub created_at: String,
|
||||
pub updated_at: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub session_id: Option<String>,
|
||||
#[serde(default)]
|
||||
pub related_tickets: Vec<RelatedTicketRef>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub(crate) enum RoleSessionOrigin {
|
||||
PreTicketIntake,
|
||||
TicketClaim,
|
||||
RoleLaunch,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, PartialOrd, Ord)]
|
||||
pub(crate) struct RelatedTicketRef {
|
||||
pub id: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub slug: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub(crate) struct TicketClaim {
|
||||
pub ticket_id: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ticket_slug: Option<String>,
|
||||
pub worker_name: String,
|
||||
pub role: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct PanelRegistrySnapshot {
|
||||
pub sessions: Vec<RoleSessionRecord>,
|
||||
pub claims: Vec<TicketClaim>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) enum TicketClaimResult {
|
||||
Claimed,
|
||||
AlreadyOwned(TicketClaim),
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) enum PanelRegistryError {
|
||||
Io(io::Error),
|
||||
Json(serde_json::Error),
|
||||
TicketAlreadyClaimed(TicketClaim),
|
||||
}
|
||||
|
||||
impl std::fmt::Display for PanelRegistryError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(error) => write!(f, "local role session registry I/O error: {error}"),
|
||||
Self::Json(error) => write!(f, "local role session registry JSON error: {error}"),
|
||||
Self::TicketAlreadyClaimed(claim) => write!(
|
||||
f,
|
||||
"Ticket {} is already claimed locally by {} ({})",
|
||||
claim.ticket_id, claim.worker_name, claim.role
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for PanelRegistryError {}
|
||||
|
||||
impl From<io::Error> for PanelRegistryError {
|
||||
fn from(error: io::Error) -> Self {
|
||||
Self::Io(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<serde_json::Error> for PanelRegistryError {
|
||||
fn from(error: serde_json::Error) -> Self {
|
||||
Self::Json(error)
|
||||
}
|
||||
}
|
||||
|
||||
impl PanelRegistryStore {
|
||||
pub(crate) fn default_for_workspace(workspace_root: &Path) -> Result<Self, PanelRegistryError> {
|
||||
let data_dir = manifest::paths::data_dir().ok_or_else(|| {
|
||||
PanelRegistryError::Io(io::Error::other("failed to resolve yoi data directory"))
|
||||
})?;
|
||||
Ok(Self::for_data_dir(data_dir, workspace_root))
|
||||
}
|
||||
|
||||
pub(crate) fn for_data_dir(data_dir: impl AsRef<Path>, workspace_root: &Path) -> Self {
|
||||
let workspace_root = normalized_workspace_key(workspace_root);
|
||||
let leaf = workspace_leaf(&workspace_root);
|
||||
let digest = fnv1a64_hex(workspace_root.as_bytes());
|
||||
Self {
|
||||
root: data_dir
|
||||
.as_ref()
|
||||
.join("panel")
|
||||
.join("workspaces")
|
||||
.join(format!("{leaf}-{digest}")),
|
||||
workspace_root: Some(workspace_root),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn from_root(root: impl Into<PathBuf>) -> Self {
|
||||
Self {
|
||||
root: root.into(),
|
||||
workspace_root: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn root(&self) -> &Path {
|
||||
&self.root
|
||||
}
|
||||
|
||||
pub(crate) fn snapshot(&self) -> Result<PanelRegistrySnapshot, PanelRegistryError> {
|
||||
let registry = self.load_registry()?;
|
||||
let claims = self.load_claims()?;
|
||||
Ok(PanelRegistrySnapshot {
|
||||
sessions: registry.sessions.into_values().collect(),
|
||||
claims,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn load_registry(&self) -> Result<RoleSessionRegistry, PanelRegistryError> {
|
||||
match fs::read(self.registry_path()) {
|
||||
Ok(bytes) => Ok(serde_json::from_slice(&bytes)?),
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(RoleSessionRegistry {
|
||||
version: REGISTRY_VERSION,
|
||||
workspace_root: self.workspace_root.clone().unwrap_or_default(),
|
||||
sessions: BTreeMap::new(),
|
||||
}),
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn record_session(
|
||||
&self,
|
||||
worker_name: impl Into<String>,
|
||||
role: impl Into<String>,
|
||||
origin: RoleSessionOrigin,
|
||||
session_id: Option<String>,
|
||||
related_tickets: impl IntoIterator<Item = RelatedTicketRef>,
|
||||
) -> Result<(), PanelRegistryError> {
|
||||
let worker_name = worker_name.into();
|
||||
let role = role.into();
|
||||
let related_tickets: Vec<RelatedTicketRef> = related_tickets.into_iter().collect();
|
||||
self.update_registry(|registry| {
|
||||
let now = now_timestamp_string();
|
||||
let mut tickets: BTreeSet<RelatedTicketRef> = registry
|
||||
.sessions
|
||||
.get(&worker_name)
|
||||
.map(|record| record.related_tickets.iter().cloned().collect())
|
||||
.unwrap_or_default();
|
||||
tickets.extend(related_tickets);
|
||||
let created_at = registry
|
||||
.sessions
|
||||
.get(&worker_name)
|
||||
.map(|record| record.created_at.clone())
|
||||
.unwrap_or_else(|| now.clone());
|
||||
registry.sessions.insert(
|
||||
worker_name.clone(),
|
||||
RoleSessionRecord {
|
||||
role,
|
||||
worker_name,
|
||||
origin,
|
||||
created_at,
|
||||
updated_at: now,
|
||||
session_id,
|
||||
related_tickets: tickets.into_iter().collect(),
|
||||
},
|
||||
);
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn claim_ticket(
|
||||
&self,
|
||||
ticket_id: &str,
|
||||
ticket_slug: Option<&str>,
|
||||
worker_name: &str,
|
||||
role: &str,
|
||||
) -> Result<TicketClaimResult, PanelRegistryError> {
|
||||
fs::create_dir_all(self.claims_dir())?;
|
||||
let claim_path = self.claim_path(ticket_id);
|
||||
let claim = TicketClaim {
|
||||
ticket_id: ticket_id.to_string(),
|
||||
ticket_slug: ticket_slug.map(ToOwned::to_owned),
|
||||
worker_name: worker_name.to_string(),
|
||||
role: role.to_string(),
|
||||
};
|
||||
match self.create_claim_file(&claim_path, &claim) {
|
||||
Ok(()) => {
|
||||
if let Err(error) = self.record_session(
|
||||
worker_name.to_string(),
|
||||
role.to_string(),
|
||||
RoleSessionOrigin::TicketClaim,
|
||||
None,
|
||||
[RelatedTicketRef {
|
||||
id: ticket_id.to_string(),
|
||||
slug: ticket_slug.map(ToOwned::to_owned),
|
||||
}],
|
||||
) {
|
||||
let _ = fs::remove_file(&claim_path);
|
||||
return Err(error);
|
||||
}
|
||||
Ok(TicketClaimResult::Claimed)
|
||||
}
|
||||
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
|
||||
let existing = self.load_claim(ticket_id)?;
|
||||
if existing.worker_name == worker_name && existing.role == role {
|
||||
Ok(TicketClaimResult::AlreadyOwned(existing))
|
||||
} else {
|
||||
Err(PanelRegistryError::TicketAlreadyClaimed(existing))
|
||||
}
|
||||
}
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn load_claim(&self, ticket_id: &str) -> Result<TicketClaim, PanelRegistryError> {
|
||||
let bytes = fs::read(self.claim_path(ticket_id))?;
|
||||
Ok(serde_json::from_slice(&bytes)?)
|
||||
}
|
||||
|
||||
pub(crate) fn claim_for_ticket(
|
||||
&self,
|
||||
ticket_id: &str,
|
||||
) -> Result<Option<TicketClaim>, PanelRegistryError> {
|
||||
match self.load_claim(ticket_id) {
|
||||
Ok(claim) => Ok(Some(claim)),
|
||||
Err(PanelRegistryError::Io(error)) if error.kind() == io::ErrorKind::NotFound => {
|
||||
Ok(None)
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn update_registry(
|
||||
&self,
|
||||
update: impl FnOnce(&mut RoleSessionRegistry) -> Result<(), PanelRegistryError>,
|
||||
) -> Result<(), PanelRegistryError> {
|
||||
fs::create_dir_all(&self.root)?;
|
||||
let _lock = self.acquire_registry_lock()?;
|
||||
let mut registry = self.load_registry()?;
|
||||
registry.version = REGISTRY_VERSION;
|
||||
if let Some(workspace_root) = self.workspace_root.as_ref() {
|
||||
registry.workspace_root = workspace_root.clone();
|
||||
}
|
||||
update(&mut registry)?;
|
||||
self.save_registry(®istry)
|
||||
}
|
||||
|
||||
fn acquire_registry_lock(&self) -> Result<RegistryLockGuard, PanelRegistryError> {
|
||||
let lock_path = self.root.join(REGISTRY_LOCK_FILE);
|
||||
for _ in 0..50 {
|
||||
match OpenOptions::new()
|
||||
.write(true)
|
||||
.create_new(true)
|
||||
.open(&lock_path)
|
||||
{
|
||||
Ok(_) => return Ok(RegistryLockGuard { path: lock_path }),
|
||||
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
|
||||
thread::sleep(Duration::from_millis(10));
|
||||
}
|
||||
Err(error) => return Err(error.into()),
|
||||
}
|
||||
}
|
||||
Err(PanelRegistryError::Io(io::Error::new(
|
||||
io::ErrorKind::WouldBlock,
|
||||
"timed out acquiring panel role session registry lock",
|
||||
)))
|
||||
}
|
||||
|
||||
fn save_registry(&self, registry: &RoleSessionRegistry) -> Result<(), PanelRegistryError> {
|
||||
let path = self.registry_path();
|
||||
let temp_path = path.with_extension(format!("json.{}.tmp", now_timestamp_string()));
|
||||
let bytes = serde_json::to_vec_pretty(registry)?;
|
||||
fs::write(&temp_path, [&bytes[..], b"\n"].concat())?;
|
||||
fs::rename(temp_path, path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn create_claim_file(&self, claim_path: &Path, claim: &TicketClaim) -> io::Result<()> {
|
||||
let temp_path = self
|
||||
.claims_dir()
|
||||
.join(format!(".{}.tmp", now_timestamp_string()));
|
||||
let bytes = serde_json::to_vec_pretty(claim).map_err(io::Error::other)?;
|
||||
fs::write(&temp_path, [&bytes[..], b"\n"].concat())?;
|
||||
let link_result = fs::hard_link(&temp_path, claim_path);
|
||||
let remove_result = fs::remove_file(&temp_path);
|
||||
match (link_result, remove_result) {
|
||||
(Ok(()), Ok(())) | (Ok(()), Err(_)) => Ok(()),
|
||||
(Err(error), _) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn load_claims(&self) -> Result<Vec<TicketClaim>, PanelRegistryError> {
|
||||
let mut claims: Vec<TicketClaim> = Vec::new();
|
||||
match fs::read_dir(self.claims_dir()) {
|
||||
Ok(entries) => {
|
||||
for entry in entries {
|
||||
let entry = entry?;
|
||||
if entry.file_type()?.is_file()
|
||||
&& entry
|
||||
.path()
|
||||
.extension()
|
||||
.is_some_and(|extension| extension == "json")
|
||||
{
|
||||
let bytes = fs::read(entry.path())?;
|
||||
claims.push(serde_json::from_slice(&bytes)?);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => {}
|
||||
Err(error) => return Err(error.into()),
|
||||
}
|
||||
claims.sort_by(|left, right| left.ticket_id.cmp(&right.ticket_id));
|
||||
Ok(claims)
|
||||
}
|
||||
|
||||
fn registry_path(&self) -> PathBuf {
|
||||
self.root.join(REGISTRY_FILE)
|
||||
}
|
||||
|
||||
fn claims_dir(&self) -> PathBuf {
|
||||
self.root.join(CLAIMS_DIR)
|
||||
}
|
||||
|
||||
fn claim_path(&self, ticket_id: &str) -> PathBuf {
|
||||
self.claims_dir()
|
||||
.join(format!("{}.json", encode_path_component(ticket_id)))
|
||||
}
|
||||
}
|
||||
|
||||
struct RegistryLockGuard {
|
||||
path: PathBuf,
|
||||
}
|
||||
|
||||
impl Drop for RegistryLockGuard {
|
||||
fn drop(&mut self) {
|
||||
let _ = fs::remove_file(&self.path);
|
||||
}
|
||||
}
|
||||
|
||||
impl PanelRegistrySnapshot {
|
||||
pub(crate) fn empty() -> Self {
|
||||
Self {
|
||||
sessions: Vec::new(),
|
||||
claims: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn claim_for_ticket(&self, ticket_id: &str) -> Option<&TicketClaim> {
|
||||
self.claims
|
||||
.iter()
|
||||
.find(|claim| claim.ticket_id == ticket_id)
|
||||
}
|
||||
}
|
||||
|
||||
fn normalized_workspace_key(path: &Path) -> String {
|
||||
path.to_string_lossy().replace('\\', "/")
|
||||
}
|
||||
|
||||
fn workspace_leaf(workspace_root: &str) -> String {
|
||||
let leaf = workspace_root
|
||||
.rsplit('/')
|
||||
.find(|part| !part.is_empty())
|
||||
.unwrap_or("workspace");
|
||||
let sanitized = leaf
|
||||
.chars()
|
||||
.map(|ch| {
|
||||
if ch.is_ascii_alphanumeric() || matches!(ch, '-' | '_') {
|
||||
ch
|
||||
} else {
|
||||
'-'
|
||||
}
|
||||
})
|
||||
.collect::<String>()
|
||||
.trim_matches('-')
|
||||
.to_string();
|
||||
if sanitized.is_empty() {
|
||||
"workspace".to_string()
|
||||
} else {
|
||||
sanitized
|
||||
}
|
||||
}
|
||||
|
||||
fn fnv1a64_hex(bytes: &[u8]) -> String {
|
||||
let mut hash = 0xcbf29ce484222325u64;
|
||||
for byte in bytes {
|
||||
hash ^= u64::from(*byte);
|
||||
hash = hash.wrapping_mul(0x100000001b3);
|
||||
}
|
||||
format!("{hash:016x}")
|
||||
}
|
||||
|
||||
fn encode_path_component(value: &str) -> String {
|
||||
let mut encoded = String::with_capacity(value.len());
|
||||
for byte in value.bytes() {
|
||||
match byte {
|
||||
b'a'..=b'z' | b'A'..=b'Z' | b'0'..=b'9' | b'-' | b'_' => encoded.push(byte as char),
|
||||
_ => encoded.push_str(&format!("%{byte:02X}")),
|
||||
}
|
||||
}
|
||||
encoded
|
||||
}
|
||||
|
||||
fn now_timestamp_string() -> String {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_nanos().to_string())
|
||||
.unwrap_or_else(|_| "0".to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use tempfile::TempDir;
|
||||
|
||||
#[test]
|
||||
fn registry_path_is_workspace_scoped_under_data_dir() {
|
||||
let data_dir = TempDir::new().unwrap();
|
||||
let store = PanelRegistryStore::for_data_dir(data_dir.path(), Path::new("/repo/yoi"));
|
||||
let other = PanelRegistryStore::for_data_dir(data_dir.path(), Path::new("/repo/other"));
|
||||
|
||||
assert!(store.root().starts_with(data_dir.path()));
|
||||
let root = store.root().to_string_lossy();
|
||||
assert!(root.contains("panel/workspaces/yoi-"));
|
||||
assert_ne!(store.root(), other.root());
|
||||
|
||||
store
|
||||
.record_session(
|
||||
"ticket-intake-preticket",
|
||||
"intake",
|
||||
RoleSessionOrigin::PreTicketIntake,
|
||||
None,
|
||||
[],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(store.load_registry().unwrap().workspace_root, "/repo/yoi");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claim_ticket_rejects_second_active_local_pod() {
|
||||
let temp = TempDir::new().unwrap();
|
||||
let store = PanelRegistryStore::from_root(temp.path().join("registry"));
|
||||
|
||||
assert!(matches!(
|
||||
store.claim_ticket("T-1", Some("ticket-one"), "ticket-one-intake", "intake"),
|
||||
Ok(TicketClaimResult::Claimed)
|
||||
));
|
||||
|
||||
let error = store
|
||||
.claim_ticket("T-1", Some("ticket-one"), "ticket-two-intake", "intake")
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, PanelRegistryError::TicketAlreadyClaimed(_)));
|
||||
let claim = store.claim_for_ticket("T-1").unwrap().unwrap();
|
||||
assert_eq!(claim.worker_name, "ticket-one-intake");
|
||||
assert_eq!(claim.ticket_slug.as_deref(), Some("ticket-one"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn intake_session_relation_is_not_one_to_one_with_tickets() {
|
||||
let temp = TempDir::new().unwrap();
|
||||
let store = PanelRegistryStore::from_root(temp.path().join("registry"));
|
||||
|
||||
store
|
||||
.record_session(
|
||||
"ticket-intake-preticket",
|
||||
"intake",
|
||||
RoleSessionOrigin::PreTicketIntake,
|
||||
None,
|
||||
[],
|
||||
)
|
||||
.unwrap();
|
||||
store
|
||||
.record_session(
|
||||
"ticket-intake-shared",
|
||||
"intake",
|
||||
RoleSessionOrigin::RoleLaunch,
|
||||
None,
|
||||
[
|
||||
RelatedTicketRef {
|
||||
id: "T-1".to_string(),
|
||||
slug: Some("one".to_string()),
|
||||
},
|
||||
RelatedTicketRef {
|
||||
id: "T-2".to_string(),
|
||||
slug: Some("two".to_string()),
|
||||
},
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let snapshot = store.snapshot().unwrap();
|
||||
let preticket = snapshot
|
||||
.sessions
|
||||
.iter()
|
||||
.find(|session| session.worker_name == "ticket-intake-preticket")
|
||||
.unwrap();
|
||||
let shared = snapshot
|
||||
.sessions
|
||||
.iter()
|
||||
.find(|session| session.worker_name == "ticket-intake-shared")
|
||||
.unwrap();
|
||||
|
||||
assert!(preticket.related_tickets.is_empty());
|
||||
assert_eq!(shared.role, "intake");
|
||||
assert_eq!(shared.origin, RoleSessionOrigin::RoleLaunch);
|
||||
assert!(!shared.created_at.is_empty());
|
||||
assert!(!shared.updated_at.is_empty());
|
||||
assert_eq!(
|
||||
shared.related_tickets,
|
||||
vec![
|
||||
RelatedTicketRef {
|
||||
id: "T-1".to_string(),
|
||||
slug: Some("one".to_string()),
|
||||
},
|
||||
RelatedTicketRef {
|
||||
id: "T-2".to_string(),
|
||||
slug: Some("two".to_string()),
|
||||
},
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -1,756 +0,0 @@
|
||||
//! Inline-viewport "spawn Worker and attach" UX.
|
||||
//!
|
||||
//! Rendered at the user's current cursor position when `yoi` is invoked
|
||||
//! with no positional argument. Uses user-configured and bundled Profile
|
||||
//! choices plus bundled profiles, defaults to the builtin profile, prompts for
|
||||
//! the Worker's name, and on confirmation launches the Worker runtime command as an
|
||||
//! independent process. Once the process reports its socket via the
|
||||
//! `YOI-READY` stderr line, the dialog hands control back so main can
|
||||
//! switch the terminal to alternate-screen mode.
|
||||
//!
|
||||
//! The viewport's last frame stays in the terminal's scrollback so the
|
||||
//! user has a record of what was spawned (or why a spawn failed).
|
||||
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Duration;
|
||||
|
||||
use client::{SpawnConfig, WorkerRuntimeCommand, spawn_worker};
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use manifest::ProfileDiscovery;
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{Frame, TerminalOptions, Viewport};
|
||||
use session_store::SegmentId;
|
||||
|
||||
const VIEWPORT_LINES: u16 = 6;
|
||||
|
||||
pub struct SpawnReady {
|
||||
pub worker_name: String,
|
||||
pub socket_path: PathBuf,
|
||||
}
|
||||
|
||||
pub enum SpawnOutcome {
|
||||
Ready(SpawnReady),
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum SpawnError {
|
||||
Io(io::Error),
|
||||
Spawn(client::SpawnError),
|
||||
}
|
||||
|
||||
impl std::fmt::Display for SpawnError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Io(e) => write!(f, "io error: {e}"),
|
||||
Self::Spawn(e) => write!(f, "{e}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for SpawnError {}
|
||||
|
||||
impl From<io::Error> for SpawnError {
|
||||
fn from(e: io::Error) -> Self {
|
||||
Self::Io(e)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<client::SpawnError> for SpawnError {
|
||||
fn from(e: client::SpawnError) -> Self {
|
||||
Self::Spawn(e)
|
||||
}
|
||||
}
|
||||
|
||||
type InlineTerminal = Terminal<CrosstermBackend<io::Stdout>>;
|
||||
|
||||
/// Source session for a resume run. `None` = fresh spawn (current
|
||||
/// behaviour); `Some(id)` swaps the dialog into "Resume Worker" mode and
|
||||
/// passes `--session <id>` to the spawned Worker runtime child.
|
||||
pub async fn run(
|
||||
resume_from: Option<SegmentId>,
|
||||
worker_name: Option<String>,
|
||||
profile: Option<String>,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<SpawnOutcome, SpawnError> {
|
||||
let defaults = load_spawn_defaults()?;
|
||||
let mut profile_choices = if resume_from.is_some() {
|
||||
Vec::new()
|
||||
} else {
|
||||
defaults.profile_choices
|
||||
};
|
||||
let profile_index = initial_profile_index(
|
||||
&mut profile_choices,
|
||||
profile.as_deref(),
|
||||
defaults.default_profile_index,
|
||||
);
|
||||
|
||||
let selected_name = worker_name.unwrap_or(defaults.default_name);
|
||||
let immediate = resume_from.is_some() || profile.is_some() && !selected_name.is_empty();
|
||||
let mut form = Form {
|
||||
cwd: defaults.cwd.clone(),
|
||||
scope_origin: defaults.scope_origin,
|
||||
name_cursor: selected_name.chars().count(),
|
||||
name: selected_name,
|
||||
message: None,
|
||||
editing: true,
|
||||
resume_from,
|
||||
profile_choices,
|
||||
profile_index,
|
||||
};
|
||||
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
|
||||
// Phase 1: confirm / cancel.
|
||||
if !immediate {
|
||||
loop {
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
match poll_event()? {
|
||||
None => continue,
|
||||
Some(Action::Submit) => {
|
||||
if form.name.trim().is_empty() {
|
||||
form.message = Some(("name is required".to_string(), MessageKind::Error));
|
||||
continue;
|
||||
}
|
||||
break;
|
||||
}
|
||||
Some(Action::Cancel) => {
|
||||
form.editing = false;
|
||||
form.message = Some(("cancelled".to_string(), MessageKind::Info));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
drop(terminal);
|
||||
return Ok(SpawnOutcome::Cancelled);
|
||||
}
|
||||
Some(Action::Char(c)) => form.insert_char(c),
|
||||
Some(Action::Backspace) => form.backspace(),
|
||||
Some(Action::Delete) => form.delete_forward(),
|
||||
Some(Action::Left) => form.move_left(),
|
||||
Some(Action::Right) => form.move_right(),
|
||||
Some(Action::Home) => form.name_cursor = 0,
|
||||
Some(Action::End) => form.name_cursor = form.name.chars().count(),
|
||||
Some(Action::ProfileNext) => form.cycle_profile_next(),
|
||||
Some(Action::ProfilePrev) => form.cycle_profile_prev(),
|
||||
}
|
||||
}
|
||||
} else if form.name.trim().is_empty() {
|
||||
return Err(SpawnError::Io(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"name is required",
|
||||
)));
|
||||
}
|
||||
|
||||
// Phase 2: launch worker and wait for ready line. Drop the cursor
|
||||
// out of the name field — subsequent frames are passive status
|
||||
// updates, not input — so the cursor doesn't end up parked there
|
||||
// when the inline terminal is finally dropped.
|
||||
form.editing = false;
|
||||
form.message = Some(("starting worker...".to_string(), MessageKind::Progress));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
|
||||
match wait_for_ready(&mut terminal, &mut form, &runtime_command).await {
|
||||
Ok(ready) => {
|
||||
form.message = Some((
|
||||
format!("ready: {} attaching...", ready.worker_name),
|
||||
MessageKind::Ok,
|
||||
));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
drop(terminal);
|
||||
Ok(SpawnOutcome::Ready(ready))
|
||||
}
|
||||
Err(e) => {
|
||||
form.message = Some((e.to_string(), MessageKind::Error));
|
||||
let _ = terminal.draw(|f| draw_form(f, &form));
|
||||
drop(terminal);
|
||||
Err(e)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Launch a Worker runtime command with `--worker <name>` without opening the name dialog. The child Worker
|
||||
/// resolves persisted Worker metadata if present, or creates a fresh same-name Worker
|
||||
/// from the default profile.
|
||||
pub async fn run_worker_name(
|
||||
worker_name: String,
|
||||
runtime_command: WorkerRuntimeCommand,
|
||||
) -> Result<SpawnOutcome, SpawnError> {
|
||||
let defaults = load_spawn_defaults()?;
|
||||
let mut form = form_for_worker_name(worker_name, defaults);
|
||||
let mut terminal = make_inline_terminal()?;
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
|
||||
match wait_for_ready(&mut terminal, &mut form, &runtime_command).await {
|
||||
Ok(ready) => {
|
||||
form.message = Some((
|
||||
format!("ready: {} attaching...", ready.worker_name),
|
||||
MessageKind::Ok,
|
||||
));
|
||||
terminal.draw(|f| draw_form(f, &form))?;
|
||||
drop(terminal);
|
||||
Ok(SpawnOutcome::Ready(ready))
|
||||
}
|
||||
Err(e) => {
|
||||
form.message = Some((e.to_string(), MessageKind::Error));
|
||||
let _ = terminal.draw(|f| draw_form(f, &form));
|
||||
drop(terminal);
|
||||
Err(e)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct SpawnDefaults {
|
||||
cwd: PathBuf,
|
||||
scope_origin: ScopeOrigin,
|
||||
default_name: String,
|
||||
default_profile_index: usize,
|
||||
profile_choices: Vec<ProfileChoice>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct ProfileChoice {
|
||||
selector: Option<String>,
|
||||
label: String,
|
||||
is_default: bool,
|
||||
}
|
||||
|
||||
fn load_spawn_defaults() -> Result<SpawnDefaults, SpawnError> {
|
||||
let cwd = std::env::current_dir().map_err(SpawnError::Io)?;
|
||||
|
||||
let default_name = cwd
|
||||
.file_name()
|
||||
.and_then(|s| s.to_str())
|
||||
.map(sanitise_default_name)
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or_else(|| "worker".to_string());
|
||||
|
||||
let (profile_choices, default_profile_index) = profile_choices_for_cwd(&cwd);
|
||||
|
||||
Ok(SpawnDefaults {
|
||||
cwd,
|
||||
scope_origin: ScopeOrigin::FromProfile,
|
||||
default_name,
|
||||
default_profile_index,
|
||||
profile_choices,
|
||||
})
|
||||
}
|
||||
|
||||
fn profile_choices_for_cwd(cwd: &Path) -> (Vec<ProfileChoice>, usize) {
|
||||
let Ok(registry) = ProfileDiscovery::for_cwd(cwd).discover() else {
|
||||
return (Vec::new(), 0);
|
||||
};
|
||||
|
||||
let mut choices = Vec::new();
|
||||
for entry in registry.entries() {
|
||||
let mut label = entry.qualified_name();
|
||||
if entry.is_default {
|
||||
label.push_str(" (default)");
|
||||
}
|
||||
if let Some(description) = entry.description.as_deref() {
|
||||
label.push_str(" — ");
|
||||
label.push_str(description);
|
||||
}
|
||||
choices.push(ProfileChoice {
|
||||
selector: Some(entry.qualified_name()),
|
||||
label,
|
||||
is_default: entry.is_default,
|
||||
});
|
||||
}
|
||||
|
||||
let default_index = choices
|
||||
.iter()
|
||||
.position(|choice| choice.is_default)
|
||||
.unwrap_or(0);
|
||||
(choices, default_index)
|
||||
}
|
||||
|
||||
fn initial_profile_index(
|
||||
choices: &mut Vec<ProfileChoice>,
|
||||
explicit_profile: Option<&str>,
|
||||
default_index: usize,
|
||||
) -> usize {
|
||||
let Some(selector) = explicit_profile else {
|
||||
return default_index.min(choices.len().saturating_sub(1));
|
||||
};
|
||||
if let Some(index) = choices
|
||||
.iter()
|
||||
.position(|choice| choice.selector.as_deref() == Some(selector))
|
||||
{
|
||||
return index;
|
||||
}
|
||||
choices.push(ProfileChoice {
|
||||
selector: Some(selector.to_string()),
|
||||
label: selector.to_string(),
|
||||
is_default: false,
|
||||
});
|
||||
choices.len() - 1
|
||||
}
|
||||
|
||||
fn form_for_worker_name(worker_name: String, defaults: SpawnDefaults) -> Form {
|
||||
Form {
|
||||
cwd: defaults.cwd,
|
||||
scope_origin: defaults.scope_origin,
|
||||
name_cursor: worker_name.chars().count(),
|
||||
name: worker_name,
|
||||
message: Some(("resuming worker...".to_string(), MessageKind::Progress)),
|
||||
editing: false,
|
||||
resume_from: None,
|
||||
profile_choices: Vec::new(),
|
||||
profile_index: 0,
|
||||
}
|
||||
}
|
||||
|
||||
fn make_inline_terminal() -> io::Result<InlineTerminal> {
|
||||
let backend = CrosstermBackend::new(io::stdout());
|
||||
Terminal::with_options(
|
||||
backend,
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(VIEWPORT_LINES),
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
enum Action {
|
||||
Submit,
|
||||
Cancel,
|
||||
Char(char),
|
||||
Backspace,
|
||||
Delete,
|
||||
Left,
|
||||
Right,
|
||||
Home,
|
||||
End,
|
||||
ProfileNext,
|
||||
ProfilePrev,
|
||||
}
|
||||
|
||||
fn poll_event() -> io::Result<Option<Action>> {
|
||||
if !event::poll(Duration::from_millis(100))? {
|
||||
return Ok(None);
|
||||
}
|
||||
match event::read()? {
|
||||
TermEvent::Key(k) if k.kind != KeyEventKind::Release => {
|
||||
let ctrl = k.modifiers.contains(KeyModifiers::CONTROL);
|
||||
Ok(match k.code {
|
||||
KeyCode::Enter => Some(Action::Submit),
|
||||
KeyCode::Esc => Some(Action::Cancel),
|
||||
KeyCode::Char('c') if ctrl => Some(Action::Cancel),
|
||||
KeyCode::Char('a') if ctrl => Some(Action::Home),
|
||||
KeyCode::Char('e') if ctrl => Some(Action::End),
|
||||
KeyCode::Char('u') if ctrl => Some(Action::Cancel),
|
||||
KeyCode::Backspace => Some(Action::Backspace),
|
||||
KeyCode::Delete => Some(Action::Delete),
|
||||
KeyCode::Left => Some(Action::Left),
|
||||
KeyCode::Right => Some(Action::Right),
|
||||
KeyCode::Up | KeyCode::BackTab => Some(Action::ProfilePrev),
|
||||
KeyCode::Down | KeyCode::Tab => Some(Action::ProfileNext),
|
||||
KeyCode::Home => Some(Action::Home),
|
||||
KeyCode::End => Some(Action::End),
|
||||
KeyCode::Char(c) if !ctrl && is_safe_name_char(c) => Some(Action::Char(c)),
|
||||
_ => None,
|
||||
})
|
||||
}
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_safe_name_char(c: char) -> bool {
|
||||
// Filesystem-safe; worker.name becomes a runtime-dir name.
|
||||
c.is_ascii_alphanumeric() || matches!(c, '-' | '_' | '.')
|
||||
}
|
||||
|
||||
fn sanitise_default_name(s: &str) -> String {
|
||||
s.chars()
|
||||
.map(|c| if is_safe_name_char(c) { c } else { '-' })
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn wait_for_ready(
|
||||
terminal: &mut InlineTerminal,
|
||||
form: &mut Form,
|
||||
runtime_command: &WorkerRuntimeCommand,
|
||||
) -> Result<SpawnReady, SpawnError> {
|
||||
let config = SpawnConfig {
|
||||
runtime_command: runtime_command.clone(),
|
||||
worker_name: form.name.clone(),
|
||||
profile: form.selected_profile_selector(),
|
||||
workspace_root: form.cwd.clone(),
|
||||
cwd: None,
|
||||
resume_from: form.resume_from,
|
||||
};
|
||||
let ready = spawn_worker(config, |line| {
|
||||
form.message = Some((line.to_string(), MessageKind::Progress));
|
||||
let _ = terminal.draw(|f| draw_form(f, form));
|
||||
})
|
||||
.await?;
|
||||
Ok(SpawnReady {
|
||||
worker_name: ready.worker_name,
|
||||
socket_path: ready.socket_path,
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum MessageKind {
|
||||
Info,
|
||||
Ok,
|
||||
Error,
|
||||
Progress,
|
||||
}
|
||||
|
||||
enum ScopeOrigin {
|
||||
FromProfile,
|
||||
}
|
||||
|
||||
struct Form {
|
||||
cwd: PathBuf,
|
||||
/// Display label for the scope row in the dialog.
|
||||
scope_origin: ScopeOrigin,
|
||||
name: String,
|
||||
/// Cursor position counted in **chars**, not bytes — `name`
|
||||
/// currently only accepts ASCII so the two coincide, but we keep
|
||||
/// char-based bookkeeping in case we relax `is_safe_name_char`.
|
||||
name_cursor: usize,
|
||||
message: Option<(String, MessageKind)>,
|
||||
/// True while the dialog is accepting name input. Drives whether
|
||||
/// the rendered frame parks the terminal cursor inside the name
|
||||
/// field — when false (post-confirm / cancel / failure frames) the
|
||||
/// cursor stays out so it does not collide with the shell prompt
|
||||
/// after the inline terminal is dropped.
|
||||
editing: bool,
|
||||
/// `Some(id)` flips the dialog into "Resume Worker" mode: the title
|
||||
/// switches, the source session is shown to the user, and the
|
||||
/// child worker is launched with `--session <id>` so it restores
|
||||
/// from `id` and appends to the same session log.
|
||||
resume_from: Option<SegmentId>,
|
||||
/// Optional profile choices passed with `--profile` for
|
||||
/// fresh spawns. This is not used for resume/attach flows because those must
|
||||
/// restore Worker state rather than re-evaluate a profile source.
|
||||
profile_choices: Vec<ProfileChoice>,
|
||||
profile_index: usize,
|
||||
}
|
||||
|
||||
impl Form {
|
||||
fn insert_char(&mut self, c: char) {
|
||||
let byte = self.char_offset_to_byte(self.name_cursor);
|
||||
self.name.insert(byte, c);
|
||||
self.name_cursor += 1;
|
||||
}
|
||||
|
||||
fn backspace(&mut self) {
|
||||
if self.name_cursor == 0 {
|
||||
return;
|
||||
}
|
||||
let end = self.char_offset_to_byte(self.name_cursor);
|
||||
let start = self.char_offset_to_byte(self.name_cursor - 1);
|
||||
self.name.replace_range(start..end, "");
|
||||
self.name_cursor -= 1;
|
||||
}
|
||||
|
||||
fn delete_forward(&mut self) {
|
||||
let total = self.name.chars().count();
|
||||
if self.name_cursor >= total {
|
||||
return;
|
||||
}
|
||||
let start = self.char_offset_to_byte(self.name_cursor);
|
||||
let end = self.char_offset_to_byte(self.name_cursor + 1);
|
||||
self.name.replace_range(start..end, "");
|
||||
}
|
||||
|
||||
fn move_left(&mut self) {
|
||||
if self.name_cursor > 0 {
|
||||
self.name_cursor -= 1;
|
||||
}
|
||||
}
|
||||
|
||||
fn move_right(&mut self) {
|
||||
let total = self.name.chars().count();
|
||||
if self.name_cursor < total {
|
||||
self.name_cursor += 1;
|
||||
}
|
||||
}
|
||||
|
||||
fn selected_profile(&self) -> Option<&ProfileChoice> {
|
||||
self.profile_choices
|
||||
.get(self.profile_index)
|
||||
.filter(|choice| choice.selector.is_some())
|
||||
}
|
||||
|
||||
fn selected_profile_selector(&self) -> Option<String> {
|
||||
self.selected_profile()
|
||||
.and_then(|choice| choice.selector.clone())
|
||||
}
|
||||
|
||||
fn cycle_profile_next(&mut self) {
|
||||
if self.profile_choices.is_empty() {
|
||||
return;
|
||||
}
|
||||
self.profile_index = (self.profile_index + 1) % self.profile_choices.len();
|
||||
self.message = None;
|
||||
}
|
||||
|
||||
fn cycle_profile_prev(&mut self) {
|
||||
if self.profile_choices.is_empty() {
|
||||
return;
|
||||
}
|
||||
self.profile_index = if self.profile_index == 0 {
|
||||
self.profile_choices.len() - 1
|
||||
} else {
|
||||
self.profile_index - 1
|
||||
};
|
||||
self.message = None;
|
||||
}
|
||||
|
||||
fn char_offset_to_byte(&self, char_off: usize) -> usize {
|
||||
self.name
|
||||
.char_indices()
|
||||
.nth(char_off)
|
||||
.map(|(b, _)| b)
|
||||
.unwrap_or(self.name.len())
|
||||
}
|
||||
}
|
||||
|
||||
fn draw_form(f: &mut Frame<'_>, form: &Form) {
|
||||
let area = f.area();
|
||||
let layout = Layout::vertical([
|
||||
Constraint::Length(1), // title
|
||||
Constraint::Length(1), // name field
|
||||
Constraint::Length(1), // context (profile or scope default)
|
||||
Constraint::Length(1), // hint
|
||||
Constraint::Length(1), // message
|
||||
Constraint::Length(1), // spacer
|
||||
])
|
||||
.split(area);
|
||||
|
||||
let title_text = match form.resume_from {
|
||||
Some(id) => format!("resume worker session: {}", short_segment(id)),
|
||||
None => "spawn worker".to_string(),
|
||||
};
|
||||
let title = Paragraph::new(Line::from(vec![Span::styled(
|
||||
title_text,
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
)]));
|
||||
f.render_widget(title, layout[0]);
|
||||
|
||||
f.render_widget(Paragraph::new(name_line(form)), layout[1]);
|
||||
f.render_widget(Paragraph::new(context_line(form)), layout[2]);
|
||||
f.render_widget(Paragraph::new(hint_line()), layout[3]);
|
||||
f.render_widget(Paragraph::new(message_line(form)), layout[4]);
|
||||
|
||||
if form.editing {
|
||||
// Place the cursor inside the name field while the user is
|
||||
// editing. Skipped on post-confirm frames so the inline
|
||||
// viewport's drop leaves the cursor at the bottom of the
|
||||
// rendered area rather than parked on the name line, which
|
||||
// would let the shell prompt (or any later eprintln) clobber
|
||||
// the rendered name field after exit.
|
||||
let cursor_col = 2 + "name: ".len() + form.name_cursor;
|
||||
f.set_cursor_position((layout[1].x + cursor_col as u16, layout[1].y));
|
||||
}
|
||||
}
|
||||
|
||||
/// First 8 hex digits of a UUID — short enough to skim, long enough
|
||||
/// to disambiguate inside a 10-row picker.
|
||||
pub(crate) fn short_segment(id: SegmentId) -> String {
|
||||
let s = id.to_string();
|
||||
s.chars().take(8).collect()
|
||||
}
|
||||
|
||||
fn name_line(form: &Form) -> Line<'_> {
|
||||
Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("name: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled(
|
||||
form.name.as_str(),
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
fn context_line(form: &Form) -> Line<'_> {
|
||||
if let Some(profile) = form.profile_choices.get(form.profile_index) {
|
||||
return Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("profile: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled(profile.label.as_str(), Style::default().fg(Color::Green)),
|
||||
Span::styled(
|
||||
" (tab/down to change)",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
),
|
||||
]);
|
||||
}
|
||||
|
||||
match form.scope_origin {
|
||||
ScopeOrigin::FromProfile => Line::from(vec![
|
||||
Span::raw(" "),
|
||||
Span::styled("scope: ", Style::default().fg(Color::DarkGray)),
|
||||
Span::styled("from selected profile", Style::default().fg(Color::Green)),
|
||||
]),
|
||||
}
|
||||
}
|
||||
|
||||
fn hint_line() -> Line<'static> {
|
||||
Line::from(vec![Span::styled(
|
||||
" enter spawn · tab/down next profile · shift-tab/up prev · esc cancel",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
)])
|
||||
}
|
||||
|
||||
fn message_line(form: &Form) -> Line<'_> {
|
||||
let Some((text, kind)) = form.message.as_ref() else {
|
||||
return Line::from("");
|
||||
};
|
||||
let style = match kind {
|
||||
MessageKind::Info => Style::default().fg(Color::DarkGray),
|
||||
MessageKind::Ok => Style::default().fg(Color::Green),
|
||||
MessageKind::Error => Style::default().fg(Color::Red),
|
||||
MessageKind::Progress => Style::default().fg(Color::Yellow),
|
||||
};
|
||||
Line::from(vec![Span::raw(" "), Span::styled(text.as_str(), style)])
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn form(name: &str) -> Form {
|
||||
Form {
|
||||
cwd: PathBuf::from("/work/example"),
|
||||
scope_origin: ScopeOrigin::FromProfile,
|
||||
name: name.to_string(),
|
||||
name_cursor: name.chars().count(),
|
||||
message: None,
|
||||
editing: true,
|
||||
resume_from: None,
|
||||
profile_choices: Vec::new(),
|
||||
profile_index: 0,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn worker_name_form_restores_or_creates_by_worker_name() {
|
||||
let defaults = SpawnDefaults {
|
||||
cwd: PathBuf::from("/work/example"),
|
||||
scope_origin: ScopeOrigin::FromProfile,
|
||||
default_name: "ignored".to_string(),
|
||||
default_profile_index: 0,
|
||||
profile_choices: Vec::new(),
|
||||
};
|
||||
let f = form_for_worker_name("agent".to_string(), defaults);
|
||||
|
||||
assert_eq!(f.name, "agent");
|
||||
assert_eq!(f.name_cursor, "agent".chars().count());
|
||||
assert_eq!(f.resume_from, None);
|
||||
assert!(!f.editing);
|
||||
assert_eq!(
|
||||
f.message,
|
||||
Some(("resuming worker...".to_string(), MessageKind::Progress))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn profile_choices_ignore_repository_local_profile_registry() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
let project = temp.path().join("project");
|
||||
let yoi = project.join(".yoi");
|
||||
std::fs::create_dir_all(&yoi).unwrap();
|
||||
std::fs::write(
|
||||
yoi.join("profiles.toml"),
|
||||
"default = \"coder\"\n[profile]\ncoder = \"profiles/coder.toml\"\n",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let (choices, default_index) = profile_choices_for_cwd(&project);
|
||||
assert_eq!(default_index, 0);
|
||||
assert!(
|
||||
choices
|
||||
.iter()
|
||||
.all(|choice| { choice.selector.as_deref() != Some("project:coder") })
|
||||
);
|
||||
assert!(
|
||||
choices
|
||||
.iter()
|
||||
.any(|choice| { choice.selector.as_deref() == Some("builtin:companion") })
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn profile_cycle_selects_only_discovered_profiles() {
|
||||
let mut form = form("coder");
|
||||
form.profile_choices = vec![
|
||||
ProfileChoice {
|
||||
selector: Some("project:coder".to_string()),
|
||||
label: "project:coder (default)".to_string(),
|
||||
is_default: true,
|
||||
},
|
||||
ProfileChoice {
|
||||
selector: Some("user:reviewer".to_string()),
|
||||
label: "user:reviewer".to_string(),
|
||||
is_default: false,
|
||||
},
|
||||
];
|
||||
form.profile_index = 0;
|
||||
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("project:coder")
|
||||
);
|
||||
form.cycle_profile_next();
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("user:reviewer")
|
||||
);
|
||||
form.cycle_profile_next();
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("project:coder")
|
||||
);
|
||||
form.cycle_profile_prev();
|
||||
assert_eq!(
|
||||
form.selected_profile_selector().as_deref(),
|
||||
Some("user:reviewer")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn initial_profile_index_adds_explicit_selector_not_in_discovery_list() {
|
||||
let mut choices = Vec::new();
|
||||
let selected = initial_profile_index(&mut choices, Some("coder"), 0);
|
||||
assert_eq!(selected, 0);
|
||||
assert_eq!(choices[0].selector.as_deref(), Some("coder"));
|
||||
assert_eq!(choices[0].label, "coder");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn name_input_handles_insert_backspace_and_cursor() {
|
||||
let mut f = form("");
|
||||
for c in "abc".chars() {
|
||||
f.insert_char(c);
|
||||
}
|
||||
assert_eq!(f.name, "abc");
|
||||
assert_eq!(f.name_cursor, 3);
|
||||
|
||||
f.move_left();
|
||||
f.move_left();
|
||||
f.insert_char('X');
|
||||
assert_eq!(f.name, "aXbc");
|
||||
|
||||
f.backspace();
|
||||
assert_eq!(f.name, "abc");
|
||||
assert_eq!(f.name_cursor, 1);
|
||||
|
||||
f.delete_forward();
|
||||
assert_eq!(f.name, "ac");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sanitise_default_name_replaces_unsafe_chars() {
|
||||
assert_eq!(sanitise_default_name("my project!"), "my-project-");
|
||||
assert_eq!(sanitise_default_name("ok-name_2.0"), "ok-name_2.0");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,166 @@
|
||||
use std::io;
|
||||
use std::time::Duration;
|
||||
|
||||
use client::{StandaloneSessionListIntent, StandaloneSessionResumeIntent, Target};
|
||||
use crossterm::event::{self, Event as TermEvent, KeyCode, KeyEventKind, KeyModifiers};
|
||||
use ratatui::Terminal;
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::layout::{Constraint, Layout};
|
||||
use ratatui::prelude::{Color, Line, Modifier, Span, Style};
|
||||
use ratatui::widgets::Paragraph;
|
||||
use ratatui::{TerminalOptions, Viewport};
|
||||
use standalone::{StandaloneListScope, StandaloneSessionRecord, StandaloneSessionStore};
|
||||
use thiserror::Error;
|
||||
|
||||
const LIMIT: usize = 100;
|
||||
|
||||
pub(crate) fn pick(
|
||||
target: &dyn Target,
|
||||
include_all: bool,
|
||||
) -> Result<Option<StandaloneSessionResumeIntent>, StandalonePickerError> {
|
||||
let intent = target
|
||||
.standalone_session_list(include_all)
|
||||
.map_err(StandalonePickerError::Target)?;
|
||||
let records = load_records(&intent)?;
|
||||
if records.is_empty() {
|
||||
return Err(StandalonePickerError::NoSessions { include_all });
|
||||
}
|
||||
let selected = run_picker(records)?;
|
||||
selected
|
||||
.map(|record| {
|
||||
target
|
||||
.standalone_session_resume(record.session_id.to_string())
|
||||
.map_err(StandalonePickerError::Target)
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn load_records(
|
||||
intent: &StandaloneSessionListIntent,
|
||||
) -> Result<Vec<StandaloneSessionRecord>, StandalonePickerError> {
|
||||
let store = StandaloneSessionStore::open(&intent.state_dir)
|
||||
.map_err(StandalonePickerError::StateStore)?;
|
||||
store
|
||||
.list(
|
||||
&intent.cwd,
|
||||
if intent.include_all {
|
||||
StandaloneListScope::All
|
||||
} else {
|
||||
StandaloneListScope::CurrentCwd
|
||||
},
|
||||
LIMIT,
|
||||
)
|
||||
.map_err(StandalonePickerError::StateStore)
|
||||
}
|
||||
|
||||
fn run_picker(
|
||||
records: Vec<StandaloneSessionRecord>,
|
||||
) -> Result<Option<StandaloneSessionRecord>, StandalonePickerError> {
|
||||
let height = u16::try_from(records.len().saturating_add(3).min(20)).unwrap_or(20);
|
||||
let mut terminal = Terminal::with_options(
|
||||
CrosstermBackend::new(io::stdout()),
|
||||
TerminalOptions {
|
||||
viewport: Viewport::Inline(height),
|
||||
},
|
||||
)
|
||||
.map_err(StandalonePickerError::Io)?;
|
||||
let mut selected = 0usize;
|
||||
loop {
|
||||
terminal
|
||||
.draw(|frame| draw(frame, &records, selected))
|
||||
.map_err(StandalonePickerError::Io)?;
|
||||
if !event::poll(Duration::from_millis(100)).map_err(StandalonePickerError::Io)? {
|
||||
continue;
|
||||
}
|
||||
let TermEvent::Key(key) = event::read().map_err(StandalonePickerError::Io)? else {
|
||||
continue;
|
||||
};
|
||||
if key.kind == KeyEventKind::Release {
|
||||
continue;
|
||||
}
|
||||
let ctrl = key.modifiers.contains(KeyModifiers::CONTROL);
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') if !ctrl => {
|
||||
selected = selected.saturating_sub(1);
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if !ctrl => {
|
||||
selected = (selected + 1).min(records.len() - 1);
|
||||
}
|
||||
KeyCode::Enter => return Ok(Some(records[selected].clone())),
|
||||
KeyCode::Esc => return Ok(None),
|
||||
KeyCode::Char('c') if ctrl => return Ok(None),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn draw(frame: &mut ratatui::Frame<'_>, records: &[StandaloneSessionRecord], selected: usize) {
|
||||
let mut constraints = vec![Constraint::Length(1)];
|
||||
constraints.extend(records.iter().map(|_| Constraint::Length(1)));
|
||||
constraints.push(Constraint::Length(1));
|
||||
let rows = Layout::vertical(constraints).split(frame.area());
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(Span::styled(
|
||||
"resume standalone session",
|
||||
Style::default().add_modifier(Modifier::BOLD),
|
||||
))),
|
||||
rows[0],
|
||||
);
|
||||
for (index, record) in records.iter().enumerate() {
|
||||
let active = index == selected;
|
||||
let marker = if active { "▶ " } else { " " };
|
||||
let style = if active {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
};
|
||||
let cwd = record.cwd.canonical_path.display();
|
||||
frame.render_widget(
|
||||
Paragraph::new(Line::from(vec![
|
||||
Span::raw(marker),
|
||||
Span::styled(record.session_id.short(), style),
|
||||
Span::raw(format!(
|
||||
" [{:?}] updated:{} {}",
|
||||
record.status, record.updated_at_unix_ms, cwd
|
||||
)),
|
||||
])),
|
||||
rows[index + 1],
|
||||
);
|
||||
}
|
||||
frame.render_widget(
|
||||
Paragraph::new(" [↑/↓] select [enter] restore [esc] cancel"),
|
||||
rows[records.len() + 1],
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub(crate) enum StandalonePickerError {
|
||||
#[error("standalone target error: {0}")]
|
||||
Target(#[source] client::TargetError),
|
||||
#[error("standalone session state is unavailable: {0}")]
|
||||
StateStore(#[source] standalone::StandaloneStoreError),
|
||||
#[error(
|
||||
"no standalone sessions found for this cwd; use `yoi --local --resume --all` to include all cwd identities"
|
||||
)]
|
||||
NoSessions { include_all: bool },
|
||||
#[error("standalone session picker I/O failed: {0}")]
|
||||
Io(#[source] io::Error),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use client::StandaloneTarget;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn empty_picker_keeps_current_cwd_as_default_scope() {
|
||||
let temp = tempfile::tempdir().expect("tempdir");
|
||||
let target = StandaloneTarget::new(temp.path());
|
||||
let error = pick(&target, false).expect_err("empty picker should fail explicitly");
|
||||
assert!(error.to_string().contains("this cwd"));
|
||||
assert!(error.to_string().contains("--all"));
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -251,14 +251,8 @@ impl DelegatingWorkdirSession {
|
||||
self.ensure_path(path, WorkdirDelegationPermission::Write)
|
||||
}
|
||||
|
||||
fn ensure_command(&self, starting: bool) -> Result<(), WorkdirError> {
|
||||
self.ensure_capability(WorkdirSessionCapability::Command, "command execution")?;
|
||||
if starting && self.has_active_write_lease() {
|
||||
return Err(WorkdirError::Denied(
|
||||
"command execution is denied while a child holds a write delegation".into(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
fn ensure_command(&self) -> Result<(), WorkdirError> {
|
||||
self.ensure_capability(WorkdirSessionCapability::Command, "command execution")
|
||||
}
|
||||
|
||||
fn ensure_parent_write_available(&self, path: &FsPath) -> Result<(), WorkdirError> {
|
||||
@@ -281,20 +275,6 @@ impl DelegatingWorkdirSession {
|
||||
}
|
||||
}
|
||||
|
||||
fn has_active_write_lease(&self) -> bool {
|
||||
let mut leases = self
|
||||
.child_write_leases
|
||||
.lock()
|
||||
.expect("workdir delegation lease mutex poisoned");
|
||||
leases.retain(|_, lease| lease.validity.upgrade().is_some_and(|v| v.is_active()));
|
||||
leases.values().any(|lease| {
|
||||
lease
|
||||
.rules
|
||||
.iter()
|
||||
.any(|rule| rule.permission == WorkdirDelegationPermission::Write)
|
||||
})
|
||||
}
|
||||
|
||||
fn validate_delegation_rules(
|
||||
&self,
|
||||
rules: &[WorkdirDelegationRule],
|
||||
@@ -311,7 +291,10 @@ impl DelegatingWorkdirSession {
|
||||
if !self.capabilities.supports(WorkdirSessionCapability::Read)
|
||||
|| (writable
|
||||
&& (!self.capabilities.supports(WorkdirSessionCapability::Write)
|
||||
|| !self.capabilities.supports(WorkdirSessionCapability::Edit)))
|
||||
|| !self.capabilities.supports(WorkdirSessionCapability::Edit)
|
||||
|| !self
|
||||
.capabilities
|
||||
.supports(WorkdirSessionCapability::Command)))
|
||||
{
|
||||
return Err(WorkdirError::Denied(
|
||||
"parent workdir session cannot delegate the requested capabilities".into(),
|
||||
@@ -342,6 +325,7 @@ impl DelegatingWorkdirSession {
|
||||
if writable {
|
||||
delegated.push(WorkdirSessionCapability::Write);
|
||||
delegated.push(WorkdirSessionCapability::Edit);
|
||||
delegated.push(WorkdirSessionCapability::Command);
|
||||
}
|
||||
Ok(WorkdirSessionCapabilities::from_capabilities(delegated))
|
||||
}
|
||||
@@ -499,12 +483,12 @@ impl WorkdirSession for DelegatingWorkdirSession {
|
||||
}
|
||||
|
||||
async fn start_command(&self, request: CommandRequest) -> Result<CommandHandle, WorkdirError> {
|
||||
self.ensure_command(true)?;
|
||||
self.ensure_command()?;
|
||||
self.source.start_command(request).await
|
||||
}
|
||||
|
||||
async fn command_status(&self, handle: CommandHandle) -> Result<CommandStatus, WorkdirError> {
|
||||
self.ensure_command(false)?;
|
||||
self.ensure_command()?;
|
||||
self.source.command_status(handle).await
|
||||
}
|
||||
|
||||
@@ -512,12 +496,12 @@ impl WorkdirSession for DelegatingWorkdirSession {
|
||||
&self,
|
||||
request: CommandOutputRequest,
|
||||
) -> Result<CommandOutput, WorkdirError> {
|
||||
self.ensure_command(false)?;
|
||||
self.ensure_command()?;
|
||||
self.source.command_output(request).await
|
||||
}
|
||||
|
||||
async fn cancel_command(&self, handle: CommandHandle) -> Result<(), WorkdirError> {
|
||||
self.ensure_command(false)?;
|
||||
self.ensure_command()?;
|
||||
self.source.cancel_command(handle).await
|
||||
}
|
||||
|
||||
@@ -749,6 +733,31 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_command(
|
||||
session: &WorkdirSessionHandle,
|
||||
command: impl Into<String>,
|
||||
tool_call_id: impl Into<String>,
|
||||
) -> CommandOutput {
|
||||
let handle = session
|
||||
.start_command(CommandRequest {
|
||||
command: command.into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 1024,
|
||||
tool_call_id: Some(tool_call_id.into()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
session
|
||||
.command_output(CommandOutputRequest {
|
||||
handle,
|
||||
cursor: 0,
|
||||
limit: 1024,
|
||||
wait: true,
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delegation_capable_session_forwards_command_telemetry() {
|
||||
let root = TempDir::new().unwrap();
|
||||
@@ -842,6 +851,18 @@ mod tests {
|
||||
);
|
||||
assert!(child.scoped_session.subscribe_command_events().is_none());
|
||||
assert!(child.scoped_session.command_snapshot().is_empty());
|
||||
assert!(matches!(
|
||||
child
|
||||
.scoped_session
|
||||
.start_command(CommandRequest {
|
||||
command: "printf denied".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 1024,
|
||||
tool_call_id: Some("read-only-command".into()),
|
||||
})
|
||||
.await,
|
||||
Err(WorkdirError::Denied(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
@@ -920,7 +941,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn write_lease_blocks_parent_region_until_release() {
|
||||
async fn write_lease_keeps_typed_parent_writes_exclusive_without_blocking_commands() {
|
||||
let root = TempDir::new().unwrap();
|
||||
fs::create_dir_all(root.path().join("leased")).unwrap();
|
||||
fs::create_dir_all(root.path().join("other")).unwrap();
|
||||
@@ -929,6 +950,30 @@ mod tests {
|
||||
.delegate(request("leased", WorkdirDelegationPermission::Write))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(
|
||||
child
|
||||
.capabilities
|
||||
.supports(WorkdirSessionCapability::Command)
|
||||
);
|
||||
let child_output = run_command(
|
||||
&child.scoped_session,
|
||||
"printf child-command",
|
||||
"delegated-child-command",
|
||||
)
|
||||
.await;
|
||||
assert_eq!(child_output.content, "child-command");
|
||||
let parent_output = run_command(
|
||||
&parent,
|
||||
"printf parent-write > leased/from-command; printf parent-command",
|
||||
"parent-command-during-child-write",
|
||||
)
|
||||
.await;
|
||||
assert_eq!(parent_output.status, CommandStatus::Completed);
|
||||
assert_eq!(parent_output.content, "parent-command");
|
||||
assert_eq!(
|
||||
fs::read_to_string(root.path().join("leased/from-command")).unwrap(),
|
||||
"parent-write"
|
||||
);
|
||||
|
||||
assert!(matches!(
|
||||
parent.write(write("leased/file", "parent")).await,
|
||||
@@ -941,6 +986,18 @@ mod tests {
|
||||
.await
|
||||
.unwrap();
|
||||
child.release();
|
||||
assert!(matches!(
|
||||
child
|
||||
.scoped_session
|
||||
.start_command(CommandRequest {
|
||||
command: "printf revoked".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 1024,
|
||||
tool_call_id: Some("revoked-child-command".into()),
|
||||
})
|
||||
.await,
|
||||
Err(WorkdirError::SessionClosed)
|
||||
));
|
||||
parent
|
||||
.write(write("leased/parent", "parent"))
|
||||
.await
|
||||
@@ -992,6 +1049,78 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn nested_write_leases_do_not_block_command_capable_ancestors() {
|
||||
let root = TempDir::new().unwrap();
|
||||
fs::create_dir_all(root.path().join("docs/sub")).unwrap();
|
||||
let root_session = session(root.path());
|
||||
let child = root_session
|
||||
.delegate(request("docs", WorkdirDelegationPermission::Write))
|
||||
.await
|
||||
.unwrap();
|
||||
let nested = child
|
||||
.scoped_session
|
||||
.delegate(request("docs/sub", WorkdirDelegationPermission::Write))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
for (session, label) in [
|
||||
(&root_session, "root"),
|
||||
(&child.scoped_session, "child"),
|
||||
(&nested.scoped_session, "nested"),
|
||||
] {
|
||||
let output = run_command(
|
||||
session,
|
||||
format!("printf {label}"),
|
||||
format!("{label}-command-during-nested-write"),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(output.status, CommandStatus::Completed);
|
||||
assert_eq!(output.content, label);
|
||||
}
|
||||
|
||||
assert!(matches!(
|
||||
root_session.write(write("docs/root", "blocked")).await,
|
||||
Err(WorkdirError::Denied(_))
|
||||
));
|
||||
assert!(matches!(
|
||||
child
|
||||
.scoped_session
|
||||
.write(write("sub/child", "blocked"))
|
||||
.await,
|
||||
Err(WorkdirError::Denied(_))
|
||||
));
|
||||
nested
|
||||
.scoped_session
|
||||
.write(write("nested", "allowed"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
nested.release();
|
||||
child.release();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reapplied_write_delegation_chain_forwards_command_lifecycle() {
|
||||
let root = TempDir::new().unwrap();
|
||||
fs::create_dir_all(root.path().join("delegated")).unwrap();
|
||||
let applied = apply_delegation_chain(
|
||||
session(root.path()),
|
||||
[request("delegated", WorkdirDelegationPermission::Write)],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let output = run_command(
|
||||
&applied.scoped_session,
|
||||
"printf reapplied",
|
||||
"reapplied-command",
|
||||
)
|
||||
.await;
|
||||
assert_eq!(output.status, CommandStatus::Completed);
|
||||
assert_eq!(output.content, "reapplied");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn applied_chain_cannot_replace_outer_provider_attenuation() {
|
||||
let root = TempDir::new().unwrap();
|
||||
@@ -1036,6 +1165,17 @@ mod tests {
|
||||
.unwrap();
|
||||
|
||||
parent.close().await.unwrap();
|
||||
assert!(matches!(
|
||||
parent
|
||||
.start_command(CommandRequest {
|
||||
command: "printf closed".into(),
|
||||
timeout_secs: 5,
|
||||
output_limit: 1024,
|
||||
tool_call_id: Some("closed-parent-command".into()),
|
||||
})
|
||||
.await,
|
||||
Err(WorkdirError::SessionClosed)
|
||||
));
|
||||
assert!(matches!(
|
||||
child.scoped_session.read(read("a")).await,
|
||||
Err(WorkdirError::SessionClosed)
|
||||
|
||||
@@ -107,6 +107,31 @@ pub enum WorkdirTransportErrorCode {
|
||||
Internal,
|
||||
}
|
||||
|
||||
impl WorkdirTransportErrorCode {
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::NotFound => "not_found",
|
||||
Self::Conflict => "conflict",
|
||||
Self::Unsupported => "unsupported",
|
||||
Self::InvalidRequest => "invalid_request",
|
||||
Self::UnknownCommand => "unknown_command",
|
||||
Self::Unavailable => "unavailable",
|
||||
Self::Internal => "internal",
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared public HTTP classification for Runtime and Workspace Workdir operation boundaries.
|
||||
pub const fn http_status(self) -> u16 {
|
||||
match self {
|
||||
Self::NotFound | Self::UnknownCommand => 404,
|
||||
Self::Conflict => 409,
|
||||
Self::Unsupported | Self::InvalidRequest => 400,
|
||||
Self::Unavailable => 503,
|
||||
Self::Internal => 500,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkdirTransportError {
|
||||
pub code: WorkdirTransportErrorCode,
|
||||
@@ -126,6 +151,9 @@ impl WorkdirTransportError {
|
||||
message: format!("Workdir capability {capability:?} is not available"),
|
||||
};
|
||||
}
|
||||
WorkdirError::UnsupportedOperation(_) => {
|
||||
(Code::Unsupported, "Workdir operation is not supported")
|
||||
}
|
||||
WorkdirError::UnknownCommand(_) => {
|
||||
(Code::UnknownCommand, "Workdir command was not found")
|
||||
}
|
||||
@@ -161,7 +189,7 @@ impl WorkdirTransportError {
|
||||
match self.code {
|
||||
Code::NotFound => WorkdirError::NotFound("<remote>".into()),
|
||||
Code::Conflict => WorkdirError::Conflict(self.message),
|
||||
Code::Unsupported => WorkdirError::Unavailable(self.message),
|
||||
Code::Unsupported => WorkdirError::UnsupportedOperation(self.message),
|
||||
Code::UnknownCommand => WorkdirError::UnknownCommand("<remote>".to_string()),
|
||||
Code::InvalidRequest => WorkdirError::InvalidArgument(self.message),
|
||||
Code::Unavailable => WorkdirError::Unavailable(self.message),
|
||||
@@ -517,13 +545,13 @@ mod client {
|
||||
.json::<WorkdirTransportError>()
|
||||
.await
|
||||
.map(WorkdirTransportError::into_workdir_error)
|
||||
.unwrap_or_else(|error| {
|
||||
WorkdirError::Unavailable(format!("Runtime HTTP error: {error}"))
|
||||
.unwrap_or_else(|_| {
|
||||
WorkdirError::Transport("Runtime Workdir error response was invalid".to_string())
|
||||
})
|
||||
}
|
||||
|
||||
fn http_unavailable(error: reqwest::Error) -> WorkdirError {
|
||||
WorkdirError::Unavailable(format!("Runtime Workdir HTTP request failed: {error}"))
|
||||
fn http_unavailable(_error: reqwest::Error) -> WorkdirError {
|
||||
WorkdirError::Transport("Runtime Workdir HTTP request failed".to_string())
|
||||
}
|
||||
|
||||
pub use self::RemoteWorkdirSession as ClientSession;
|
||||
@@ -536,6 +564,57 @@ pub use client::{ClientSession as RemoteWorkdirSession, WorkdirHttpAuthorization
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn transport_error_round_trip_keeps_public_classification() {
|
||||
for (code, expected_status, expected_error) in [
|
||||
(
|
||||
WorkdirTransportErrorCode::InvalidRequest,
|
||||
400,
|
||||
"invalid argument",
|
||||
),
|
||||
(WorkdirTransportErrorCode::NotFound, 404, "file not found"),
|
||||
(
|
||||
WorkdirTransportErrorCode::UnknownCommand,
|
||||
404,
|
||||
"unknown Workdir session command",
|
||||
),
|
||||
(
|
||||
WorkdirTransportErrorCode::Conflict,
|
||||
409,
|
||||
"modified externally",
|
||||
),
|
||||
(WorkdirTransportErrorCode::Unsupported, 400, "unsupported"),
|
||||
(WorkdirTransportErrorCode::Unavailable, 503, "unavailable"),
|
||||
(WorkdirTransportErrorCode::Internal, 500, "transport failed"),
|
||||
] {
|
||||
let transport = WorkdirTransportError {
|
||||
code,
|
||||
message: "safe provider message".to_string(),
|
||||
};
|
||||
assert_eq!(code.http_status(), expected_status);
|
||||
let workdir_error = transport.clone().into_workdir_error();
|
||||
assert!(workdir_error.to_string().contains(expected_error));
|
||||
assert_eq!(
|
||||
WorkdirTransportError::from_workdir_error(&workdir_error).code,
|
||||
code
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn local_validation_errors_share_invalid_request_classification() {
|
||||
for error in [
|
||||
WorkdirError::InvalidGlob("[".to_string()),
|
||||
WorkdirError::InvalidRegex("(".to_string()),
|
||||
WorkdirError::InvalidArgument("limit must be positive".to_string()),
|
||||
] {
|
||||
let transport = WorkdirTransportError::from_workdir_error(&error);
|
||||
assert_eq!(transport.code, WorkdirTransportErrorCode::InvalidRequest);
|
||||
assert_eq!(transport.code.http_status(), 400);
|
||||
assert_eq!(transport.message, "Workdir operation request is invalid");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_failure_remains_distinct_from_session_unavailable() {
|
||||
let transport = WorkdirTransportError::from_workdir_error(&WorkdirError::Transport(
|
||||
|
||||
@@ -28,7 +28,9 @@ pub use fs_operation::{
|
||||
GlobResult, GrepOutputMode, GrepRequest, GrepResult, ListEntry, ListRequest, ListResult,
|
||||
ReadRequest, ReadResult, StatRequest, StatResult, WriteRequest, WriteResult,
|
||||
};
|
||||
pub use local::{LocalWorkdirSession, SymlinkInfo, direct_symlink, first_symlink};
|
||||
pub use local::{
|
||||
LocalWorkdirSession, SymlinkInfo, WorkdirSessionResource, direct_symlink, first_symlink,
|
||||
};
|
||||
pub use operation::*;
|
||||
|
||||
/// Persistent, opaque identity of one materialized Workdir.
|
||||
@@ -223,6 +225,9 @@ pub enum WorkdirError {
|
||||
#[error("Workdir session does not support {0:?}")]
|
||||
Unsupported(WorkdirSessionCapability),
|
||||
|
||||
#[error("Workdir operation is unsupported: {0}")]
|
||||
UnsupportedOperation(String),
|
||||
|
||||
#[error("invalid Workdir path: {0}")]
|
||||
InvalidPath(String),
|
||||
|
||||
|
||||
@@ -8,7 +8,8 @@
|
||||
//! `LocalWorkdirSession` is cheap to clone (`Arc` inside). Tool-specific session
|
||||
//! state, such as read-before-edit tracking, remains owned by the tool layer.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
use std::fmt::Debug;
|
||||
#[cfg(test)]
|
||||
use std::io::Write as _;
|
||||
use std::io::{Read as _, Seek as _, SeekFrom};
|
||||
@@ -228,6 +229,8 @@ struct LocalWorkdirSessionInner {
|
||||
next_command_id: AtomicU64,
|
||||
commands: Mutex<HashMap<String, LocalCommand>>,
|
||||
command_telemetry: CommandTelemetry,
|
||||
command_environment: BTreeMap<String, String>,
|
||||
resources: StdMutex<Vec<Arc<dyn WorkdirSessionResource>>>,
|
||||
}
|
||||
|
||||
impl Drop for LocalWorkdirSessionInner {
|
||||
@@ -242,6 +245,9 @@ impl Drop for LocalWorkdirSessionInner {
|
||||
}
|
||||
}
|
||||
|
||||
pub trait WorkdirSessionResource: Debug + Send + Sync {}
|
||||
impl<T> WorkdirSessionResource for T where T: Debug + Send + Sync {}
|
||||
|
||||
/// Scope-aware filesystem handle. Clone-cheap (`Arc` inside).
|
||||
///
|
||||
/// The wrapped [`SharedScope`] is shared with every clone of this
|
||||
@@ -318,6 +324,26 @@ impl LocalWorkdirSession {
|
||||
cwd: PathBuf,
|
||||
scope: SharedScope,
|
||||
capabilities: WorkdirSessionCapabilities,
|
||||
) -> Self {
|
||||
Self::materialized_bound_with_environment(
|
||||
workdir,
|
||||
root,
|
||||
cwd,
|
||||
scope,
|
||||
capabilities,
|
||||
BTreeMap::new(),
|
||||
Vec::new(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn materialized_bound_with_environment(
|
||||
workdir: Workdir,
|
||||
root: PathBuf,
|
||||
cwd: PathBuf,
|
||||
scope: SharedScope,
|
||||
capabilities: WorkdirSessionCapabilities,
|
||||
command_environment: BTreeMap<String, String>,
|
||||
resources: Vec<Arc<dyn WorkdirSessionResource>>,
|
||||
) -> Self {
|
||||
Self {
|
||||
inner: Arc::new(LocalWorkdirSessionInner {
|
||||
@@ -331,6 +357,8 @@ impl LocalWorkdirSession {
|
||||
next_command_id: AtomicU64::new(1),
|
||||
commands: Mutex::new(HashMap::new()),
|
||||
command_telemetry: CommandTelemetry::new(),
|
||||
command_environment,
|
||||
resources: StdMutex::new(resources),
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -669,9 +697,18 @@ impl WorkdirSession for LocalWorkdirSession {
|
||||
let (completion_tx, completion) = watch::channel(false);
|
||||
let command_id = handle.0.clone();
|
||||
let telemetry = self.inner.command_telemetry.clone();
|
||||
let command_environment = self.inner.command_environment.clone();
|
||||
let (cancel, cancel_rx) = watch::channel(false);
|
||||
let task = tokio::spawn(async move {
|
||||
let output = run_command(cwd, request, command_id, telemetry, cancel_rx).await;
|
||||
let output = run_command(
|
||||
cwd,
|
||||
request,
|
||||
command_id,
|
||||
telemetry,
|
||||
command_environment,
|
||||
cancel_rx,
|
||||
)
|
||||
.await;
|
||||
let _ = completion_tx.send(true);
|
||||
output
|
||||
});
|
||||
@@ -840,6 +877,9 @@ impl WorkdirSession for LocalWorkdirSession {
|
||||
LocalCommand::Completed(_) => {}
|
||||
}
|
||||
}
|
||||
if let Ok(mut resources) = self.inner.resources.lock() {
|
||||
resources.clear();
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -909,6 +949,7 @@ async fn run_command(
|
||||
request: CommandRequest,
|
||||
command_id: String,
|
||||
telemetry: CommandTelemetry,
|
||||
command_environment: BTreeMap<String, String>,
|
||||
mut cancel: watch::Receiver<bool>,
|
||||
) -> Result<CommandOutput, WorkdirError> {
|
||||
let stdout = tempfile::NamedTempFile::new().map_err(|error| WorkdirError::io(&cwd, error))?;
|
||||
@@ -925,6 +966,7 @@ async fn run_command(
|
||||
.arg("-c")
|
||||
.arg(&request.command)
|
||||
.current_dir(&cwd)
|
||||
.envs(command_environment)
|
||||
.stdin(Stdio::null())
|
||||
.stdout(Stdio::from(stdout_file))
|
||||
.stderr(Stdio::from(stderr_file))
|
||||
@@ -1901,9 +1943,9 @@ mod tests {
|
||||
&workdir,
|
||||
GrepRequest {
|
||||
pattern: "NEEDLE".into(),
|
||||
path: WorkdirPath::root(),
|
||||
glob: None,
|
||||
file_type: None,
|
||||
path: WorkdirPath::new("src/main.rs").unwrap(),
|
||||
glob: Some("src/*.rs".into()),
|
||||
file_type: Some("rust".into()),
|
||||
case_insensitive: false,
|
||||
before_context: 0,
|
||||
after_context: 0,
|
||||
@@ -2319,6 +2361,31 @@ mod tests {
|
||||
assert_eq!(terminal, Some((handle.0, CommandStatus::TimedOut, None)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn closing_session_releases_runtime_resources() {
|
||||
#[derive(Debug)]
|
||||
struct Resource(Arc<AtomicBool>);
|
||||
impl Drop for Resource {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::Release);
|
||||
}
|
||||
}
|
||||
|
||||
let dir = TempDir::new().unwrap();
|
||||
let released = Arc::new(AtomicBool::new(false));
|
||||
let session = LocalWorkdirSession::materialized_bound_with_environment(
|
||||
Workdir::new("resource-session"),
|
||||
dir.path().to_path_buf(),
|
||||
dir.path().to_path_buf(),
|
||||
SharedScope::new(Scope::writable(dir.path()).unwrap()),
|
||||
WorkdirSessionCapabilities::ALL,
|
||||
BTreeMap::from([("SSH_AUTH_SOCK".to_string(), "test-socket".to_string())]),
|
||||
vec![Arc::new(Resource(released.clone()))],
|
||||
);
|
||||
WorkdirSession::close(&session).await.unwrap();
|
||||
assert!(released.load(Ordering::Acquire));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_cancels_active_command() {
|
||||
let dir = TempDir::new().unwrap();
|
||||
|
||||
@@ -30,6 +30,8 @@ impl RuntimeWorkerRef {
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum MaterializerKind {
|
||||
#[default]
|
||||
RuntimeGitCache,
|
||||
/// Legacy persisted value from the pre-cache local `git worktree` materializer.
|
||||
LocalGitWorktree,
|
||||
}
|
||||
|
||||
@@ -109,6 +111,8 @@ pub struct WorkingDirectoryProvenance {
|
||||
pub creation_selector: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub creation_ref: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub creation_tree: Option<String>,
|
||||
pub materializer_kind: MaterializerKind,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cleanup_target: Option<WorkingDirectoryCleanupTarget>,
|
||||
@@ -122,6 +126,10 @@ pub struct WorkingDirectoryCurrentObservation {
|
||||
pub current_selector: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_ref: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_tree: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub observed_at_epoch_seconds: Option<u64>,
|
||||
pub status: WorkingDirectoryStatusKind,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cleanliness: Option<String>,
|
||||
@@ -141,9 +149,15 @@ pub struct WorkingDirectorySummary {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub creation_ref: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub creation_tree: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_selector: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_ref: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub current_tree: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub observed_at_epoch_seconds: Option<u64>,
|
||||
pub materializer_kind: MaterializerKind,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cleanup_target: Option<WorkingDirectoryCleanupTarget>,
|
||||
@@ -166,6 +180,7 @@ impl WorkingDirectorySummary {
|
||||
WorkingDirectoryProvenance {
|
||||
creation_selector: self.creation_selector.clone(),
|
||||
creation_ref: self.creation_ref.clone(),
|
||||
creation_tree: self.creation_tree.clone(),
|
||||
materializer_kind: self.materializer_kind.clone(),
|
||||
cleanup_target: self.cleanup_target.clone(),
|
||||
}
|
||||
@@ -175,6 +190,8 @@ impl WorkingDirectorySummary {
|
||||
WorkingDirectoryCurrentObservation {
|
||||
current_selector: self.current_selector.clone(),
|
||||
current_ref: self.current_ref.clone(),
|
||||
current_tree: self.current_tree.clone(),
|
||||
observed_at_epoch_seconds: self.observed_at_epoch_seconds,
|
||||
status: self.status.clone(),
|
||||
cleanliness: self.cleanliness.clone(),
|
||||
primary_worker_id: self.primary_worker_id.clone(),
|
||||
@@ -211,6 +228,7 @@ pub struct WorkingDirectoryListResponse {
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct WorkingDirectoryDetailResponse {
|
||||
pub workspace_id: String,
|
||||
pub runtime_id: String,
|
||||
pub item: WorkingDirectorySummary,
|
||||
pub diagnostics: Vec<WorkingDirectoryDiagnostic>,
|
||||
}
|
||||
@@ -246,8 +264,11 @@ mod tests {
|
||||
repository_id: "repo".to_string(),
|
||||
creation_selector: Some("develop".to_string()),
|
||||
creation_ref: Some("abc123".to_string()),
|
||||
creation_tree: Some("tree123".to_string()),
|
||||
current_selector: Some("work/ticket".to_string()),
|
||||
current_ref: Some("def456".to_string()),
|
||||
current_tree: Some("tree456".to_string()),
|
||||
observed_at_epoch_seconds: Some(1_777_777_777),
|
||||
materializer_kind: MaterializerKind::LocalGitWorktree,
|
||||
cleanup_target: Some(WorkingDirectoryCleanupTarget {
|
||||
kind: "git_worktree".to_string(),
|
||||
@@ -268,8 +289,11 @@ mod tests {
|
||||
repository_id: "repo".to_string(),
|
||||
creation_selector: None,
|
||||
creation_ref: None,
|
||||
creation_tree: None,
|
||||
current_selector: None,
|
||||
current_ref: Some("987fed".to_string()),
|
||||
current_tree: None,
|
||||
observed_at_epoch_seconds: None,
|
||||
materializer_kind: MaterializerKind::LocalGitWorktree,
|
||||
cleanup_target: None,
|
||||
status: WorkingDirectoryStatusKind::Active,
|
||||
@@ -306,6 +330,7 @@ mod tests {
|
||||
|
||||
let detail = WorkingDirectoryDetailResponse {
|
||||
workspace_id: decoded.workspace_id.clone(),
|
||||
runtime_id: "arcadia".to_string(),
|
||||
item: decoded.items[0].clone(),
|
||||
diagnostics: decoded.diagnostics.clone(),
|
||||
};
|
||||
|
||||
@@ -41,9 +41,12 @@ tar.workspace = true
|
||||
thiserror = { workspace = true }
|
||||
tokio = { workspace = true, features = ["net", "rt", "sync", "time"] }
|
||||
toml.workspace = true
|
||||
url.workspace = true
|
||||
uuid = { workspace = true, features = ["v7"] }
|
||||
zeroize.workspace = true
|
||||
tower = { workspace = true, features = ["util"], optional = true }
|
||||
worker.workspace = true
|
||||
workspace-api = { path = "../workspace-api" }
|
||||
workdir.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
|
||||
@@ -3,6 +3,7 @@ use base64::engine::general_purpose::URL_SAFE_NO_PAD;
|
||||
use ring::rand::{SecureRandom, SystemRandom};
|
||||
use ring::signature::{ED25519, Ed25519KeyPair, KeyPair, UnparsedPublicKey};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::fmt;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
@@ -14,6 +15,12 @@ pub const WORKER_MUTATION_SOURCE_PROOF_HEADER: &str = "x-yoi-worker-mutation-pro
|
||||
const WORKER_MUTATION_SOURCE_PROOF_PREFIX: &str = "yoi-worker-source-v1";
|
||||
const WORKER_MUTATION_SOURCE_SIGNING_INPUT_PREFIX: &str = "yoi-worker-source-v1.";
|
||||
pub const WORKER_REMOVE_PERMISSION: &str = "workspace:worker-remove";
|
||||
pub const RUNTIME_REQUEST_SOURCE_PROOF_HEADER: &str = "x-yoi-runtime-request-proof";
|
||||
pub const WORKSPACE_REQUEST_PERMISSION: &str = "workspace:request";
|
||||
pub const WORKSPACE_WORKER_DISCOVERY_PERMISSION: &str = "workspace:worker-discovery";
|
||||
pub const BACKEND_RESOURCE_FETCH_PERMISSION: &str = "workspace:resource-fetch";
|
||||
const RUNTIME_REQUEST_SOURCE_PROOF_PREFIX: &str = "yoi-runtime-request-v1";
|
||||
const RUNTIME_REQUEST_SOURCE_SIGNING_INPUT_PREFIX: &str = "yoi-runtime-request-v1.";
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum RuntimeAuthError {
|
||||
@@ -33,6 +40,10 @@ pub enum RuntimeAuthError {
|
||||
InvalidTokenFormat,
|
||||
#[error("malformed capability token claims: {0}")]
|
||||
MalformedClaims(#[from] serde_json::Error),
|
||||
#[error("runtime request proof contains an invalid `{0}` claim")]
|
||||
InvalidClaim(&'static str),
|
||||
#[error("runtime request proof does not match the HTTP request")]
|
||||
ClaimMismatch,
|
||||
#[error("unknown token issuer `{0}`")]
|
||||
UnknownIssuer(String),
|
||||
#[error("invalid token signature")]
|
||||
@@ -224,6 +235,162 @@ pub fn verify_capability_token(
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RuntimeRequestSourceClaims {
|
||||
pub iss: String,
|
||||
pub aud: String,
|
||||
pub workspace_id: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub worker_id: Option<String>,
|
||||
pub permission: String,
|
||||
pub method: String,
|
||||
pub path: String,
|
||||
pub body_digest: String,
|
||||
pub iat: i64,
|
||||
pub exp: i64,
|
||||
pub jti: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct RuntimeRequestSourceSigner {
|
||||
identity_id: String,
|
||||
private_key: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct RuntimeRequestSourceExpectation<'a> {
|
||||
pub identity_id: &'a str,
|
||||
pub audience: &'a str,
|
||||
pub workspace_id: &'a str,
|
||||
pub worker_id: Option<&'a str>,
|
||||
pub permission: &'a str,
|
||||
pub method: &'a str,
|
||||
pub path: &'a str,
|
||||
pub body_digest: &'a str,
|
||||
pub now_unix: i64,
|
||||
}
|
||||
|
||||
pub fn request_body_digest(body: &[u8]) -> String {
|
||||
URL_SAFE_NO_PAD.encode(Sha256::digest(body))
|
||||
}
|
||||
|
||||
impl RuntimeRequestSourceSigner {
|
||||
pub fn from_identity(identity: &RuntimeIdentityMaterial) -> Self {
|
||||
Self {
|
||||
identity_id: identity.identity_id.clone(),
|
||||
private_key: identity.private_key.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn issue(
|
||||
&self,
|
||||
audience: &str,
|
||||
workspace_id: &str,
|
||||
worker_id: Option<&str>,
|
||||
permission: &str,
|
||||
method: &str,
|
||||
path: &str,
|
||||
body: &[u8],
|
||||
now_unix: i64,
|
||||
ttl_seconds: u64,
|
||||
) -> Result<String, RuntimeAuthError> {
|
||||
for (name, value) in [
|
||||
("audience", audience),
|
||||
("workspace_id", workspace_id),
|
||||
("permission", permission),
|
||||
("method", method),
|
||||
("path", path),
|
||||
] {
|
||||
if value.trim().is_empty() {
|
||||
return Err(RuntimeAuthError::InvalidClaim(name));
|
||||
}
|
||||
}
|
||||
if worker_id.is_some_and(str::is_empty) {
|
||||
return Err(RuntimeAuthError::InvalidClaim("worker_id"));
|
||||
}
|
||||
let ttl_seconds = i64::try_from(ttl_seconds).unwrap_or(i64::MAX);
|
||||
let claims = RuntimeRequestSourceClaims {
|
||||
iss: self.identity_id.clone(),
|
||||
aud: audience.to_owned(),
|
||||
workspace_id: workspace_id.to_owned(),
|
||||
worker_id: worker_id.map(str::to_owned),
|
||||
permission: permission.to_owned(),
|
||||
method: method.to_owned(),
|
||||
path: path.to_owned(),
|
||||
body_digest: request_body_digest(body),
|
||||
iat: now_unix,
|
||||
exp: now_unix.saturating_add(ttl_seconds),
|
||||
jti: new_token_id()?,
|
||||
};
|
||||
let payload = serde_json::to_vec(&claims)?;
|
||||
let payload = URL_SAFE_NO_PAD.encode(payload);
|
||||
let signing_input = format!("{RUNTIME_REQUEST_SOURCE_SIGNING_INPUT_PREFIX}{payload}");
|
||||
let private = decode_private_key(&self.private_key)?;
|
||||
let key_pair = Ed25519KeyPair::from_pkcs8(&private)
|
||||
.map_err(|_| RuntimeAuthError::InvalidPrivateKey)?;
|
||||
let signature = URL_SAFE_NO_PAD.encode(key_pair.sign(signing_input.as_bytes()).as_ref());
|
||||
Ok(format!(
|
||||
"{RUNTIME_REQUEST_SOURCE_PROOF_PREFIX}.{payload}.{signature}"
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decode_runtime_request_source_claims(
|
||||
proof: &str,
|
||||
) -> Result<RuntimeRequestSourceClaims, RuntimeAuthError> {
|
||||
let (prefix, payload, _signature) = split_runtime_request_source_proof(proof)?;
|
||||
if prefix != RUNTIME_REQUEST_SOURCE_PROOF_PREFIX {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
let payload = URL_SAFE_NO_PAD.decode(payload)?;
|
||||
serde_json::from_slice(&payload).map_err(RuntimeAuthError::from)
|
||||
}
|
||||
|
||||
pub fn verify_runtime_request_source(
|
||||
proof: &str,
|
||||
public_key: &str,
|
||||
expected: &RuntimeRequestSourceExpectation<'_>,
|
||||
) -> Result<RuntimeRequestSourceClaims, RuntimeAuthError> {
|
||||
let (prefix, payload, signature) = split_runtime_request_source_proof(proof)?;
|
||||
if prefix != RUNTIME_REQUEST_SOURCE_PROOF_PREFIX {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
let signature = URL_SAFE_NO_PAD.decode(signature)?;
|
||||
let signing_input = format!("{RUNTIME_REQUEST_SOURCE_SIGNING_INPUT_PREFIX}{payload}");
|
||||
let public_key = decode_public_key(public_key)?;
|
||||
UnparsedPublicKey::new(&ED25519, public_key)
|
||||
.verify(signing_input.as_bytes(), &signature)
|
||||
.map_err(|_| RuntimeAuthError::InvalidSignature)?;
|
||||
let claims = decode_runtime_request_source_claims(proof)?;
|
||||
if claims.iss != expected.identity_id
|
||||
|| claims.aud != expected.audience
|
||||
|| claims.workspace_id != expected.workspace_id
|
||||
|| claims.worker_id.as_deref() != expected.worker_id
|
||||
|| claims.permission != expected.permission
|
||||
|| claims.method != expected.method
|
||||
|| claims.path != expected.path
|
||||
|| claims.body_digest != expected.body_digest
|
||||
{
|
||||
return Err(RuntimeAuthError::ClaimMismatch);
|
||||
}
|
||||
if claims.iat > expected.now_unix || claims.exp < expected.now_unix {
|
||||
return Err(RuntimeAuthError::Expired);
|
||||
}
|
||||
Ok(claims)
|
||||
}
|
||||
|
||||
fn split_runtime_request_source_proof(proof: &str) -> Result<(&str, &str, &str), RuntimeAuthError> {
|
||||
let mut parts = proof.split('.');
|
||||
let prefix = parts.next().unwrap_or_default();
|
||||
let payload = parts.next().unwrap_or_default();
|
||||
let signature = parts.next().unwrap_or_default();
|
||||
if prefix.is_empty() || payload.is_empty() || signature.is_empty() || parts.next().is_some() {
|
||||
return Err(RuntimeAuthError::InvalidTokenFormat);
|
||||
}
|
||||
Ok((prefix, payload, signature))
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkerMutationSourceClaims {
|
||||
pub iss: String,
|
||||
@@ -592,6 +759,99 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_request_source_proof_binds_request_and_rejects_spoofed_signature() {
|
||||
let trusted = RuntimeIdentityMaterial::generate("runtime-main").unwrap();
|
||||
let signer = RuntimeRequestSourceSigner::from_identity(&trusted);
|
||||
let body = br#"{"ticket":"T-1"}"#;
|
||||
let proof = signer
|
||||
.issue(
|
||||
"server-main",
|
||||
"workspace-a",
|
||||
Some("worker-7"),
|
||||
WORKSPACE_REQUEST_PERMISSION,
|
||||
"POST",
|
||||
"/api/w/workspace-a/tickets/comment",
|
||||
body,
|
||||
90,
|
||||
10,
|
||||
)
|
||||
.unwrap();
|
||||
let expected = RuntimeRequestSourceExpectation {
|
||||
identity_id: "runtime-main",
|
||||
audience: "server-main",
|
||||
workspace_id: "workspace-a",
|
||||
worker_id: Some("worker-7"),
|
||||
permission: WORKSPACE_REQUEST_PERMISSION,
|
||||
method: "POST",
|
||||
path: "/api/w/workspace-a/tickets/comment",
|
||||
body_digest: &request_body_digest(body),
|
||||
now_unix: 99,
|
||||
};
|
||||
let claims = verify_runtime_request_source(&proof, &trusted.public_key, &expected).unwrap();
|
||||
assert_eq!(claims.iss, "runtime-main");
|
||||
let changed_body = RuntimeRequestSourceExpectation {
|
||||
body_digest: &request_body_digest(br#"{"ticket":"T-2"}"#),
|
||||
..expected.clone()
|
||||
};
|
||||
assert!(matches!(
|
||||
verify_runtime_request_source(&proof, &trusted.public_key, &changed_body),
|
||||
Err(RuntimeAuthError::ClaimMismatch)
|
||||
));
|
||||
let spoofed = RuntimeIdentityMaterial::generate("runtime-main").unwrap();
|
||||
assert!(matches!(
|
||||
verify_runtime_request_source(&proof, &spoofed.public_key, &expected),
|
||||
Err(RuntimeAuthError::InvalidSignature)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn runtime_request_source_proof_rejects_wrong_scope_and_expiry() {
|
||||
let runtime = RuntimeIdentityMaterial::generate("runtime-main").unwrap();
|
||||
let proof = RuntimeRequestSourceSigner::from_identity(&runtime)
|
||||
.issue(
|
||||
"server-main",
|
||||
"workspace-a",
|
||||
None,
|
||||
BACKEND_RESOURCE_FETCH_PERMISSION,
|
||||
"POST",
|
||||
"/api/runtime/v1/workspaces/workspace-a/resources/fetch",
|
||||
b"{}",
|
||||
90,
|
||||
10,
|
||||
)
|
||||
.unwrap();
|
||||
let digest = request_body_digest(b"{}");
|
||||
let expected = RuntimeRequestSourceExpectation {
|
||||
identity_id: "runtime-main",
|
||||
audience: "server-main",
|
||||
workspace_id: "workspace-a",
|
||||
worker_id: None,
|
||||
permission: BACKEND_RESOURCE_FETCH_PERMISSION,
|
||||
method: "POST",
|
||||
path: "/api/runtime/v1/workspaces/workspace-a/resources/fetch",
|
||||
body_digest: &digest,
|
||||
now_unix: 99,
|
||||
};
|
||||
assert!(verify_runtime_request_source(&proof, &runtime.public_key, &expected).is_ok());
|
||||
let wrong_workspace = RuntimeRequestSourceExpectation {
|
||||
workspace_id: "workspace-b",
|
||||
..expected.clone()
|
||||
};
|
||||
assert!(matches!(
|
||||
verify_runtime_request_source(&proof, &runtime.public_key, &wrong_workspace),
|
||||
Err(RuntimeAuthError::ClaimMismatch)
|
||||
));
|
||||
let expired = RuntimeRequestSourceExpectation {
|
||||
now_unix: 101,
|
||||
..expected
|
||||
};
|
||||
assert!(matches!(
|
||||
verify_runtime_request_source(&proof, &runtime.public_key, &expired),
|
||||
Err(RuntimeAuthError::Expired)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capability_token_verifies_signature_audience_expiry_and_permission() {
|
||||
let server = RuntimeIdentityMaterial::generate("server-main").unwrap();
|
||||
|
||||
@@ -2,7 +2,6 @@ use crate::identity::{RuntimeWorkerRef, WorkerId, WorkerRef};
|
||||
use crate::interaction::WorkerInput;
|
||||
use crate::profile_archive::{ProfileSourceArchive, ProfileSourceArchiveRef};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::path::PathBuf;
|
||||
|
||||
fn is_false(value: &bool) -> bool {
|
||||
!*value
|
||||
@@ -85,9 +84,9 @@ impl std::ops::Deref for RepositorySelector {
|
||||
pub struct WorkingDirectoryRepository {
|
||||
pub id: String,
|
||||
pub provider: String,
|
||||
pub uri: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub local_path: Option<PathBuf>,
|
||||
pub source: workspace_api::RepositorySource,
|
||||
pub source_revision: u64,
|
||||
pub source_fingerprint: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub selector: Option<RepositorySelector>,
|
||||
}
|
||||
@@ -98,6 +97,74 @@ pub use workdir::workspace::{
|
||||
WorkingDirectorySummary,
|
||||
};
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct SensitiveString(String);
|
||||
|
||||
impl SensitiveString {
|
||||
pub fn new(value: impl Into<String>) -> Self {
|
||||
Self(value.into())
|
||||
}
|
||||
|
||||
pub fn expose(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for SensitiveString {
|
||||
fn drop(&mut self) {
|
||||
zeroize::Zeroize::zeroize(&mut self.0);
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for SensitiveString {
|
||||
fn default() -> Self {
|
||||
Self(String::new())
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for SensitiveString {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("[REDACTED]")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RepositorySshMaterializationAccess {
|
||||
pub credential_id: String,
|
||||
pub credential_revision: u64,
|
||||
pub host_trust_id: String,
|
||||
pub host_trust_revision: u64,
|
||||
pub access: workspace_api::RepositoryAccessMode,
|
||||
pub expires_at_epoch_seconds: u64,
|
||||
pub repository_id: String,
|
||||
pub repository_source_fingerprint: String,
|
||||
pub repository_uri: String,
|
||||
pub secret_resource: crate::resource::BackendResourceHandle,
|
||||
#[serde(skip, default)]
|
||||
pub private_key: SensitiveString,
|
||||
#[serde(skip, default)]
|
||||
pub known_hosts_entry: SensitiveString,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RepositoryMaterializationContext {
|
||||
pub workspace_id: String,
|
||||
pub runtime_id: String,
|
||||
pub operation_id: String,
|
||||
pub config_revision: u64,
|
||||
pub config_projection_digest: String,
|
||||
#[serde(default)]
|
||||
pub cache_generation: u64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ssh: Option<RepositorySshMaterializationAccess>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkingDirectoryRepositoryAccessRequest {
|
||||
pub working_directory_id: String,
|
||||
pub materialization: RepositoryMaterializationContext,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct WorkingDirectoryRequest {
|
||||
pub repository: WorkingDirectoryRepository,
|
||||
@@ -107,6 +174,9 @@ pub struct WorkingDirectoryRequest {
|
||||
/// Backend can create canonical registry rows before materialization.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub backend_workdir_id: Option<String>,
|
||||
/// Backend-authored, operation-scoped repository access and cache identity.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub materialization: Option<RepositoryMaterializationContext>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
use crate::catalog::{WorkingDirectoryRequest, WorkingDirectoryStatus};
|
||||
use crate::catalog::{
|
||||
WorkingDirectoryRepositoryAccessRequest, WorkingDirectoryRequest, WorkingDirectoryStatus,
|
||||
};
|
||||
use crate::config_bundle::ConfigBundle;
|
||||
use crate::error::RuntimeError;
|
||||
use crate::identity::WorkerRef;
|
||||
@@ -319,6 +321,16 @@ pub trait WorkerExecutionBackend: Send + Sync + 'static {
|
||||
))
|
||||
}
|
||||
|
||||
fn authorize_working_directory_repository_access(
|
||||
&self,
|
||||
_request: &WorkingDirectoryRepositoryAccessRequest,
|
||||
) -> Result<(), WorkingDirectoryDiagnostic> {
|
||||
Err(WorkingDirectoryDiagnostic::rejected(
|
||||
"working_directory_repository_access_unsupported",
|
||||
"Worker execution backend does not support Repository access authorization",
|
||||
))
|
||||
}
|
||||
|
||||
fn list_working_directories(&self) -> Vec<WorkingDirectoryStatus> {
|
||||
Vec::new()
|
||||
}
|
||||
@@ -454,6 +466,14 @@ impl WorkerExecutionBackendRef {
|
||||
self.backend.create_working_directory(request)
|
||||
}
|
||||
|
||||
pub(crate) fn authorize_working_directory_repository_access(
|
||||
&self,
|
||||
request: &WorkingDirectoryRepositoryAccessRequest,
|
||||
) -> Result<(), WorkingDirectoryDiagnostic> {
|
||||
self.backend
|
||||
.authorize_working_directory_repository_access(request)
|
||||
}
|
||||
|
||||
pub(crate) fn list_working_directories(&self) -> Vec<WorkingDirectoryStatus> {
|
||||
self.backend.list_working_directories()
|
||||
}
|
||||
|
||||
@@ -12,7 +12,8 @@ use crate::auth::{
|
||||
};
|
||||
use crate::catalog::{
|
||||
ConfigBundleRef, CreateWorkerRequest, WorkerDetail, WorkerLifecycleAck, WorkerSummary,
|
||||
WorkingDirectoryRequest, WorkingDirectoryStatus, WorkspaceApiRef,
|
||||
WorkingDirectoryRepositoryAccessRequest, WorkingDirectoryRequest, WorkingDirectoryStatus,
|
||||
WorkspaceApiRef,
|
||||
};
|
||||
use crate::config_bundle::{ConfigBundle, ConfigBundleAvailability, ConfigBundleSummary};
|
||||
use crate::error::RuntimeError;
|
||||
@@ -203,6 +204,10 @@ fn runtime_http_router_with_optional_auth(
|
||||
"/v1/working-directories",
|
||||
get(list_working_directories).post(create_working_directory),
|
||||
)
|
||||
.route(
|
||||
"/v1/working-directories/repository-access",
|
||||
post(authorize_working_directory_repository_access),
|
||||
)
|
||||
.route(
|
||||
"/v1/working-directories/{working_directory_id}/sessions",
|
||||
post(open_workdir_session),
|
||||
@@ -335,6 +340,11 @@ pub struct RuntimeHttpWorkingDirectoriesResponse {
|
||||
pub working_directories: Vec<WorkingDirectoryStatus>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RuntimeHttpRepositoryAccessResponse {
|
||||
pub authorized: bool,
|
||||
}
|
||||
|
||||
/// Working directory response used by create/detail/delete endpoints.
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RuntimeHttpWorkingDirectoryResponse {
|
||||
@@ -513,6 +523,29 @@ async fn list_workers(
|
||||
Ok(Json(RuntimeHttpWorkersResponse { workers }))
|
||||
}
|
||||
|
||||
async fn authorize_working_directory_repository_access(
|
||||
State(state): State<RuntimeHttpState>,
|
||||
Extension(auth): Extension<RuntimeAuthContext>,
|
||||
body: Result<Json<WorkingDirectoryRepositoryAccessRequest>, JsonRejection>,
|
||||
) -> RestResult<RuntimeHttpRepositoryAccessResponse> {
|
||||
let Json(request) = body.map_err(RuntimeHttpRestError::json_rejection)?;
|
||||
if request.materialization.workspace_id != auth.workspace_id {
|
||||
return Err(RuntimeHttpRestError::new(
|
||||
StatusCode::FORBIDDEN,
|
||||
"working_directory_materialization_workspace_mismatch",
|
||||
"Repository access authority does not match the authenticated Workspace",
|
||||
));
|
||||
}
|
||||
state
|
||||
.runtime
|
||||
.authorize_working_directory_repository_access_from_resource(request)
|
||||
.await
|
||||
.map_err(RuntimeHttpRestError::runtime)?;
|
||||
Ok(Json(RuntimeHttpRepositoryAccessResponse {
|
||||
authorized: true,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn list_working_directories(
|
||||
State(state): State<RuntimeHttpState>,
|
||||
) -> RestResult<RuntimeHttpWorkingDirectoriesResponse> {
|
||||
@@ -527,12 +560,23 @@ async fn list_working_directories(
|
||||
|
||||
async fn create_working_directory(
|
||||
State(state): State<RuntimeHttpState>,
|
||||
Extension(auth): Extension<RuntimeAuthContext>,
|
||||
body: Result<Json<WorkingDirectoryRequest>, JsonRejection>,
|
||||
) -> RestResult<RuntimeHttpWorkingDirectoryResponse> {
|
||||
let Json(request) = body.map_err(RuntimeHttpRestError::json_rejection)?;
|
||||
if let Some(materialization) = request.materialization.as_ref()
|
||||
&& materialization.workspace_id != auth.workspace_id
|
||||
{
|
||||
return Err(RuntimeHttpRestError::new(
|
||||
StatusCode::FORBIDDEN,
|
||||
"working_directory_materialization_workspace_mismatch",
|
||||
"Repository materialization authority does not match the authenticated Workspace",
|
||||
));
|
||||
}
|
||||
let working_directory = state
|
||||
.runtime
|
||||
.create_working_directory(request)
|
||||
.create_working_directory_from_resource(request)
|
||||
.await
|
||||
.map_err(RuntimeHttpRestError::runtime)?;
|
||||
Ok(Json(RuntimeHttpWorkingDirectoryResponse {
|
||||
working_directory,
|
||||
@@ -1559,6 +1603,9 @@ fn required_runtime_permission(method: &Method, path: &str) -> Option<&'static s
|
||||
if path == "/v1/workers" && *method == Method::POST {
|
||||
return Some("workers:create");
|
||||
}
|
||||
if path == "/v1/working-directories/repository-access" && *method == Method::POST {
|
||||
return Some("workdirs:operate");
|
||||
}
|
||||
if path.starts_with("/v1/workdir-sessions")
|
||||
|| (path.starts_with("/v1/working-directories/") && path.ends_with("/sessions"))
|
||||
{
|
||||
@@ -1674,17 +1721,8 @@ impl RuntimeHttpWorkdirError {
|
||||
impl From<workdir::WorkdirError> for RuntimeHttpWorkdirError {
|
||||
fn from(error: workdir::WorkdirError) -> Self {
|
||||
let payload = WorkdirTransportError::from_workdir_error(&error);
|
||||
let status = match payload.code {
|
||||
WorkdirTransportErrorCode::NotFound | WorkdirTransportErrorCode::UnknownCommand => {
|
||||
StatusCode::NOT_FOUND
|
||||
}
|
||||
WorkdirTransportErrorCode::Conflict => StatusCode::CONFLICT,
|
||||
WorkdirTransportErrorCode::Unsupported | WorkdirTransportErrorCode::InvalidRequest => {
|
||||
StatusCode::BAD_REQUEST
|
||||
}
|
||||
WorkdirTransportErrorCode::Unavailable => StatusCode::SERVICE_UNAVAILABLE,
|
||||
WorkdirTransportErrorCode::Internal => StatusCode::INTERNAL_SERVER_ERROR,
|
||||
};
|
||||
let status = StatusCode::from_u16(payload.code.http_status())
|
||||
.expect("Workdir transport error status is valid");
|
||||
Self { status, payload }
|
||||
}
|
||||
}
|
||||
@@ -1839,8 +1877,8 @@ mod tests {
|
||||
use manifest::{Scope, SharedScope};
|
||||
use tower::ServiceExt;
|
||||
use workdir::{
|
||||
LocalWorkdirSession, ReadRequest, StatRequest, Workdir, WorkdirPath,
|
||||
WorkdirSessionCapabilities,
|
||||
GrepOutputMode, GrepRequest, LocalWorkdirSession, ReadRequest, StatRequest, Workdir,
|
||||
WorkdirPath, WorkdirSessionCapabilities,
|
||||
};
|
||||
|
||||
fn test_bundle(profile: ProfileSelector) -> ConfigBundle {
|
||||
@@ -2220,6 +2258,10 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn workdir_routes_require_dedicated_operation_permission() {
|
||||
assert_eq!(
|
||||
required_runtime_permission(&Method::POST, "/v1/working-directories/repository-access",),
|
||||
Some("workdirs:operate")
|
||||
);
|
||||
assert_eq!(
|
||||
required_runtime_permission(&Method::POST, "/v1/working-directories/wd-1/sessions"),
|
||||
Some("workdirs:operate")
|
||||
@@ -2297,6 +2339,40 @@ mod tests {
|
||||
.expect("owned operation");
|
||||
assert!(matches!(result, WorkdirSessionOperationResult::Stat(_)));
|
||||
|
||||
let grep = WorkdirSessionOperationRequest {
|
||||
delegations: Vec::new(),
|
||||
operation: WorkdirSessionOperation::Grep(GrepRequest {
|
||||
pattern: "hello".into(),
|
||||
path: WorkdirPath::new("hello.txt").unwrap(),
|
||||
glob: Some("*.txt".into()),
|
||||
file_type: Some("txt".into()),
|
||||
case_insensitive: false,
|
||||
before_context: 0,
|
||||
after_context: 0,
|
||||
multiline: false,
|
||||
output_mode: GrepOutputMode::Content,
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
}),
|
||||
};
|
||||
let Json(result) = run_workdir_session_operation(
|
||||
State(state.clone()),
|
||||
Path("session-1".to_string()),
|
||||
Some(Extension(auth.clone())),
|
||||
Ok(Json(grep)),
|
||||
)
|
||||
.await
|
||||
.expect("grep direct file through provider operation");
|
||||
match result {
|
||||
WorkdirSessionOperationResult::Grep(result) => {
|
||||
assert_eq!(result.match_count, 1);
|
||||
assert_eq!(result.matched_files, 1);
|
||||
assert!(result.output.starts_with("hello.txt\n"));
|
||||
assert!(result.output.contains("> 1 │ hello"));
|
||||
}
|
||||
other => panic!("unexpected workdir grep result: {other:?}"),
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
{
|
||||
let delegated_visible = WorkdirSessionOperationRequest {
|
||||
|
||||
@@ -23,10 +23,21 @@ use worker_runtime::http_server::{
|
||||
RuntimeHttpServerConfig, RuntimeHttpServerError, RuntimeHttpStoreSelection,
|
||||
};
|
||||
use worker_runtime::worker_backend::{ProfileRuntimeWorkerFactory, WorkerRuntimeExecutionBackend};
|
||||
use worker_runtime::working_directory::LocalGitWorktreeMaterializer;
|
||||
use worker_runtime::working_directory::RuntimeGitCacheMaterializer;
|
||||
use worker_runtime::{Runtime, RuntimeOptions};
|
||||
|
||||
fn main() -> ExitCode {
|
||||
let mut arguments = std::env::args().skip(1).collect::<Vec<_>>();
|
||||
if arguments.first().map(String::as_str) == Some("__repository-ssh") {
|
||||
arguments.remove(0);
|
||||
return match worker_runtime::working_directory::run_repository_ssh_client(&arguments) {
|
||||
Ok(status) => ExitCode::from(u8::try_from(status).unwrap_or(1)),
|
||||
Err(error) => {
|
||||
eprintln!("{error}");
|
||||
ExitCode::from(1)
|
||||
}
|
||||
};
|
||||
}
|
||||
match run() {
|
||||
Ok(()) => ExitCode::SUCCESS,
|
||||
Err(error) => {
|
||||
@@ -160,29 +171,52 @@ fn build_runtime(config: &ProcessConfig) -> Result<Runtime, ProcessError> {
|
||||
};
|
||||
let mut factory = ProfileRuntimeWorkerFactory::new(fs_paths.worker_dir.join("worker-root"))
|
||||
.with_runtime_store_dir(runtime_store_dir);
|
||||
if let Some(identity) = read_runtime_auth_file(&runtime_auth_path(config))?.identity {
|
||||
factory = factory.with_remote_worker_mutation_identity(identity);
|
||||
let runtime_auth = read_runtime_auth_file(&runtime_auth_path(config))?;
|
||||
if let Some(identity) = runtime_auth.identity.clone() {
|
||||
if let [trusted_server] = runtime_auth.trusted_servers.as_slice() {
|
||||
factory =
|
||||
factory.with_runtime_request_identity(identity, trusted_server.server_id.clone());
|
||||
} else {
|
||||
factory = factory.with_remote_worker_mutation_identity(identity);
|
||||
}
|
||||
}
|
||||
let mut backend_resource_client: Option<
|
||||
Arc<dyn worker_runtime::resource::BackendResourceClient>,
|
||||
> = None;
|
||||
if let Some(endpoint) = config.backend_resource_endpoint.clone() {
|
||||
factory = factory.with_resource_client(Arc::new(
|
||||
let identity = runtime_auth.identity.as_ref().ok_or_else(|| {
|
||||
ProcessError::Auth(
|
||||
"--backend-resource-endpoint requires a configured Runtime identity".to_owned(),
|
||||
)
|
||||
})?;
|
||||
let [trusted_server] = runtime_auth.trusted_servers.as_slice() else {
|
||||
return Err(ProcessError::Auth(
|
||||
"--backend-resource-endpoint requires exactly one trusted Server identity"
|
||||
.to_owned(),
|
||||
));
|
||||
};
|
||||
let client = Arc::new(
|
||||
worker_runtime::resource::HttpBackendResourceClient::new(
|
||||
endpoint,
|
||||
config.backend_resource_token.clone(),
|
||||
),
|
||||
));
|
||||
)
|
||||
.with_runtime_request_source(identity, trusted_server.server_id.clone()),
|
||||
);
|
||||
factory = factory.with_resource_client(client.clone());
|
||||
backend_resource_client = Some(client);
|
||||
}
|
||||
let backend = Arc::new(
|
||||
WorkerRuntimeExecutionBackend::new(factory)
|
||||
.map_err(ProcessError::WorkerAdapter)?
|
||||
.with_working_directory_materializer(LocalGitWorktreeMaterializer::new(
|
||||
.with_working_directory_materializer(RuntimeGitCacheMaterializer::new(
|
||||
fs_paths.workdir_target.clone(),
|
||||
)),
|
||||
);
|
||||
|
||||
match &config.http.store {
|
||||
let runtime = match &config.http.store {
|
||||
RuntimeHttpStoreSelection::Memory => {
|
||||
Runtime::with_execution_backend(runtime_options_from_http(&config.http), backend)
|
||||
.map_err(ProcessError::Runtime)
|
||||
.map_err(ProcessError::Runtime)?
|
||||
}
|
||||
RuntimeHttpStoreSelection::Fs { root } => {
|
||||
let mut options = FsRuntimeStoreOptions::new(root.clone()).with_runtime_id(
|
||||
@@ -195,12 +229,20 @@ fn build_runtime(config: &ProcessConfig) -> Result<Runtime, ProcessError> {
|
||||
);
|
||||
options.display_name = config.http.display_name.clone();
|
||||
Runtime::with_fs_store_and_execution_backend(options, backend)
|
||||
.map_err(ProcessError::Runtime)
|
||||
.map_err(ProcessError::Runtime)?
|
||||
}
|
||||
_ => Err(ProcessError::usage(
|
||||
"unsupported Runtime catalog store selection".to_string(),
|
||||
)),
|
||||
_ => {
|
||||
return Err(ProcessError::usage(
|
||||
"unsupported Runtime catalog store selection".to_string(),
|
||||
));
|
||||
}
|
||||
};
|
||||
if let Some(client) = backend_resource_client {
|
||||
runtime
|
||||
.install_backend_resource_client(client)
|
||||
.map_err(ProcessError::Runtime)?;
|
||||
}
|
||||
Ok(runtime)
|
||||
}
|
||||
|
||||
fn runtime_options_from_http(config: &RuntimeHttpServerConfig) -> RuntimeOptions {
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
use crate::auth::{
|
||||
BACKEND_RESOURCE_FETCH_PERMISSION, RUNTIME_REQUEST_SOURCE_PROOF_HEADER,
|
||||
RuntimeIdentityMaterial, RuntimeRequestSourceSigner, unix_now_seconds,
|
||||
};
|
||||
use crate::identity::WorkerId;
|
||||
use crate::profile_archive::{ProfileSourceArchive, ProfileSourceArchiveRef, sha256_hex};
|
||||
use async_trait::async_trait;
|
||||
@@ -7,18 +11,46 @@ use std::sync::Mutex;
|
||||
|
||||
pub const PROFILE_SOURCE_ARCHIVE_CONTENT_TYPE: &str =
|
||||
"application/vnd.yoi.profile-source-archive+tar";
|
||||
pub const REPOSITORY_SSH_ACCESS_CONTENT_TYPE: &str =
|
||||
"application/vnd.yoi.repository-ssh-access+json";
|
||||
pub const DEFAULT_PROFILE_SOURCE_ARCHIVE_MAX_BYTES: u64 = 2 * 1024 * 1024;
|
||||
pub const DEFAULT_REPOSITORY_SSH_ACCESS_MAX_BYTES: u64 = 64 * 1024;
|
||||
|
||||
#[derive(Clone, Serialize, Deserialize)]
|
||||
pub struct RepositorySshAccessSecret {
|
||||
pub private_key: String,
|
||||
pub known_hosts_entry: String,
|
||||
}
|
||||
|
||||
impl Drop for RepositorySshAccessSecret {
|
||||
fn drop(&mut self) {
|
||||
zeroize::Zeroize::zeroize(&mut self.private_key);
|
||||
zeroize::Zeroize::zeroize(&mut self.known_hosts_entry);
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for RepositorySshAccessSecret {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("RepositorySshAccessSecret")
|
||||
.field("private_key", &"[REDACTED]")
|
||||
.field("known_hosts_entry", &"[REDACTED]")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum BackendResourceKind {
|
||||
ProfileSourceArchive,
|
||||
RepositorySshAccess,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum BackendResourceOperation {
|
||||
FetchArchive,
|
||||
FetchOnce,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
@@ -62,7 +94,7 @@ pub struct BackendResourceFetchRequest {
|
||||
pub audit_correlation_id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct BackendResourceFetchResponse {
|
||||
pub kind: BackendResourceKind,
|
||||
pub resource_id: String,
|
||||
@@ -72,6 +104,29 @@ pub struct BackendResourceFetchResponse {
|
||||
pub audit_correlation_id: String,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for BackendResourceFetchResponse {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("BackendResourceFetchResponse")
|
||||
.field("kind", &self.kind)
|
||||
.field("resource_id", &self.resource_id)
|
||||
.field("digest", &self.digest)
|
||||
.field("content_type", &self.content_type)
|
||||
.field(
|
||||
"bytes",
|
||||
&format_args!("[REDACTED; {} bytes]", self.bytes.len()),
|
||||
)
|
||||
.field("audit_correlation_id", &self.audit_correlation_id)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for BackendResourceFetchResponse {
|
||||
fn drop(&mut self) {
|
||||
self.bytes.fill(0);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, thiserror::Error)]
|
||||
#[serde(tag = "code", rename_all = "snake_case")]
|
||||
pub enum BackendResourceError {
|
||||
@@ -108,6 +163,8 @@ pub trait BackendResourceClient: Send + Sync + 'static {
|
||||
pub struct HttpBackendResourceClient {
|
||||
endpoint: String,
|
||||
bearer_token: Option<String>,
|
||||
request_source_signer: Option<RuntimeRequestSourceSigner>,
|
||||
request_source_audience: Option<String>,
|
||||
client: reqwest::Client,
|
||||
}
|
||||
|
||||
@@ -117,9 +174,21 @@ impl HttpBackendResourceClient {
|
||||
Self {
|
||||
endpoint: endpoint.into(),
|
||||
bearer_token,
|
||||
request_source_signer: None,
|
||||
request_source_audience: None,
|
||||
client: reqwest::Client::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_runtime_request_source(
|
||||
mut self,
|
||||
identity: &RuntimeIdentityMaterial,
|
||||
audience: impl Into<String>,
|
||||
) -> Self {
|
||||
self.request_source_signer = Some(RuntimeRequestSourceSigner::from_identity(identity));
|
||||
self.request_source_audience = Some(audience.into());
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "http-server")]
|
||||
@@ -129,7 +198,44 @@ impl BackendResourceClient for HttpBackendResourceClient {
|
||||
&self,
|
||||
request: BackendResourceFetchRequest,
|
||||
) -> Result<BackendResourceFetchResponse, BackendResourceError> {
|
||||
let builder = self.client.post(&self.endpoint).json(&request);
|
||||
let body = serde_json::to_vec(&request).map_err(|error| {
|
||||
BackendResourceError::InvalidResponse {
|
||||
message: error.to_string(),
|
||||
}
|
||||
})?;
|
||||
let endpoint = reqwest::Url::parse(&self.endpoint).map_err(|error| {
|
||||
BackendResourceError::Transport {
|
||||
message: error.to_string(),
|
||||
}
|
||||
})?;
|
||||
let mut builder = self
|
||||
.client
|
||||
.post(endpoint.clone())
|
||||
.header(reqwest::header::CONTENT_TYPE, "application/json")
|
||||
.body(body.clone());
|
||||
if let Some(signer) = self.request_source_signer.as_ref() {
|
||||
let audience = self.request_source_audience.as_deref().ok_or_else(|| {
|
||||
BackendResourceError::Unauthorized {
|
||||
message: "Runtime request proof audience is unavailable".to_owned(),
|
||||
}
|
||||
})?;
|
||||
let proof = signer
|
||||
.issue(
|
||||
audience,
|
||||
&request.handle.workspace_id,
|
||||
None,
|
||||
BACKEND_RESOURCE_FETCH_PERMISSION,
|
||||
"POST",
|
||||
endpoint.path(),
|
||||
&body,
|
||||
i64::try_from(unix_now_seconds()).unwrap_or(i64::MAX),
|
||||
30,
|
||||
)
|
||||
.map_err(|error| BackendResourceError::Unauthorized {
|
||||
message: error.to_string(),
|
||||
})?;
|
||||
builder = builder.header(RUNTIME_REQUEST_SOURCE_PROOF_HEADER, proof);
|
||||
}
|
||||
let builder = if let Some(token) = self.bearer_token.as_deref() {
|
||||
builder.bearer_auth(token)
|
||||
} else {
|
||||
@@ -193,7 +299,7 @@ pub fn build_profile_source_archive_fetch_request(
|
||||
|
||||
pub fn profile_source_archive_from_response(
|
||||
handle: &BackendResourceHandle,
|
||||
response: BackendResourceFetchResponse,
|
||||
mut response: BackendResourceFetchResponse,
|
||||
) -> Result<ProfileSourceArchive, BackendResourceError> {
|
||||
if handle.kind != BackendResourceKind::ProfileSourceArchive
|
||||
|| response.kind != BackendResourceKind::ProfileSourceArchive
|
||||
@@ -208,7 +314,7 @@ pub fn profile_source_archive_from_response(
|
||||
if response.content_type != handle.content_type {
|
||||
return Err(BackendResourceError::ContentTypeMismatch {
|
||||
expected: handle.content_type.clone(),
|
||||
actual: response.content_type,
|
||||
actual: response.content_type.clone(),
|
||||
});
|
||||
}
|
||||
let actual_bytes = response.bytes.len() as u64;
|
||||
@@ -223,7 +329,7 @@ pub fn profile_source_archive_from_response(
|
||||
return Err(BackendResourceError::DigestMismatch {
|
||||
expected: handle.digest.clone(),
|
||||
actual: if response.digest != handle.digest {
|
||||
response.digest
|
||||
response.digest.clone()
|
||||
} else {
|
||||
actual_digest
|
||||
},
|
||||
@@ -241,7 +347,7 @@ pub fn profile_source_archive_from_response(
|
||||
}
|
||||
})?,
|
||||
},
|
||||
content: response.bytes,
|
||||
content: std::mem::take(&mut response.bytes),
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use crate::catalog::{
|
||||
ConfigBundleRef, CreateWorkerRequest, ProfileSelector, WorkerDetail, WorkerLifecycleAck,
|
||||
WorkerStatus, WorkerSummary, WorkingDirectoryRequest,
|
||||
WorkerStatus, WorkerSummary, WorkingDirectoryRepositoryAccessRequest, WorkingDirectoryRequest,
|
||||
WorkingDirectoryStatus as CatalogWorkingDirectoryStatus, WorkspaceApiRef,
|
||||
};
|
||||
use crate::config_bundle::{
|
||||
@@ -26,6 +26,10 @@ use crate::management::{
|
||||
};
|
||||
#[cfg(feature = "ws-server")]
|
||||
use crate::observation::{WorkerObservationCursor, WorkerObservationEvent};
|
||||
use crate::resource::{
|
||||
BackendResourceClient, BackendResourceError, BackendResourceFetchRequest, BackendResourceKind,
|
||||
REPOSITORY_SSH_ACCESS_CONTENT_TYPE, RepositorySshAccessSecret,
|
||||
};
|
||||
#[cfg(feature = "fs-store")]
|
||||
use crate::retention::{
|
||||
FsWorkerRetentionProvider, WorkerRetentionExecutionRequest, WorkerRetentionExecutionResult,
|
||||
@@ -172,6 +176,14 @@ impl Runtime {
|
||||
Ok(runtime)
|
||||
}
|
||||
|
||||
pub fn install_backend_resource_client(
|
||||
&self,
|
||||
client: Arc<dyn BackendResourceClient>,
|
||||
) -> Result<(), RuntimeError> {
|
||||
self.lock()?.backend_resource_client = Some(BackendResourceClientRef(client));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Create or restore a filesystem-backed Runtime.
|
||||
///
|
||||
/// The store is scoped by `options.root`; if the directory already exists,
|
||||
@@ -366,6 +378,103 @@ impl Runtime {
|
||||
.map_err(RuntimeError::from)
|
||||
}
|
||||
|
||||
pub async fn create_working_directory_from_resource(
|
||||
&self,
|
||||
mut request: WorkingDirectoryRequest,
|
||||
) -> Result<CatalogWorkingDirectoryStatus, RuntimeError> {
|
||||
if let Some(ssh) = request
|
||||
.materialization
|
||||
.as_mut()
|
||||
.and_then(|materialization| materialization.ssh.as_mut())
|
||||
{
|
||||
self.resolve_repository_access_resource(ssh).await?;
|
||||
}
|
||||
self.create_working_directory(request)
|
||||
}
|
||||
|
||||
pub fn authorize_working_directory_repository_access(
|
||||
&self,
|
||||
request: WorkingDirectoryRepositoryAccessRequest,
|
||||
) -> Result<(), RuntimeError> {
|
||||
let backend = {
|
||||
let state = self.lock()?;
|
||||
state.ensure_running()?;
|
||||
state.execution_backend.clone().ok_or_else(|| {
|
||||
RuntimeError::ExecutionBackendUnavailable {
|
||||
message: "working directory Repository access requires an execution backend"
|
||||
.to_string(),
|
||||
}
|
||||
})?
|
||||
};
|
||||
backend
|
||||
.authorize_working_directory_repository_access(&request)
|
||||
.map_err(RuntimeError::from)
|
||||
}
|
||||
|
||||
async fn resolve_repository_access_resource(
|
||||
&self,
|
||||
ssh: &mut crate::catalog::RepositorySshMaterializationAccess,
|
||||
) -> Result<(), RuntimeError> {
|
||||
if !ssh.private_key.expose().is_empty() && !ssh.known_hosts_entry.expose().is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let (client, runtime_id) = {
|
||||
let state = self.lock()?;
|
||||
let client = state.backend_resource_client.clone().ok_or_else(|| {
|
||||
RuntimeError::InvalidRequest(
|
||||
"Backend Repository access resource client is unavailable".to_string(),
|
||||
)
|
||||
})?;
|
||||
let runtime_id = state.runtime_identity.clone().ok_or_else(|| {
|
||||
RuntimeError::InvalidRequest("Runtime identity is unavailable".to_string())
|
||||
})?;
|
||||
(client, runtime_id)
|
||||
};
|
||||
let mut response = client
|
||||
.0
|
||||
.fetch_resource(BackendResourceFetchRequest {
|
||||
handle: ssh.secret_resource.clone(),
|
||||
runtime_id,
|
||||
worker_id: None,
|
||||
audit_correlation_id: ssh.secret_resource.audit_correlation_id.clone(),
|
||||
})
|
||||
.await
|
||||
.map_err(repository_resource_error)?;
|
||||
if response.kind != BackendResourceKind::RepositorySshAccess
|
||||
|| response.content_type != REPOSITORY_SSH_ACCESS_CONTENT_TYPE
|
||||
|| response.resource_id != ssh.secret_resource.resource_id
|
||||
|| response.digest != ssh.secret_resource.digest
|
||||
|| response.bytes.len() as u64 > ssh.secret_resource.max_bytes
|
||||
{
|
||||
return Err(RuntimeError::InvalidRequest(
|
||||
"Backend Repository SSH access resource response was invalid".to_string(),
|
||||
));
|
||||
}
|
||||
let secret = serde_json::from_slice::<RepositorySshAccessSecret>(&response.bytes);
|
||||
response.bytes.fill(0);
|
||||
let mut secret = secret.map_err(|_| {
|
||||
RuntimeError::InvalidRequest(
|
||||
"Backend Repository SSH access resource payload was invalid".to_string(),
|
||||
)
|
||||
})?;
|
||||
ssh.private_key =
|
||||
crate::catalog::SensitiveString::new(std::mem::take(&mut secret.private_key));
|
||||
ssh.known_hosts_entry =
|
||||
crate::catalog::SensitiveString::new(std::mem::take(&mut secret.known_hosts_entry));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn authorize_working_directory_repository_access_from_resource(
|
||||
&self,
|
||||
mut request: WorkingDirectoryRepositoryAccessRequest,
|
||||
) -> Result<(), RuntimeError> {
|
||||
let ssh = request.materialization.ssh.as_mut().ok_or_else(|| {
|
||||
RuntimeError::InvalidRequest("Repository SSH access metadata is missing".to_string())
|
||||
})?;
|
||||
self.resolve_repository_access_resource(ssh).await?;
|
||||
self.authorize_working_directory_repository_access(request)
|
||||
}
|
||||
|
||||
/// List Runtime-owned working directories through the attached execution backend.
|
||||
pub fn list_working_directories(
|
||||
&self,
|
||||
@@ -566,12 +675,13 @@ impl Runtime {
|
||||
let worker_id = request.worker_id;
|
||||
let worker_ref = WorkerRef::new(worker_id);
|
||||
|
||||
let durable_request = durable_create_worker_request(&request);
|
||||
let record = WorkerRecord {
|
||||
worker_ref: worker_ref.clone(),
|
||||
worker_id: worker_id.clone(),
|
||||
status: WorkerStatus::Stopped,
|
||||
workspace_id: scope.map(|scope| scope.workspace_id.clone()),
|
||||
request: request.clone(),
|
||||
request: durable_request,
|
||||
run_generation: 1,
|
||||
working_directory: None,
|
||||
execution_handle: None,
|
||||
@@ -1420,7 +1530,9 @@ impl Runtime {
|
||||
}
|
||||
}
|
||||
Ok(protocol::Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
greeting: protocol::Greeting {
|
||||
worker_name: worker_ref.worker_id.to_string(),
|
||||
cwd: String::new(),
|
||||
@@ -1842,6 +1954,15 @@ struct SubscriptionSink {
|
||||
lagged: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct BackendResourceClientRef(Arc<dyn BackendResourceClient>);
|
||||
|
||||
impl std::fmt::Debug for BackendResourceClientRef {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("BackendResourceClientRef(..)")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct RuntimeState {
|
||||
display_name: Option<String>,
|
||||
@@ -1853,6 +1974,7 @@ struct RuntimeState {
|
||||
persistence: RuntimePersistence,
|
||||
status: RuntimeStatus,
|
||||
execution_backend: Option<WorkerExecutionBackendRef>,
|
||||
backend_resource_client: Option<BackendResourceClientRef>,
|
||||
#[cfg(feature = "fs-store")]
|
||||
next_diagnostic_id: u64,
|
||||
workers: BTreeMap<WorkerId, WorkerRecord>,
|
||||
@@ -1880,6 +2002,7 @@ impl RuntimeState {
|
||||
persistence: RuntimePersistence::Memory,
|
||||
status: RuntimeStatus::Running,
|
||||
execution_backend: None,
|
||||
backend_resource_client: None,
|
||||
#[cfg(feature = "fs-store")]
|
||||
next_diagnostic_id: 1,
|
||||
workers: BTreeMap::new(),
|
||||
@@ -1908,6 +2031,7 @@ impl RuntimeState {
|
||||
persistence: RuntimePersistence::Fs(store),
|
||||
status: RuntimeStatus::Running,
|
||||
execution_backend: None,
|
||||
backend_resource_client: None,
|
||||
#[cfg(feature = "fs-store")]
|
||||
next_diagnostic_id: 1,
|
||||
workers: BTreeMap::new(),
|
||||
@@ -1959,6 +2083,7 @@ impl RuntimeState {
|
||||
persistence: RuntimePersistence::Fs(store),
|
||||
status: persisted.status,
|
||||
execution_backend: None,
|
||||
backend_resource_client: None,
|
||||
next_diagnostic_id,
|
||||
workers,
|
||||
config_bundles: BTreeMap::new(),
|
||||
@@ -2619,6 +2744,7 @@ impl RuntimeState {
|
||||
protocol::WorkerStatus::Running => Some(WorkerStatus::Running),
|
||||
protocol::WorkerStatus::Idle => Some(WorkerStatus::Idle),
|
||||
protocol::WorkerStatus::Paused => Some(WorkerStatus::Paused),
|
||||
protocol::WorkerStatus::Stopped => Some(WorkerStatus::Stopped),
|
||||
},
|
||||
protocol::Event::RunEnd { result } => match result {
|
||||
protocol::RunResult::Finished | protocol::RunResult::RolledBack => {
|
||||
@@ -2714,6 +2840,33 @@ fn worker_status_from_run_state(run_state: WorkerExecutionRunState) -> WorkerSta
|
||||
}
|
||||
}
|
||||
|
||||
fn repository_resource_error(error: BackendResourceError) -> RuntimeError {
|
||||
let category = match error {
|
||||
BackendResourceError::Expired => "expired",
|
||||
BackendResourceError::Unauthorized { .. } => "unauthorized",
|
||||
BackendResourceError::UnsupportedKind => "unsupported_kind",
|
||||
BackendResourceError::MissingResource => "missing_resource",
|
||||
BackendResourceError::Oversized { .. } => "oversized",
|
||||
BackendResourceError::DigestMismatch { .. } => "digest_mismatch",
|
||||
BackendResourceError::ContentTypeMismatch { .. } => "content_type_mismatch",
|
||||
BackendResourceError::InvalidResponse { .. } => "invalid_response",
|
||||
BackendResourceError::Transport { .. } => "transport",
|
||||
};
|
||||
RuntimeError::InvalidRequest(format!(
|
||||
"Backend Repository SSH access resource fetch failed: {category}"
|
||||
))
|
||||
}
|
||||
|
||||
fn durable_create_worker_request(request: &CreateWorkerRequest) -> CreateWorkerRequest {
|
||||
let mut durable = request.clone();
|
||||
if let Some(working_directory) = durable.working_directory_request.as_mut()
|
||||
&& let Some(materialization) = working_directory.materialization.as_mut()
|
||||
{
|
||||
materialization.ssh = None;
|
||||
}
|
||||
durable
|
||||
}
|
||||
|
||||
fn requested_primary_workdir_id(request: &CreateWorkerRequest) -> Option<&str> {
|
||||
request
|
||||
.working_directory
|
||||
@@ -2884,7 +3037,9 @@ fn subscription_worker_state(status: WorkerStatus) -> SubscriptionWorkerState {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::catalog::{
|
||||
ConfigBundleRef, ProfileSelector, WorkingDirectoryClaim, WorkspaceApiRef,
|
||||
ConfigBundleRef, MaterializerKind, ProfileSelector, RepositoryMaterializationContext,
|
||||
RepositorySshMaterializationAccess, SensitiveString, WorkingDirectoryClaim,
|
||||
WorkingDirectoryRepository, WorkingDirectoryRequest, WorkspaceApiRef,
|
||||
};
|
||||
use crate::config_bundle::{
|
||||
ConfigBundle, ConfigBundleMetadata, ConfigBundleProvenance, ConfigDeclaration,
|
||||
@@ -2894,6 +3049,8 @@ mod tests {
|
||||
WorkerExecutionBackend, WorkerExecutionContext, WorkerExecutionHandle,
|
||||
WorkerExecutionRestoreRequest, WorkerExecutionRunState,
|
||||
};
|
||||
use crate::working_directory::WorkingDirectoryDiagnostic;
|
||||
use async_trait::async_trait;
|
||||
use std::collections::BTreeMap;
|
||||
#[cfg(feature = "fs-store")]
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
@@ -2950,7 +3107,7 @@ mod tests {
|
||||
&mut activity,
|
||||
&internal_worker_status_event(
|
||||
internal_worker_ref("child-b", None),
|
||||
protocol::WorkerStatus::Idle,
|
||||
protocol::WorkerStatus::Stopped,
|
||||
),
|
||||
));
|
||||
}
|
||||
@@ -2997,7 +3154,9 @@ mod tests {
|
||||
),
|
||||
);
|
||||
let snapshot = protocol::Event::Snapshot {
|
||||
entries: Vec::new(),
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: Vec::new(),
|
||||
},
|
||||
greeting: protocol::Greeting {
|
||||
worker_name: "parent".to_string(),
|
||||
cwd: "/tmp".to_string(),
|
||||
@@ -3115,6 +3274,243 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn durable_worker_request_omits_repository_credentials() {
|
||||
let mut request = task_request("worker-secret-redaction");
|
||||
request.working_directory_request = Some(WorkingDirectoryRequest {
|
||||
repository: WorkingDirectoryRepository {
|
||||
id: "repository-1".to_string(),
|
||||
provider: "git".to_string(),
|
||||
source: workspace_api::RepositorySource {
|
||||
kind: workspace_api::RepositorySourceKind::Ssh,
|
||||
uri: "ssh://git@example.test/repo.git".to_string(),
|
||||
},
|
||||
source_revision: 1,
|
||||
source_fingerprint: "sha256:source".to_string(),
|
||||
selector: None,
|
||||
},
|
||||
materializer: MaterializerKind::RuntimeGitCache,
|
||||
backend_workdir_id: Some("working-directory-1".to_string()),
|
||||
materialization: Some(RepositoryMaterializationContext {
|
||||
workspace_id: "workspace-1".to_string(),
|
||||
runtime_id: "runtime-1".to_string(),
|
||||
operation_id: "operation-1".to_string(),
|
||||
config_revision: 1,
|
||||
config_projection_digest: "sha256:projection".to_string(),
|
||||
cache_generation: 0,
|
||||
ssh: Some(RepositorySshMaterializationAccess {
|
||||
credential_id: "credential-1".to_string(),
|
||||
credential_revision: 1,
|
||||
host_trust_id: "host-trust-1".to_string(),
|
||||
host_trust_revision: 1,
|
||||
access: workspace_api::RepositoryAccessMode::ReadOnly,
|
||||
expires_at_epoch_seconds: u64::MAX,
|
||||
repository_id: "repository-1".to_string(),
|
||||
repository_source_fingerprint: "sha256:source".to_string(),
|
||||
repository_uri: "ssh://git@example.test/repo.git".to_string(),
|
||||
secret_resource: repository_resource_handle(),
|
||||
private_key: SensitiveString::new("private-key-bytes"),
|
||||
known_hosts_entry: SensitiveString::new("known-hosts-entry"),
|
||||
}),
|
||||
}),
|
||||
});
|
||||
|
||||
let durable = durable_create_worker_request(&request);
|
||||
|
||||
assert!(
|
||||
request
|
||||
.working_directory_request
|
||||
.as_ref()
|
||||
.and_then(|working_directory| working_directory.materialization.as_ref())
|
||||
.and_then(|materialization| materialization.ssh.as_ref())
|
||||
.is_some()
|
||||
);
|
||||
assert!(
|
||||
durable
|
||||
.working_directory_request
|
||||
.as_ref()
|
||||
.and_then(|working_directory| working_directory.materialization.as_ref())
|
||||
.and_then(|materialization| materialization.ssh.as_ref())
|
||||
.is_none()
|
||||
);
|
||||
let serialized = serde_json::to_string(&durable).unwrap();
|
||||
assert!(!serialized.contains("private-key-bytes"));
|
||||
assert!(!serialized.contains("known-hosts-entry"));
|
||||
}
|
||||
|
||||
fn repository_resource_handle() -> crate::resource::BackendResourceHandle {
|
||||
crate::resource::BackendResourceHandle {
|
||||
kind: crate::resource::BackendResourceKind::RepositorySshAccess,
|
||||
workspace_id: "workspace-1".to_string(),
|
||||
scope_id: Some("repository-ssh-access".to_string()),
|
||||
runtime_id: Some("runtime-1".to_string()),
|
||||
worker_id: None,
|
||||
resource_id: "repository-access-1".to_string(),
|
||||
digest: "opaque:repository-access-1".to_string(),
|
||||
operation: crate::resource::BackendResourceOperation::FetchOnce,
|
||||
expires_at_unix_seconds: i64::MAX,
|
||||
nonce: "repository-access-1".to_string(),
|
||||
revision: "1".to_string(),
|
||||
generation: None,
|
||||
max_bytes: crate::resource::DEFAULT_REPOSITORY_SSH_ACCESS_MAX_BYTES,
|
||||
content_type: crate::resource::REPOSITORY_SSH_ACCESS_CONTENT_TYPE.to_string(),
|
||||
redaction: crate::resource::ResourceRedactionPolicy::RuntimeInternalOnly,
|
||||
audit_correlation_id: "repository-access-1".to_string(),
|
||||
profile_source_graph: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_access_resource_is_fetched_before_provider_authorization() {
|
||||
let (runtime, backend) = runtime_and_backend();
|
||||
backend
|
||||
.repository_access_available
|
||||
.store(true, Ordering::SeqCst);
|
||||
runtime.bind_runtime_identity("runtime-1").unwrap();
|
||||
let handle = repository_resource_handle();
|
||||
runtime
|
||||
.install_backend_resource_client(Arc::new(TestRepositoryResourceClient {
|
||||
response: Mutex::new(Some(crate::resource::BackendResourceFetchResponse {
|
||||
kind: crate::resource::BackendResourceKind::RepositorySshAccess,
|
||||
resource_id: handle.resource_id.clone(),
|
||||
digest: handle.digest.clone(),
|
||||
content_type: crate::resource::REPOSITORY_SSH_ACCESS_CONTENT_TYPE.to_string(),
|
||||
bytes: serde_json::to_vec(&RepositorySshAccessSecret {
|
||||
private_key: "private-key-bytes".to_string(),
|
||||
known_hosts_entry: "known-hosts-entry".to_string(),
|
||||
})
|
||||
.unwrap(),
|
||||
audit_correlation_id: handle.audit_correlation_id.clone(),
|
||||
})),
|
||||
}))
|
||||
.unwrap();
|
||||
let request = WorkingDirectoryRepositoryAccessRequest {
|
||||
working_directory_id: "working-directory-1".to_string(),
|
||||
materialization: RepositoryMaterializationContext {
|
||||
workspace_id: "workspace-1".to_string(),
|
||||
runtime_id: "runtime-1".to_string(),
|
||||
operation_id: "operation-1".to_string(),
|
||||
config_revision: 1,
|
||||
config_projection_digest: "sha256:projection".to_string(),
|
||||
cache_generation: 0,
|
||||
ssh: Some(RepositorySshMaterializationAccess {
|
||||
credential_id: "credential-1".to_string(),
|
||||
credential_revision: 1,
|
||||
host_trust_id: "host-trust-1".to_string(),
|
||||
host_trust_revision: 1,
|
||||
access: workspace_api::RepositoryAccessMode::ReadOnly,
|
||||
expires_at_epoch_seconds: u64::MAX,
|
||||
repository_id: "repository-1".to_string(),
|
||||
repository_source_fingerprint: "sha256:source".to_string(),
|
||||
repository_uri: "ssh://git@example.test/repo.git".to_string(),
|
||||
secret_resource: handle,
|
||||
private_key: SensitiveString::default(),
|
||||
known_hosts_entry: SensitiveString::default(),
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
let replay = request.clone();
|
||||
runtime
|
||||
.authorize_working_directory_repository_access_from_resource(request)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(
|
||||
runtime
|
||||
.authorize_working_directory_repository_access_from_resource(replay)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
|
||||
let accesses = backend.repository_accesses.lock().unwrap();
|
||||
assert_eq!(accesses.len(), 1);
|
||||
let access = accesses[0].materialization.ssh.as_ref().unwrap();
|
||||
assert_eq!(access.private_key.expose(), "private-key-bytes");
|
||||
assert_eq!(access.known_hosts_entry.expose(), "known-hosts-entry");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn working_directory_create_fetches_repository_access_before_provider_call() {
|
||||
let (runtime, backend) = runtime_and_backend();
|
||||
backend
|
||||
.repository_access_available
|
||||
.store(true, Ordering::SeqCst);
|
||||
runtime.bind_runtime_identity("runtime-1").unwrap();
|
||||
let handle = repository_resource_handle();
|
||||
runtime
|
||||
.install_backend_resource_client(Arc::new(TestRepositoryResourceClient {
|
||||
response: Mutex::new(Some(crate::resource::BackendResourceFetchResponse {
|
||||
kind: crate::resource::BackendResourceKind::RepositorySshAccess,
|
||||
resource_id: handle.resource_id.clone(),
|
||||
digest: handle.digest.clone(),
|
||||
content_type: crate::resource::REPOSITORY_SSH_ACCESS_CONTENT_TYPE.to_string(),
|
||||
bytes: serde_json::to_vec(&RepositorySshAccessSecret {
|
||||
private_key: "create-private-key-bytes".to_string(),
|
||||
known_hosts_entry: "create-known-hosts-entry".to_string(),
|
||||
})
|
||||
.unwrap(),
|
||||
audit_correlation_id: handle.audit_correlation_id.clone(),
|
||||
})),
|
||||
}))
|
||||
.unwrap();
|
||||
let request = WorkingDirectoryRequest {
|
||||
repository: WorkingDirectoryRepository {
|
||||
id: "repository-1".to_string(),
|
||||
provider: "git".to_string(),
|
||||
source: workspace_api::RepositorySource {
|
||||
kind: workspace_api::RepositorySourceKind::Ssh,
|
||||
uri: "ssh://git@example.test/repo.git".to_string(),
|
||||
},
|
||||
source_revision: 1,
|
||||
source_fingerprint: "sha256:source".to_string(),
|
||||
selector: None,
|
||||
},
|
||||
materializer: MaterializerKind::RuntimeGitCache,
|
||||
backend_workdir_id: Some("working-directory-1".to_string()),
|
||||
materialization: Some(RepositoryMaterializationContext {
|
||||
workspace_id: "workspace-1".to_string(),
|
||||
runtime_id: "runtime-1".to_string(),
|
||||
operation_id: "operation-create".to_string(),
|
||||
config_revision: 1,
|
||||
config_projection_digest: "sha256:projection".to_string(),
|
||||
cache_generation: 0,
|
||||
ssh: Some(RepositorySshMaterializationAccess {
|
||||
credential_id: "credential-1".to_string(),
|
||||
credential_revision: 1,
|
||||
host_trust_id: "host-trust-1".to_string(),
|
||||
host_trust_revision: 1,
|
||||
access: workspace_api::RepositoryAccessMode::ReadOnly,
|
||||
expires_at_epoch_seconds: u64::MAX,
|
||||
repository_id: "repository-1".to_string(),
|
||||
repository_source_fingerprint: "sha256:source".to_string(),
|
||||
repository_uri: "ssh://git@example.test/repo.git".to_string(),
|
||||
secret_resource: handle,
|
||||
private_key: SensitiveString::default(),
|
||||
known_hosts_entry: SensitiveString::default(),
|
||||
}),
|
||||
}),
|
||||
};
|
||||
|
||||
assert!(
|
||||
runtime
|
||||
.create_working_directory_from_resource(request)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
|
||||
let requests = backend.working_directory_requests.lock().unwrap();
|
||||
let access = requests[0]
|
||||
.materialization
|
||||
.as_ref()
|
||||
.and_then(|materialization| materialization.ssh.as_ref())
|
||||
.unwrap();
|
||||
assert_eq!(access.private_key.expose(), "create-private-key-bytes");
|
||||
assert_eq!(
|
||||
access.known_hosts_entry.expose(),
|
||||
"create-known-hosts-entry"
|
||||
);
|
||||
}
|
||||
|
||||
fn scoped_task_request(objective: &str, workspace_id: &str) -> CreateWorkerRequest {
|
||||
let mut request = task_request(objective);
|
||||
request.workspace_api = Some(WorkspaceApiRef {
|
||||
@@ -3196,6 +3592,9 @@ mod tests {
|
||||
config_bundles: Mutex<Vec<Option<ConfigBundle>>>,
|
||||
contexts: Mutex<BTreeMap<WorkerId, WorkerExecutionContext>>,
|
||||
dispatched_inputs: Mutex<Vec<WorkerInput>>,
|
||||
repository_accesses: Mutex<Vec<WorkingDirectoryRepositoryAccessRequest>>,
|
||||
repository_access_available: AtomicBool,
|
||||
working_directory_requests: Mutex<Vec<WorkingDirectoryRequest>>,
|
||||
preserve_commit_ack_submission_id: AtomicBool,
|
||||
#[cfg(feature = "ws-server")]
|
||||
snapshots: Mutex<BTreeMap<WorkerId, protocol::Event>>,
|
||||
@@ -3236,6 +3635,38 @@ mod tests {
|
||||
"test-execution-backend"
|
||||
}
|
||||
|
||||
fn create_working_directory(
|
||||
&self,
|
||||
request: &WorkingDirectoryRequest,
|
||||
) -> Result<CatalogWorkingDirectoryStatus, WorkingDirectoryDiagnostic> {
|
||||
self.working_directory_requests
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(request.clone());
|
||||
Err(WorkingDirectoryDiagnostic::rejected(
|
||||
"working_directory_unsupported",
|
||||
"Worker execution backend does not support working directory materialization",
|
||||
))
|
||||
}
|
||||
|
||||
fn authorize_working_directory_repository_access(
|
||||
&self,
|
||||
request: &WorkingDirectoryRepositoryAccessRequest,
|
||||
) -> Result<(), WorkingDirectoryDiagnostic> {
|
||||
self.repository_accesses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(request.clone());
|
||||
if self.repository_access_available.load(Ordering::SeqCst) {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(WorkingDirectoryDiagnostic::rejected(
|
||||
"working_directory_repository_access_unsupported",
|
||||
"Worker execution backend does not support Repository access authorization",
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
fn spawn_worker(&self, request: WorkerExecutionSpawnRequest) -> WorkerExecutionSpawnResult {
|
||||
self.run_generations
|
||||
.lock()
|
||||
@@ -3343,6 +3774,24 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
struct TestRepositoryResourceClient {
|
||||
response: Mutex<Option<crate::resource::BackendResourceFetchResponse>>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl BackendResourceClient for TestRepositoryResourceClient {
|
||||
async fn fetch_resource(
|
||||
&self,
|
||||
_request: BackendResourceFetchRequest,
|
||||
) -> Result<crate::resource::BackendResourceFetchResponse, BackendResourceError> {
|
||||
self.response
|
||||
.lock()
|
||||
.unwrap()
|
||||
.take()
|
||||
.ok_or(BackendResourceError::MissingResource)
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_with_backend() -> Runtime {
|
||||
let runtime = Runtime::with_execution_backend(
|
||||
RuntimeOptions::default(),
|
||||
@@ -4136,7 +4585,17 @@ mod tests {
|
||||
backend.set_worker_snapshot(
|
||||
&detail.worker_ref,
|
||||
protocol::Event::Snapshot {
|
||||
entries: vec![expected_entry.clone()],
|
||||
session: protocol::SessionSnapshot {
|
||||
entries: vec![protocol::SessionSnapshotEntry {
|
||||
entry_id: "restored-log-entry".to_owned(),
|
||||
timestamp: 1,
|
||||
provenance: protocol::SessionEntryProvenance::LegacyUnknown,
|
||||
derived_from: Vec::new(),
|
||||
data: protocol::SessionSnapshotEntryData::RunError {
|
||||
message: expected_entry.to_string(),
|
||||
},
|
||||
}],
|
||||
},
|
||||
greeting: protocol::Greeting {
|
||||
worker_name: "live-worker".to_string(),
|
||||
cwd: "/tmp/live".to_string(),
|
||||
@@ -4161,12 +4620,13 @@ mod tests {
|
||||
.unwrap();
|
||||
match snapshot {
|
||||
protocol::Event::Snapshot {
|
||||
entries,
|
||||
session,
|
||||
greeting,
|
||||
status,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(entries, vec![expected_entry]);
|
||||
assert_eq!(session.entries.len(), 1);
|
||||
assert_eq!(session.entries[0].entry_id, "restored-log-entry");
|
||||
assert_eq!(greeting.worker_name, "live-worker");
|
||||
assert_eq!(status, protocol::WorkerStatus::Running);
|
||||
}
|
||||
|
||||
@@ -14,10 +14,13 @@ use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::{Arc, Mutex, mpsc};
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::auth::RuntimeIdentityMaterial;
|
||||
use crate::auth::{
|
||||
BACKEND_RESOURCE_FETCH_PERMISSION, RUNTIME_REQUEST_SOURCE_PROOF_HEADER,
|
||||
RuntimeIdentityMaterial, RuntimeRequestSourceSigner, unix_now_seconds,
|
||||
};
|
||||
use crate::catalog::{
|
||||
CreateWorkerRequest, ProfileSourceArchiveHttpRef, ProfileSourceArchiveSource,
|
||||
WorkingDirectoryRequest, WorkingDirectoryStatus,
|
||||
WorkingDirectoryRepositoryAccessRequest, WorkingDirectoryRequest, WorkingDirectoryStatus,
|
||||
};
|
||||
use crate::execution::{
|
||||
WorkerExecutionBackend, WorkerExecutionHandle, WorkerExecutionOperation,
|
||||
@@ -34,10 +37,8 @@ use crate::working_directory::{
|
||||
WorkingDirectoryBinding, WorkingDirectoryDiagnostic, WorkingDirectoryMaterializer,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use protocol::{Event, Method, Segment, WorkerStatus};
|
||||
use session_store::{
|
||||
CombinedStore, LogEntry, WorkerAggregateStore, WorkerSessionStore, collect_state,
|
||||
};
|
||||
use protocol::{ErrorCode, Event, Method, Segment, WorkerStatus};
|
||||
use session_store::{CombinedStore, LogEntry, WorkerAggregateStore, WorkerSessionStore};
|
||||
#[cfg(test)]
|
||||
use session_store::{FsStore, FsWorkerStore};
|
||||
use tokio::runtime::Runtime;
|
||||
@@ -45,6 +46,8 @@ use tokio::runtime::Runtime;
|
||||
use tokio::sync::broadcast;
|
||||
use workdir::{LocalWorkdirSession, Workdir, WorkdirSessionCapabilities, WorkdirSessionHandle};
|
||||
|
||||
#[cfg(test)]
|
||||
use worker::WorkerController;
|
||||
use worker::feature::builtin::{
|
||||
CompositeWorkerObservationProvider, WorkerObservationError, WorkerObservationProvider,
|
||||
WorkerObservationSubject, WorkerObservationSubjectRef, WorkerSessionCapture,
|
||||
@@ -53,9 +56,10 @@ use worker::feature::builtin::{
|
||||
#[cfg(feature = "ws-server")]
|
||||
use worker::ipc::protocol_session::{live_log_entry_event, subscribe_worker_protocol_session};
|
||||
use worker::{
|
||||
PromptCatalogSource, SegmentLogSink, WORKER_INPUT_SUBMISSION_EXTENSION_DOMAIN, Worker,
|
||||
WorkerController, WorkerControllerTransport, WorkerError, WorkerFilesystemAuthority,
|
||||
WorkerHandle, WorkerSharedState, WorkerWorkspaceContext, WorkspaceClient, WorkspaceId,
|
||||
PreparedWorker, PromptCatalogSource, SegmentLogSink, WORKER_INPUT_SUBMISSION_EXTENSION_DOMAIN,
|
||||
Worker, WorkerBootstrap, WorkerBootstrapError, WorkerBootstrapLayout,
|
||||
WorkerControllerTransport, WorkerError, WorkerFilesystemAuthority, WorkerHandle,
|
||||
WorkerSharedState, WorkerWorkspaceContext, WorkspaceClient, WorkspaceId,
|
||||
};
|
||||
|
||||
const DEFAULT_BACKEND_ID: &str = "worker-crate";
|
||||
@@ -65,8 +69,9 @@ const RUNTIME_TASK_TIMEOUT: Duration = Duration::from_secs(10);
|
||||
const USER_INPUT_COMMIT_TIMEOUT: Duration = Duration::from_secs(9);
|
||||
|
||||
fn user_input_has_submission(entry: &LogEntry, submission_id: &str) -> bool {
|
||||
let LogEntry::UserInput { extensions, .. } = entry else {
|
||||
return false;
|
||||
let extensions = match entry {
|
||||
LogEntry::AnnotatedUserInput { extensions, .. } => extensions,
|
||||
_ => return false,
|
||||
};
|
||||
extensions.iter().any(|extension| {
|
||||
extension.domain == WORKER_INPUT_SUBMISSION_EXTENSION_DOMAIN
|
||||
@@ -209,11 +214,11 @@ impl WorkerObservationProvider for RuntimeGrantedWorkerObservationProvider {
|
||||
return Err(WorkerObservationError::NotFound);
|
||||
}
|
||||
let entries = sink.subscribe_with_snapshot().0;
|
||||
let state = collect_state(&entries);
|
||||
Ok(WorkerSessionCapture {
|
||||
segment_id: format!("runtime:{runtime_id}:worker:{worker_id}"),
|
||||
items: state.history,
|
||||
})
|
||||
WorkerSessionCapture::from_log_entries(
|
||||
format!("runtime:{runtime_id}:worker:{worker_id}"),
|
||||
&entries,
|
||||
)
|
||||
.map_err(WorkerObservationError::Unavailable)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -295,6 +300,7 @@ pub struct ProfileRuntimeWorkerFactory {
|
||||
prompt_projection_cache: Arc<WorkspacePromptProjectionCache>,
|
||||
runtime_id: Option<String>,
|
||||
worker_mutation_identity: Option<RuntimeIdentityMaterial>,
|
||||
runtime_request_audience: Option<String>,
|
||||
embedded_worker_mutation_dispatcher: Option<Arc<dyn EmbeddedWorkerMutationDispatcher>>,
|
||||
controller_transport: WorkerControllerTransport,
|
||||
}
|
||||
@@ -311,6 +317,7 @@ impl ProfileRuntimeWorkerFactory {
|
||||
prompt_projection_cache: Arc::new(WorkspacePromptProjectionCache::default()),
|
||||
runtime_id: None,
|
||||
worker_mutation_identity: None,
|
||||
runtime_request_audience: None,
|
||||
embedded_worker_mutation_dispatcher: None,
|
||||
controller_transport: WorkerControllerTransport::UnixSocket,
|
||||
}
|
||||
@@ -331,6 +338,17 @@ impl ProfileRuntimeWorkerFactory {
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_runtime_request_identity(
|
||||
mut self,
|
||||
identity: RuntimeIdentityMaterial,
|
||||
audience: impl Into<String>,
|
||||
) -> Self {
|
||||
self.runtime_id = Some(identity.identity_id.clone());
|
||||
self.worker_mutation_identity = Some(identity);
|
||||
self.runtime_request_audience = Some(audience.into());
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_embedded_worker_mutation_dispatcher(
|
||||
mut self,
|
||||
runtime_id: impl Into<String>,
|
||||
@@ -406,7 +424,7 @@ impl ProfileRuntimeWorkerFactory {
|
||||
fn restore_fallback_manifest(
|
||||
worker_name: &str,
|
||||
) -> Result<(manifest::WorkerManifest, PromptCatalogSource), String> {
|
||||
let mut config = manifest::WorkerManifestConfig::builtin_defaults();
|
||||
let mut config = manifest::WorkerManifestConfig::resolution_defaults();
|
||||
config.worker.name = Some(worker_name.to_string());
|
||||
let manifest = manifest::WorkerManifest::try_from(config)
|
||||
.map_err(|err| format!("failed to build restore fallback manifest: {err}"))?;
|
||||
@@ -457,13 +475,15 @@ impl ProfileRuntimeWorkerFactory {
|
||||
async fn resolve_profile_source_archive(
|
||||
&self,
|
||||
source: &ProfileSourceArchiveSource,
|
||||
request_audience: Option<&str>,
|
||||
) -> Result<crate::profile_archive::VerifiedProfileSourceArchive, String> {
|
||||
match source {
|
||||
ProfileSourceArchiveSource::Embedded { archive } => archive
|
||||
.verify()
|
||||
.map_err(|err| format!("failed to verify embedded profile source archive: {err}")),
|
||||
ProfileSourceArchiveSource::Http { location } => {
|
||||
self.fetch_profile_source_archive(location).await
|
||||
self.fetch_profile_source_archive(location, request_audience)
|
||||
.await
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -471,10 +491,18 @@ impl ProfileRuntimeWorkerFactory {
|
||||
async fn fetch_profile_source_archive(
|
||||
&self,
|
||||
location: &ProfileSourceArchiveHttpRef,
|
||||
request_audience: Option<&str>,
|
||||
) -> Result<crate::profile_archive::VerifiedProfileSourceArchive, String> {
|
||||
if let Some(cached) = self.profile_archive_cache.get(&location.archive.digest) {
|
||||
let response =
|
||||
fetch_profile_source_archive_http(location, Some(&location.archive.digest)).await?;
|
||||
let response = fetch_profile_source_archive_http(
|
||||
location,
|
||||
Some(&location.archive.digest),
|
||||
self.worker_mutation_identity.as_ref(),
|
||||
self.runtime_request_audience
|
||||
.as_deref()
|
||||
.or(request_audience),
|
||||
)
|
||||
.await?;
|
||||
if let Some(fetched) = response {
|
||||
self.profile_archive_cache.insert(fetched.clone());
|
||||
fetched.verify().map_err(|err| {
|
||||
@@ -486,12 +514,19 @@ impl ProfileRuntimeWorkerFactory {
|
||||
.map_err(|err| format!("failed to verify cached profile source archive: {err}"))
|
||||
}
|
||||
} else {
|
||||
let archive = fetch_profile_source_archive_http(location, None)
|
||||
.await?
|
||||
.ok_or_else(|| {
|
||||
"profile source archive HTTP revalidation returned 304 without a cached archive"
|
||||
.to_string()
|
||||
})?;
|
||||
let archive = fetch_profile_source_archive_http(
|
||||
location,
|
||||
None,
|
||||
self.worker_mutation_identity.as_ref(),
|
||||
self.runtime_request_audience
|
||||
.as_deref()
|
||||
.or(request_audience),
|
||||
)
|
||||
.await?
|
||||
.ok_or_else(|| {
|
||||
"profile source archive HTTP revalidation returned 304 without a cached archive"
|
||||
.to_string()
|
||||
})?;
|
||||
self.profile_archive_cache.insert(archive.clone());
|
||||
archive
|
||||
.verify()
|
||||
@@ -527,6 +562,7 @@ impl RuntimeWorkspaceBackendRef {
|
||||
worker_ref: &WorkerRef,
|
||||
workspace_scope: Option<&crate::runtime::RuntimeWorkspaceScope>,
|
||||
mutation_identity: Option<&RuntimeIdentityMaterial>,
|
||||
runtime_request_audience: Option<&str>,
|
||||
embedded_dispatcher: Option<&Arc<dyn EmbeddedWorkerMutationDispatcher>>,
|
||||
prompt_projection_cache: Option<Arc<WorkspacePromptProjectionCache>>,
|
||||
) -> WorkerWorkspaceContext {
|
||||
@@ -546,6 +582,13 @@ impl RuntimeWorkspaceBackendRef {
|
||||
if let Some(cache) = prompt_projection_cache {
|
||||
client = client.with_prompt_projection_cache(cache);
|
||||
}
|
||||
if let Some(identity) = mutation_identity {
|
||||
let audience = runtime_request_audience
|
||||
.or_else(|| workspace_scope.map(|scope| scope.server_id.as_str()));
|
||||
if let Some(audience) = audience {
|
||||
client = client.with_runtime_request_source(identity, audience.to_owned());
|
||||
}
|
||||
}
|
||||
if let (Some(scope), Some(identity)) = (workspace_scope, mutation_identity) {
|
||||
client = client.with_worker_remove(RuntimeWorkerMutationForwarder::remote(
|
||||
identity,
|
||||
@@ -576,9 +619,40 @@ impl RuntimeWorkspaceBackendRef {
|
||||
async fn fetch_profile_source_archive_http(
|
||||
location: &ProfileSourceArchiveHttpRef,
|
||||
cached_digest: Option<&str>,
|
||||
identity: Option<&RuntimeIdentityMaterial>,
|
||||
audience: Option<&str>,
|
||||
) -> Result<Option<crate::profile_archive::ProfileSourceArchive>, String> {
|
||||
let client = reqwest::Client::new();
|
||||
let mut request = client.get(&location.url);
|
||||
let url = reqwest::Url::parse(&location.url)
|
||||
.map_err(|error| format!("profile source archive URL is invalid: {error}"))?;
|
||||
let path = url.path().to_owned();
|
||||
let workspace_id = path
|
||||
.split('/')
|
||||
.collect::<Vec<_>>()
|
||||
.windows(2)
|
||||
.find_map(|parts| (parts[0] == "w").then_some(parts[1]))
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| "profile source archive URL is not workspace-scoped".to_owned())?;
|
||||
let mut request = client.get(url);
|
||||
if let Some(identity) = identity {
|
||||
let audience = audience.ok_or_else(|| {
|
||||
"profile source archive request proof audience is unavailable".to_owned()
|
||||
})?;
|
||||
let proof = RuntimeRequestSourceSigner::from_identity(identity)
|
||||
.issue(
|
||||
audience,
|
||||
workspace_id,
|
||||
None,
|
||||
BACKEND_RESOURCE_FETCH_PERMISSION,
|
||||
"GET",
|
||||
&path,
|
||||
b"",
|
||||
i64::try_from(unix_now_seconds()).unwrap_or(i64::MAX),
|
||||
30,
|
||||
)
|
||||
.map_err(|error| error.to_string())?;
|
||||
request = request.header(RUNTIME_REQUEST_SOURCE_PROOF_HEADER, proof);
|
||||
}
|
||||
if cached_digest == Some(location.archive.digest.as_str()) {
|
||||
if let Some(etag) = location.etag.as_deref() {
|
||||
request = request.header(reqwest::header::IF_NONE_MATCH, etag);
|
||||
@@ -620,6 +694,8 @@ async fn fetch_profile_source_archive_http(
|
||||
async fn fetch_profile_source_archive_http(
|
||||
_location: &ProfileSourceArchiveHttpRef,
|
||||
_cached_digest: Option<&str>,
|
||||
_identity: Option<&RuntimeIdentityMaterial>,
|
||||
_audience: Option<&str>,
|
||||
) -> Result<Option<crate::profile_archive::ProfileSourceArchive>, String> {
|
||||
Err(
|
||||
"HTTP profile source archive fetch requires the worker-runtime http-server feature"
|
||||
@@ -632,13 +708,17 @@ fn runtime_local_workdir_session(
|
||||
root: &Path,
|
||||
cwd: &Path,
|
||||
scope: manifest::SharedScope,
|
||||
command_environment: std::collections::BTreeMap<String, String>,
|
||||
resources: Vec<Arc<dyn workdir::WorkdirSessionResource>>,
|
||||
) -> WorkdirSessionHandle {
|
||||
Arc::new(LocalWorkdirSession::materialized_bound(
|
||||
Arc::new(LocalWorkdirSession::materialized_bound_with_environment(
|
||||
Workdir::new(workdir_id),
|
||||
root.to_path_buf(),
|
||||
cwd.to_path_buf(),
|
||||
scope,
|
||||
WorkdirSessionCapabilities::ALL,
|
||||
command_environment,
|
||||
resources,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -743,12 +823,19 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory {
|
||||
&request.worker_ref,
|
||||
request.workspace_scope.as_ref(),
|
||||
self.worker_mutation_identity.as_ref(),
|
||||
self.runtime_request_audience.as_deref(),
|
||||
self.embedded_worker_mutation_dispatcher.as_ref(),
|
||||
Some(self.prompt_projection_cache.clone()),
|
||||
);
|
||||
let selector = profile.as_ref();
|
||||
let archive = self
|
||||
.resolve_profile_source_archive(&request.request.profile_source)
|
||||
.resolve_profile_source_archive(
|
||||
&request.request.profile_source,
|
||||
request
|
||||
.workspace_scope
|
||||
.as_ref()
|
||||
.map(|scope| scope.server_id.as_str()),
|
||||
)
|
||||
.await?;
|
||||
let (mut manifest, mut loader) = {
|
||||
let manifest = archive
|
||||
@@ -796,15 +883,31 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory {
|
||||
)?;
|
||||
let store = CombinedStore::new(session_store, worker_metadata_store);
|
||||
|
||||
let mut worker = Worker::from_manifest_with_context(
|
||||
let run_dir = worker_aggregate_dir
|
||||
.join("runs")
|
||||
.join(request.run_generation.to_string());
|
||||
let mut prepared = WorkerBootstrap::new(
|
||||
manifest,
|
||||
store,
|
||||
loader,
|
||||
workspace_context,
|
||||
filesystem_authority,
|
||||
WorkerBootstrapLayout::RuntimeManagedRun {
|
||||
run_dir: run_dir.clone(),
|
||||
},
|
||||
self.controller_transport,
|
||||
)
|
||||
.prepare()
|
||||
.await
|
||||
.map_err(|err| format!("failed to create Worker from profile: {err}"))?;
|
||||
.map_err(|error| match error {
|
||||
WorkerBootstrapError::Worker(source) => {
|
||||
format!("failed to create Worker from profile: {source}")
|
||||
}
|
||||
WorkerBootstrapError::Controller { source, .. } => {
|
||||
format!("failed to prepare Worker controller: {source}")
|
||||
}
|
||||
})?;
|
||||
let worker = prepared.worker_mut();
|
||||
validate_worker_memory_settings(worker.manifest(), &request.request)?;
|
||||
if let Some(binding) = request.working_directory.as_ref() {
|
||||
worker.bind_workdir_session(Some(runtime_local_workdir_session(
|
||||
@@ -812,6 +915,8 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory {
|
||||
binding.root(),
|
||||
binding.cwd(),
|
||||
worker.scope().clone(),
|
||||
binding.command_environment(),
|
||||
binding.session_resources(),
|
||||
)));
|
||||
} else {
|
||||
worker.bind_workdir_session(None);
|
||||
@@ -848,21 +953,16 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory {
|
||||
}
|
||||
|
||||
let workspace_client = worker.workspace_client_handle();
|
||||
let run_dir = worker_aggregate_dir
|
||||
.join("runs")
|
||||
.join(request.run_generation.to_string());
|
||||
let (handle, shutdown_rx) = WorkerController::spawn_runtime_managed_run_with_transport(
|
||||
worker,
|
||||
&run_dir,
|
||||
self.controller_transport,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
format!(
|
||||
"failed to spawn Worker controller in {}: {err}",
|
||||
let started = prepared.start().await.map_err(|error| match error {
|
||||
WorkerBootstrapError::Worker(source) => {
|
||||
format!("failed to prepare Worker before controller start: {source}")
|
||||
}
|
||||
WorkerBootstrapError::Controller { source, .. } => format!(
|
||||
"failed to spawn Worker controller in {}: {source}",
|
||||
run_dir.display()
|
||||
)
|
||||
),
|
||||
})?;
|
||||
let (handle, shutdown_rx) = (started.handle, started.shutdown);
|
||||
if flow_transition_enabled {
|
||||
handle.shared_state.enable_flow_transition();
|
||||
}
|
||||
@@ -909,6 +1009,7 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory {
|
||||
&request.worker_ref,
|
||||
request.workspace_scope.as_ref(),
|
||||
self.worker_mutation_identity.as_ref(),
|
||||
self.runtime_request_audience.as_deref(),
|
||||
self.embedded_worker_mutation_dispatcher.as_ref(),
|
||||
Some(self.prompt_projection_cache.clone()),
|
||||
);
|
||||
@@ -989,6 +1090,8 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory {
|
||||
binding.root(),
|
||||
binding.cwd(),
|
||||
worker.scope().clone(),
|
||||
binding.command_environment(),
|
||||
binding.session_resources(),
|
||||
)));
|
||||
} else {
|
||||
worker.bind_workdir_session(None);
|
||||
@@ -1028,18 +1131,25 @@ impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory {
|
||||
let run_dir = worker_aggregate_dir
|
||||
.join("runs")
|
||||
.join(request.run_generation.to_string());
|
||||
let (handle, shutdown_rx) = WorkerController::spawn_runtime_managed_run_with_transport(
|
||||
let started = PreparedWorker::new(
|
||||
worker,
|
||||
&run_dir,
|
||||
WorkerBootstrapLayout::RuntimeManagedRun {
|
||||
run_dir: run_dir.clone(),
|
||||
},
|
||||
self.controller_transport,
|
||||
)
|
||||
.start()
|
||||
.await
|
||||
.map_err(|err| {
|
||||
format!(
|
||||
"failed to spawn restored Worker controller in {}: {err}",
|
||||
.map_err(|error| match error {
|
||||
WorkerBootstrapError::Worker(source) => {
|
||||
format!("failed to prepare restored Worker: {source}")
|
||||
}
|
||||
WorkerBootstrapError::Controller { source, .. } => format!(
|
||||
"failed to spawn restored Worker controller in {}: {source}",
|
||||
run_dir.display()
|
||||
)
|
||||
),
|
||||
})?;
|
||||
let (handle, shutdown_rx) = (started.handle, started.shutdown);
|
||||
if flow_transition_enabled {
|
||||
handle.shared_state.enable_flow_transition();
|
||||
}
|
||||
@@ -1353,7 +1463,6 @@ where
|
||||
let streams = subscribe_worker_protocol_session(&handle);
|
||||
let mut events = streams.events;
|
||||
let mut entry_events = streams.log_entries;
|
||||
let bridge_handle = handle.clone();
|
||||
let bridge_busy = busy.clone();
|
||||
if let Err(message) = self.spawn_on_adapter_runtime(async move {
|
||||
loop {
|
||||
@@ -1361,12 +1470,28 @@ where
|
||||
event = events.recv() => {
|
||||
match event {
|
||||
Ok(event) => {
|
||||
let next_busy = match &event {
|
||||
Event::InvokeStart { .. }
|
||||
| Event::Status {
|
||||
status: WorkerStatus::Running,
|
||||
} => Some(true),
|
||||
Event::RunEnd { .. }
|
||||
| Event::Error {
|
||||
code: ErrorCode::NotPaused,
|
||||
..
|
||||
}
|
||||
| Event::Status {
|
||||
status:
|
||||
WorkerStatus::Idle
|
||||
| WorkerStatus::Paused
|
||||
| WorkerStatus::Stopped,
|
||||
}
|
||||
| Event::Shutdown => Some(false),
|
||||
_ => None,
|
||||
};
|
||||
let _ = bridge_context.publish_protocol_event(event);
|
||||
if matches!(
|
||||
bridge_handle.shared_state.get_status(),
|
||||
WorkerStatus::Idle | WorkerStatus::Paused
|
||||
) {
|
||||
bridge_busy.store(false, Ordering::SeqCst);
|
||||
if let Some(next_busy) = next_busy {
|
||||
bridge_busy.store(next_busy, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
Err(broadcast::error::RecvError::Lagged(_)) => continue,
|
||||
@@ -1456,7 +1581,9 @@ fn accepted_notify_run_state(status: WorkerStatus, auto_run: bool) -> WorkerExec
|
||||
match status {
|
||||
WorkerStatus::Running => WorkerExecutionRunState::Busy,
|
||||
WorkerStatus::Idle if auto_run => WorkerExecutionRunState::Busy,
|
||||
WorkerStatus::Idle | WorkerStatus::Paused => WorkerExecutionRunState::Idle,
|
||||
WorkerStatus::Idle | WorkerStatus::Paused | WorkerStatus::Stopped => {
|
||||
WorkerExecutionRunState::Idle
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1500,6 +1627,19 @@ where
|
||||
Ok(materializer.create(request)?.status())
|
||||
}
|
||||
|
||||
fn authorize_working_directory_repository_access(
|
||||
&self,
|
||||
request: &WorkingDirectoryRepositoryAccessRequest,
|
||||
) -> Result<(), WorkingDirectoryDiagnostic> {
|
||||
let materializer = self.working_directory_materializer.as_ref().ok_or_else(|| {
|
||||
WorkingDirectoryDiagnostic::rejected(
|
||||
"working_directory_materializer_unavailable",
|
||||
"working directory Repository access requested, but no materializer is configured for this runtime backend",
|
||||
)
|
||||
})?;
|
||||
materializer.authorize_repository_access(request)
|
||||
}
|
||||
|
||||
fn list_working_directories(&self) -> Vec<WorkingDirectoryStatus> {
|
||||
self.working_directory_materializer
|
||||
.as_ref()
|
||||
@@ -1542,6 +1682,8 @@ where
|
||||
binding.root(),
|
||||
binding.cwd(),
|
||||
manifest::SharedScope::new(scope),
|
||||
binding.command_environment(),
|
||||
binding.session_resources(),
|
||||
))
|
||||
}
|
||||
|
||||
@@ -2060,7 +2202,7 @@ mod tests {
|
||||
use crate::identity::WorkerRef;
|
||||
use crate::management::RuntimeOptions;
|
||||
use crate::observation::WorkerObservationCursor;
|
||||
use crate::working_directory::LocalGitWorktreeMaterializer;
|
||||
use crate::working_directory::RuntimeGitCacheMaterializer;
|
||||
use agen::Engine;
|
||||
use agen::llm_client::event::{Event as LlmEvent, ResponseStatus, StatusEvent};
|
||||
use agen::llm_client::{ClientError, LlmClient, Request};
|
||||
@@ -2187,12 +2329,18 @@ mod tests {
|
||||
let scope = crate::runtime::RuntimeWorkspaceScope::new("workspace-a", "server-main");
|
||||
|
||||
let before_restart =
|
||||
backend.worker_context(&worker_ref, Some(&scope), Some(&identity), None, None);
|
||||
backend.worker_context(&worker_ref, Some(&scope), Some(&identity), None, None, None);
|
||||
let adapter = WorkerRuntimeExecutionBackend::new(FailingFactory).unwrap();
|
||||
let (after_restore_kind, after_restore_workspace_id) = adapter
|
||||
.run_on_adapter_runtime(async move {
|
||||
let after_restore =
|
||||
backend.worker_context(&worker_ref, Some(&scope), Some(&identity), None, None);
|
||||
let after_restore = backend.worker_context(
|
||||
&worker_ref,
|
||||
Some(&scope),
|
||||
Some(&identity),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
let client = after_restore.client_handle();
|
||||
Ok((
|
||||
client.kind().to_string(),
|
||||
@@ -2383,6 +2531,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
let workspace_client = workspace_context.client_handle();
|
||||
self.observed_workspace_clients.lock().unwrap().push((
|
||||
@@ -2393,7 +2542,9 @@ mod tests {
|
||||
let scope = Scope::writable(&scope_root).map_err(|err| err.to_string())?;
|
||||
let worker = Worker::new(
|
||||
manifest,
|
||||
Engine::new(self.client.clone()),
|
||||
Engine::<_, agen::state::Mutable, worker::SessionHistoryMetadata>::new_annotated(
|
||||
self.client.clone(),
|
||||
),
|
||||
store,
|
||||
workspace_context,
|
||||
filesystem_authority,
|
||||
@@ -2450,18 +2601,22 @@ mod tests {
|
||||
) {
|
||||
let deadline = std::time::Instant::now() + Duration::from_secs(5);
|
||||
loop {
|
||||
let matches = {
|
||||
let observed = {
|
||||
let workers = backend.workers.lock().unwrap();
|
||||
let execution = workers.get(worker_ref).expect("live Worker execution");
|
||||
execution.handle.shared_state.get_status() == expected_status
|
||||
&& execution.busy.load(Ordering::SeqCst) == expected_busy
|
||||
(
|
||||
execution.handle.shared_state.get_status(),
|
||||
execution.busy.load(Ordering::SeqCst),
|
||||
)
|
||||
};
|
||||
if matches {
|
||||
if observed == (expected_status, expected_busy) {
|
||||
return;
|
||||
}
|
||||
assert!(
|
||||
std::time::Instant::now() < deadline,
|
||||
"timed out waiting for adapter state {expected_status:?}, busy={expected_busy}"
|
||||
"timed out waiting for adapter state {expected_status:?}, busy={expected_busy}; last observed status={:?}, busy={}",
|
||||
observed.0,
|
||||
observed.1,
|
||||
);
|
||||
std::thread::sleep(Duration::from_millis(10));
|
||||
}
|
||||
@@ -2631,12 +2786,17 @@ mod tests {
|
||||
repository: WorkingDirectoryRepository {
|
||||
id: "repo-main".to_string(),
|
||||
provider: "git".to_string(),
|
||||
uri: ".".to_string(),
|
||||
local_path: Some(repo.to_path_buf()),
|
||||
source: workspace_api::RepositorySource {
|
||||
kind: workspace_api::RepositorySourceKind::LocalPath,
|
||||
uri: repo.display().to_string(),
|
||||
},
|
||||
source_revision: 1,
|
||||
source_fingerprint: "sha256:test".to_string(),
|
||||
selector: Some(RepositorySelector::from("HEAD")),
|
||||
},
|
||||
materializer: MaterializerKind::LocalGitWorktree,
|
||||
materializer: MaterializerKind::RuntimeGitCache,
|
||||
backend_workdir_id: None,
|
||||
materialization: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2760,12 +2920,16 @@ mod tests {
|
||||
root.path(),
|
||||
root.path(),
|
||||
manifest::SharedScope::new(Scope::writable(root.path()).unwrap()),
|
||||
Default::default(),
|
||||
Vec::new(),
|
||||
);
|
||||
let restored = runtime_local_workdir_session(
|
||||
"working-directory-42",
|
||||
root.path(),
|
||||
root.path(),
|
||||
manifest::SharedScope::new(Scope::writable(root.path()).unwrap()),
|
||||
Default::default(),
|
||||
Vec::new(),
|
||||
);
|
||||
|
||||
assert_eq!(spawned.workdir().id().as_str(), "working-directory-42");
|
||||
@@ -2781,7 +2945,7 @@ mod tests {
|
||||
archive: bundle.profile_source_archive.clone().unwrap(),
|
||||
};
|
||||
factory
|
||||
.resolve_profile_source_archive(&source)
|
||||
.resolve_profile_source_archive(&source, None)
|
||||
.await
|
||||
.expect("embedded archive should resolve without Backend resource client");
|
||||
}
|
||||
@@ -2966,9 +3130,68 @@ mod tests {
|
||||
assert!(!socket_path.exists());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn profile_runtime_factory_uses_shared_worker_bootstrap_seams() {
|
||||
let source = include_str!("worker_backend.rs");
|
||||
let production = source
|
||||
.split_once("#[cfg(test)]\nmod tests")
|
||||
.map(|(production, _)| production)
|
||||
.expect("worker backend test module marker");
|
||||
let factory = production
|
||||
.split_once("impl RuntimeWorkerFactory for ProfileRuntimeWorkerFactory")
|
||||
.map(|(_, factory)| factory)
|
||||
.expect("profile runtime factory implementation");
|
||||
let (fresh, restore) = factory
|
||||
.split_once("async fn restore_controller")
|
||||
.expect("fresh and restore factory paths");
|
||||
let assert_in_order = |path: &str, markers: &[&str]| {
|
||||
let mut offset = 0;
|
||||
for marker in markers {
|
||||
let relative = path[offset..]
|
||||
.find(marker)
|
||||
.unwrap_or_else(|| panic!("missing ordered factory marker {marker}"));
|
||||
offset += relative + marker.len();
|
||||
}
|
||||
};
|
||||
assert_in_order(
|
||||
fresh,
|
||||
&[
|
||||
"WorkerBootstrap::new(",
|
||||
".prepare()",
|
||||
"worker.bind_workdir_session(",
|
||||
"worker.bind_worker_observation_provider(",
|
||||
"install_runtime_flow_transition_feature()",
|
||||
"prepared.start()",
|
||||
],
|
||||
);
|
||||
assert_in_order(
|
||||
restore,
|
||||
&[
|
||||
"Worker::restore_from_worker_metadata_with_context(",
|
||||
"worker.bind_workdir_session(",
|
||||
"worker.bind_worker_observation_provider(",
|
||||
"install_runtime_flow_transition_feature()",
|
||||
"PreparedWorker::new(",
|
||||
".start()",
|
||||
],
|
||||
);
|
||||
assert!(
|
||||
production.contains("WorkerBootstrap::new("),
|
||||
"fresh runtime Workers must use the shared construction bootstrap"
|
||||
);
|
||||
assert!(
|
||||
production.contains("PreparedWorker::new("),
|
||||
"restored runtime Workers must use the shared pre-exposure lifecycle"
|
||||
);
|
||||
assert!(
|
||||
!production.contains("WorkerController::spawn_runtime_managed_run_with_transport"),
|
||||
"runtime factory paths must not bypass the shared controller lifecycle"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[serial_test::serial(worker_allocation)]
|
||||
fn in_process_runtime_reopens_persisted_worker_without_overlong_unix_socket() {
|
||||
fn shared_bootstrap_preserves_in_process_transport_for_fresh_and_restored_runtime_workers() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let long_component = "embedded-workspace-store-segment".repeat(4);
|
||||
let runtime_store_dir = root.path().join(long_component);
|
||||
@@ -3119,15 +3342,16 @@ mod tests {
|
||||
assert!(entries.iter().any(|entry| {
|
||||
matches!(
|
||||
entry,
|
||||
LogEntry::UserInput { segments, .. }
|
||||
LogEntry::AnnotatedUserInput { segments, .. }
|
||||
if segments == &vec![Segment::text("start the ticket")]
|
||||
)
|
||||
}));
|
||||
let submission_id = entries
|
||||
.iter()
|
||||
.find_map(|entry| {
|
||||
let LogEntry::UserInput { extensions, .. } = entry else {
|
||||
return None;
|
||||
let extensions = match entry {
|
||||
LogEntry::AnnotatedUserInput { extensions, .. } => extensions,
|
||||
_ => return None,
|
||||
};
|
||||
extensions
|
||||
.iter()
|
||||
@@ -3237,7 +3461,7 @@ mod tests {
|
||||
};
|
||||
let backend = WorkerRuntimeExecutionBackend::new(factory)
|
||||
.unwrap()
|
||||
.with_working_directory_materializer(LocalGitWorktreeMaterializer::new(
|
||||
.with_working_directory_materializer(RuntimeGitCacheMaterializer::new(
|
||||
runtime_base.path(),
|
||||
));
|
||||
let runtime =
|
||||
@@ -3392,7 +3616,7 @@ mod tests {
|
||||
};
|
||||
let backend = WorkerRuntimeExecutionBackend::new(factory)
|
||||
.unwrap()
|
||||
.with_working_directory_materializer(LocalGitWorktreeMaterializer::new(
|
||||
.with_working_directory_materializer(RuntimeGitCacheMaterializer::new(
|
||||
runtime_base.path(),
|
||||
));
|
||||
let runtime =
|
||||
@@ -3431,7 +3655,7 @@ mod tests {
|
||||
let repo = create_clean_repo();
|
||||
let backend = WorkerRuntimeExecutionBackend::new(FailingFactory)
|
||||
.unwrap()
|
||||
.with_working_directory_materializer(LocalGitWorktreeMaterializer::new(
|
||||
.with_working_directory_materializer(RuntimeGitCacheMaterializer::new(
|
||||
runtime_base.path(),
|
||||
));
|
||||
let runtime =
|
||||
@@ -3467,7 +3691,7 @@ mod tests {
|
||||
let repo = create_clean_repo();
|
||||
let backend = WorkerRuntimeExecutionBackend::new(FailingFactory)
|
||||
.unwrap()
|
||||
.with_working_directory_materializer(LocalGitWorktreeMaterializer::new(
|
||||
.with_working_directory_materializer(RuntimeGitCacheMaterializer::new(
|
||||
runtime_base.path(),
|
||||
));
|
||||
let runtime =
|
||||
@@ -3481,9 +3705,15 @@ mod tests {
|
||||
|
||||
assert!(format!("{error:?}").contains("spawn failed"));
|
||||
let working_directories_root = runtime_base.path();
|
||||
let remaining_entries = fs::read_dir(working_directories_root)
|
||||
.map(|entries| entries.count())
|
||||
let remaining_workdirs = fs::read_dir(working_directories_root)
|
||||
.map(|entries| {
|
||||
entries
|
||||
.flatten()
|
||||
.filter(|entry| !entry.file_name().to_string_lossy().starts_with('.'))
|
||||
.count()
|
||||
})
|
||||
.unwrap_or(0);
|
||||
assert_eq!(remaining_entries, 0);
|
||||
assert_eq!(remaining_workdirs, 0);
|
||||
assert!(working_directories_root.join(".repository-cache").is_dir());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,9 +7,10 @@ use worker::{
|
||||
};
|
||||
|
||||
use crate::auth::{
|
||||
RuntimeAuthError, RuntimeIdentityMaterial, RuntimeWorkerMutationSourceSigner,
|
||||
WORKER_REMOVE_PERMISSION, WorkerMutationActorKind, WorkerMutationOperation,
|
||||
WorkerMutationSourceClaims, new_token_id,
|
||||
RUNTIME_REQUEST_SOURCE_PROOF_HEADER, RuntimeAuthError, RuntimeIdentityMaterial,
|
||||
RuntimeRequestSourceSigner, RuntimeWorkerMutationSourceSigner, WORKER_REMOVE_PERMISSION,
|
||||
WORKSPACE_REQUEST_PERMISSION, WORKSPACE_WORKER_DISCOVERY_PERMISSION, WorkerMutationActorKind,
|
||||
WorkerMutationOperation, WorkerMutationSourceClaims, new_token_id,
|
||||
};
|
||||
use crate::runtime::RuntimeWorkspaceScope;
|
||||
use crate::worker_backend::WorkspacePromptProjectionCache;
|
||||
@@ -289,6 +290,8 @@ pub struct RuntimeOwnedWorkspaceClient {
|
||||
worker_id: String,
|
||||
request_timeout: Option<Duration>,
|
||||
worker_remove: Option<RuntimeWorkerMutationForwarder>,
|
||||
request_source_signer: Option<RuntimeRequestSourceSigner>,
|
||||
request_source_audience: Option<String>,
|
||||
prompt_projection_cache: Option<Arc<WorkspacePromptProjectionCache>>,
|
||||
}
|
||||
|
||||
@@ -306,6 +309,8 @@ impl RuntimeOwnedWorkspaceClient {
|
||||
worker_id: worker_id.into(),
|
||||
request_timeout: None,
|
||||
worker_remove: None,
|
||||
request_source_signer: None,
|
||||
request_source_audience: None,
|
||||
prompt_projection_cache: None,
|
||||
}
|
||||
}
|
||||
@@ -315,6 +320,16 @@ impl RuntimeOwnedWorkspaceClient {
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_runtime_request_source(
|
||||
mut self,
|
||||
identity: &RuntimeIdentityMaterial,
|
||||
audience: impl Into<String>,
|
||||
) -> Self {
|
||||
self.request_source_signer = Some(RuntimeRequestSourceSigner::from_identity(identity));
|
||||
self.request_source_audience = Some(audience.into());
|
||||
self
|
||||
}
|
||||
|
||||
pub(crate) fn with_prompt_projection_cache(
|
||||
mut self,
|
||||
cache: Arc<WorkspacePromptProjectionCache>,
|
||||
@@ -328,6 +343,51 @@ impl RuntimeOwnedWorkspaceClient {
|
||||
self.request_timeout = request_timeout;
|
||||
self
|
||||
}
|
||||
|
||||
fn execute_with_permission(
|
||||
&self,
|
||||
request: WorkspaceRequest,
|
||||
permission: &'static str,
|
||||
) -> Result<WorkspaceResponse, WorkspaceClientError> {
|
||||
let base_url = self.base_url.clone();
|
||||
let workspace_id = self.workspace_id.clone();
|
||||
let runtime_id = self.runtime_id.clone();
|
||||
let worker_id = self.worker_id.clone();
|
||||
let request_source_signer = self.request_source_signer.clone();
|
||||
let request_source_audience = self.request_source_audience.clone();
|
||||
let request_timeout = self.request_timeout;
|
||||
if tokio::runtime::Handle::try_current().is_ok() {
|
||||
std::thread::spawn(move || {
|
||||
execute_runtime_owned_workspace_http(
|
||||
&base_url,
|
||||
&workspace_id,
|
||||
&runtime_id,
|
||||
&worker_id,
|
||||
request_source_signer.as_ref(),
|
||||
request_source_audience.as_deref(),
|
||||
request_timeout,
|
||||
permission,
|
||||
request,
|
||||
)
|
||||
})
|
||||
.join()
|
||||
.map_err(|_| {
|
||||
WorkspaceClientError::Request("workspace request thread panicked".to_string())
|
||||
})?
|
||||
} else {
|
||||
execute_runtime_owned_workspace_http(
|
||||
&self.base_url,
|
||||
&self.workspace_id,
|
||||
&self.runtime_id,
|
||||
&self.worker_id,
|
||||
self.request_source_signer.as_ref(),
|
||||
self.request_source_audience.as_deref(),
|
||||
self.request_timeout,
|
||||
permission,
|
||||
request,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for RuntimeOwnedWorkspaceClient {
|
||||
@@ -362,33 +422,40 @@ impl WorkspaceClient for RuntimeOwnedWorkspaceClient {
|
||||
&self,
|
||||
request: WorkspaceRequest,
|
||||
) -> Result<WorkspaceResponse, WorkspaceClientError> {
|
||||
let base_url = self.base_url.clone();
|
||||
let runtime_id = self.runtime_id.clone();
|
||||
let worker_id = self.worker_id.clone();
|
||||
let request_timeout = self.request_timeout;
|
||||
if tokio::runtime::Handle::try_current().is_ok() {
|
||||
std::thread::spawn(move || {
|
||||
execute_runtime_owned_workspace_http(
|
||||
&base_url,
|
||||
&runtime_id,
|
||||
&worker_id,
|
||||
request_timeout,
|
||||
request,
|
||||
)
|
||||
})
|
||||
.join()
|
||||
.map_err(|_| {
|
||||
WorkspaceClientError::Request("workspace request thread panicked".to_string())
|
||||
})?
|
||||
} else {
|
||||
execute_runtime_owned_workspace_http(
|
||||
&self.base_url,
|
||||
&self.runtime_id,
|
||||
&self.worker_id,
|
||||
self.request_timeout,
|
||||
request,
|
||||
)
|
||||
self.execute_with_permission(request, WORKSPACE_REQUEST_PERMISSION)
|
||||
}
|
||||
|
||||
fn list_workspace_workers(
|
||||
&self,
|
||||
request: worker::WorkspaceWorkerDiscoveryRequest,
|
||||
) -> Result<workspace_api::WorkspaceWorkerDiscoveryPage, WorkspaceClientError> {
|
||||
let mut path = format!(
|
||||
"/api/w/{}/worker-discovery/workers?limit={}",
|
||||
self.workspace_id, request.limit
|
||||
);
|
||||
if let Some(cursor) = request.cursor.as_deref() {
|
||||
path.push_str("&cursor=");
|
||||
path.push_str(&percent_encode_query(cursor));
|
||||
}
|
||||
if let Some(query) = request.query.as_deref() {
|
||||
path.push_str("&query=");
|
||||
path.push_str(&percent_encode_query(query));
|
||||
}
|
||||
let response = self.execute_with_permission(
|
||||
WorkspaceRequest::get(path),
|
||||
WORKSPACE_WORKER_DISCOVERY_PERMISSION,
|
||||
)?;
|
||||
if !(200..300).contains(&response.status) {
|
||||
return Err(WorkspaceClientError::Request(format!(
|
||||
"Workspace Worker discovery failed with HTTP {}: {}",
|
||||
response.status, response.body
|
||||
)));
|
||||
}
|
||||
serde_json::from_str(&response.body).map_err(|error| {
|
||||
WorkspaceClientError::Request(format!(
|
||||
"invalid Workspace Worker discovery response: {error}"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn current_prompt_projection(
|
||||
@@ -482,11 +549,28 @@ impl WorkspaceClient for RuntimeOwnedWorkspaceClient {
|
||||
}
|
||||
}
|
||||
|
||||
fn percent_encode_query(value: &str) -> String {
|
||||
let mut encoded = String::with_capacity(value.len());
|
||||
for byte in value.bytes() {
|
||||
if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') {
|
||||
encoded.push(char::from(byte));
|
||||
} else {
|
||||
use std::fmt::Write as _;
|
||||
let _ = write!(encoded, "%{byte:02X}");
|
||||
}
|
||||
}
|
||||
encoded
|
||||
}
|
||||
|
||||
fn execute_runtime_owned_workspace_http(
|
||||
base_url: &str,
|
||||
workspace_id: &str,
|
||||
runtime_id: &str,
|
||||
worker_id: &str,
|
||||
request_source_signer: Option<&RuntimeRequestSourceSigner>,
|
||||
request_source_audience: Option<&str>,
|
||||
request_timeout: Option<Duration>,
|
||||
permission: &'static str,
|
||||
request: WorkspaceRequest,
|
||||
) -> Result<WorkspaceResponse, WorkspaceClientError> {
|
||||
if !request.path.starts_with('/') || request.path.starts_with("//") {
|
||||
@@ -510,11 +594,33 @@ fn execute_runtime_owned_workspace_http(
|
||||
))
|
||||
})?;
|
||||
let request_label = format!("{method} {}", request.path);
|
||||
let body = request.body.unwrap_or_default();
|
||||
let mut request_builder = client
|
||||
.request(method, url)
|
||||
.request(method.clone(), url)
|
||||
.header("x-yoi-runtime-id", runtime_id)
|
||||
.header("x-yoi-worker-id", worker_id);
|
||||
if let Some(body) = request.body {
|
||||
if let Some(signer) = request_source_signer {
|
||||
let audience = request_source_audience.ok_or_else(|| {
|
||||
WorkspaceClientError::Request(
|
||||
"runtime request proof audience is unavailable".to_owned(),
|
||||
)
|
||||
})?;
|
||||
let proof = signer
|
||||
.issue(
|
||||
audience,
|
||||
workspace_id,
|
||||
Some(worker_id),
|
||||
permission,
|
||||
method.as_str(),
|
||||
&request.path,
|
||||
body.as_bytes(),
|
||||
i64::try_from(unix_now_seconds()).unwrap_or(i64::MAX),
|
||||
30,
|
||||
)
|
||||
.map_err(|error| WorkspaceClientError::Request(error.to_string()))?;
|
||||
request_builder = request_builder.header(RUNTIME_REQUEST_SOURCE_PROOF_HEADER, proof);
|
||||
}
|
||||
if !body.is_empty() {
|
||||
request_builder = request_builder
|
||||
.header(reqwest::header::CONTENT_TYPE, "application/json")
|
||||
.body(body);
|
||||
@@ -590,8 +696,8 @@ fn unix_now_seconds() -> u64 {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::auth::{
|
||||
WorkerMutationSourceExpectation, decode_worker_mutation_source_claims,
|
||||
verify_worker_mutation_source_proof,
|
||||
WorkerMutationSourceExpectation, decode_runtime_request_source_claims,
|
||||
decode_worker_mutation_source_claims, verify_worker_mutation_source_proof,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -797,7 +903,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ordinary_workspace_forwarding_stamps_legacy_source_only_inside_runtime() {
|
||||
fn ordinary_workspace_forwarding_stamps_runtime_identity_and_signs_path_and_query() {
|
||||
use std::io::{Read, Write};
|
||||
use std::net::TcpListener;
|
||||
use std::sync::Mutex;
|
||||
@@ -817,21 +923,120 @@ mod tests {
|
||||
.unwrap();
|
||||
});
|
||||
|
||||
let identity = RuntimeIdentityMaterial::generate("runtime-a").unwrap();
|
||||
let client = RuntimeOwnedWorkspaceClient::new(
|
||||
"workspace-a",
|
||||
format!("http://{address}"),
|
||||
"runtime-a",
|
||||
"worker-a",
|
||||
);
|
||||
)
|
||||
.with_runtime_request_source(&identity, "server-a");
|
||||
let response = client
|
||||
.execute(WorkspaceRequest::get("/api/w/workspace-a/tickets/search"))
|
||||
.execute(WorkspaceRequest::get(
|
||||
"/api/w/workspace-a/tickets/search?state=planning&limit=20",
|
||||
))
|
||||
.unwrap();
|
||||
assert_eq!(response.status, 200);
|
||||
server.join().unwrap();
|
||||
let request = received.lock().unwrap().to_ascii_lowercase();
|
||||
assert!(request.contains("x-yoi-runtime-id: runtime-a"));
|
||||
assert!(request.contains("x-yoi-worker-id: worker-a"));
|
||||
assert!(!request.contains("authorization:"));
|
||||
let request = received.lock().unwrap().clone();
|
||||
let lowercase_request = request.to_ascii_lowercase();
|
||||
assert!(lowercase_request.contains("x-yoi-runtime-id: runtime-a"));
|
||||
assert!(lowercase_request.contains("x-yoi-worker-id: worker-a"));
|
||||
assert!(lowercase_request.contains("x-yoi-runtime-request-proof: yoi-runtime-request-v1."));
|
||||
assert!(!lowercase_request.contains("authorization:"));
|
||||
let token = request
|
||||
.lines()
|
||||
.find_map(|line| {
|
||||
line.split_once(':').and_then(|(name, value)| {
|
||||
name.eq_ignore_ascii_case(RUNTIME_REQUEST_SOURCE_PROOF_HEADER)
|
||||
.then(|| value.trim())
|
||||
})
|
||||
})
|
||||
.expect("runtime proof header");
|
||||
let claims = decode_runtime_request_source_claims(token).unwrap();
|
||||
assert_eq!(
|
||||
claims.path,
|
||||
"/api/w/workspace-a/tickets/search?state=planning&limit=20"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn workspace_worker_discovery_signs_dedicated_permission_and_encoded_query() {
|
||||
use std::io::{Read, Write};
|
||||
use std::net::TcpListener;
|
||||
use std::sync::Mutex;
|
||||
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let received = Arc::new(Mutex::new(String::new()));
|
||||
let received_for_server = received.clone();
|
||||
let body = serde_json::json!({
|
||||
"workers": [{
|
||||
"subject": {
|
||||
"kind": "runtime_worker",
|
||||
"runtime_id": "runtime-b",
|
||||
"worker_id": "worker-b"
|
||||
},
|
||||
"resource_key": "W-2",
|
||||
"display_name": "coder two",
|
||||
"profile": "builtin:coder",
|
||||
"status": "idle"
|
||||
}],
|
||||
"next_cursor": "v1:1"
|
||||
})
|
||||
.to_string();
|
||||
let server = std::thread::spawn(move || {
|
||||
let (mut stream, _) = listener.accept().unwrap();
|
||||
let mut bytes = [0_u8; 4096];
|
||||
let count = stream.read(&mut bytes).unwrap();
|
||||
*received_for_server.lock().unwrap() =
|
||||
String::from_utf8_lossy(&bytes[..count]).into_owned();
|
||||
write!(
|
||||
stream,
|
||||
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}",
|
||||
body.len(),
|
||||
body
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
|
||||
let identity = RuntimeIdentityMaterial::generate("runtime-a").unwrap();
|
||||
let client = RuntimeOwnedWorkspaceClient::new(
|
||||
"workspace-a",
|
||||
format!("http://{address}"),
|
||||
"runtime-a",
|
||||
"worker-a",
|
||||
)
|
||||
.with_runtime_request_source(&identity, "server-a");
|
||||
let page = client
|
||||
.list_workspace_workers(worker::WorkspaceWorkerDiscoveryRequest {
|
||||
cursor: Some("v1:0".to_string()),
|
||||
limit: 1,
|
||||
query: Some("coder two".to_string()),
|
||||
})
|
||||
.unwrap();
|
||||
assert_eq!(page.workers[0].resource_key, "W-2");
|
||||
server.join().unwrap();
|
||||
|
||||
let request = received.lock().unwrap().clone();
|
||||
assert!(request.contains(
|
||||
"GET /api/w/workspace-a/worker-discovery/workers?limit=1&cursor=v1%3A0&query=coder%20two "
|
||||
));
|
||||
let token = request
|
||||
.lines()
|
||||
.find_map(|line| {
|
||||
line.split_once(':').and_then(|(name, value)| {
|
||||
name.eq_ignore_ascii_case(RUNTIME_REQUEST_SOURCE_PROOF_HEADER)
|
||||
.then(|| value.trim())
|
||||
})
|
||||
})
|
||||
.unwrap();
|
||||
let claims = decode_runtime_request_source_claims(token).unwrap();
|
||||
assert_eq!(claims.permission, WORKSPACE_WORKER_DISCOVERY_PERMISSION);
|
||||
assert_eq!(
|
||||
claims.path,
|
||||
"/api/w/workspace-a/worker-discovery/workers?limit=1&cursor=v1%3A0&query=coder%20two"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user