mirror of
https://github.com/RightNow-AI/openfang.git
synced 2026-08-14 08:52:02 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
acf2587e46 | ||
|
|
4583157b49 | ||
|
|
7185ea8808 | ||
|
|
4aa1508f54 | ||
|
|
5447bf7f1d | ||
|
|
b77ebfb897 | ||
|
|
2323cd5e67 | ||
|
|
df29e5e8b9 | ||
|
|
8b411c21c4 | ||
|
|
b00af5eddd | ||
|
|
36177425c4 | ||
|
|
6bed6c04ff | ||
|
|
4c496be02f | ||
|
|
7f7b071528 | ||
|
|
2d1fb8171c | ||
|
|
e683acc565 | ||
|
|
5396889ff1 | ||
|
|
c9e8d31571 | ||
|
|
d8ad91572e | ||
|
|
c27bfebd17 | ||
|
|
f05ba5e42f | ||
|
|
d9e72abb4b | ||
|
|
505a8e8080 | ||
|
|
5cc865e6e6 | ||
|
|
6a1ce40d86 | ||
|
|
fbb7936234 | ||
|
|
569e76c79a | ||
|
|
88ad029999 | ||
|
|
838836b29c | ||
|
|
8564872181 | ||
|
|
efbefa1682 | ||
|
|
bdcd440cd6 | ||
|
|
25516c7f87 | ||
|
|
ae2706bdab | ||
|
|
32299bb506 | ||
|
|
6f8463fc91 | ||
|
|
6b03cb2e9d | ||
|
|
8cb7541678 | ||
|
|
6b5b7674d3 | ||
|
|
247dca508b | ||
|
|
90d16e52be | ||
|
|
68bde60fac | ||
|
|
a422058049 | ||
|
|
6ba0bfb7ef | ||
|
|
7699b86037 | ||
|
|
e31216d5ec | ||
|
|
538e943d3d | ||
|
|
94fca22124 | ||
|
|
37e2043ed7 | ||
|
|
31eb833cdf | ||
|
|
15da248faf | ||
|
|
c27a6f3609 | ||
|
|
f792f1a14b | ||
|
|
5e228336e4 | ||
|
|
8b10930e40 | ||
|
|
5c1b1508a2 | ||
|
|
701fcd8e2e | ||
|
|
118eacea64 | ||
|
|
aaad1fdf32 | ||
|
|
dd8c53026e | ||
|
|
218f2dba1f | ||
|
|
24aca4e31d | ||
|
|
3cce1eb3fb | ||
|
|
c89958b66d | ||
|
|
b0a92456bf | ||
|
|
67bbcc623d | ||
|
|
a91bfc0e9c | ||
|
|
948117d5de | ||
|
|
2dedab2a8b | ||
|
|
8642c4d442 | ||
|
|
46a6eb33d9 | ||
|
|
99b4ce2931 | ||
|
|
87932f5da0 | ||
|
|
9130811433 | ||
|
|
325734c6aa | ||
|
|
faf2cf9211 | ||
|
|
15ed29c667 | ||
|
|
1d1bf0fb09 | ||
|
|
d3363142b2 | ||
|
|
76929a41aa | ||
|
|
fe34a37e6f | ||
|
|
4b63eb18cc | ||
|
|
fe21d4b4df | ||
|
|
f52bc53e47 | ||
|
|
10f7ee1885 | ||
|
|
7bc6591338 | ||
|
|
53f2066945 | ||
|
|
c69dd84184 | ||
|
|
c1356fc95d | ||
|
|
aabf83b351 | ||
|
|
da6b567ac3 | ||
|
|
ccdd7943a2 | ||
|
|
7fe87babe6 | ||
|
|
9c0e1637a5 | ||
|
|
fd450dbfc2 | ||
|
|
81176dc626 | ||
|
|
ef9096f7c5 | ||
|
|
fbb5bb1ae9 | ||
|
|
c435a6adcd | ||
|
|
f67c4e8754 | ||
|
|
3b237ac526 | ||
|
|
96c572df32 | ||
|
|
17e0d519ca | ||
|
|
37c233d489 | ||
|
|
92f7e996de | ||
|
|
b1c4061247 | ||
|
|
79aa34c77a | ||
|
|
4ae2961b1c | ||
|
|
6a90aa08df | ||
|
|
356500bb1e | ||
|
|
40bd7e2c11 | ||
|
|
a7197d7b97 | ||
|
|
bc26d5e8c3 | ||
|
|
84d90ad342 | ||
|
|
9fee63d58c | ||
|
|
40903cceee | ||
|
|
5a86141677 | ||
|
|
e97eb6fff3 | ||
|
|
93b57bdd52 | ||
|
|
0227ff1790 | ||
|
|
525d7d844a | ||
|
|
f2587995a2 | ||
|
|
45e7ea7948 | ||
|
|
de8a692036 | ||
|
|
8d3d77dd99 | ||
|
|
c9701627a9 | ||
|
|
643a22b295 |
@@ -91,7 +91,9 @@ jobs:
|
||||
- uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
components: rustfmt
|
||||
- run: cargo fmt --check
|
||||
# Gate every workspace crate on rustfmt to keep `cargo fmt --all --check` clean.
|
||||
# See issue #1121.
|
||||
- run: cargo fmt --all -- --check
|
||||
|
||||
audit:
|
||||
name: Security Audit
|
||||
|
||||
@@ -204,7 +204,7 @@ jobs:
|
||||
$hash = (Get-FileHash "openfang-${{ matrix.target }}.zip" -Algorithm SHA256).Hash.ToLower()
|
||||
"$hash openfang-${{ matrix.target }}.zip" | Out-File -Encoding ASCII "openfang-${{ matrix.target }}.zip.sha256"
|
||||
- name: Upload to GitHub Release
|
||||
uses: softprops/action-gh-release@v2
|
||||
uses: softprops/action-gh-release@v3
|
||||
with:
|
||||
files: openfang-${{ matrix.target }}.*
|
||||
env:
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
## Health Stack
|
||||
|
||||
- typecheck: cargo build --workspace --lib
|
||||
- lint: cargo clippy --workspace --all-targets -- -D warnings
|
||||
- test: cargo test --workspace
|
||||
- shell: shellcheck scripts/install.sh
|
||||
Generated
+95
-90
@@ -1029,27 +1029,27 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-assembler-x64"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "046d4b584c3bb9b5eb500c8f29549bec36be11000f1ba2a927cef3d1a9875691"
|
||||
checksum = "adc822414b18d1f5b1b33ce1441534e311e62fef86ebb5b9d382af857d0272c9"
|
||||
dependencies = [
|
||||
"cranelift-assembler-x64-meta",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-assembler-x64-meta"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b9b194a7870becb1490366fc0ae392ccd188065ff35f8391e77ac659db6fb977"
|
||||
checksum = "8c646808b06f4532478d8d6057d74f15c3322f10d995d9486e7dcea405bf521a"
|
||||
dependencies = [
|
||||
"cranelift-srcgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-bforest"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bb6a4ab44c6b371e661846b97dab687387a60ac4e2f864e2d4257284aad9e889"
|
||||
checksum = "7b5996f01a686b2349cdb379083ec5ad3e8cb8767fb2d495d3a4f2ee4163a18d"
|
||||
dependencies = [
|
||||
"cranelift-entity",
|
||||
"wasmtime-internal-core",
|
||||
@@ -1057,9 +1057,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-bitset"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b8b7a44150c2f471a94023482bda1902710746e4bed9f9973d60c5a94319b06d"
|
||||
checksum = "523fea83273f6a985520f57788809a4de2165794d9ab00fb1254fceb4f5aa00c"
|
||||
dependencies = [
|
||||
"serde",
|
||||
"serde_derive",
|
||||
@@ -1068,9 +1068,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-codegen"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "01b06598133b1dd76758b8b95f8d6747c124124aade50cea96a3d88b962da9fa"
|
||||
checksum = "d73d1e372730b5f64ed1a2bd9f01fe4686c8ec14a28034e3084e530c8d951878"
|
||||
dependencies = [
|
||||
"bumpalo",
|
||||
"cranelift-assembler-x64",
|
||||
@@ -1096,9 +1096,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-codegen-meta"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6190e2e7bcf0a678da2f715363d34ed530fedf7a2f0ab75edaefef72a70465ff"
|
||||
checksum = "b0319c18165e93dc1ebf78946a8da0b1c341c95b4a39729a69574671639bdb5f"
|
||||
dependencies = [
|
||||
"cranelift-assembler-x64-meta",
|
||||
"cranelift-codegen-shared",
|
||||
@@ -1109,24 +1109,24 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-codegen-shared"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f583cf203d1aa8b79560e3b01f929bdacf9070b015eec4ea9c46e22a3f83e4a0"
|
||||
checksum = "9195cd8aeecb55e401aa96b2eaa55921636e8246c127ed7908f7ef7e0d40f270"
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-control"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "803159df35cc398ae54473c150b16d6c77e92ab2948be638488de126a3328fbc"
|
||||
checksum = "8976c2154b74136322befc74222ab5c7249edd7e2604f8cbef2b94975541ffb9"
|
||||
dependencies = [
|
||||
"arbitrary",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-entity"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3109e417257082d88087f5bcce677525bdaa8322b88dd7f175ed1a1fd41d546c"
|
||||
checksum = "6038b3147c7982f4951150d5f96c7c06c1e7214b99d4b4a98607aadf8ded89d1"
|
||||
dependencies = [
|
||||
"cranelift-bitset",
|
||||
"serde",
|
||||
@@ -1136,9 +1136,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-frontend"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "14db6b0e0e4994c581092df78d837be2072578f7cb2528f96a6cf895e56dee63"
|
||||
checksum = "4cbd294abe236e23cc3d907b0936226b6a8342db7636daa9c7c72be1e323420e"
|
||||
dependencies = [
|
||||
"cranelift-codegen",
|
||||
"log",
|
||||
@@ -1148,15 +1148,15 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-isle"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ec66ea5025c7317383699778282ac98741d68444f956e3b1d7b62f12b7216e67"
|
||||
checksum = "b5a90b6ed3aba84189352a87badeb93b2126d3724225a42dc67fdce53d1b139c"
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-native"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "373ade56438e6232619d85678477d0a88a31b3581936e0503e61e96b546b0800"
|
||||
checksum = "c3ec0cc1a54e22925eacf4fc3dc815f907734d3b377899d19d52bec04863e853"
|
||||
dependencies = [
|
||||
"cranelift-codegen",
|
||||
"libc",
|
||||
@@ -1165,9 +1165,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "cranelift-srcgen"
|
||||
version = "0.130.1"
|
||||
version = "0.130.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ef53619d3cd5c78fd998c6d9420547af26b72e6456f94c2a8a2334cb76b42baa"
|
||||
checksum = "948865622f87f30907bb46fbb081b235ae63c1896a99a83c26a003305c1fa82d"
|
||||
|
||||
[[package]]
|
||||
name = "crc32fast"
|
||||
@@ -2743,7 +2743,7 @@ dependencies = [
|
||||
"libc",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"socket2 0.5.10",
|
||||
"socket2 0.6.3",
|
||||
"tokio",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
@@ -3255,9 +3255,9 @@ checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2"
|
||||
|
||||
[[package]]
|
||||
name = "lettre"
|
||||
version = "0.11.20"
|
||||
version = "0.11.21"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "471816f3e24b85e820dee02cde962379ea1a669e5242f19c61bcbcffedf4c4fb"
|
||||
checksum = "dabda5859ee7c06b995b9d1165aa52c39110e079ef609db97178d86aeb051fa7"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"base64 0.22.1",
|
||||
@@ -3320,9 +3320,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "libc"
|
||||
version = "0.2.183"
|
||||
version = "0.2.185"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d"
|
||||
checksum = "52ff2c0fe9bc6cb6b14a0592c2ff4fa9ceb83eea9db979b0487cd054946a2b8f"
|
||||
|
||||
[[package]]
|
||||
name = "libloading"
|
||||
@@ -3947,9 +3947,9 @@ checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381"
|
||||
|
||||
[[package]]
|
||||
name = "open"
|
||||
version = "5.3.3"
|
||||
version = "5.3.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "43bb73a7fa3799b198970490a51174027ba0d4ec504b03cd08caf513d40024bc"
|
||||
checksum = "9f3bab717c29a857abf75fcef718d441ec7cb2725f937343c734740a985d37fd"
|
||||
dependencies = [
|
||||
"dunce",
|
||||
"is-wsl",
|
||||
@@ -3959,7 +3959,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-api"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"argon2",
|
||||
"async-trait",
|
||||
@@ -4001,7 +4001,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-channels"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"aes",
|
||||
"async-trait",
|
||||
@@ -4040,7 +4040,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-cli"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"clap",
|
||||
"clap_complete",
|
||||
@@ -4068,7 +4068,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-desktop"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"open",
|
||||
@@ -4094,7 +4094,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-extensions"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"argon2",
|
||||
@@ -4122,14 +4122,16 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-hands"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"chrono",
|
||||
"dashmap",
|
||||
"dirs 6.0.0",
|
||||
"hex",
|
||||
"openfang-types",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
"tokio-test",
|
||||
@@ -4140,7 +4142,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-kernel"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"chrono",
|
||||
@@ -4165,6 +4167,7 @@ dependencies = [
|
||||
"rustls",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2",
|
||||
"subtle",
|
||||
"tempfile",
|
||||
"thiserror 2.0.18",
|
||||
@@ -4179,7 +4182,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-memory"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"chrono",
|
||||
@@ -4199,7 +4202,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-migrate"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"chrono",
|
||||
"dirs 6.0.0",
|
||||
@@ -4218,7 +4221,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-runtime"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"async-trait",
|
||||
@@ -4254,11 +4257,13 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-skills"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"chrono",
|
||||
"ed25519-dalek",
|
||||
"hex",
|
||||
"openfang-types",
|
||||
"rand 0.8.5",
|
||||
"reqwest 0.12.28",
|
||||
"serde",
|
||||
"serde_json",
|
||||
@@ -4277,7 +4282,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-types"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"bitflags 2.11.0",
|
||||
@@ -4297,7 +4302,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openfang-wire"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
dependencies = [
|
||||
"async-trait",
|
||||
"chrono",
|
||||
@@ -5019,7 +5024,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "27c6023962132f4b30eb4c172c91ce92d933da334c59c23cddee82358ddafb0b"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"itertools 0.14.0",
|
||||
"itertools 0.13.0",
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
@@ -5027,9 +5032,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "pulley-interpreter"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "010dec3755eb61b2f1051ecb3611b718460b7a74c131e474de2af20a845938af"
|
||||
checksum = "7ec12fe19a9588315a49fe5704502a9c02d6a198303314b0c7c86123b06d29e5"
|
||||
dependencies = [
|
||||
"cranelift-bitset",
|
||||
"log",
|
||||
@@ -5039,9 +5044,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "pulley-macros"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ad360c32e85ca4b083ac0e2b6856e8f11c3d5060dafa7d5dc57b370857fa3018"
|
||||
checksum = "36f7d5ef31ebf1b46cd7e722ffef934e670d7e462f49aa01cde07b9b76dca580"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
@@ -5100,7 +5105,7 @@ dependencies = [
|
||||
"quinn-udp",
|
||||
"rustc-hash",
|
||||
"rustls",
|
||||
"socket2 0.5.10",
|
||||
"socket2 0.6.3",
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"tracing",
|
||||
@@ -5138,7 +5143,7 @@ dependencies = [
|
||||
"cfg_aliases",
|
||||
"libc",
|
||||
"once_cell",
|
||||
"socket2 0.5.10",
|
||||
"socket2 0.6.3",
|
||||
"tracing",
|
||||
"windows-sys 0.60.2",
|
||||
]
|
||||
@@ -5709,9 +5714,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "rustls"
|
||||
version = "0.23.37"
|
||||
version = "0.23.39"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "758025cb5fccfd3bc2fd74708fd4682be41d99e5dff73c377c0646c6012c73a4"
|
||||
checksum = "7c2c118cb077cca2822033836dfb1b975355dfb784b5e8da48f7b6c5db74e60e"
|
||||
dependencies = [
|
||||
"aws-lc-rs",
|
||||
"log",
|
||||
@@ -5774,9 +5779,9 @@ checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f"
|
||||
|
||||
[[package]]
|
||||
name = "rustls-webpki"
|
||||
version = "0.103.10"
|
||||
version = "0.103.13"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "df33b2b81ac578cabaf06b89b0631153a3f416b0a886e8a7a1707fb51abbd1ef"
|
||||
checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e"
|
||||
dependencies = [
|
||||
"aws-lc-rs",
|
||||
"ring",
|
||||
@@ -7781,9 +7786,9 @@ checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821"
|
||||
|
||||
[[package]]
|
||||
name = "uuid"
|
||||
version = "1.23.0"
|
||||
version = "1.23.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5ac8b6f42ead25368cf5b098aeb3dc8a1a2c05a3eee8a9a1a68c640edbfc79d9"
|
||||
checksum = "ddd74a9687298c6858e9b88ec8935ec45d22e8fd5e6394fa1bd4e99a87789c76"
|
||||
dependencies = [
|
||||
"getrandom 0.4.2",
|
||||
"js-sys",
|
||||
@@ -8061,9 +8066,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ce205cd643d661b5ba5ba4717e13730262e8cdbc8f2eacbc7b906d45c1a74026"
|
||||
checksum = "efb1ed5899dde98357cfdcf647a4614498798719793898245b4b34e663addabf"
|
||||
dependencies = [
|
||||
"addr2line",
|
||||
"async-trait",
|
||||
@@ -8114,9 +8119,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-environ"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0b8b78abf3677d4a0a5db82e5015b4d085ff3a1b8b472cbb8c70d4b769f019ce"
|
||||
checksum = "4172382dcc785c31d0e862c6780a18f5dd437914d22c4691351f965ef751c821"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"cpp_demangle",
|
||||
@@ -8145,9 +8150,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-cache"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8e4fd4103ba413c0da2e636f73490c6c8e446d708cbde7573703941bc3d6a448"
|
||||
checksum = "4ed398988226d7aa0505ac6bb576e09532ad722d702ec4e66365d78ed695c95f"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"directories-next",
|
||||
@@ -8165,9 +8170,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-component-macro"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0d3d6914f34be2f9d78d8ee9f422e834dfc204e71ccce697205fae95fed87892"
|
||||
checksum = "ae5ec9fff073ff13b81732d56a9515d761c245750bcda09093827f84130ebc25"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"proc-macro2",
|
||||
@@ -8180,15 +8185,15 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-component-util"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3751b0616b914fdd87fe1bf804694a078f321b000338e6476bc48a4d6e454f21"
|
||||
checksum = "935d9ab293ba27d1ec9aa7bc1b3a43993dbe961af2a8f23f90a11e1331b4c13f"
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-core"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "22632b187e1b0716f1b9ac57ad29013bed33175fcb19e10bb6896126f82fac67"
|
||||
checksum = "9a3820b174f477d2a7083209d1ad5353fcdb11eaea434b2137b8681029460dd3"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"hashbrown 0.16.1",
|
||||
@@ -8198,9 +8203,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-cranelift"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8b3ca07b3e0bb3429674b173b5800577719d600774dd81bff58f775c0aaa64ee"
|
||||
checksum = "d1679d205caf9766c6aa309d45bb3e7c634d7725e3164404df33824b9f7c4fb7"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cranelift-codegen",
|
||||
@@ -8225,9 +8230,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-fiber"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "20c8b2c9704eb1f33ead025ec16038277ccb63d0a14c31e99d5b765d7c36da55"
|
||||
checksum = "f1e505254058be5b0df458d670ee42d9eafe2349d04c1296e9dc01071dc20a85"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"cfg-if",
|
||||
@@ -8240,9 +8245,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-jit-debug"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d950310d07391d34369f62c48336ebb14eacbd4d6f772bb5f349c24e838e0664"
|
||||
checksum = "1c2e05b345f1773e59c20e6ad7298fd6857cdea245023d88bb659c96d8f0ea72"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"object",
|
||||
@@ -8252,9 +8257,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-jit-icache-coherence"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3606662c156962d096be3127b8b8ae8ee2f8be3f896dad29259ff01ddb64abfd"
|
||||
checksum = "b86701b234a4643e3f111869aa792b3a05a06e02d486ee9cb6c04dae16b52dab"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"libc",
|
||||
@@ -8264,9 +8269,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-unwinder"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "75eef0747e52dc545b075f64fd0e0cc237ae738e641266b1970e07e2d744bc32"
|
||||
checksum = "f63558d801beb83dde9b336eb4ae049019aee26627926edb32cd119d7e4c83cd"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cranelift-codegen",
|
||||
@@ -8277,9 +8282,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-versioned-export-macros"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d8b0a5dab02a8fb527f547855ecc0e05f9fdc3d5bd57b8b080349408f9a6cece"
|
||||
checksum = "737c4d956fc3a848541a064afb683dd2771132a6b125be5baaf95c4379aa47df"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
@@ -8288,9 +8293,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-winch"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8007342bd12ff400293a817973f7ecd6f1d9a8549a53369a9c1af357166f1f1e"
|
||||
checksum = "f599b79545e3bba0b7913406055ebede5bb0dabee9ba2015ef25a9f4c9f47807"
|
||||
dependencies = [
|
||||
"cranelift-codegen",
|
||||
"gimli",
|
||||
@@ -8305,9 +8310,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "wasmtime-internal-wit-bindgen"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7900c3e3c1d6e475bc225d73b02d6d5484815f260022e6964dca9558e50dd01a"
|
||||
checksum = "2192a77a00b9a67800c2b4e1c70fb6abca79d6b529e53a2ef9dcdcc36090330d"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"bitflags 2.11.0",
|
||||
@@ -8501,9 +8506,9 @@ checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f"
|
||||
|
||||
[[package]]
|
||||
name = "winch-codegen"
|
||||
version = "43.0.1"
|
||||
version = "43.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "eb9f45f7172a2628c8317766e427babc0a400f9d10b1c0f0b0617c5ed5b79de6"
|
||||
checksum = "52dbb0cf07b0dfe7b7a1ca8efb8f94ba98bd0fb144c411ea1665c78f0449e958"
|
||||
dependencies = [
|
||||
"cranelift-assembler-x64",
|
||||
"cranelift-codegen",
|
||||
@@ -9231,7 +9236,7 @@ checksum = "b9cc00251562a284751c9973bace760d86c0276c471b4be569fe6b068ee97a56"
|
||||
|
||||
[[package]]
|
||||
name = "xtask"
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
|
||||
[[package]]
|
||||
name = "yoke"
|
||||
|
||||
+1
-1
@@ -18,7 +18,7 @@ members = [
|
||||
]
|
||||
|
||||
[workspace.package]
|
||||
version = "0.6.0"
|
||||
version = "0.6.9"
|
||||
edition = "2021"
|
||||
license = "Apache-2.0 OR MIT"
|
||||
repository = "https://github.com/RightNow-AI/openfang"
|
||||
|
||||
@@ -19,8 +19,8 @@
|
||||
<p align="center">
|
||||
<img src="https://img.shields.io/badge/language-Rust-orange?style=flat-square" alt="Rust" />
|
||||
<img src="https://img.shields.io/badge/license-MIT-blue?style=flat-square" alt="MIT" />
|
||||
<img src="https://img.shields.io/badge/version-0.5.10-green?style=flat-square" alt="v0.5.10" />
|
||||
<img src="https://img.shields.io/badge/tests-1,767%2B%20passing-brightgreen?style=flat-square" alt="Tests" />
|
||||
<img src="https://img.shields.io/badge/version-0.6.9-green?style=flat-square" alt="v0.6.9" />
|
||||
<img src="https://img.shields.io/badge/tests-2,696%2B%20passing-brightgreen?style=flat-square" alt="Tests" />
|
||||
<img src="https://img.shields.io/badge/clippy-0%20warnings-brightgreen?style=flat-square" alt="Clippy" />
|
||||
<a href="https://www.buymeacoffee.com/openfang" target="_blank"><img src="https://img.shields.io/badge/Buy%20Me%20a%20Coffee-FFDD00?style=flat-square&logo=buy-me-a-coffee&logoColor=black" alt="Buy Me A Coffee" /></a>
|
||||
</p>
|
||||
|
||||
@@ -1166,11 +1166,12 @@ pub async fn start_channel_bridge_with_config(
|
||||
if let Some(ref tg_config) = config.telegram {
|
||||
if let Some(token) = read_token(&tg_config.bot_token_env, "Telegram") {
|
||||
let poll_interval = Duration::from_secs(tg_config.poll_interval_secs);
|
||||
let adapter = Arc::new(TelegramAdapter::new(
|
||||
let adapter = Arc::new(TelegramAdapter::with_thread_routes(
|
||||
token,
|
||||
tg_config.allowed_users.clone(),
|
||||
poll_interval,
|
||||
tg_config.api_url.clone(),
|
||||
tg_config.thread_routes.clone(),
|
||||
));
|
||||
adapters.push((adapter, tg_config.default_agent.clone()));
|
||||
}
|
||||
@@ -1185,6 +1186,7 @@ pub async fn start_channel_bridge_with_config(
|
||||
dc_config.allowed_users.clone(),
|
||||
dc_config.ignore_bots,
|
||||
dc_config.intents,
|
||||
dc_config.auto_thread.clone(),
|
||||
));
|
||||
adapters.push((adapter, dc_config.default_agent.clone()));
|
||||
}
|
||||
@@ -1249,10 +1251,17 @@ pub async fn start_channel_bridge_with_config(
|
||||
// Matrix
|
||||
if let Some(ref mx_config) = config.matrix {
|
||||
if let Some(token) = read_token(&mx_config.access_token_env, "Matrix") {
|
||||
let adapter = Arc::new(MatrixAdapter::new(
|
||||
// MSC2918 refresh-token support: optional env var, when present the
|
||||
// adapter auto-recovers from M_UNKNOWN_TOKEN 401s.
|
||||
let refresh = mx_config
|
||||
.refresh_token_env
|
||||
.as_deref()
|
||||
.and_then(|env| read_token(env, "Matrix refresh"));
|
||||
let adapter = Arc::new(MatrixAdapter::with_refresh_token(
|
||||
mx_config.homeserver_url.clone(),
|
||||
mx_config.user_id.clone(),
|
||||
token,
|
||||
refresh,
|
||||
mx_config.allowed_rooms.clone(),
|
||||
mx_config.auto_accept_invites,
|
||||
));
|
||||
@@ -1477,9 +1486,10 @@ pub async fn start_channel_bridge_with_config(
|
||||
encrypt_key,
|
||||
fs_config.bot_names.clone(),
|
||||
)),
|
||||
FeishuMode::Websocket => Arc::new(FeishuAdapter::new_websocket(
|
||||
FeishuMode::Websocket => Arc::new(FeishuAdapter::new_websocket_with_region(
|
||||
fs_config.app_id.clone(),
|
||||
secret,
|
||||
region,
|
||||
)),
|
||||
};
|
||||
adapters.push((adapter, fs_config.default_agent.clone()));
|
||||
|
||||
@@ -216,7 +216,7 @@ pub async fn auth(
|
||||
|
||||
// Check session cookie (dashboard login sessions)
|
||||
if auth_state.auth_enabled {
|
||||
if let Some(token) = extract_session_cookie(&request) {
|
||||
if let Some(token) = crate::session_auth::extract_session_cookie(request.headers()) {
|
||||
if crate::session_auth::verify_session_token(&token, &auth_state.session_secret)
|
||||
.is_some()
|
||||
{
|
||||
@@ -242,21 +242,6 @@ pub async fn auth(
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
/// Extract the `openfang_session` cookie value from a request.
|
||||
fn extract_session_cookie(request: &Request<Body>) -> Option<String> {
|
||||
request
|
||||
.headers()
|
||||
.get("cookie")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|cookies| {
|
||||
cookies.split(';').find_map(|c| {
|
||||
c.trim()
|
||||
.strip_prefix("openfang_session=")
|
||||
.map(|v| v.to_string())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/// Security headers middleware — applied to ALL API responses.
|
||||
pub async fn security_headers(request: Request<Body>, next: Next) -> Response<Body> {
|
||||
let mut response = next.run(request).await;
|
||||
|
||||
@@ -235,7 +235,12 @@ fn convert_messages(oai_messages: &[OaiMessage]) -> Vec<Message> {
|
||||
OaiContent::Null => return None,
|
||||
};
|
||||
|
||||
Some(Message { role, content })
|
||||
Some(Message {
|
||||
msg_id: uuid::Uuid::new_v4().to_string(),
|
||||
provider_msg_id: None,
|
||||
role,
|
||||
content,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
@@ -215,6 +215,12 @@ pub async fn list_agents(State(state): State<Arc<AppState>>) -> impl IntoRespons
|
||||
let ready = matches!(e.state, openfang_types::agent::AgentState::Running)
|
||||
&& auth_status != "missing";
|
||||
|
||||
// Issue #1026: surface which agents are currently calling the LLM
|
||||
// so the dashboard can render a live "inferencing" indicator.
|
||||
// A running task in the kernel's `running_tasks` map means the
|
||||
// agent loop is in flight (LLM call + tool dispatch).
|
||||
let is_inferencing = state.kernel.running_tasks.contains_key(&e.id);
|
||||
|
||||
serde_json::json!({
|
||||
"id": e.id.to_string(),
|
||||
"name": e.name,
|
||||
@@ -227,6 +233,7 @@ pub async fn list_agents(State(state): State<Arc<AppState>>) -> impl IntoRespons
|
||||
"model_tier": tier,
|
||||
"auth_status": auth_status,
|
||||
"ready": ready,
|
||||
"is_inferencing": is_inferencing,
|
||||
"profile": e.manifest.profile,
|
||||
"identity": {
|
||||
"emoji": e.identity.emoji,
|
||||
@@ -321,6 +328,7 @@ pub fn inject_attachments_into_session(
|
||||
session.messages.push(Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Blocks(image_blocks),
|
||||
..Default::default()
|
||||
});
|
||||
|
||||
if let Err(e) = kernel.memory.save_session(&session) {
|
||||
@@ -684,6 +692,98 @@ pub async fn kill_agent(
|
||||
}
|
||||
}
|
||||
|
||||
/// DELETE /api/agents/{id}/uninstall — Permanently uninstall an agent.
|
||||
///
|
||||
/// Issue #1163: in addition to killing the agent (registry + memory + cron),
|
||||
/// this also removes the on-disk `~/.openfang/agents/<name>/` directory so
|
||||
/// the agent does not auto-respawn on the next daemon start.
|
||||
pub async fn uninstall_agent(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Path(id): Path<String>,
|
||||
) -> impl IntoResponse {
|
||||
let agent_id: AgentId = match id.parse() {
|
||||
Ok(id) => id,
|
||||
Err(_) => {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(serde_json::json!({"error": "Invalid agent ID"})),
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
// Capture the agent name BEFORE killing — registry entry is gone after.
|
||||
let agent_name = match state.kernel.registry.get(agent_id) {
|
||||
Some(entry) => entry.name.clone(),
|
||||
None => {
|
||||
return (
|
||||
StatusCode::NOT_FOUND,
|
||||
Json(serde_json::json!({"error": "Agent not found"})),
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
// Step 1: kill the agent (registry, memory, cron, triggers, caps).
|
||||
if let Err(e) = state.kernel.kill_agent(agent_id) {
|
||||
tracing::warn!("kill_agent failed during uninstall for {id}: {e}");
|
||||
return (
|
||||
StatusCode::NOT_FOUND,
|
||||
Json(serde_json::json!({"error": "Agent not found or already terminated"})),
|
||||
);
|
||||
}
|
||||
|
||||
// Step 2: remove ~/.openfang/agents/<name>/ so the agent does NOT
|
||||
// auto-respawn from disk on the next daemon start.
|
||||
let agents_dir = state.kernel.config.home_dir.join("agents");
|
||||
let agent_dir = agents_dir.join(&agent_name);
|
||||
|
||||
let dir_removed = if agent_dir.is_dir() {
|
||||
// Safety: only allow removal if the parent is exactly the agents root.
|
||||
let parent_ok = agent_dir
|
||||
.parent()
|
||||
.map(|p| p == agents_dir.as_path())
|
||||
.unwrap_or(false);
|
||||
if !parent_ok {
|
||||
tracing::warn!(
|
||||
agent = %agent_name,
|
||||
path = %agent_dir.display(),
|
||||
"Refusing to remove agent dir outside agents root"
|
||||
);
|
||||
false
|
||||
} else {
|
||||
match std::fs::remove_dir_all(&agent_dir) {
|
||||
Ok(()) => {
|
||||
tracing::info!(
|
||||
agent = %agent_name,
|
||||
path = %agent_dir.display(),
|
||||
"Removed agent directory on uninstall (#1163)"
|
||||
);
|
||||
true
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
agent = %agent_name,
|
||||
path = %agent_dir.display(),
|
||||
"Failed to remove agent directory: {e}"
|
||||
);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
false
|
||||
};
|
||||
|
||||
(
|
||||
StatusCode::OK,
|
||||
Json(serde_json::json!({
|
||||
"status": "uninstalled",
|
||||
"agent_id": id,
|
||||
"name": agent_name,
|
||||
"dir_removed": dir_removed,
|
||||
})),
|
||||
)
|
||||
}
|
||||
|
||||
/// POST /api/agents/{id}/restart — Restart a crashed/stuck agent.
|
||||
///
|
||||
/// Cancels any active task, resets agent state to Running, and updates last_active.
|
||||
@@ -1409,11 +1509,15 @@ pub async fn get_agent(
|
||||
"network": entry.manifest.capabilities.network,
|
||||
},
|
||||
"description": entry.manifest.description,
|
||||
"system_prompt": entry.manifest.model.system_prompt,
|
||||
"tags": entry.manifest.tags,
|
||||
"identity": {
|
||||
"emoji": entry.identity.emoji,
|
||||
"avatar_url": entry.identity.avatar_url,
|
||||
"color": entry.identity.color,
|
||||
"archetype": entry.identity.archetype,
|
||||
"vibe": entry.identity.vibe,
|
||||
"greeting_style": entry.identity.greeting_style,
|
||||
},
|
||||
"skills": entry.manifest.skills,
|
||||
"skills_mode": if entry.manifest.skills.is_empty() { "all" } else { "allowlist" },
|
||||
@@ -3599,7 +3703,14 @@ pub async fn install_skill(
|
||||
let config = openfang_skills::marketplace::MarketplaceConfig::default();
|
||||
let client = openfang_skills::marketplace::MarketplaceClient::new(config);
|
||||
|
||||
match client.install(&req.name, &skills_dir).await {
|
||||
let opts = openfang_skills::installer::InstallOptions {
|
||||
require_signed: req.require_signed,
|
||||
allowed_signer_keys: req.allowed_signer_keys.clone(),
|
||||
};
|
||||
match client
|
||||
.install_with_options(&req.name, &skills_dir, &opts)
|
||||
.await
|
||||
{
|
||||
Ok(version) => {
|
||||
// Hot-reload so agents see the new skill immediately
|
||||
state.kernel.reload_skills();
|
||||
@@ -3656,6 +3767,112 @@ pub async fn reload_skills(State(state): State<Arc<AppState>>) -> impl IntoRespo
|
||||
Json(serde_json::json!({"status": "reloaded"}))
|
||||
}
|
||||
|
||||
/// POST /api/audit/append — Append an entry to the Merkle hash chain audit
|
||||
/// trail on behalf of an external (instance-side) wrapper (issue #1174).
|
||||
///
|
||||
/// RBAC: gated by the same bearer-token middleware as POST /api/skills/install
|
||||
/// (see `middleware::auth_middleware`). When `api_key` is configured every
|
||||
/// caller must present `Authorization: Bearer <key>` — wrappers running in the
|
||||
/// same trust boundary as the daemon are expected to share that key.
|
||||
pub async fn audit_append(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(req): Json<AuditAppendRequest>,
|
||||
) -> impl IntoResponse {
|
||||
use openfang_runtime::audit::AuditAction;
|
||||
|
||||
// SECURITY: bound input sizes so a wrapper cannot wedge the chain with
|
||||
// unbounded strings. The audit table stores TEXT columns and the chain
|
||||
// hash is computed over the same bytes — keep it sane.
|
||||
const MAX_FIELD: usize = 16 * 1024;
|
||||
if req.event_type.len() > MAX_FIELD
|
||||
|| req.agent_id.len() > MAX_FIELD
|
||||
|| req.detail.len() > MAX_FIELD
|
||||
|| req
|
||||
.outcome
|
||||
.as_ref()
|
||||
.map(|s| s.len() > MAX_FIELD)
|
||||
.unwrap_or(false)
|
||||
|| req
|
||||
.signing_context
|
||||
.as_ref()
|
||||
.map(|s| s.len() > MAX_FIELD)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return (
|
||||
StatusCode::PAYLOAD_TOO_LARGE,
|
||||
Json(serde_json::json!({"error": "field exceeds 16KB limit"})),
|
||||
);
|
||||
}
|
||||
|
||||
// Map operator-supplied event_type → AuditAction (case-insensitive).
|
||||
let action = match req.event_type.trim().to_ascii_lowercase().as_str() {
|
||||
"toolinvoke" | "tool_invoke" | "tool" => AuditAction::ToolInvoke,
|
||||
"capabilitycheck" | "capability_check" | "capability" => AuditAction::CapabilityCheck,
|
||||
"agentspawn" | "agent_spawn" | "spawn" => AuditAction::AgentSpawn,
|
||||
"agentkill" | "agent_kill" | "kill" => AuditAction::AgentKill,
|
||||
"agentmessage" | "agent_message" | "message" => AuditAction::AgentMessage,
|
||||
"memoryaccess" | "memory_access" | "memory" => AuditAction::MemoryAccess,
|
||||
"fileaccess" | "file_access" | "file" => AuditAction::FileAccess,
|
||||
"networkaccess" | "network_access" | "network" => AuditAction::NetworkAccess,
|
||||
"shellexec" | "shell_exec" | "shell" => AuditAction::ShellExec,
|
||||
"authattempt" | "auth_attempt" | "auth" => AuditAction::AuthAttempt,
|
||||
"wireconnect" | "wire_connect" | "wire" => AuditAction::WireConnect,
|
||||
"configchange" | "config_change" | "config" => AuditAction::ConfigChange,
|
||||
other => {
|
||||
tracing::warn!(
|
||||
"audit_append: unknown event_type {other:?}, falling back to ToolInvoke"
|
||||
);
|
||||
AuditAction::ToolInvoke
|
||||
}
|
||||
};
|
||||
|
||||
// Compose a detail string that preserves the operator's free-form detail
|
||||
// plus optional signing context and structured payload, so wrappers can
|
||||
// attach context without changing the on-chain schema.
|
||||
let mut detail = req.detail.clone();
|
||||
if let Some(ctx) = req.signing_context.as_ref().filter(|s| !s.is_empty()) {
|
||||
if !detail.is_empty() {
|
||||
detail.push_str(" | ");
|
||||
}
|
||||
detail.push_str("signer=");
|
||||
detail.push_str(ctx);
|
||||
}
|
||||
if let Some(payload) = req.payload.as_ref() {
|
||||
let serialised = serde_json::to_string(payload)
|
||||
.unwrap_or_else(|_| String::from("<unserialisable payload>"));
|
||||
// Cap payload contribution so a huge JSON blob cannot blow the entry.
|
||||
let truncated: String = serialised.chars().take(8 * 1024).collect();
|
||||
if !detail.is_empty() {
|
||||
detail.push_str(" | ");
|
||||
}
|
||||
detail.push_str("payload=");
|
||||
detail.push_str(&truncated);
|
||||
}
|
||||
|
||||
let agent_id = if req.agent_id.trim().is_empty() {
|
||||
"external-wrapper".to_string()
|
||||
} else {
|
||||
req.agent_id.clone()
|
||||
};
|
||||
let outcome = req.outcome.clone().unwrap_or_else(|| "ok".to_string());
|
||||
|
||||
let hash = state
|
||||
.kernel
|
||||
.audit_log
|
||||
.record(agent_id, action, detail, outcome);
|
||||
let seq = state.kernel.audit_log.len().saturating_sub(1) as u64;
|
||||
|
||||
(
|
||||
StatusCode::OK,
|
||||
Json(serde_json::json!({
|
||||
"status": "appended",
|
||||
"seq": seq,
|
||||
"hash": hash,
|
||||
"tip": state.kernel.audit_log.tip_hash(),
|
||||
})),
|
||||
)
|
||||
}
|
||||
|
||||
/// GET /api/marketplace/search — Search the FangHub marketplace.
|
||||
pub async fn marketplace_search(
|
||||
Query(params): Query<HashMap<String, String>>,
|
||||
@@ -6287,7 +6504,7 @@ pub async fn list_providers(State(state): State<Arc<AppState>>) -> impl IntoResp
|
||||
// Index probe results by provider list position for O(1) lookup
|
||||
let mut probe_map: HashMap<usize, openfang_runtime::provider_health::ProbeResult> =
|
||||
HashMap::with_capacity(local_providers.len());
|
||||
for ((idx, _, _), result) in local_providers.iter().zip(probe_results.into_iter()) {
|
||||
for ((idx, _, _), result) in local_providers.iter().zip(probe_results) {
|
||||
probe_map.insert(*idx, result);
|
||||
}
|
||||
|
||||
@@ -7069,6 +7286,10 @@ pub async fn compact_session(
|
||||
}
|
||||
|
||||
/// POST /api/agents/{id}/stop — Cancel an agent's current LLM run.
|
||||
///
|
||||
/// If the agent is owned by an active hand instance, the hand instance is
|
||||
/// also deactivated. Otherwise the hand stays registered as `Active` and the
|
||||
/// user cannot re-activate it via the wizard (issue #1164).
|
||||
pub async fn stop_agent(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Path(id): Path<String>,
|
||||
@@ -7082,6 +7303,33 @@ pub async fn stop_agent(
|
||||
)
|
||||
}
|
||||
};
|
||||
|
||||
// If this agent is the agent of an active hand instance, deactivate the
|
||||
// hand entirely — which also kills the agent and cancels the run. This
|
||||
// matches what users expect when they click Stop on a hand-owned agent.
|
||||
if let Some(instance) = state.kernel.hand_registry.find_by_agent(agent_id) {
|
||||
match state.kernel.deactivate_hand(instance.instance_id) {
|
||||
Ok(()) => {
|
||||
return (
|
||||
StatusCode::OK,
|
||||
Json(serde_json::json!({
|
||||
"status": "ok",
|
||||
"message": "Hand deactivated",
|
||||
"hand_deactivated": true,
|
||||
"hand_id": instance.hand_id,
|
||||
"instance_id": instance.instance_id,
|
||||
})),
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
return (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": format!("{e}")})),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
match state.kernel.stop_agent_run(agent_id) {
|
||||
Ok(true) => (
|
||||
StatusCode::OK,
|
||||
@@ -7531,6 +7779,7 @@ pub async fn set_provider_key(
|
||||
model: model_id,
|
||||
api_key_env: env_var.clone(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let mut guard = state
|
||||
.kernel
|
||||
@@ -7698,6 +7947,7 @@ pub async fn test_provider(
|
||||
Some(base_url)
|
||||
},
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
|
||||
match openfang_runtime::drivers::create_driver(&driver_config) {
|
||||
@@ -8030,7 +8280,9 @@ fn build_skill_config_snapshot(
|
||||
.read()
|
||||
.unwrap_or_else(|e| e.into_inner());
|
||||
let source: &std::collections::HashMap<String, std::collections::HashMap<String, String>> =
|
||||
override_guard.as_ref().unwrap_or(&state.kernel.config.skills);
|
||||
override_guard
|
||||
.as_ref()
|
||||
.unwrap_or(&state.kernel.config.skills);
|
||||
source.get(skill_name).cloned().unwrap_or_default()
|
||||
};
|
||||
|
||||
@@ -8392,7 +8644,9 @@ fn remove_skill_config_var(
|
||||
|
||||
let mut remove_skill = false;
|
||||
if let Some(skills_table) = root.get_mut("skills").and_then(|v| v.as_table_mut()) {
|
||||
if let Some(skill_section) = skills_table.get_mut(skill_name).and_then(|v| v.as_table_mut())
|
||||
if let Some(skill_section) = skills_table
|
||||
.get_mut(skill_name)
|
||||
.and_then(|v| v.as_table_mut())
|
||||
{
|
||||
skill_section.remove(var_name);
|
||||
if skill_section.is_empty() {
|
||||
@@ -9019,9 +9273,8 @@ pub async fn create_schedule(
|
||||
}
|
||||
if let Some(arr) = delivery_targets_raw.as_array() {
|
||||
for (idx, t) in arr.iter().enumerate() {
|
||||
if let Err(e) = serde_json::from_value::<
|
||||
openfang_types::scheduler::CronDeliveryTarget,
|
||||
>(t.clone())
|
||||
if let Err(e) =
|
||||
serde_json::from_value::<openfang_types::scheduler::CronDeliveryTarget>(t.clone())
|
||||
{
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
@@ -9152,9 +9405,8 @@ pub async fn update_schedule(
|
||||
let mut parsed: Vec<openfang_types::scheduler::CronDeliveryTarget> =
|
||||
Vec::with_capacity(arr.len());
|
||||
for (idx, t) in arr.iter().enumerate() {
|
||||
match serde_json::from_value::<openfang_types::scheduler::CronDeliveryTarget>(
|
||||
t.clone(),
|
||||
) {
|
||||
match serde_json::from_value::<openfang_types::scheduler::CronDeliveryTarget>(t.clone())
|
||||
{
|
||||
Ok(dt) => parsed.push(dt),
|
||||
Err(e) => {
|
||||
return (
|
||||
@@ -9693,12 +9945,49 @@ pub async fn patch_agent_config(
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Request body for cloning an agent.
|
||||
///
|
||||
/// `overrides` is a free-form JSON object that is deep-merged onto the cloned
|
||||
/// manifest before spawning. This lets callers tweak fields like `description`,
|
||||
/// `tags`, `model`, etc. without re-stating the entire manifest. The `name`
|
||||
/// field on `overrides` is ignored — `new_name` always wins.
|
||||
#[derive(serde::Deserialize)]
|
||||
pub struct CloneAgentRequest {
|
||||
pub new_name: String,
|
||||
#[serde(default)]
|
||||
pub overrides: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// POST /api/agents/{id}/clone — Clone an agent with its workspace files.
|
||||
/// Workspace files that contain accumulated memory or per-session state and
|
||||
/// MUST NOT be copied when cloning an agent. The cloned agent starts with
|
||||
/// independent memory by design (see issue #868).
|
||||
const MEMORY_FILES: &[&str] = &["MEMORY.md", "HEARTBEAT.md"];
|
||||
|
||||
/// Deep-merge `overrides` onto `base` JSON. Object fields are recursively
|
||||
/// merged; arrays and scalars are replaced.
|
||||
fn deep_merge_json(base: &mut serde_json::Value, overrides: serde_json::Value) {
|
||||
match (base, overrides) {
|
||||
(serde_json::Value::Object(base_map), serde_json::Value::Object(over_map)) => {
|
||||
for (k, v) in over_map {
|
||||
match base_map.get_mut(&k) {
|
||||
Some(existing) => deep_merge_json(existing, v),
|
||||
None => {
|
||||
base_map.insert(k, v);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
(slot, value) => {
|
||||
*slot = value;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// POST /api/agents/{template_id}/clone — Clone a template agent into a new
|
||||
/// independent agent with its own workspace and memory.
|
||||
///
|
||||
/// Body: `{ "new_name": "user-42", "overrides": { ... } }`
|
||||
///
|
||||
/// Returns: `{ "agent_id": "...", "name": "...", "manifest": { ... } }`
|
||||
pub async fn clone_agent(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Path(id): Path<String>,
|
||||
@@ -9709,7 +9998,7 @@ pub async fn clone_agent(
|
||||
Err(_) => {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(serde_json::json!({"error": "Invalid agent ID"})),
|
||||
Json(serde_json::json!({"error": "Invalid template agent ID"})),
|
||||
);
|
||||
}
|
||||
};
|
||||
@@ -9721,29 +10010,90 @@ pub async fn clone_agent(
|
||||
);
|
||||
}
|
||||
|
||||
if req.new_name.trim().is_empty() {
|
||||
let new_name = req.new_name.trim().to_string();
|
||||
if new_name.is_empty() {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(serde_json::json!({"error": "new_name cannot be empty"})),
|
||||
);
|
||||
}
|
||||
|
||||
// Reject names with path separators / control chars to keep workspace dir naming safe.
|
||||
if new_name
|
||||
.chars()
|
||||
.any(|c| c == '/' || c == '\\' || c == '\0' || c.is_control())
|
||||
{
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(serde_json::json!({"error": "new_name contains invalid characters"})),
|
||||
);
|
||||
}
|
||||
|
||||
// Reject if template doesn't exist.
|
||||
let source = match state.kernel.registry.get(agent_id) {
|
||||
Some(e) => e,
|
||||
None => {
|
||||
return (
|
||||
StatusCode::NOT_FOUND,
|
||||
Json(serde_json::json!({"error": "Agent not found"})),
|
||||
Json(serde_json::json!({"error": "Template agent not found"})),
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
// Deep-clone manifest with new name
|
||||
let mut cloned_manifest = source.manifest.clone();
|
||||
cloned_manifest.name = req.new_name.clone();
|
||||
cloned_manifest.workspace = None; // Let kernel assign a new workspace
|
||||
// Reject if new_name collides with an existing agent.
|
||||
if state.kernel.registry.find_by_name(&new_name).is_some() {
|
||||
return (
|
||||
StatusCode::CONFLICT,
|
||||
Json(serde_json::json!({
|
||||
"error": format!("An agent named '{}' already exists", new_name)
|
||||
})),
|
||||
);
|
||||
}
|
||||
|
||||
// Spawn the cloned agent
|
||||
// Deep-clone manifest and apply overrides.
|
||||
let mut cloned_manifest = source.manifest.clone();
|
||||
if let Some(overrides) = req.overrides.clone() {
|
||||
if !overrides.is_object() {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(serde_json::json!({"error": "overrides must be a JSON object"})),
|
||||
);
|
||||
}
|
||||
// Serialize manifest to JSON, merge overrides, deserialize back.
|
||||
// This lets callers patch arbitrary nested fields without exhaustive enumeration.
|
||||
let mut manifest_json = match serde_json::to_value(&cloned_manifest) {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
return (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(
|
||||
serde_json::json!({"error": format!("Failed to serialize manifest: {e}")}),
|
||||
),
|
||||
);
|
||||
}
|
||||
};
|
||||
deep_merge_json(&mut manifest_json, overrides);
|
||||
cloned_manifest = match serde_json::from_value(manifest_json) {
|
||||
Ok(m) => m,
|
||||
Err(e) => {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(
|
||||
serde_json::json!({"error": format!("Invalid overrides — manifest no longer valid: {e}")}),
|
||||
),
|
||||
);
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
// new_name always wins over any name field in overrides.
|
||||
cloned_manifest.name = new_name.clone();
|
||||
// Let the kernel assign fresh workspace and state directory paths.
|
||||
// Never share private state with the template (see issue #868, #1097).
|
||||
cloned_manifest.workspace = None;
|
||||
cloned_manifest.state_dir = None;
|
||||
|
||||
// Spawn the cloned agent.
|
||||
let new_id = match state.kernel.spawn_agent(cloned_manifest) {
|
||||
Ok(id) => id,
|
||||
Err(e) => {
|
||||
@@ -9754,43 +10104,103 @@ pub async fn clone_agent(
|
||||
}
|
||||
};
|
||||
|
||||
// Copy workspace files from source to destination
|
||||
// Copy non-memory identity files from source state_dir to destination
|
||||
// state_dir. MEMORY.md and HEARTBEAT.md are intentionally skipped — the
|
||||
// cloned agent must start with independent memory (issue #868). Identity
|
||||
// files live in state_dir per #1097; fall back to legacy workspace for
|
||||
// older agents.
|
||||
let new_entry = state.kernel.registry.get(new_id);
|
||||
if let (Some(ref src_ws), Some(ref new_entry)) = (source.manifest.workspace, new_entry) {
|
||||
let src_state = source
|
||||
.manifest
|
||||
.state_dir
|
||||
.as_ref()
|
||||
.or(source.manifest.workspace.as_ref());
|
||||
let dst_state = new_entry.as_ref().and_then(|e| {
|
||||
e.manifest
|
||||
.state_dir
|
||||
.as_ref()
|
||||
.or(e.manifest.workspace.as_ref())
|
||||
});
|
||||
if let (Some(src_ws), Some(dst_ws)) = (src_state, dst_state) {
|
||||
if let (Ok(src_can), Ok(dst_can)) = (src_ws.canonicalize(), dst_ws.canonicalize()) {
|
||||
for &fname in KNOWN_IDENTITY_FILES {
|
||||
if MEMORY_FILES.contains(&fname) {
|
||||
continue;
|
||||
}
|
||||
let src_file = src_can.join(fname);
|
||||
let dst_file = dst_can.join(fname);
|
||||
if src_file.exists() {
|
||||
let _ = std::fs::copy(&src_file, &dst_file);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Copy the user-facing `skills/` subdirectory so curated skills travel
|
||||
// with the template. These live in the workspace, not the state dir.
|
||||
if let (Some(ref src_ws), Some(ref new_entry)) = (&source.manifest.workspace, &new_entry) {
|
||||
if let Some(ref dst_ws) = new_entry.manifest.workspace {
|
||||
// Security: canonicalize both paths
|
||||
if let (Ok(src_can), Ok(dst_can)) = (src_ws.canonicalize(), dst_ws.canonicalize()) {
|
||||
for &fname in KNOWN_IDENTITY_FILES {
|
||||
let src_file = src_can.join(fname);
|
||||
let dst_file = dst_can.join(fname);
|
||||
if src_file.exists() {
|
||||
let _ = std::fs::copy(&src_file, &dst_file);
|
||||
}
|
||||
let src_skills = src_can.join("skills");
|
||||
let dst_skills = dst_can.join("skills");
|
||||
if src_skills.is_dir() {
|
||||
let _ = copy_dir_recursive(&src_skills, &dst_skills);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Copy identity from source
|
||||
// Copy visual identity from source so the clone looks like the template.
|
||||
let _ = state
|
||||
.kernel
|
||||
.registry
|
||||
.update_identity(new_id, source.identity.clone());
|
||||
|
||||
// Register in channel router so binding resolution finds the cloned agent
|
||||
// Register in channel router so binding resolution finds the cloned agent.
|
||||
if let Some(ref mgr) = *state.bridge_manager.lock().await {
|
||||
mgr.router().register_agent(req.new_name.clone(), new_id);
|
||||
mgr.router().register_agent(new_name.clone(), new_id);
|
||||
}
|
||||
|
||||
// Return the freshly-spawned manifest (after kernel assigned workspace etc).
|
||||
let manifest_value = new_entry
|
||||
.as_ref()
|
||||
.and_then(|e| serde_json::to_value(&e.manifest).ok())
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
|
||||
(
|
||||
StatusCode::CREATED,
|
||||
Json(serde_json::json!({
|
||||
"agent_id": new_id.to_string(),
|
||||
"name": req.new_name,
|
||||
"name": new_name,
|
||||
"manifest": manifest_value,
|
||||
})),
|
||||
)
|
||||
}
|
||||
|
||||
/// Recursively copy a directory tree. Skips any file whose name appears in
|
||||
/// `MEMORY_FILES`. Best-effort: errors on individual entries are swallowed so
|
||||
/// a clone can still partially succeed (the caller will see the new agent ID).
|
||||
fn copy_dir_recursive(src: &std::path::Path, dst: &std::path::Path) -> std::io::Result<()> {
|
||||
std::fs::create_dir_all(dst)?;
|
||||
for entry in std::fs::read_dir(src)? {
|
||||
let entry = entry?;
|
||||
let file_type = entry.file_type()?;
|
||||
let name = entry.file_name();
|
||||
if let Some(name_str) = name.to_str() {
|
||||
if MEMORY_FILES.contains(&name_str) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
let src_path = entry.path();
|
||||
let dst_path = dst.join(&name);
|
||||
if file_type.is_dir() {
|
||||
let _ = copy_dir_recursive(&src_path, &dst_path);
|
||||
} else if file_type.is_file() {
|
||||
let _ = std::fs::copy(&src_path, &dst_path);
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workspace File Editor endpoints
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -9832,8 +10242,16 @@ pub async fn list_agent_files(
|
||||
}
|
||||
};
|
||||
|
||||
let workspace = match entry.manifest.workspace {
|
||||
Some(ref ws) => ws.clone(),
|
||||
// Identity files live in the agent's private state directory (see #1097).
|
||||
// Fall back to the legacy workspace location for agents created before the
|
||||
// split so existing on-disk files remain reachable.
|
||||
let workspace = match entry
|
||||
.manifest
|
||||
.state_dir
|
||||
.as_ref()
|
||||
.or(entry.manifest.workspace.as_ref())
|
||||
{
|
||||
Some(ws) => ws.clone(),
|
||||
None => {
|
||||
return (
|
||||
StatusCode::NOT_FOUND,
|
||||
@@ -9894,8 +10312,15 @@ pub async fn get_agent_file(
|
||||
}
|
||||
};
|
||||
|
||||
let workspace = match entry.manifest.workspace {
|
||||
Some(ref ws) => ws.clone(),
|
||||
// Identity files live in the agent's private state directory (see #1097).
|
||||
// Fall back to legacy workspace for agents created before the split.
|
||||
let workspace = match entry
|
||||
.manifest
|
||||
.state_dir
|
||||
.as_ref()
|
||||
.or(entry.manifest.workspace.as_ref())
|
||||
{
|
||||
Some(ws) => ws.clone(),
|
||||
None => {
|
||||
return (
|
||||
StatusCode::NOT_FOUND,
|
||||
@@ -10001,8 +10426,15 @@ pub async fn set_agent_file(
|
||||
}
|
||||
};
|
||||
|
||||
let workspace = match entry.manifest.workspace {
|
||||
Some(ref ws) => ws.clone(),
|
||||
// Identity files live in the agent's private state directory (see #1097).
|
||||
// Fall back to legacy workspace for agents created before the split.
|
||||
let workspace = match entry
|
||||
.manifest
|
||||
.state_dir
|
||||
.as_ref()
|
||||
.or(entry.manifest.workspace.as_ref())
|
||||
{
|
||||
Some(ws) => ws.clone(),
|
||||
None => {
|
||||
return (
|
||||
StatusCode::NOT_FOUND,
|
||||
@@ -12212,17 +12644,7 @@ pub async fn auth_check(
|
||||
};
|
||||
|
||||
// Check session cookie
|
||||
let session_user = request
|
||||
.headers()
|
||||
.get("cookie")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|cookies| {
|
||||
cookies.split(';').find_map(|c| {
|
||||
c.trim()
|
||||
.strip_prefix("openfang_session=")
|
||||
.map(|v| v.to_string())
|
||||
})
|
||||
})
|
||||
let session_user = crate::session_auth::extract_session_cookie(request.headers())
|
||||
.and_then(|token| crate::session_auth::verify_session_token(&token, &secret));
|
||||
|
||||
if let Some(username) = session_user {
|
||||
@@ -12444,3 +12866,110 @@ mod skill_config_tests {
|
||||
assert_eq!(back, doc);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod uninstall_agent_tests {
|
||||
//! Issue #1163 — directory-removal portion of the uninstall flow.
|
||||
//!
|
||||
//! These tests exercise the same logic the route handler runs after
|
||||
//! `kernel.kill_agent()`: locate `<home>/agents/<name>/`, verify it is
|
||||
//! directly under the agents root, and remove it. Live end-to-end
|
||||
//! coverage (real HTTP + kernel) belongs in `tests/api_integration_test.rs`.
|
||||
use std::path::Path;
|
||||
|
||||
/// Mirror of the dir-removal logic in `uninstall_agent`. Kept in sync
|
||||
/// with the route handler so the rules can be unit-tested without a
|
||||
/// running kernel. Returns whether the directory was removed.
|
||||
fn remove_agent_dir(home_dir: &Path, agent_name: &str) -> bool {
|
||||
let agents_dir = home_dir.join("agents");
|
||||
let agent_dir = agents_dir.join(agent_name);
|
||||
if !agent_dir.is_dir() {
|
||||
return false;
|
||||
}
|
||||
let parent_ok = agent_dir
|
||||
.parent()
|
||||
.map(|p| p == agents_dir.as_path())
|
||||
.unwrap_or(false);
|
||||
if !parent_ok {
|
||||
return false;
|
||||
}
|
||||
std::fs::remove_dir_all(&agent_dir).is_ok()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn removes_agent_directory_under_agents_root() {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let home = tmp.path().to_path_buf();
|
||||
let agents = home.join("agents");
|
||||
std::fs::create_dir_all(agents.join("trash-agent")).unwrap();
|
||||
std::fs::write(
|
||||
agents.join("trash-agent").join("agent.toml"),
|
||||
"name = \"trash-agent\"\n",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(agents.join("trash-agent").is_dir());
|
||||
let removed = remove_agent_dir(&home, "trash-agent");
|
||||
assert!(removed, "agent directory must be removed");
|
||||
assert!(!agents.join("trash-agent").exists());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_false_when_no_directory_exists() {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let home = tmp.path().to_path_buf();
|
||||
std::fs::create_dir_all(home.join("agents")).unwrap();
|
||||
|
||||
let removed = remove_agent_dir(&home, "ghost-agent");
|
||||
assert!(!removed, "no dir => false, but uninstall still succeeds");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn does_not_touch_siblings() {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let home = tmp.path().to_path_buf();
|
||||
let agents = home.join("agents");
|
||||
std::fs::create_dir_all(agents.join("trash-agent")).unwrap();
|
||||
std::fs::create_dir_all(agents.join("keep-me")).unwrap();
|
||||
std::fs::write(
|
||||
agents.join("trash-agent").join("agent.toml"),
|
||||
"name = \"trash-agent\"\n",
|
||||
)
|
||||
.unwrap();
|
||||
std::fs::write(
|
||||
agents.join("keep-me").join("agent.toml"),
|
||||
"name = \"keep-me\"\n",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert!(remove_agent_dir(&home, "trash-agent"));
|
||||
assert!(!agents.join("trash-agent").exists());
|
||||
assert!(
|
||||
agents.join("keep-me").is_dir(),
|
||||
"sibling agent dirs must not be touched by uninstall"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_path_traversal_attempt() {
|
||||
// A name like "../escape" would join to a path whose parent is the
|
||||
// agents root only if the file system resolves it that way — but
|
||||
// `parent()` on a non-canonicalized Path returns the textual parent,
|
||||
// which for `<home>/agents/../escape` is `<home>/agents/..`, not
|
||||
// `<home>/agents`. The check rejects it.
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let home = tmp.path().to_path_buf();
|
||||
std::fs::create_dir_all(home.join("agents")).unwrap();
|
||||
// Create a sibling dir outside agents/ that an attacker might want
|
||||
// to delete.
|
||||
std::fs::create_dir_all(home.join("escape")).unwrap();
|
||||
std::fs::write(home.join("escape").join("secret.toml"), "x = 1\n").unwrap();
|
||||
|
||||
let removed = remove_agent_dir(&home, "../escape");
|
||||
assert!(!removed, "must reject path-traversal names");
|
||||
assert!(
|
||||
home.join("escape").is_dir(),
|
||||
"sibling dir outside agents/ must NOT be deleted"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -54,6 +54,10 @@ pub async fn build_router(
|
||||
budget_config: Arc::new(tokio::sync::RwLock::new(kernel.config.budget.clone())),
|
||||
});
|
||||
|
||||
// Start WS cron broadcaster — subscribes to kernel event bus and pushes
|
||||
// cron job results to all connected WebSocket clients in real-time.
|
||||
ws::start_ws_cron_broadcaster(kernel.clone());
|
||||
|
||||
// CORS: allow localhost origins by default. If API key is set, the API
|
||||
// is protected anyway. For development, permissive CORS is convenient.
|
||||
let cors = if state.kernel.config.api_key.trim().is_empty() {
|
||||
@@ -182,6 +186,10 @@ pub async fn build_router(
|
||||
.delete(routes::kill_agent)
|
||||
.patch(routes::patch_agent),
|
||||
)
|
||||
.route(
|
||||
"/api/agents/{id}/uninstall",
|
||||
axum::routing::delete(routes::uninstall_agent),
|
||||
)
|
||||
.route(
|
||||
"/api/agents/{id}/mode",
|
||||
axum::routing::put(routes::set_agent_mode),
|
||||
@@ -195,6 +203,12 @@ pub async fn build_router(
|
||||
"/api/agents/{id}/start",
|
||||
axum::routing::post(routes::restart_agent),
|
||||
)
|
||||
.route(
|
||||
// Issue #890 — alias so dashboards and external orchestrators can
|
||||
// wake an inactive agent via a verb that matches the agent_activate tool.
|
||||
"/api/agents/{id}/activate",
|
||||
axum::routing::post(routes::restart_agent),
|
||||
)
|
||||
.route(
|
||||
"/api/agents/{id}/message",
|
||||
axum::routing::post(routes::send_message),
|
||||
@@ -380,6 +394,11 @@ pub async fn build_router(
|
||||
"/api/skills/reload",
|
||||
axum::routing::post(routes::reload_skills),
|
||||
)
|
||||
// Audit trail (issue #1174 — instance-side wrapper integration)
|
||||
.route(
|
||||
"/api/audit/append",
|
||||
axum::routing::post(routes::audit_append),
|
||||
)
|
||||
.route(
|
||||
"/api/skills/{id}/config",
|
||||
axum::routing::get(routes::get_skill_config).put(routes::put_skill_config),
|
||||
|
||||
@@ -17,6 +17,25 @@ pub fn create_session_token(username: &str, secret: &str, ttl_hours: u64) -> Str
|
||||
base64::engine::general_purpose::STANDARD.encode(format!("{payload}:{signature}"))
|
||||
}
|
||||
|
||||
/// Extract the `openfang_session` cookie value from a `Cookie` header string.
|
||||
///
|
||||
/// Returns `None` if the header is absent or the cookie is not present.
|
||||
/// Used by both the HTTP auth middleware and the WebSocket upgrade handler so
|
||||
/// that browser sessions established via `sessionLogin()` are honored on both
|
||||
/// surfaces (issue #1085).
|
||||
pub fn extract_session_cookie(headers: &axum::http::HeaderMap) -> Option<String> {
|
||||
headers
|
||||
.get("cookie")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|cookies| {
|
||||
cookies.split(';').find_map(|c| {
|
||||
c.trim()
|
||||
.strip_prefix("openfang_session=")
|
||||
.map(|v| v.to_string())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/// Verify a session token. Returns the username if valid and not expired.
|
||||
pub fn verify_session_token(token: &str, secret: &str) -> Option<String> {
|
||||
use base64::Engine;
|
||||
@@ -141,4 +160,36 @@ mod tests {
|
||||
// Starts with $argon2 but is not a valid PHC string.
|
||||
assert!(!verify_password("x", "$argon2id$garbage"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_session_cookie_present() {
|
||||
let mut h = axum::http::HeaderMap::new();
|
||||
h.insert(
|
||||
"cookie",
|
||||
"foo=bar; openfang_session=abc.def.ghi; baz=qux"
|
||||
.parse()
|
||||
.unwrap(),
|
||||
);
|
||||
assert_eq!(extract_session_cookie(&h).as_deref(), Some("abc.def.ghi"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_session_cookie_absent() {
|
||||
let mut h = axum::http::HeaderMap::new();
|
||||
h.insert("cookie", "foo=bar; baz=qux".parse().unwrap());
|
||||
assert_eq!(extract_session_cookie(&h), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_session_cookie_no_header() {
|
||||
let h = axum::http::HeaderMap::new();
|
||||
assert_eq!(extract_session_cookie(&h), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_session_cookie_only_value() {
|
||||
let mut h = axum::http::HeaderMap::new();
|
||||
h.insert("cookie", "openfang_session=lonely".parse().unwrap());
|
||||
assert_eq!(extract_session_cookie(&h).as_deref(), Some("lonely"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -65,6 +65,15 @@ pub struct MessageResponse {
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct SkillInstallRequest {
|
||||
pub name: String,
|
||||
/// When true, reject the install unless the bundle ships a valid
|
||||
/// Ed25519 SignedManifest envelope bound to the on-disk manifest.
|
||||
/// Maps to `InstallOptions::require_signed` (issue #1170).
|
||||
#[serde(default)]
|
||||
pub require_signed: bool,
|
||||
/// Optional hex-encoded allow-list of acceptable signer public keys.
|
||||
/// Empty = TOFU (any valid signature accepted).
|
||||
#[serde(default)]
|
||||
pub allowed_signer_keys: Vec<String>,
|
||||
}
|
||||
|
||||
/// Request to uninstall a skill.
|
||||
@@ -115,3 +124,96 @@ pub struct CommandsQuery {
|
||||
#[serde(default)]
|
||||
pub surface: Option<String>,
|
||||
}
|
||||
|
||||
/// Request body for `POST /api/audit/append` (issue #1174).
|
||||
///
|
||||
/// Lets external (instance-side) wrappers append entries to the Merkle hash
|
||||
/// chain audit log. The handler maps `event_type` to an `AuditAction` and
|
||||
/// records the entry through `kernel.audit_log`.
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct AuditAppendRequest {
|
||||
/// Operator-supplied event category. Case-insensitive, matched against the
|
||||
/// `AuditAction` enum variants (e.g. `tool_invoke`, `ConfigChange`,
|
||||
/// `agent_message`). Unknown values fall back to `ToolInvoke`.
|
||||
pub event_type: String,
|
||||
/// Agent or wrapper identifier responsible for the event. When empty,
|
||||
/// recorded as `"external-wrapper"`.
|
||||
#[serde(default)]
|
||||
pub agent_id: String,
|
||||
/// Free-form detail string (e.g. tool name, URL, file path).
|
||||
#[serde(default)]
|
||||
pub detail: String,
|
||||
/// Optional arbitrary payload. When present it is serialised to JSON and
|
||||
/// appended onto the entry's detail so the wrapper retains structured
|
||||
/// context without changing the on-chain schema.
|
||||
#[serde(default)]
|
||||
pub payload: Option<serde_json::Value>,
|
||||
/// Optional outcome string (`"ok"`, `"denied"`, or an error). Defaults to
|
||||
/// `"ok"` when omitted.
|
||||
#[serde(default)]
|
||||
pub outcome: Option<String>,
|
||||
/// Optional operator-supplied signing context (e.g. wrapper identity, key
|
||||
/// fingerprint). Mixed into the detail when present so the chain captures
|
||||
/// who attested to the event.
|
||||
#[serde(default)]
|
||||
pub signing_context: Option<String>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn skill_install_request_defaults_back_compat() {
|
||||
// Existing callers send `{"name": "..."}` only. New optional fields
|
||||
// must default cleanly (issue #1170).
|
||||
let req: SkillInstallRequest = serde_json::from_str(r#"{"name":"github-helper"}"#).unwrap();
|
||||
assert_eq!(req.name, "github-helper");
|
||||
assert!(!req.require_signed);
|
||||
assert!(req.allowed_signer_keys.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn skill_install_request_parses_require_signed() {
|
||||
let req: SkillInstallRequest = serde_json::from_str(
|
||||
r#"{"name":"x","require_signed":true,"allowed_signer_keys":["abc123"]}"#,
|
||||
)
|
||||
.unwrap();
|
||||
assert!(req.require_signed);
|
||||
assert_eq!(req.allowed_signer_keys, vec!["abc123".to_string()]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn audit_append_request_required_only() {
|
||||
// Only `event_type` is required; everything else must default.
|
||||
let req: AuditAppendRequest =
|
||||
serde_json::from_str(r#"{"event_type":"ToolInvoke"}"#).unwrap();
|
||||
assert_eq!(req.event_type, "ToolInvoke");
|
||||
assert!(req.agent_id.is_empty());
|
||||
assert!(req.detail.is_empty());
|
||||
assert!(req.payload.is_none());
|
||||
assert!(req.outcome.is_none());
|
||||
assert!(req.signing_context.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn audit_append_request_full_payload() {
|
||||
let body = r#"{
|
||||
"event_type": "config_change",
|
||||
"agent_id": "wrapper-1",
|
||||
"detail": "rotated key",
|
||||
"payload": {"key_id": "k-42", "ts": 1700000000},
|
||||
"outcome": "ok",
|
||||
"signing_context": "ed25519:deadbeef"
|
||||
}"#;
|
||||
let req: AuditAppendRequest = serde_json::from_str(body).unwrap();
|
||||
assert_eq!(req.event_type, "config_change");
|
||||
assert_eq!(req.agent_id, "wrapper-1");
|
||||
assert_eq!(req.detail, "rotated key");
|
||||
assert_eq!(req.outcome.as_deref(), Some("ok"));
|
||||
assert_eq!(req.signing_context.as_deref(), Some("ed25519:deadbeef"));
|
||||
let payload = req.payload.expect("payload present");
|
||||
assert_eq!(payload["key_id"], "k-42");
|
||||
assert_eq!(payload["ts"], 1_700_000_000);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -90,11 +90,11 @@ pub async fn webchat_page() -> impl IntoResponse {
|
||||
let html = WEBCHAT_HTML.replace(NONCE_PLACEHOLDER, &nonce);
|
||||
let csp = format!(
|
||||
"default-src 'self'; \
|
||||
script-src 'self' 'nonce-{nonce}' 'unsafe-eval'; \
|
||||
style-src 'self' 'unsafe-inline' https://fonts.googleapis.com https://fonts.gstatic.com; \
|
||||
script-src 'self' 'nonce-{nonce}' 'unsafe-eval' https://cdn.jsdelivr.net; \
|
||||
style-src 'self' 'unsafe-inline' https://fonts.googleapis.com https://fonts.gstatic.com https://cdn.jsdelivr.net; \
|
||||
img-src 'self' data: blob:; \
|
||||
connect-src 'self' ws://localhost:* ws://127.0.0.1:* wss://localhost:* wss://127.0.0.1:*; \
|
||||
font-src 'self' https://fonts.gstatic.com; \
|
||||
connect-src 'self' ws://localhost:* ws://127.0.0.1:* wss://localhost:* wss://127.0.0.1:* https://cdn.jsdelivr.net; \
|
||||
font-src 'self' https://fonts.gstatic.com https://cdn.jsdelivr.net; \
|
||||
media-src 'self' blob:; \
|
||||
frame-src 'self' blob:; \
|
||||
object-src 'none'; \
|
||||
@@ -120,6 +120,7 @@ pub async fn webchat_page() -> impl IntoResponse {
|
||||
/// All vendor libraries (Alpine.js, marked.js, highlight.js) are bundled
|
||||
/// locally — no CDN dependency. Alpine.js is included LAST because it
|
||||
/// immediately processes x-data directives and fires alpine:init on load.
|
||||
/// KaTeX is loaded dynamically from jsdelivr CDN when needed for LaTeX rendering.
|
||||
const WEBCHAT_HTML: &str = concat!(
|
||||
include_str!("../static/index_head.html"),
|
||||
"<style>\n",
|
||||
|
||||
+644
-51
@@ -19,6 +19,7 @@ use axum::response::IntoResponse;
|
||||
use dashmap::DashMap;
|
||||
use futures::stream::SplitSink;
|
||||
use futures::{SinkExt, StreamExt};
|
||||
use openfang_kernel::OpenFangKernel;
|
||||
use openfang_runtime::kernel_handle::KernelHandle;
|
||||
use openfang_runtime::llm_driver::StreamEvent;
|
||||
use openfang_runtime::llm_errors;
|
||||
@@ -30,7 +31,7 @@ use std::net::{IpAddr, SocketAddr};
|
||||
use std::sync::atomic::{AtomicU8, AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
/// Per-IP WebSocket connection tracker.
|
||||
@@ -98,6 +99,62 @@ fn ws_tracker() -> &'static DashMap<IpAddr, AtomicUsize> {
|
||||
TRACKER.get_or_init(DashMap::new)
|
||||
}
|
||||
|
||||
/// Per-agent WebSocket sender entry.
|
||||
struct WsSender {
|
||||
sender: Arc<Mutex<SplitSink<WebSocket, Message>>>,
|
||||
}
|
||||
|
||||
/// Global registry: agent_id → active WebSocket senders.
|
||||
/// Uses RwLock for fine-grained read/write access to the sender list.
|
||||
fn ws_agent_connections() -> &'static DashMap<AgentId, RwLock<Vec<WsSender>>> {
|
||||
static REGISTRY: std::sync::OnceLock<DashMap<AgentId, RwLock<Vec<WsSender>>>> =
|
||||
std::sync::OnceLock::new();
|
||||
REGISTRY.get_or_init(DashMap::new)
|
||||
}
|
||||
|
||||
/// Register a WebSocket connection for an agent (async).
|
||||
pub async fn register_ws_connection(
|
||||
agent_id: AgentId,
|
||||
sender: Arc<Mutex<SplitSink<WebSocket, Message>>>,
|
||||
) {
|
||||
let entry = ws_agent_connections().entry(agent_id).or_default();
|
||||
let mut senders = entry.value().write().await;
|
||||
senders.push(WsSender { sender });
|
||||
}
|
||||
|
||||
/// Deregister a WebSocket connection for an agent.
|
||||
/// Returns the number of remaining connections for this agent.
|
||||
pub async fn deregister_ws_connection(
|
||||
agent_id: AgentId,
|
||||
sender: &Arc<Mutex<SplitSink<WebSocket, Message>>>,
|
||||
) -> usize {
|
||||
let entry = match ws_agent_connections().get(&agent_id) {
|
||||
Some(e) => e,
|
||||
None => return 0,
|
||||
};
|
||||
let mut senders = entry.value().write().await;
|
||||
senders.retain(|s| !Arc::ptr_eq(&s.sender, sender));
|
||||
senders.len()
|
||||
}
|
||||
|
||||
/// Broadcast a JSON message to all active WebSocket connections for an agent.
|
||||
/// Returns the number of connections the message was sent to.
|
||||
pub async fn broadcast_to_ws(agent_id: AgentId, msg: serde_json::Value) -> usize {
|
||||
let entry = match ws_agent_connections().get(&agent_id) {
|
||||
Some(e) => e,
|
||||
None => return 0,
|
||||
};
|
||||
let senders = entry.value().read().await;
|
||||
let mut success_count = 0;
|
||||
for ws_sender in senders.iter() {
|
||||
let sender = &ws_sender.sender;
|
||||
if send_json(sender, &msg).await.is_ok() {
|
||||
success_count += 1;
|
||||
}
|
||||
}
|
||||
success_count
|
||||
}
|
||||
|
||||
/// RAII guard that decrements the connection count on drop.
|
||||
struct WsConnectionGuard {
|
||||
ip: IpAddr,
|
||||
@@ -133,11 +190,121 @@ fn try_acquire_ws_slot(ip: IpAddr) -> Option<WsConnectionGuard> {
|
||||
// WS Upgrade Handler
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Parameters for [`check_ws_auth`]. Kept as a struct so the auth gate stays
|
||||
/// pure and unit-testable without an `AppState` or live socket.
|
||||
pub(crate) struct WsAuthCtx<'a> {
|
||||
/// Trimmed API key from kernel config. Empty string means no key configured.
|
||||
pub api_key: &'a str,
|
||||
/// Whether dashboard session login is enabled in config.
|
||||
pub auth_enabled: bool,
|
||||
/// Secret used to verify session cookies (api_key when set, else password hash).
|
||||
pub session_secret: &'a str,
|
||||
/// Whether the request originated from a loopback address.
|
||||
pub is_loopback: bool,
|
||||
/// True iff `OPENFANG_ALLOW_NO_AUTH=1` is set (loose mode for LAN binds).
|
||||
pub allow_no_auth: bool,
|
||||
pub headers: &'a axum::http::HeaderMap,
|
||||
pub uri: &'a axum::http::Uri,
|
||||
}
|
||||
|
||||
/// Pure auth gate for WebSocket upgrades.
|
||||
///
|
||||
/// Returns `Ok(())` if the request should be allowed through, or
|
||||
/// `Err(StatusCode::UNAUTHORIZED)` otherwise. Accepts:
|
||||
/// 1. `Authorization: Bearer <api_key>` header
|
||||
/// 2. `?token=<api_key>` query parameter
|
||||
/// 3. `openfang_session=<token>` cookie when dashboard auth is enabled
|
||||
/// 4. Loopback origin when no api_key is configured
|
||||
/// 5. Any origin when `OPENFANG_ALLOW_NO_AUTH=1`
|
||||
///
|
||||
/// Fix for issue #1085: previously only (1), (2), and (4) were honored, so
|
||||
/// dashboard users logged in via session cookie saw "No active connection"
|
||||
/// because the WS upgrade rejected them even though HTTP requests succeeded.
|
||||
pub(crate) fn check_ws_auth(ctx: &WsAuthCtx<'_>) -> Result<(), axum::http::StatusCode> {
|
||||
use axum::http::StatusCode;
|
||||
|
||||
// No api_key configured: behavior depends on whether dashboard auth is on.
|
||||
//
|
||||
// Issue #1189: previously this path allowed any loopback request through
|
||||
// when api_key was empty, EVEN IF dashboard auth was enabled. That diverged
|
||||
// from the HTTP middleware (which only opens the loopback no-auth path
|
||||
// when api_key is empty AND auth.enabled is false). A local attacker with
|
||||
// loopback access could chat with agents over WS even when the operator
|
||||
// had configured dashboard credentials. Now mirror HTTP exactly.
|
||||
if ctx.api_key.is_empty() {
|
||||
// When dashboard auth is configured, require a valid session cookie
|
||||
// regardless of bind address. Loopback no longer bypasses login.
|
||||
if ctx.auth_enabled {
|
||||
if !ctx.session_secret.is_empty() {
|
||||
if let Some(token) = crate::session_auth::extract_session_cookie(ctx.headers) {
|
||||
if crate::session_auth::verify_session_token(&token, ctx.session_secret)
|
||||
.is_some()
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
}
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
// No api_key AND dashboard auth disabled: keep the dev convenience
|
||||
// path (loopback or explicit OPENFANG_ALLOW_NO_AUTH=1).
|
||||
if ctx.is_loopback || ctx.allow_no_auth {
|
||||
return Ok(());
|
||||
}
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
// SECURITY: constant-time comparison to prevent timing attacks on API key.
|
||||
let ct_eq = |token: &str, key: &str| -> bool {
|
||||
use subtle::ConstantTimeEq;
|
||||
if token.len() != key.len() {
|
||||
return false;
|
||||
}
|
||||
token.as_bytes().ct_eq(key.as_bytes()).into()
|
||||
};
|
||||
|
||||
let header_auth = ctx
|
||||
.headers
|
||||
.get("authorization")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|v| v.strip_prefix("Bearer "))
|
||||
.map(|token| ct_eq(token, ctx.api_key))
|
||||
.unwrap_or(false);
|
||||
if header_auth {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let query_auth = ctx
|
||||
.uri
|
||||
.query()
|
||||
.and_then(|q| q.split('&').find_map(|pair| pair.strip_prefix("token=")))
|
||||
.map(crate::percent_decode)
|
||||
.map(|token| ct_eq(&token, ctx.api_key))
|
||||
.unwrap_or(false);
|
||||
if query_auth {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Dashboard session cookie (issue #1085). When auth_enabled is on the
|
||||
// session_secret is set by server.rs to either the api_key or the
|
||||
// configured password hash, mirroring the HTTP auth middleware.
|
||||
if ctx.auth_enabled && !ctx.session_secret.is_empty() {
|
||||
if let Some(token) = crate::session_auth::extract_session_cookie(ctx.headers) {
|
||||
if crate::session_auth::verify_session_token(&token, ctx.session_secret).is_some() {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Err(StatusCode::UNAUTHORIZED)
|
||||
}
|
||||
|
||||
/// GET /api/agents/:id/ws — Upgrade to WebSocket for real-time chat.
|
||||
///
|
||||
/// SECURITY: Authenticates via Bearer token in Authorization header
|
||||
/// or `?token=` query parameter (for browser WebSocket clients that
|
||||
/// cannot set custom headers).
|
||||
/// SECURITY: Authenticates via Bearer token in Authorization header,
|
||||
/// `?token=` query parameter (for browser WebSocket clients that cannot
|
||||
/// set custom headers), or the `openfang_session` cookie set by the
|
||||
/// dashboard's session login flow (issue #1085).
|
||||
pub async fn agent_ws(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<Arc<AppState>>,
|
||||
@@ -152,48 +319,37 @@ pub async fn agent_ws(
|
||||
let api_key_raw = &state.kernel.config.api_key;
|
||||
let api_key = api_key_raw.trim();
|
||||
let is_loopback = addr.ip().is_loopback();
|
||||
let allow_no_auth = std::env::var("OPENFANG_ALLOW_NO_AUTH")
|
||||
.map(|v| matches!(v.trim(), "1" | "true" | "TRUE" | "yes" | "on"))
|
||||
.unwrap_or(false);
|
||||
|
||||
if api_key.is_empty() {
|
||||
// No key configured. Only allow loopback, unless the operator has
|
||||
// explicitly opted in to running open via OPENFANG_ALLOW_NO_AUTH=1.
|
||||
let allow_no_auth = std::env::var("OPENFANG_ALLOW_NO_AUTH")
|
||||
.map(|v| matches!(v.trim(), "1" | "true" | "TRUE" | "yes" | "on"))
|
||||
.unwrap_or(false);
|
||||
if !is_loopback && !allow_no_auth {
|
||||
warn!(
|
||||
ip = %addr.ip(),
|
||||
"WebSocket upgrade rejected: no api_key configured and origin is not loopback"
|
||||
);
|
||||
return axum::http::StatusCode::UNAUTHORIZED.into_response();
|
||||
}
|
||||
// Mirror the session_secret derivation in server.rs::AuthState so cookies
|
||||
// issued by /api/auth/login verify the same way over HTTP and WS.
|
||||
let auth_enabled = state.kernel.config.auth.enabled;
|
||||
let session_secret_owned: String = if !api_key.is_empty() {
|
||||
api_key.to_string()
|
||||
} else if auth_enabled {
|
||||
state.kernel.config.auth.password_hash.clone()
|
||||
} else {
|
||||
// SECURITY: Use constant-time comparison to prevent timing attacks on API key
|
||||
let ct_eq = |token: &str, key: &str| -> bool {
|
||||
use subtle::ConstantTimeEq;
|
||||
if token.len() != key.len() {
|
||||
return false;
|
||||
}
|
||||
token.as_bytes().ct_eq(key.as_bytes()).into()
|
||||
};
|
||||
String::new()
|
||||
};
|
||||
|
||||
let header_auth = headers
|
||||
.get("authorization")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|v| v.strip_prefix("Bearer "))
|
||||
.map(|token| ct_eq(token, api_key))
|
||||
.unwrap_or(false);
|
||||
let auth_ctx = WsAuthCtx {
|
||||
api_key,
|
||||
auth_enabled,
|
||||
session_secret: &session_secret_owned,
|
||||
is_loopback,
|
||||
allow_no_auth,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
|
||||
let query_auth = uri
|
||||
.query()
|
||||
.and_then(|q| q.split('&').find_map(|pair| pair.strip_prefix("token=")))
|
||||
.map(crate::percent_decode)
|
||||
.map(|token| ct_eq(&token, api_key))
|
||||
.unwrap_or(false);
|
||||
|
||||
if !header_auth && !query_auth {
|
||||
warn!("WebSocket upgrade rejected: invalid auth");
|
||||
return axum::http::StatusCode::UNAUTHORIZED.into_response();
|
||||
}
|
||||
if let Err(status) = check_ws_auth(&auth_ctx) {
|
||||
warn!(
|
||||
ip = %addr.ip(),
|
||||
"WebSocket upgrade rejected: no valid Bearer token, ?token=, or openfang_session cookie"
|
||||
);
|
||||
return status.into_response();
|
||||
}
|
||||
|
||||
// SECURITY: Enforce per-IP WebSocket connection limit
|
||||
@@ -264,6 +420,9 @@ async fn handle_agent_ws(
|
||||
let (sender, mut receiver) = socket.split();
|
||||
let sender = Arc::new(Mutex::new(sender));
|
||||
|
||||
// Register this connection in the global agent-WS registry
|
||||
register_ws_connection(agent_id, Arc::clone(&sender)).await;
|
||||
|
||||
// Per-connection verbose level (default: Full)
|
||||
let verbose = Arc::new(AtomicU8::new(VerboseLevel::Full as u8));
|
||||
|
||||
@@ -416,7 +575,8 @@ async fn handle_agent_ws(
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup
|
||||
// Cleanup: deregister from agent-WS registry and abort background tasks
|
||||
deregister_ws_connection(agent_id, &sender).await;
|
||||
update_handle.abort();
|
||||
info!(agent_id = %id_str, "WebSocket disconnected");
|
||||
}
|
||||
@@ -878,15 +1038,34 @@ async fn handle_command(
|
||||
serde_json::json!({"type": "error", "content": format!("Compaction failed: {e}")})
|
||||
}
|
||||
},
|
||||
"stop" => match state.kernel.stop_agent_run(agent_id) {
|
||||
Ok(true) => {
|
||||
serde_json::json!({"type": "command_result", "command": cmd, "message": "Run cancelled."})
|
||||
"stop" => {
|
||||
// If this agent is owned by an active hand instance, deactivate the
|
||||
// hand entirely so the user can re-activate it (issue #1164).
|
||||
if let Some(instance) = state.kernel.hand_registry.find_by_agent(agent_id) {
|
||||
match state.kernel.deactivate_hand(instance.instance_id) {
|
||||
Ok(()) => serde_json::json!({
|
||||
"type": "command_result",
|
||||
"command": cmd,
|
||||
"message": format!("Hand '{}' deactivated.", instance.hand_id),
|
||||
}),
|
||||
Err(e) => {
|
||||
serde_json::json!({"type": "error", "content": format!("Stop failed: {e}")})
|
||||
}
|
||||
}
|
||||
} else {
|
||||
match state.kernel.stop_agent_run(agent_id) {
|
||||
Ok(true) => {
|
||||
serde_json::json!({"type": "command_result", "command": cmd, "message": "Run cancelled."})
|
||||
}
|
||||
Ok(false) => {
|
||||
serde_json::json!({"type": "command_result", "command": cmd, "message": "No active run to cancel."})
|
||||
}
|
||||
Err(e) => {
|
||||
serde_json::json!({"type": "error", "content": format!("Stop failed: {e}")})
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(false) => {
|
||||
serde_json::json!({"type": "command_result", "command": cmd, "message": "No active run to cancel."})
|
||||
}
|
||||
Err(e) => serde_json::json!({"type": "error", "content": format!("Stop failed: {e}")}),
|
||||
},
|
||||
}
|
||||
"model" => {
|
||||
if args.is_empty() {
|
||||
if let Some(entry) = state.kernel.registry.get(agent_id) {
|
||||
@@ -1333,6 +1512,110 @@ pub fn strip_think_tags(text: &str) -> String {
|
||||
result
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Cron Job WS Broadcasting
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Start a background task that subscribes to the kernel's event bus and
|
||||
/// broadcasts cron job results to all connected WebSocket clients for the
|
||||
/// relevant agent.
|
||||
///
|
||||
/// This runs independently of the channel bridge — it uses the kernel's
|
||||
/// event bus to receive `CronJobExecuted` events and pushes them to WS.
|
||||
pub fn start_ws_cron_broadcaster(kernel: Arc<OpenFangKernel>) {
|
||||
tokio::spawn(async move {
|
||||
let mut rx = kernel.event_bus.subscribe_all();
|
||||
loop {
|
||||
let event = rx.recv().await;
|
||||
match event {
|
||||
Ok(event) => {
|
||||
if let openfang_types::event::EventPayload::System(
|
||||
openfang_types::event::SystemEvent::CronJobExecuted {
|
||||
agent_id,
|
||||
job_id,
|
||||
job_name,
|
||||
trigger_message,
|
||||
response,
|
||||
delivered_to_channel: _,
|
||||
},
|
||||
) = event.payload
|
||||
{
|
||||
// Build the trigger message (synthetic user message from cron)
|
||||
let trigger_msg = serde_json::json!({
|
||||
"type": "message",
|
||||
"content": trigger_message,
|
||||
"source": "cron",
|
||||
"job_id": job_id,
|
||||
"job_name": job_name
|
||||
});
|
||||
let _ = broadcast_to_ws(agent_id, trigger_msg).await;
|
||||
|
||||
// Send typing start
|
||||
let _ = broadcast_to_ws(
|
||||
agent_id,
|
||||
serde_json::json!({"state": "start", "type": "typing"}),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Send streaming phase
|
||||
let _ = broadcast_to_ws(
|
||||
agent_id,
|
||||
serde_json::json!({"detail": null, "phase": "streaming", "type": "phase"}),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Send text delta (full response since we don't have streaming chunks)
|
||||
let text_delta = serde_json::json!({
|
||||
"content": response,
|
||||
"type": "text_delta"
|
||||
});
|
||||
let _ = broadcast_to_ws(agent_id, text_delta).await;
|
||||
|
||||
// Send done phase
|
||||
let _ = broadcast_to_ws(
|
||||
agent_id,
|
||||
serde_json::json!({"detail": null, "phase": "done", "type": "phase"}),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Send typing stop
|
||||
let _ = broadcast_to_ws(
|
||||
agent_id,
|
||||
serde_json::json!({"state": "stop", "type": "typing"}),
|
||||
)
|
||||
.await;
|
||||
|
||||
// Send final response (mimics the format from agent_loop)
|
||||
let response_msg = serde_json::json!({
|
||||
"type": "response",
|
||||
"content": response,
|
||||
"context_pressure": "low",
|
||||
"cost_usd": null,
|
||||
"input_tokens": 0,
|
||||
"iterations": 0,
|
||||
"output_tokens": 0
|
||||
});
|
||||
let _ = broadcast_to_ws(agent_id, response_msg).await;
|
||||
|
||||
info!(
|
||||
agent_id = %agent_id,
|
||||
job_id = %job_id,
|
||||
"Cron job result broadcast to WS"
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
|
||||
warn!(lagged_messages = n, "WS cron broadcaster lagged, skipping");
|
||||
}
|
||||
Err(tokio::sync::broadcast::error::RecvError::Closed) => {
|
||||
info!("WS cron broadcaster channel closed, stopping");
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -1432,4 +1715,314 @@ mod tests {
|
||||
assert_eq!(strip_think_tags("No thinking here"), "No thinking here");
|
||||
assert_eq!(strip_think_tags("<think>all thinking</think>"), "");
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// WebSocket auth gate (issue #1085)
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
fn empty_uri() -> axum::http::Uri {
|
||||
"/api/agents/x/ws".parse().unwrap()
|
||||
}
|
||||
|
||||
fn uri_with_token(tok: &str) -> axum::http::Uri {
|
||||
format!("/api/agents/x/ws?token={tok}").parse().unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_accepts_bearer_token() {
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert("authorization", "Bearer secret".parse().unwrap());
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "secret",
|
||||
auth_enabled: false,
|
||||
session_secret: "secret",
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(check_ws_auth(&ctx).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_accepts_query_token() {
|
||||
let headers = axum::http::HeaderMap::new();
|
||||
let uri = uri_with_token("secret");
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "secret",
|
||||
auth_enabled: false,
|
||||
session_secret: "secret",
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(check_ws_auth(&ctx).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_accepts_session_cookie() {
|
||||
// Issue #1085: the dashboard logs in via cookie, so WS must accept it.
|
||||
let secret = "shared-secret";
|
||||
let token = crate::session_auth::create_session_token("alice", secret, 1);
|
||||
let cookie = format!("foo=bar; openfang_session={token}");
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert("cookie", cookie.parse().unwrap());
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: secret,
|
||||
auth_enabled: true,
|
||||
session_secret: secret,
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(
|
||||
check_ws_auth(&ctx).is_ok(),
|
||||
"valid session cookie should authorize WS upgrade"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_session_cookie_rejected_when_auth_disabled() {
|
||||
// If dashboard auth is off, cookies must not grant access.
|
||||
let secret = "shared-secret";
|
||||
let token = crate::session_auth::create_session_token("alice", secret, 1);
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"cookie",
|
||||
format!("openfang_session={token}").parse().unwrap(),
|
||||
);
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: secret,
|
||||
auth_enabled: false,
|
||||
session_secret: secret,
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert_eq!(
|
||||
check_ws_auth(&ctx).unwrap_err(),
|
||||
axum::http::StatusCode::UNAUTHORIZED
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_rejects_wrong_session_cookie() {
|
||||
// Cookie signed with the wrong secret must fail.
|
||||
let bad = crate::session_auth::create_session_token("alice", "other-secret", 1);
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert("cookie", format!("openfang_session={bad}").parse().unwrap());
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "secret",
|
||||
auth_enabled: true,
|
||||
session_secret: "secret",
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert_eq!(
|
||||
check_ws_auth(&ctx).unwrap_err(),
|
||||
axum::http::StatusCode::UNAUTHORIZED
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_rejects_when_no_credentials() {
|
||||
let headers = axum::http::HeaderMap::new();
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "secret",
|
||||
auth_enabled: true,
|
||||
session_secret: "secret",
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert_eq!(
|
||||
check_ws_auth(&ctx).unwrap_err(),
|
||||
axum::http::StatusCode::UNAUTHORIZED
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_rejects_wrong_bearer() {
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert("authorization", "Bearer wrong".parse().unwrap());
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "secret",
|
||||
auth_enabled: false,
|
||||
session_secret: "secret",
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert_eq!(
|
||||
check_ws_auth(&ctx).unwrap_err(),
|
||||
axum::http::StatusCode::UNAUTHORIZED
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_empty_key_loopback_ok() {
|
||||
let headers = axum::http::HeaderMap::new();
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "",
|
||||
auth_enabled: false,
|
||||
session_secret: "",
|
||||
is_loopback: true,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(check_ws_auth(&ctx).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_empty_key_non_loopback_rejected() {
|
||||
// Issue #1034 B2 regression guard.
|
||||
let headers = axum::http::HeaderMap::new();
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "",
|
||||
auth_enabled: false,
|
||||
session_secret: "",
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert_eq!(
|
||||
check_ws_auth(&ctx).unwrap_err(),
|
||||
axum::http::StatusCode::UNAUTHORIZED
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_empty_key_allow_no_auth_opens() {
|
||||
let headers = axum::http::HeaderMap::new();
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "",
|
||||
auth_enabled: false,
|
||||
session_secret: "",
|
||||
is_loopback: false,
|
||||
allow_no_auth: true,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(check_ws_auth(&ctx).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_empty_key_session_cookie_grants_non_loopback() {
|
||||
// When only dashboard login is configured (no api_key, auth_enabled=true),
|
||||
// a valid session cookie must allow non-loopback WS upgrades.
|
||||
let secret = "password-hash-style-secret";
|
||||
let token = crate::session_auth::create_session_token("admin", secret, 1);
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"cookie",
|
||||
format!("openfang_session={token}").parse().unwrap(),
|
||||
);
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "",
|
||||
auth_enabled: true,
|
||||
session_secret: secret,
|
||||
is_loopback: false,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(check_ws_auth(&ctx).is_ok());
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Issue #1189: WS auth must mirror HTTP middleware. When dashboard auth
|
||||
// is enabled, loopback + empty api_key + no cookie must NOT bypass.
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn ws_auth_dashboard_on_loopback_empty_key_no_cookie_rejected() {
|
||||
// Issue #1189 regression guard: previously this returned Ok(()) because
|
||||
// the empty-api_key branch allowed any loopback request through, even
|
||||
// when dashboard credentials were configured. HTTP middleware rejects
|
||||
// this path; WS must too.
|
||||
let secret = "password-hash-style-secret";
|
||||
let headers = axum::http::HeaderMap::new();
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "",
|
||||
auth_enabled: true,
|
||||
session_secret: secret,
|
||||
is_loopback: true,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert_eq!(
|
||||
check_ws_auth(&ctx).unwrap_err(),
|
||||
axum::http::StatusCode::UNAUTHORIZED,
|
||||
"loopback must not bypass dashboard auth when api_key is empty"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_dashboard_on_loopback_valid_cookie_accepted() {
|
||||
// With dashboard auth on, a valid session cookie is the supported
|
||||
// credential and must upgrade successfully from loopback too.
|
||||
let secret = "password-hash-style-secret";
|
||||
let token = crate::session_auth::create_session_token("admin", secret, 1);
|
||||
let mut headers = axum::http::HeaderMap::new();
|
||||
headers.insert(
|
||||
"cookie",
|
||||
format!("openfang_session={token}").parse().unwrap(),
|
||||
);
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "",
|
||||
auth_enabled: true,
|
||||
session_secret: secret,
|
||||
is_loopback: true,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(
|
||||
check_ws_auth(&ctx).is_ok(),
|
||||
"valid session cookie should authorize loopback WS upgrade"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn ws_auth_dashboard_off_loopback_empty_key_accepted() {
|
||||
// Preserve the development convenience path: when dashboard auth is
|
||||
// NOT configured AND api_key is empty, loopback still upgrades.
|
||||
let headers = axum::http::HeaderMap::new();
|
||||
let uri = empty_uri();
|
||||
let ctx = WsAuthCtx {
|
||||
api_key: "",
|
||||
auth_enabled: false,
|
||||
session_secret: "",
|
||||
is_loopback: true,
|
||||
allow_no_auth: false,
|
||||
headers: &headers,
|
||||
uri: &uri,
|
||||
};
|
||||
assert!(
|
||||
check_ws_auth(&ctx).is_ok(),
|
||||
"loopback dev path must work when dashboard auth is disabled"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -312,6 +312,12 @@ tr:hover td { background: var(--surface2); }
|
||||
|
||||
@keyframes pulse { 0%, 100% { opacity: 1; } 50% { opacity: 0.4; } }
|
||||
|
||||
/* Issue #1026: live indicator for agents currently calling the LLM */
|
||||
@keyframes agent-inferencing-pulse {
|
||||
0%, 100% { transform: scale(1); opacity: 1; box-shadow: 0 0 0 0 var(--accent); }
|
||||
50% { transform: scale(1.25); opacity: 0.85; box-shadow: 0 0 0 4px rgba(255, 92, 0, 0); }
|
||||
}
|
||||
|
||||
.message.user {
|
||||
flex-direction: row-reverse;
|
||||
}
|
||||
@@ -3253,7 +3259,7 @@ mark.search-highlight {
|
||||
═══════════════════════════════════════════════════════════════════════════ */
|
||||
|
||||
.trader-dashboard {
|
||||
background: var(--bg-card);
|
||||
background: var(--surface);
|
||||
border: 1px solid var(--border);
|
||||
border-radius: 12px;
|
||||
width: 96vw;
|
||||
@@ -3270,7 +3276,7 @@ mark.search-highlight {
|
||||
border-bottom: 1px solid var(--border);
|
||||
position: sticky;
|
||||
top: 0;
|
||||
background: var(--bg-card);
|
||||
background: var(--surface);
|
||||
z-index: 10;
|
||||
border-radius: 12px 12px 0 0;
|
||||
}
|
||||
@@ -3330,6 +3336,7 @@ mark.search-highlight {
|
||||
border-radius: 8px;
|
||||
padding: 14px 16px;
|
||||
min-width: 0;
|
||||
position: relative;
|
||||
}
|
||||
.trader-chart-title {
|
||||
font-size: 0.75rem;
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
/* OpenFang Layout — Grid + Sidebar + Responsive */
|
||||
|
||||
/* Firefox compat: hide x-cloak elements until Alpine.js initializes.
|
||||
Without this, the sidebar flashes hidden in Firefox while Alpine
|
||||
processes the nested x-data scopes for nav sections. */
|
||||
[x-cloak] { display: none !important; }
|
||||
|
||||
.app-layout {
|
||||
display: flex;
|
||||
height: 100vh;
|
||||
|
||||
@@ -27,8 +27,8 @@
|
||||
</div>
|
||||
|
||||
<div class="app-layout" :class="{ 'focus-mode': $store.app.focusMode }">
|
||||
<!-- Sidebar -->
|
||||
<nav class="sidebar" :class="{ collapsed: sidebarCollapsed, 'mobile-open': mobileMenuOpen }">
|
||||
<!-- Sidebar — x-cloak prevents Firefox flash-hidden during Alpine init -->
|
||||
<nav class="sidebar" x-cloak :class="{ collapsed: sidebarCollapsed, 'mobile-open': mobileMenuOpen }">
|
||||
<div class="sidebar-header">
|
||||
<div class="sidebar-header-text">
|
||||
<div class="sidebar-logo">
|
||||
@@ -68,8 +68,8 @@
|
||||
<span class="nav-label">Agents</span>
|
||||
<span class="nav-section-chevron" :style="collapsed ? '' : 'transform:rotate(90deg)'">›</span>
|
||||
</div>
|
||||
<template x-if="!collapsed">
|
||||
<div x-transition>
|
||||
<!-- x-show + x-cloak: Firefox-safe replacement for nested <template x-if> which has render quirks. -->
|
||||
<div x-show="!collapsed" x-cloak x-transition>
|
||||
<a class="nav-item" :class="{ active: page === 'agents' }" @click="navigate('agents')" :aria-current="page === 'agents' ? 'page' : false">
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><path d="M21 15a2 2 0 0 1-2 2H7l-4 4V5a2 2 0 0 1 2-2h14a2 2 0 0 1 2 2z"/></svg></span>
|
||||
<span class="nav-label">Chat</span>
|
||||
@@ -87,8 +87,7 @@
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><path d="M21 11.5a8.38 8.38 0 01-.9 3.8 8.5 8.5 0 01-7.6 4.7 8.38 8.38 0 01-3.8-.9L3 21l1.9-5.7a8.38 8.38 0 01-.9-3.8 8.5 8.5 0 014.7-7.6 8.38 8.38 0 013.8-.9h.5a8.48 8.48 0 018 8v.5z"/></svg></span>
|
||||
<span class="nav-label">Comms</span>
|
||||
</a>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Automation -->
|
||||
@@ -97,8 +96,7 @@
|
||||
<span class="nav-label">Automation</span>
|
||||
<span class="nav-section-chevron" :style="collapsed ? '' : 'transform:rotate(90deg)'">›</span>
|
||||
</div>
|
||||
<template x-if="!collapsed">
|
||||
<div x-transition>
|
||||
<div x-show="!collapsed" x-cloak x-transition>
|
||||
<a class="nav-item" :class="{ active: page === 'workflows' }" @click="navigate('workflows')" :aria-current="page === 'workflows' ? 'page' : false">
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><path d="M6 3v12M18 9a9 9 0 0 1-9 9"/><circle cx="18" cy="6" r="3"/><circle cx="6" cy="18" r="3"/></svg></span>
|
||||
<span class="nav-label">Workflows</span>
|
||||
@@ -107,8 +105,7 @@
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><circle cx="12" cy="12" r="10"/><path d="M12 6v6l4 2"/></svg></span>
|
||||
<span class="nav-label">Scheduler</span>
|
||||
</a>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Extensions -->
|
||||
@@ -117,8 +114,7 @@
|
||||
<span class="nav-label">Extensions</span>
|
||||
<span class="nav-section-chevron" :style="collapsed ? '' : 'transform:rotate(90deg)'">›</span>
|
||||
</div>
|
||||
<template x-if="!collapsed">
|
||||
<div x-transition>
|
||||
<div x-show="!collapsed" x-cloak x-transition>
|
||||
<a class="nav-item" :class="{ active: page === 'channels' }" @click="navigate('channels')" :aria-current="page === 'channels' ? 'page' : false">
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><path d="M4 9h16M4 15h16M10 3l-2 18M16 3l-2 18"/></svg></span>
|
||||
<span class="nav-label">Channels</span>
|
||||
@@ -131,8 +127,7 @@
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><path d="M18 11V6a2 2 0 0 0-2-2 2 2 0 0 0-2 2"/><path d="M14 10V4a2 2 0 0 0-2-2 2 2 0 0 0-2 2v6"/><path d="M10 10.5V6a2 2 0 0 0-2-2 2 2 0 0 0-2 2v8"/><path d="M18 8a2 2 0 1 1 4 0v6a8 8 0 0 1-8 8h-2c-2.8 0-4.5-.9-5.7-2.4L3.4 16a2 2 0 0 1 3.2-2.4L8 15"/></svg></span>
|
||||
<span class="nav-label">Hands</span>
|
||||
</a>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Monitor -->
|
||||
@@ -141,8 +136,7 @@
|
||||
<span class="nav-label">Monitor</span>
|
||||
<span class="nav-section-chevron" :style="collapsed ? '' : 'transform:rotate(90deg)'">›</span>
|
||||
</div>
|
||||
<template x-if="!collapsed">
|
||||
<div x-transition>
|
||||
<div x-show="!collapsed" x-cloak x-transition>
|
||||
<a class="nav-item" :class="{ active: page === 'analytics' }" @click="navigate('analytics')" :aria-current="page === 'analytics' ? 'page' : false">
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><path d="M18 20V10M12 20V4M6 20v-6"/></svg></span>
|
||||
<span class="nav-label">Analytics</span>
|
||||
@@ -151,8 +145,7 @@
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><path d="m4 17 6-6-6-6"/><path d="M12 19h8"/></svg></span>
|
||||
<span class="nav-label">Logs</span>
|
||||
</a>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- System -->
|
||||
@@ -161,8 +154,7 @@
|
||||
<span class="nav-label">System</span>
|
||||
<span class="nav-section-chevron" :style="collapsed ? '' : 'transform:rotate(90deg)'">›</span>
|
||||
</div>
|
||||
<template x-if="!collapsed">
|
||||
<div x-transition>
|
||||
<div x-show="!collapsed" x-cloak x-transition>
|
||||
<a class="nav-item" :class="{ active: page === 'runtime' }" @click="navigate('runtime')" :aria-current="page === 'runtime' ? 'page' : false">
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><rect x="2" y="3" width="20" height="14" rx="2"/><path d="M8 21h8M12 17v4"/></svg></span>
|
||||
<span class="nav-label">Runtime</span>
|
||||
@@ -171,8 +163,7 @@
|
||||
<span class="nav-icon"><svg viewBox="0 0 24 24"><path d="M4 21v-7M4 10V3M12 21v-9M12 8V3M20 21v-5M20 12V3"/><path d="M1 14h6M9 8h6M17 16h6"/></svg></span>
|
||||
<span class="nav-label">Settings</span>
|
||||
</a>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -186,7 +177,7 @@
|
||||
<div class="sidebar-toggle" @click="toggleSidebar()" x-text="sidebarCollapsed ? '\u276F' : '\u276E'"></div>
|
||||
</nav>
|
||||
|
||||
<div class="sidebar-overlay" @click="mobileMenuOpen = false"></div>
|
||||
<div class="sidebar-overlay" x-cloak @click="mobileMenuOpen = false"></div>
|
||||
|
||||
<!-- Main Content -->
|
||||
<main class="main-content">
|
||||
@@ -589,6 +580,7 @@
|
||||
<svg width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><g x-show="!$store.app.focusMode"><path d="M8 3H5a2 2 0 0 0-2 2v3"/><path d="M21 8V5a2 2 0 0 0-2-2h-3"/><path d="M3 16v3a2 2 0 0 0 2 2h3"/><path d="M16 21h3a2 2 0 0 0 2-2v-3"/></g><g x-show="$store.app.focusMode"><path d="M8 3v3a2 2 0 0 1-2 2H3"/><path d="M21 8h-3a2 2 0 0 1-2-2V3"/><path d="M3 16h3a2 2 0 0 1 2 2v3"/><path d="M16 21v-3a2 2 0 0 1 2-2h3"/></g></svg>
|
||||
</button>
|
||||
<button class="btn btn-danger btn-sm" @click="killAgent()">Stop</button>
|
||||
<button class="btn btn-danger btn-sm" @click="uninstallAgent()" title="Stop and remove agent files from workspace">Uninstall</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -748,7 +740,7 @@
|
||||
<span class="text-xs" style="color:var(--danger)" x-text="formatRecordingTime()"></span>
|
||||
</div>
|
||||
<textarea id="msg-input" rows="1" :placeholder="recording ? 'Recording... release to send' : 'Message OpenFang... (/ for commands)'"
|
||||
@keydown.enter.prevent="if(!$event.isComposing && $event.keyCode !== 229 && !$event.shiftKey){if(showModelPicker && filteredModelPicker.length){pickModel(filteredModelPicker[modelPickerIdx].id)}else if(showSlashMenu && filteredSlashCommands.length){executeSlashCommand(filteredSlashCommands[slashIdx].cmd)}else{sendMessage()}}"
|
||||
@keydown.enter="if(!$event.isComposing && $event.keyCode !== 229 && !$event.shiftKey){$event.preventDefault();if(showModelPicker && filteredModelPicker.length){pickModel(filteredModelPicker[modelPickerIdx].id)}else if(showSlashMenu && filteredSlashCommands.length){executeSlashCommand(filteredSlashCommands[slashIdx].cmd)}else{sendMessage()}}"
|
||||
@keydown.escape="showSlashMenu = false; showModelPicker = false"
|
||||
@keydown.arrow-up.prevent="if(showModelPicker){modelPickerIdx = Math.max(0, modelPickerIdx - 1)}else if(showSlashMenu){slashIdx = Math.max(0, slashIdx - 1)}"
|
||||
@keydown.arrow-down.prevent="if(showModelPicker){modelPickerIdx = Math.min(filteredModelPicker.length - 1, modelPickerIdx + 1)}else if(showSlashMenu){slashIdx = Math.min(filteredSlashCommands.length - 1, slashIdx + 1)}"
|
||||
@@ -856,8 +848,15 @@
|
||||
<svg x-show="!agent.identity || !agent.identity.emoji" width="18" height="18" viewBox="0 0 24 24" fill="none" stroke="var(--accent)" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><path d="M21 15a2 2 0 0 1-2 2H7l-4 4V5a2 2 0 0 1 2-2h14a2 2 0 0 1 2 2z"/></svg>
|
||||
</div>
|
||||
<div style="min-width:0;flex:1">
|
||||
<div class="font-bold" style="font-size:13px" x-text="agent.name"></div>
|
||||
<div class="text-xs text-dim font-mono" style="font-size:11px" x-text="agent.model_name"></div>
|
||||
<div class="font-bold" style="font-size:13px">
|
||||
<span x-text="agent.name"></span>
|
||||
<!-- Issue #1026: live inferencing indicator -->
|
||||
<span x-show="agent.is_inferencing" class="agent-inferencing-dot" title="Agent is calling the LLM right now" style="display:inline-block;width:8px;height:8px;border-radius:50%;background:var(--accent);margin-left:6px;vertical-align:middle;animation:agent-inferencing-pulse 1.2s ease-in-out infinite"></span>
|
||||
</div>
|
||||
<div class="text-xs text-dim font-mono" style="font-size:11px">
|
||||
<span x-show="!agent.is_inferencing" x-text="agent.model_name"></span>
|
||||
<span x-show="agent.is_inferencing" style="color:var(--accent);font-weight:600">Inferencing…</span>
|
||||
</div>
|
||||
</div>
|
||||
<span class="badge" :class="'badge-' + agent.state.toLowerCase()" x-text="agent.state" style="font-size:10px"></span>
|
||||
<button class="agent-chip-config-btn" @click.stop="showDetail(agent)" title="Agent settings" style="display:flex;align-items:center;justify-content:center;width:28px;height:28px;border-radius:50%;border:1px solid var(--border);background:transparent;cursor:pointer;color:var(--text-dim);transition:all 0.15s;flex-shrink:0" @mouseenter="$el.style.borderColor='var(--accent)';$el.style.color='var(--accent)';$el.style.background='var(--surface2)'" @mouseleave="$el.style.borderColor='var(--border)';$el.style.color='var(--text-dim)';$el.style.background='transparent'">
|
||||
@@ -990,6 +989,7 @@
|
||||
<button class="btn btn-ghost" @click="cloneAgent(detailAgent)">Clone</button>
|
||||
<button class="btn btn-ghost" @click="clearHistory(detailAgent)">Clear History</button>
|
||||
<button class="btn btn-danger" @click="killAgent(detailAgent)">Stop</button>
|
||||
<button class="btn btn-danger" @click="uninstallAgent(detailAgent)" title="Stop and remove agent files from workspace">Uninstall</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -3573,11 +3573,47 @@ args = ["-y", "@modelcontextprotocol/server-filesystem", "/path"]</pre>
|
||||
<div x-show="tab === 'providers'">
|
||||
<div class="info-card">
|
||||
<h4>LLM Providers</h4>
|
||||
<p>OpenFang supports 12 LLM providers out of the box. Configure API keys to unlock models from each provider. Set environment variables and restart, or use the form below to save keys directly.</p>
|
||||
<p>OpenFang ships with <span x-text="providers.length"></span> built-in providers and you can add unlimited custom ones. <span x-text="configuredProviderCount"></span> currently configured. Filter, search, or jump to a category below — only providers with a saved key (or that need none) light up models for your agents.</p>
|
||||
</div>
|
||||
<div class="card-grid">
|
||||
<template x-for="p in providers" :key="p.id">
|
||||
<div class="card provider-card" :class="providerCardClass(p)">
|
||||
<!-- Filter toolbar -->
|
||||
<div class="flex gap-2 mb-4" style="flex-wrap:wrap;align-items:center">
|
||||
<div class="search-input" style="flex:1;min-width:200px">
|
||||
<span style="color:var(--text-muted)"><svg viewBox="0 0 24 24" width="14" height="14" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><circle cx="11" cy="11" r="8"/><path d="m21 21-4.35-4.35"/></svg></span>
|
||||
<input placeholder="Search providers..." x-model="providerSearch">
|
||||
</div>
|
||||
<select class="form-select" style="width:170px" x-model="providerStatusFilter">
|
||||
<option value="">All Statuses</option>
|
||||
<option value="configured">Configured</option>
|
||||
<option value="unconfigured">Needs Key</option>
|
||||
</select>
|
||||
<select class="form-select" style="width:200px" x-model="providerCategoryFilter">
|
||||
<option value="">All Categories</option>
|
||||
<option value="frontier">Frontier</option>
|
||||
<option value="oss">Open-Weight Hosts</option>
|
||||
<option value="aggregator">Aggregators</option>
|
||||
<option value="regional">Regional / China</option>
|
||||
<option value="local">Local / Self-Hosted</option>
|
||||
<option value="other">Other</option>
|
||||
</select>
|
||||
<button class="btn btn-ghost btn-sm" @click="clearProviderFilters()" x-show="providerSearch || providerStatusFilter || providerCategoryFilter">Clear</button>
|
||||
</div>
|
||||
<div class="text-xs text-dim mb-2" x-text="filteredProviders.length + ' of ' + providers.length + ' providers'"></div>
|
||||
<!-- Empty state for filters -->
|
||||
<div x-show="!filteredProviders.length && providers.length" style="text-align:center;padding:32px 16px">
|
||||
<h3 style="margin:0 0 4px;font-size:14px">No providers match your filters</h3>
|
||||
<p class="text-xs text-dim">Try a different search term or category.</p>
|
||||
<button class="btn btn-ghost btn-sm mt-2" @click="clearProviderFilters()">Clear Filters</button>
|
||||
</div>
|
||||
<!-- Grouped provider sections -->
|
||||
<template x-for="group in providersGrouped" :key="group.category">
|
||||
<div style="margin-bottom:1.25rem">
|
||||
<div class="card-header" style="display:flex;align-items:center;gap:8px;margin-bottom:6px;font-size:11px;text-transform:uppercase;letter-spacing:0.5px;color:var(--text-muted)">
|
||||
<span x-text="group.label"></span>
|
||||
<span class="text-xs text-dim" style="font-weight:normal;text-transform:none;letter-spacing:0" x-text="'(' + group.items.length + ')'"></span>
|
||||
</div>
|
||||
<div class="card-grid">
|
||||
<template x-for="p in group.items" :key="p.id">
|
||||
<div class="card provider-card" :class="providerCardClass(p)">
|
||||
<div class="flex justify-between items-center mb-2">
|
||||
<div class="card-header" style="margin:0" x-text="p.display_name"></div>
|
||||
<span class="badge" :class="providerAuthClass(p)" x-text="providerAuthText(p)"></span>
|
||||
@@ -3636,9 +3672,11 @@ args = ["-y", "@modelcontextprotocol/server-filesystem", "/path"]</pre>
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
<!-- Add Custom Provider -->
|
||||
<div class="info-card mt-4" style="border:1px solid var(--border)">
|
||||
<h4 style="margin-top:0">Add Custom Provider</h4>
|
||||
|
||||
@@ -18,7 +18,7 @@ if (typeof marked !== 'undefined') {
|
||||
function escapeHtml(text) {
|
||||
var div = document.createElement('div');
|
||||
div.textContent = text || '';
|
||||
return div.innerHTML;
|
||||
return div.innerHTML.replace(/\n/g, '<br>');
|
||||
}
|
||||
|
||||
function renderMarkdown(text) {
|
||||
|
||||
@@ -337,29 +337,38 @@ function agentsPage() {
|
||||
OpenFangAPI.wsDisconnect();
|
||||
},
|
||||
|
||||
buildConfigForm(agent) {
|
||||
var identity = (agent && agent.identity) || {};
|
||||
return {
|
||||
name: (agent && agent.name) || '',
|
||||
system_prompt: (agent && agent.system_prompt) || '',
|
||||
emoji: identity.emoji || '',
|
||||
color: identity.color || '#FF5C00',
|
||||
archetype: identity.archetype || '',
|
||||
vibe: identity.vibe || ''
|
||||
};
|
||||
},
|
||||
|
||||
async showDetail(agent) {
|
||||
this.detailAgent = agent;
|
||||
this.detailAgent._fallbacks = [];
|
||||
this.detailTab = 'info';
|
||||
this.agentFiles = [];
|
||||
this.editingFile = null;
|
||||
this.fileContent = '';
|
||||
this.editingFallback = false;
|
||||
this.newFallbackValue = '';
|
||||
this.configForm = {
|
||||
name: agent.name || '',
|
||||
system_prompt: agent.system_prompt || '',
|
||||
emoji: (agent.identity && agent.identity.emoji) || '',
|
||||
color: (agent.identity && agent.identity.color) || '#FF5C00',
|
||||
archetype: (agent.identity && agent.identity.archetype) || '',
|
||||
vibe: (agent.identity && agent.identity.vibe) || ''
|
||||
};
|
||||
this.showDetailModal = true;
|
||||
// Fetch full agent detail to get fallback_models
|
||||
// Load the full detail payload before opening the modal so editable
|
||||
// fields such as system_prompt and identity metadata are hydrated.
|
||||
var detail = agent;
|
||||
try {
|
||||
var full = await OpenFangAPI.get('/api/agents/' + agent.id);
|
||||
this.detailAgent._fallbacks = full.fallback_models || [];
|
||||
} catch(e) { /* ignore */ }
|
||||
detail = Object.assign({}, agent, full, {
|
||||
identity: Object.assign({}, (agent && agent.identity) || {}, (full && full.identity) || {})
|
||||
});
|
||||
} catch(e) { /* fall back to list payload */ }
|
||||
this.detailAgent = detail;
|
||||
this.detailAgent._fallbacks = detail.fallback_models || [];
|
||||
this.configForm = this.buildConfigForm(detail);
|
||||
this.showDetailModal = true;
|
||||
},
|
||||
|
||||
killAgent(agent) {
|
||||
@@ -376,6 +385,29 @@ function agentsPage() {
|
||||
});
|
||||
},
|
||||
|
||||
// Issue #1163: uninstall an agent (kill + remove ~/.openfang/agents/<name>/).
|
||||
uninstallAgent(agent) {
|
||||
var self = this;
|
||||
OpenFangToast.confirm(
|
||||
'Uninstall Agent',
|
||||
'Uninstall agent "' + agent.name + '"? This stops the agent AND deletes its files from your workspace. This cannot be undone.',
|
||||
async function() {
|
||||
try {
|
||||
var res = await OpenFangAPI.del('/api/agents/' + agent.id + '/uninstall');
|
||||
var msg = 'Agent "' + agent.name + '" uninstalled';
|
||||
if (res && res.dir_removed === false) {
|
||||
msg += ' (no on-disk files found)';
|
||||
}
|
||||
OpenFangToast.success(msg);
|
||||
self.showDetailModal = false;
|
||||
await Alpine.store('app').refreshAgents();
|
||||
} catch(e) {
|
||||
OpenFangToast.error('Failed to uninstall agent: ' + e.message);
|
||||
}
|
||||
}
|
||||
);
|
||||
},
|
||||
|
||||
killAllAgents() {
|
||||
var list = this.filteredAgents;
|
||||
if (!list.length) return;
|
||||
|
||||
@@ -143,6 +143,29 @@ function chatPage() {
|
||||
// Fetch dynamic commands from server
|
||||
this.fetchCommands();
|
||||
|
||||
// Observe DOM for new messages and render LaTeX
|
||||
this._latexObserver = new MutationObserver(function(mutations) {
|
||||
mutations.forEach(function(mutation) {
|
||||
mutation.addedNodes.forEach(function(node) {
|
||||
if (node.nodeType === Node.ELEMENT_NODE) {
|
||||
var bubbles = node.querySelector ? node.querySelectorAll('.message-bubble') : [];
|
||||
if (node.classList && node.classList.contains('message-bubble')) {
|
||||
bubbles = [node];
|
||||
}
|
||||
bubbles.forEach(function(bubble) {
|
||||
if (bubble.textContent && hasLatexDelimiters(bubble.textContent)) {
|
||||
renderLatex(bubble);
|
||||
}
|
||||
});
|
||||
}
|
||||
});
|
||||
});
|
||||
});
|
||||
this._latexObserver.observe(document.getElementById('messages') || document.body, {
|
||||
childList: true,
|
||||
subtree: true
|
||||
});
|
||||
|
||||
// Ctrl+/ keyboard shortcut
|
||||
document.addEventListener('keydown', function(e) {
|
||||
if ((e.ctrlKey || e.metaKey) && e.key === '/') {
|
||||
@@ -175,6 +198,10 @@ function chatPage() {
|
||||
if (store.pendingAgent) {
|
||||
self.selectAgent(store.pendingAgent);
|
||||
store.pendingAgent = null;
|
||||
} else {
|
||||
// Restore previously active agent after page refresh (#1179).
|
||||
// The agent list may not be loaded yet, so resolve once it appears.
|
||||
self._restoreActiveAgent();
|
||||
}
|
||||
|
||||
// Watch for future pending agent selections (e.g., user clicks agent while on chat)
|
||||
@@ -185,6 +212,13 @@ function chatPage() {
|
||||
}
|
||||
});
|
||||
|
||||
// Re-attempt restore once the agent list arrives from the server
|
||||
this.$watch('$store.app.agents', function(agents) {
|
||||
if (!self.currentAgent && agents && agents.length) {
|
||||
self._restoreActiveAgent();
|
||||
}
|
||||
});
|
||||
|
||||
// Watch for slash commands + model autocomplete
|
||||
this.$watch('inputText', function(val) {
|
||||
var modelMatch = val.match(/^\/model\s+(.*)$/i);
|
||||
@@ -473,7 +507,7 @@ function chatPage() {
|
||||
if (self.currentAgent && OpenFangAPI.isWsConnected()) {
|
||||
OpenFangAPI.wsSend({ type: 'command', command: 'context', args: '' });
|
||||
} else {
|
||||
self.messages.push({ id: ++msgId, role: 'system', text: 'Not connected. Connect to an agent first.', meta: '', tools: [] });
|
||||
self.messages.push({ id: ++msgId, role: 'system', text: 'Not connected (' + (OpenFangAPI.getConnectionState ? OpenFangAPI.getConnectionState() : 'unknown') + '). Pick an agent or check that your session is still valid.', meta: '', tools: [] });
|
||||
self.scrollToBottom();
|
||||
}
|
||||
break;
|
||||
@@ -481,7 +515,7 @@ function chatPage() {
|
||||
if (self.currentAgent && OpenFangAPI.isWsConnected()) {
|
||||
OpenFangAPI.wsSend({ type: 'command', command: 'verbose', args: cmdArgs });
|
||||
} else {
|
||||
self.messages.push({ id: ++msgId, role: 'system', text: 'Not connected. Connect to an agent first.', meta: '', tools: [] });
|
||||
self.messages.push({ id: ++msgId, role: 'system', text: 'Not connected (' + (OpenFangAPI.getConnectionState ? OpenFangAPI.getConnectionState() : 'unknown') + '). Pick an agent or check that your session is still valid.', meta: '', tools: [] });
|
||||
self.scrollToBottom();
|
||||
}
|
||||
break;
|
||||
@@ -489,7 +523,7 @@ function chatPage() {
|
||||
if (self.currentAgent && OpenFangAPI.isWsConnected()) {
|
||||
OpenFangAPI.wsSend({ type: 'command', command: 'queue', args: '' });
|
||||
} else {
|
||||
self.messages.push({ id: ++msgId, role: 'system', text: 'Not connected.', meta: '', tools: [] });
|
||||
self.messages.push({ id: ++msgId, role: 'system', text: 'Not connected (' + (OpenFangAPI.getConnectionState ? OpenFangAPI.getConnectionState() : 'unknown') + ').', meta: '', tools: [] });
|
||||
self.scrollToBottom();
|
||||
}
|
||||
break;
|
||||
@@ -528,6 +562,7 @@ function chatPage() {
|
||||
self._wsAgent = null;
|
||||
self.currentAgent = null;
|
||||
self.messages = [];
|
||||
try { localStorage.removeItem('of-active-agent'); } catch(e) { /* ignore */ }
|
||||
window.dispatchEvent(new Event('close-chat'));
|
||||
break;
|
||||
case '/budget':
|
||||
@@ -563,9 +598,27 @@ function chatPage() {
|
||||
}
|
||||
},
|
||||
|
||||
// Restore the previously-active agent (set in selectAgent) after a page
|
||||
// refresh, so the WebSocket re-attaches to the same session and any
|
||||
// in-flight tool output streams back into the chat (#1179).
|
||||
_restoreActiveAgent: function() {
|
||||
var storedId = null;
|
||||
try { storedId = localStorage.getItem('of-active-agent'); } catch(e) { /* ignore */ }
|
||||
if (!storedId) return;
|
||||
var agents = (Alpine.store('app') && Alpine.store('app').agents) || [];
|
||||
var match = null;
|
||||
for (var i = 0; i < agents.length; i++) {
|
||||
if (agents[i] && agents[i].id === storedId) { match = agents[i]; break; }
|
||||
}
|
||||
if (match) {
|
||||
this.selectAgent(match);
|
||||
}
|
||||
},
|
||||
|
||||
selectAgent(agent) {
|
||||
this.currentAgent = agent;
|
||||
this.messages = [];
|
||||
try { localStorage.setItem('of-active-agent', agent.id); } catch(e) { /* ignore */ }
|
||||
this.connectWs(agent.id);
|
||||
var t = typeof window.t === 'function' ? window.t : function(s) { return s; };
|
||||
// Show welcome tips on first use
|
||||
@@ -695,6 +748,15 @@ function chatPage() {
|
||||
switch (data.type) {
|
||||
case 'connected': break;
|
||||
|
||||
// Incoming message from server (e.g., cron trigger) — display as user message
|
||||
case 'message':
|
||||
if (data.content) {
|
||||
var meta = data.source === 'cron' ? '[Scheduled: ' + (data.job_name || data.job_id || '') + ']' : '';
|
||||
this.messages.push({ id: ++msgId, role: 'user', text: data.content, meta: meta, tools: [], images: [], ts: Date.now() });
|
||||
this.scrollToBottom();
|
||||
}
|
||||
break;
|
||||
|
||||
// Legacy thinking event (backward compat)
|
||||
case 'thinking':
|
||||
if (!this.messages.length || !this.messages[this.messages.length - 1].thinking) {
|
||||
@@ -1141,6 +1203,7 @@ function chatPage() {
|
||||
self._wsAgent = null;
|
||||
self.currentAgent = null;
|
||||
self.messages = [];
|
||||
try { localStorage.removeItem('of-active-agent'); } catch(e) { /* ignore */ }
|
||||
OpenFangToast.success(t('chat.agent_stopped') + ' "' + name + '"');
|
||||
Alpine.store('app').refreshAgents();
|
||||
} catch(e) {
|
||||
@@ -1149,6 +1212,37 @@ function chatPage() {
|
||||
});
|
||||
},
|
||||
|
||||
// Permanently uninstall the agent: kill + remove ~/.openfang/agents/<name>/
|
||||
// Issue #1163.
|
||||
uninstallAgent: function() {
|
||||
if (!this.currentAgent) return;
|
||||
var self = this;
|
||||
var name = this.currentAgent.name;
|
||||
var agentId = this.currentAgent.id;
|
||||
OpenFangToast.confirm(
|
||||
'Uninstall Agent',
|
||||
'Uninstall agent "' + name + '"? This stops the agent AND deletes its files from your workspace. This cannot be undone.',
|
||||
async function() {
|
||||
try {
|
||||
var res = await OpenFangAPI.del('/api/agents/' + agentId + '/uninstall');
|
||||
OpenFangAPI.wsDisconnect();
|
||||
self._wsAgent = null;
|
||||
self.currentAgent = null;
|
||||
self.messages = [];
|
||||
try { localStorage.removeItem('of-active-agent'); } catch(e) { /* ignore */ }
|
||||
var msg = 'Agent "' + name + '" uninstalled';
|
||||
if (res && res.dir_removed === false) {
|
||||
msg += ' (no on-disk files found)';
|
||||
}
|
||||
OpenFangToast.success(msg);
|
||||
Alpine.store('app').refreshAgents();
|
||||
} catch(e) {
|
||||
OpenFangToast.error('Failed to uninstall agent: ' + e.message);
|
||||
}
|
||||
}
|
||||
);
|
||||
},
|
||||
|
||||
_latexTimer: null,
|
||||
scrollToBottom() {
|
||||
var self = this;
|
||||
|
||||
@@ -25,6 +25,9 @@ function settingsPage() {
|
||||
providerUrlSaving: {},
|
||||
providerTesting: {},
|
||||
providerTestResults: {},
|
||||
providerSearch: '',
|
||||
providerStatusFilter: '',
|
||||
providerCategoryFilter: '',
|
||||
copilotOAuth: { polling: false, userCode: '', verificationUri: '', pollId: '', interval: 5 },
|
||||
customProviderName: '',
|
||||
customProviderUrl: '',
|
||||
@@ -338,6 +341,94 @@ function settingsPage() {
|
||||
return Object.keys(seen).sort();
|
||||
},
|
||||
|
||||
/// Coarse category for a provider used to group the Providers tab.
|
||||
/// Returns: 'frontier' | 'oss' | 'local' | 'aggregator' | 'regional' | 'other'.
|
||||
providerCategory(p) {
|
||||
if (!p) return 'other';
|
||||
if (p.is_local || p.key_required === false) return 'local';
|
||||
var id = (p.id || '').toLowerCase();
|
||||
var FRONTIER = ['anthropic','openai','gemini','google','xai','bedrock','azure','vertex'];
|
||||
var OSS = ['groq','together','fireworks','cerebras','sambanova','deepseek','mistral','perplexity','cohere','ai21','huggingface','replicate','nvidia','venice','novita','chutes'];
|
||||
var AGG = ['openrouter','litellm','github-copilot','claude-code'];
|
||||
var REGIONAL = ['qwen','minimax','zhipu','zai','moonshot','qianfan','volcengine','kimi'];
|
||||
if (FRONTIER.indexOf(id) !== -1) return 'frontier';
|
||||
if (REGIONAL.indexOf(id) !== -1) return 'regional';
|
||||
if (AGG.indexOf(id) !== -1) return 'aggregator';
|
||||
if (OSS.indexOf(id) !== -1) return 'oss';
|
||||
return 'other';
|
||||
},
|
||||
|
||||
providerCategoryLabel(cat) {
|
||||
switch (cat) {
|
||||
case 'frontier': return 'Frontier (Anthropic, OpenAI, Google, xAI, Bedrock)';
|
||||
case 'oss': return 'Open-Weight Hosts (Groq, Together, Fireworks, DeepSeek, etc.)';
|
||||
case 'aggregator': return 'Aggregators & Gateways (OpenRouter, GitHub Copilot)';
|
||||
case 'regional': return 'Regional / China (Qwen, Zhipu, Moonshot, MiniMax)';
|
||||
case 'local': return 'Local / Self-Hosted (Ollama, vLLM, LM Studio, Lemonade)';
|
||||
default: return 'Other Providers';
|
||||
}
|
||||
},
|
||||
|
||||
/// Stable category order for grouped rendering.
|
||||
get providerCategoriesOrdered() {
|
||||
return ['frontier', 'oss', 'aggregator', 'regional', 'local', 'other'];
|
||||
},
|
||||
|
||||
/// Returns filter-matched providers grouped by category, preserving order.
|
||||
/// Each entry: { category, label, items: [...] }. Empty groups are omitted.
|
||||
get providersGrouped() {
|
||||
var self = this;
|
||||
var filtered = this.filteredProviders;
|
||||
var by = {};
|
||||
filtered.forEach(function(p) {
|
||||
var c = self.providerCategory(p);
|
||||
if (!by[c]) by[c] = [];
|
||||
by[c].push(p);
|
||||
});
|
||||
// Sort each group: configured first, then alphabetical
|
||||
Object.keys(by).forEach(function(c) {
|
||||
by[c].sort(function(a, b) {
|
||||
var ac = a.auth_status === 'configured' ? 0 : 1;
|
||||
var bc = b.auth_status === 'configured' ? 0 : 1;
|
||||
if (ac !== bc) return ac - bc;
|
||||
return (a.display_name || a.id).localeCompare(b.display_name || b.id);
|
||||
});
|
||||
});
|
||||
var out = [];
|
||||
this.providerCategoriesOrdered.forEach(function(c) {
|
||||
if (by[c] && by[c].length) {
|
||||
out.push({ category: c, label: self.providerCategoryLabel(c), items: by[c] });
|
||||
}
|
||||
});
|
||||
return out;
|
||||
},
|
||||
|
||||
get filteredProviders() {
|
||||
var self = this;
|
||||
return this.providers.filter(function(p) {
|
||||
if (self.providerStatusFilter === 'configured' && p.auth_status !== 'configured') return false;
|
||||
if (self.providerStatusFilter === 'unconfigured' && p.auth_status === 'configured') return false;
|
||||
if (self.providerCategoryFilter && self.providerCategory(p) !== self.providerCategoryFilter) return false;
|
||||
if (self.providerSearch) {
|
||||
var q = self.providerSearch.toLowerCase();
|
||||
if ((p.display_name || '').toLowerCase().indexOf(q) === -1 &&
|
||||
(p.id || '').toLowerCase().indexOf(q) === -1 &&
|
||||
(p.api_key_env || '').toLowerCase().indexOf(q) === -1) return false;
|
||||
}
|
||||
return true;
|
||||
});
|
||||
},
|
||||
|
||||
get configuredProviderCount() {
|
||||
return this.providers.filter(function(p) { return p.auth_status === 'configured'; }).length;
|
||||
},
|
||||
|
||||
clearProviderFilters() {
|
||||
this.providerSearch = '';
|
||||
this.providerStatusFilter = '';
|
||||
this.providerCategoryFilter = '';
|
||||
},
|
||||
|
||||
get uniqueTiers() {
|
||||
var seen = {};
|
||||
this.models.forEach(function(m) { if (m.tier) seen[m.tier] = true; });
|
||||
|
||||
@@ -61,6 +61,7 @@ async fn start_test_server_with_provider(
|
||||
model: model.to_string(),
|
||||
api_key_env: api_key_env.to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
};
|
||||
@@ -101,6 +102,10 @@ async fn start_test_server_with_provider(
|
||||
"/api/agents/{id}",
|
||||
axum::routing::delete(routes::kill_agent),
|
||||
)
|
||||
.route(
|
||||
"/api/agents/{id}/clone",
|
||||
axum::routing::post(routes::clone_agent),
|
||||
)
|
||||
.route(
|
||||
"/api/triggers",
|
||||
axum::routing::get(routes::list_triggers).post(routes::create_trigger),
|
||||
@@ -301,6 +306,86 @@ async fn test_spawn_list_kill_agent() {
|
||||
assert_eq!(agents[0]["name"], "assistant");
|
||||
}
|
||||
|
||||
/// Regression test for issue #1026: GET /api/agents returns `is_inferencing`
|
||||
/// reflecting whether the agent has an in-flight LLM task. This drives the
|
||||
/// live dashboard indicator that shows which agents are calling the LLM.
|
||||
#[tokio::test]
|
||||
async fn test_list_agents_includes_inferencing_flag() {
|
||||
let server = start_test_server().await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
// Spawn a test agent.
|
||||
let resp = client
|
||||
.post(format!("{}/api/agents", server.base_url))
|
||||
.json(&serde_json::json!({"manifest_toml": TEST_MANIFEST}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 201);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
let agent_id_str = body["agent_id"].as_str().unwrap().to_string();
|
||||
let agent_id: openfang_types::agent::AgentId = agent_id_str.parse().unwrap();
|
||||
|
||||
// Baseline: idle agent must report is_inferencing = false.
|
||||
let resp = client
|
||||
.get(format!("{}/api/agents", server.base_url))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 200);
|
||||
let agents: Vec<serde_json::Value> = resp.json().await.unwrap();
|
||||
let test_agent = agents
|
||||
.iter()
|
||||
.find(|a| a["id"] == agent_id_str)
|
||||
.expect("spawned agent should appear in list");
|
||||
assert_eq!(
|
||||
test_agent["is_inferencing"], false,
|
||||
"freshly spawned agent should not be inferencing"
|
||||
);
|
||||
|
||||
// Simulate an in-flight LLM call by inserting a real AbortHandle into
|
||||
// the kernel's running_tasks map. This is exactly what the agent loop
|
||||
// does when it starts processing a message.
|
||||
let handle = tokio::spawn(async {
|
||||
// Long-lived task we will abort at end of test.
|
||||
tokio::time::sleep(std::time::Duration::from_secs(60)).await;
|
||||
});
|
||||
server
|
||||
.state
|
||||
.kernel
|
||||
.running_tasks
|
||||
.insert(agent_id, handle.abort_handle());
|
||||
|
||||
// Now list_agents should report is_inferencing = true for that agent.
|
||||
let resp = client
|
||||
.get(format!("{}/api/agents", server.base_url))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 200);
|
||||
let agents: Vec<serde_json::Value> = resp.json().await.unwrap();
|
||||
let test_agent = agents
|
||||
.iter()
|
||||
.find(|a| a["id"] == agent_id_str)
|
||||
.expect("spawned agent should still appear in list");
|
||||
assert_eq!(
|
||||
test_agent["is_inferencing"], true,
|
||||
"agent with an entry in running_tasks must be flagged is_inferencing"
|
||||
);
|
||||
|
||||
// Other agents (the default assistant) must NOT be flagged.
|
||||
if let Some(other) = agents.iter().find(|a| a["id"] != agent_id_str) {
|
||||
assert_eq!(
|
||||
other["is_inferencing"], false,
|
||||
"agents without a running task must not be flagged"
|
||||
);
|
||||
}
|
||||
|
||||
// Cleanup so the spawned future does not outlive the test.
|
||||
server.state.kernel.running_tasks.remove(&agent_id);
|
||||
handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_agent_session_empty() {
|
||||
let server = start_test_server().await;
|
||||
@@ -379,6 +464,7 @@ async fn test_agent_session_filters_system_messages() {
|
||||
content: openfang_types::message::MessageContent::Text(
|
||||
"INTERNAL SYSTEM PROMPT — must not leak to UI".to_string(),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
Message::user("hello"),
|
||||
Message::assistant("hi there"),
|
||||
@@ -820,6 +906,7 @@ async fn start_test_server_with_auth(api_key: &str) -> TestServer {
|
||||
model: "test-model".to_string(),
|
||||
api_key_env: "OLLAMA_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
};
|
||||
@@ -874,6 +961,10 @@ async fn start_test_server_with_auth(api_key: &str) -> TestServer {
|
||||
"/api/agents/{id}",
|
||||
axum::routing::delete(routes::kill_agent),
|
||||
)
|
||||
.route(
|
||||
"/api/agents/{id}/clone",
|
||||
axum::routing::post(routes::clone_agent),
|
||||
)
|
||||
.route(
|
||||
"/api/triggers",
|
||||
axum::routing::get(routes::list_triggers).post(routes::create_trigger),
|
||||
@@ -1156,7 +1247,10 @@ async fn test_commands_invalid_surface_400() {
|
||||
assert_eq!(resp.status(), 400);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
let err = body["error"].as_str().unwrap_or_default();
|
||||
assert!(err.contains("bogus"), "error should mention the bad value: {err}");
|
||||
assert!(
|
||||
err.contains("bogus"),
|
||||
"error should mention the bad value: {err}"
|
||||
);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -1211,7 +1305,10 @@ async fn test_schedules_delivery_targets_roundtrip() {
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 201);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
let sched_id = body["id"].as_str().expect("created schedule id").to_string();
|
||||
let sched_id = body["id"]
|
||||
.as_str()
|
||||
.expect("created schedule id")
|
||||
.to_string();
|
||||
let got = body["delivery_targets"]
|
||||
.as_array()
|
||||
.expect("response must include delivery_targets");
|
||||
@@ -1291,7 +1388,9 @@ async fn test_schedules_delivery_targets_update() {
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert_eq!(body["status"], "updated");
|
||||
let echoed = &body["schedule"]["delivery_targets"];
|
||||
let arr = echoed.as_array().expect("schedule.delivery_targets must be array");
|
||||
let arr = echoed
|
||||
.as_array()
|
||||
.expect("schedule.delivery_targets must be array");
|
||||
assert_eq!(arr.len(), 2);
|
||||
assert_eq!(arr[0]["type"], "webhook");
|
||||
assert_eq!(arr[1]["type"], "local_file");
|
||||
@@ -1406,10 +1505,7 @@ async fn test_schedules_delivery_log_endpoint() {
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 201);
|
||||
let sched_id = resp
|
||||
.json::<serde_json::Value>()
|
||||
.await
|
||||
.unwrap()["id"]
|
||||
let sched_id = resp.json::<serde_json::Value>().await.unwrap()["id"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.to_string();
|
||||
@@ -1510,3 +1606,195 @@ async fn test_cron_jobs_delivery_targets_roundtrip() {
|
||||
assert_eq!(targets[1]["type"], "webhook");
|
||||
assert_eq!(targets[1]["url"], "http://example.com/pulse");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Clone agent endpoint tests (issue #868)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Happy path: clone an existing template agent into a new agent with a
|
||||
/// distinct name. The clone must get a fresh ID, fresh workspace path, and
|
||||
/// inherit non-name manifest fields from the template.
|
||||
#[tokio::test]
|
||||
async fn test_clone_agent_happy_path() {
|
||||
let server = start_test_server().await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
// Spawn a template agent.
|
||||
let resp = client
|
||||
.post(format!("{}/api/agents", server.base_url))
|
||||
.json(&serde_json::json!({"manifest_toml": TEST_MANIFEST}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 201);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
let template_id = body["agent_id"].as_str().unwrap().to_string();
|
||||
|
||||
// Clone it.
|
||||
let resp = client
|
||||
.post(format!(
|
||||
"{}/api/agents/{}/clone",
|
||||
server.base_url, template_id
|
||||
))
|
||||
.json(&serde_json::json!({
|
||||
"new_name": "cloned-user-1",
|
||||
"overrides": {
|
||||
"description": "Cloned for user 1",
|
||||
"tags": ["clone", "user-1"]
|
||||
}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 201, "clone should succeed");
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
|
||||
let new_id = body["agent_id"].as_str().unwrap();
|
||||
assert_ne!(new_id, template_id, "clone must have a fresh agent ID");
|
||||
assert_eq!(body["name"], "cloned-user-1");
|
||||
|
||||
// The full manifest should be returned and reflect the new name + overrides.
|
||||
let manifest = &body["manifest"];
|
||||
assert!(manifest.is_object(), "manifest must be returned");
|
||||
assert_eq!(manifest["name"], "cloned-user-1");
|
||||
assert_eq!(manifest["description"], "Cloned for user 1");
|
||||
assert_eq!(
|
||||
manifest["tags"].as_array().unwrap(),
|
||||
&vec![serde_json::json!("clone"), serde_json::json!("user-1"),]
|
||||
);
|
||||
// Inherited from template — the system_prompt should match.
|
||||
assert_eq!(
|
||||
manifest["model"]["system_prompt"],
|
||||
"You are a test agent. Reply concisely."
|
||||
);
|
||||
|
||||
// The agent list should now contain both template and clone.
|
||||
let resp = client
|
||||
.get(format!("{}/api/agents", server.base_url))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
let agents: Vec<serde_json::Value> = resp.json().await.unwrap();
|
||||
let names: Vec<&str> = agents.iter().map(|a| a["name"].as_str().unwrap()).collect();
|
||||
assert!(names.contains(&"test-agent"));
|
||||
assert!(names.contains(&"cloned-user-1"));
|
||||
}
|
||||
|
||||
/// Cloning into a name that's already taken must fail with 409 Conflict.
|
||||
#[tokio::test]
|
||||
async fn test_clone_agent_name_collision() {
|
||||
let server = start_test_server().await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
// Spawn a template agent named "test-agent".
|
||||
let resp = client
|
||||
.post(format!("{}/api/agents", server.base_url))
|
||||
.json(&serde_json::json!({"manifest_toml": TEST_MANIFEST}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 201);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
let template_id = body["agent_id"].as_str().unwrap().to_string();
|
||||
|
||||
// First clone — succeeds.
|
||||
let resp = client
|
||||
.post(format!(
|
||||
"{}/api/agents/{}/clone",
|
||||
server.base_url, template_id
|
||||
))
|
||||
.json(&serde_json::json!({"new_name": "duplicate-name"}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 201);
|
||||
|
||||
// Second clone with the same name — must be rejected.
|
||||
let resp = client
|
||||
.post(format!(
|
||||
"{}/api/agents/{}/clone",
|
||||
server.base_url, template_id
|
||||
))
|
||||
.json(&serde_json::json!({"new_name": "duplicate-name"}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
resp.status(),
|
||||
409,
|
||||
"duplicate name must return 409 Conflict"
|
||||
);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert!(body["error"].as_str().unwrap().contains("already exists"));
|
||||
|
||||
// Cloning into the template's own name must also be rejected.
|
||||
let resp = client
|
||||
.post(format!(
|
||||
"{}/api/agents/{}/clone",
|
||||
server.base_url, template_id
|
||||
))
|
||||
.json(&serde_json::json!({"new_name": "test-agent"}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 409);
|
||||
}
|
||||
|
||||
/// Cloning a non-existent template must return 404.
|
||||
#[tokio::test]
|
||||
async fn test_clone_agent_template_not_found() {
|
||||
let server = start_test_server().await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
// Random valid UUID that does not match any agent.
|
||||
let bogus_id = "00000000-0000-0000-0000-000000000000";
|
||||
let resp = client
|
||||
.post(format!("{}/api/agents/{}/clone", server.base_url, bogus_id))
|
||||
.json(&serde_json::json!({"new_name": "ghost-clone"}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 404);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert!(body["error"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("Template agent not found"));
|
||||
|
||||
// Malformed agent id → 400.
|
||||
let resp = client
|
||||
.post(format!("{}/api/agents/not-a-uuid/clone", server.base_url))
|
||||
.json(&serde_json::json!({"new_name": "ghost-clone"}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 400);
|
||||
}
|
||||
|
||||
/// Empty new_name must be rejected with 400.
|
||||
#[tokio::test]
|
||||
async fn test_clone_agent_empty_name_rejected() {
|
||||
let server = start_test_server().await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
// Spawn a template agent.
|
||||
let resp = client
|
||||
.post(format!("{}/api/agents", server.base_url))
|
||||
.json(&serde_json::json!({"manifest_toml": TEST_MANIFEST}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
let template_id = body["agent_id"].as_str().unwrap().to_string();
|
||||
|
||||
let resp = client
|
||||
.post(format!(
|
||||
"{}/api/agents/{}/clone",
|
||||
server.base_url, template_id
|
||||
))
|
||||
.json(&serde_json::json!({"new_name": " "}))
|
||||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 400);
|
||||
}
|
||||
|
||||
@@ -98,6 +98,7 @@ async fn test_full_daemon_lifecycle() {
|
||||
model: "test".to_string(),
|
||||
api_key_env: "OLLAMA_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
};
|
||||
@@ -225,6 +226,7 @@ async fn test_server_immediate_responsiveness() {
|
||||
model: "test".to_string(),
|
||||
api_key_env: "OLLAMA_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
};
|
||||
|
||||
@@ -42,6 +42,7 @@ async fn start_test_server() -> TestServer {
|
||||
model: "test-model".to_string(),
|
||||
api_key_env: "OLLAMA_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
};
|
||||
|
||||
@@ -76,6 +76,7 @@ async fn start_test_server() -> TestServer {
|
||||
model: "test-model".to_string(),
|
||||
api_key_env: "OLLAMA_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
};
|
||||
@@ -166,7 +167,9 @@ async fn get_config_returns_declared_and_resolved() {
|
||||
body["resolved"]["github_token"]["source"], "unresolved",
|
||||
"github_token should be unresolved without env"
|
||||
);
|
||||
assert!(body["resolved"]["github_token"]["is_secret"].as_bool().unwrap());
|
||||
assert!(body["resolved"]["github_token"]["is_secret"]
|
||||
.as_bool()
|
||||
.unwrap());
|
||||
|
||||
// default_branch falls back to default "main".
|
||||
assert_eq!(body["resolved"]["default_branch"]["source"], "default");
|
||||
@@ -255,10 +258,7 @@ async fn put_rejects_unknown_variable() {
|
||||
.unwrap();
|
||||
assert_eq!(resp.status(), 400);
|
||||
let body: serde_json::Value = resp.json().await.unwrap();
|
||||
assert!(body["error"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("nonexistent_var"));
|
||||
assert!(body["error"].as_str().unwrap().contains("nonexistent_var"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -377,12 +377,7 @@ async fn put_reloads_registry_so_agents_see_change() {
|
||||
.unwrap();
|
||||
|
||||
// The kernel's live override map must now hold the new values.
|
||||
let guard = server
|
||||
.state
|
||||
.kernel
|
||||
.skill_config_overrides
|
||||
.read()
|
||||
.unwrap();
|
||||
let guard = server.state.kernel.skill_config_overrides.read().unwrap();
|
||||
let overrides = guard.as_ref().expect("override map set after PUT");
|
||||
let skill_cfg = overrides.get("test-config-skill").expect("skill present");
|
||||
assert_eq!(skill_cfg.get("github_token").unwrap(), "ghp_new");
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
//! `BridgeManager` which owns running adapters and dispatches messages.
|
||||
|
||||
use crate::formatter;
|
||||
use crate::router::AgentRouter;
|
||||
use crate::router::{AgentRouter, BindingContext};
|
||||
use crate::types::{
|
||||
default_phase_emoji, AgentPhase, ChannelAdapter, ChannelContent, ChannelMessage, ChannelUser,
|
||||
LifecycleReaction,
|
||||
@@ -421,6 +421,26 @@ impl BridgeManager {
|
||||
adapter: Arc<dyn ChannelAdapter>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let stream = adapter.start().await?;
|
||||
|
||||
// Migration note for Discord/Slack: prior versions keyed `/agent <name>`
|
||||
// selections on the channel ID rather than the user. `user_defaults` is
|
||||
// in-memory only, so the daemon restart that loads this binary already
|
||||
// wipes any stale entries — but log a one-line nudge so users know to
|
||||
// re-run `/agent <name>` if their previous selection appears to have
|
||||
// gone away. See `set_user_default` call sites in `dispatch_message`
|
||||
// and `handle_command` for the keying fix.
|
||||
match adapter.name() {
|
||||
"discord" | "slack" => {
|
||||
info!(
|
||||
adapter = adapter.name(),
|
||||
"Channel adapter starting: per-user `/agent <name>` defaults are \
|
||||
in-memory and reset on daemon restart. If a previous selection \
|
||||
no longer takes effect, re-run `/agent <name>` once."
|
||||
);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
let handle = self.handle.clone();
|
||||
let router = self.router.clone();
|
||||
let rate_limiter = self.rate_limiter.clone();
|
||||
@@ -659,6 +679,36 @@ fn sender_user_id(message: &ChannelMessage) -> &str {
|
||||
.unwrap_or(&message.sender.platform_id)
|
||||
}
|
||||
|
||||
/// Build a `BindingContext` for routing the given inbound message.
|
||||
///
|
||||
/// Populates `channel_id` so per-channel bindings (e.g. `channel_id = "<discord_channel>"`)
|
||||
/// can route to dedicated agents. The channel ID source is delegated to
|
||||
/// [`ChannelMessage::channel_id`] — the single source of truth shared with
|
||||
/// config validation (see `CHANNELS_WITH_PLATFORM_ID_AS_CHANNEL` in
|
||||
/// `openfang-types::config`). `peer_id` uses the resolved user ID, not
|
||||
/// `sender.platform_id`, so user-scoped bindings still match correctly on
|
||||
/// Discord/Slack/etc. where `platform_id` holds the channel.
|
||||
///
|
||||
/// This replaces the earlier heuristic `sender_channel_id()` (which inferred
|
||||
/// "platform_id is the channel" from "metadata has `sender_user_id`"). The
|
||||
/// allowlist is explicit, the metadata-fallback path is documented, and
|
||||
/// adapters can be added or removed in one place (`openfang-types::config`)
|
||||
/// without touching this file.
|
||||
fn binding_context_for(message: &ChannelMessage) -> BindingContext {
|
||||
BindingContext {
|
||||
channel: channel_type_str(&message.channel).to_string(),
|
||||
account_id: None,
|
||||
peer_id: sender_user_id(message).to_string(),
|
||||
channel_id: message.channel_id(),
|
||||
guild_id: message
|
||||
.metadata
|
||||
.get("guild_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from),
|
||||
roles: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// If an error contains "Agent not found", try to re-resolve the channel's default agent
|
||||
/// by name (the name stored at bridge startup). Returns `Some(new_id)` on success.
|
||||
async fn try_reresolution(
|
||||
@@ -723,12 +773,20 @@ async fn dispatch_message(
|
||||
.as_ref()
|
||||
.map(|o| o.lifecycle_reactions)
|
||||
.unwrap_or(true);
|
||||
let thread_id = if threading_enabled {
|
||||
message.thread_id.as_deref()
|
||||
|
||||
// --- Auto-thread: decide intent now, but create AFTER all policy guards ---
|
||||
let auto_thread_name = if !threading_enabled && message.thread_id.is_none() {
|
||||
adapter.should_auto_thread(message).await
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// thread_id is resolved later, after all guards pass.
|
||||
// Always propagate an existing thread_id (message arrived inside a thread),
|
||||
// regardless of threading_enabled — that flag controls explicit threading config,
|
||||
// not auto-detected thread context.
|
||||
let mut effective_thread_id: Option<String> = message.thread_id.clone();
|
||||
|
||||
// --- DM/Group policy check ---
|
||||
if let Some(ref ov) = overrides {
|
||||
if message.is_group {
|
||||
@@ -789,19 +847,144 @@ async fn dispatch_message(
|
||||
if let Err(msg) =
|
||||
rate_limiter.check(ct_str, sender_user_id(message), ov.rate_limit_per_user)
|
||||
{
|
||||
send_response(adapter, &message.sender, msg, thread_id, output_format).await;
|
||||
// Rate-limit rejection: don't create a thread, use existing thread if any
|
||||
send_response(
|
||||
adapter,
|
||||
&message.sender,
|
||||
msg,
|
||||
message.thread_id.as_deref(),
|
||||
output_format,
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Create auto-thread NOW (after all policy guards have passed) ---
|
||||
if let Some(ref thread_name) = auto_thread_name {
|
||||
match adapter
|
||||
.create_thread(&message.sender, &message.platform_message_id, thread_name)
|
||||
.await
|
||||
{
|
||||
Ok(new_thread_id) => {
|
||||
info!(
|
||||
"Created auto-thread {} for message {}",
|
||||
thread_name, message.platform_message_id
|
||||
);
|
||||
effective_thread_id = Some(new_thread_id);
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("Failed to create auto-thread: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Resolve final thread_id reference used by all downstream send_response calls
|
||||
let thread_id = effective_thread_id.as_deref();
|
||||
|
||||
// Handle commands first (early return)
|
||||
if let ChannelContent::Command { ref name, ref args } = message.content {
|
||||
let result = handle_command(name, args, handle, router, &message.sender).await;
|
||||
let result = handle_command(
|
||||
name,
|
||||
args,
|
||||
handle,
|
||||
router,
|
||||
&message.sender,
|
||||
sender_user_id(message),
|
||||
)
|
||||
.await;
|
||||
send_response(adapter, &message.sender, result, thread_id, output_format).await;
|
||||
return;
|
||||
}
|
||||
|
||||
// Multipart: flatten children into LLM content blocks. If any image
|
||||
// succeeds, dispatch as multimodal; otherwise fall through to the text
|
||||
// path (Multipart arm in the match below builds the combined descriptor).
|
||||
if let ChannelContent::Multipart(parts) = &message.content {
|
||||
let mut blocks: Vec<ContentBlock> = Vec::new();
|
||||
for part in parts {
|
||||
debug_assert!(
|
||||
!matches!(part, ChannelContent::Multipart(_)),
|
||||
"nested Multipart in ChannelContent — adapters should produce flat lists"
|
||||
);
|
||||
match part {
|
||||
ChannelContent::Text(t) => blocks.push(ContentBlock::Text {
|
||||
text: t.clone(),
|
||||
provider_metadata: None,
|
||||
}),
|
||||
ChannelContent::Image { url, caption } => {
|
||||
let mut img = download_image_to_blocks(url, caption.as_deref()).await;
|
||||
blocks.append(&mut img);
|
||||
}
|
||||
ChannelContent::File { url, filename, .. } => {
|
||||
blocks.push(ContentBlock::Text {
|
||||
text: format!("[User sent a file ({filename}): {url}]"),
|
||||
provider_metadata: None,
|
||||
});
|
||||
}
|
||||
ChannelContent::Voice {
|
||||
url,
|
||||
duration_seconds,
|
||||
} => {
|
||||
blocks.push(ContentBlock::Text {
|
||||
text: format!("[User sent a voice message ({duration_seconds}s): {url}]"),
|
||||
provider_metadata: None,
|
||||
});
|
||||
}
|
||||
ChannelContent::Location { lat, lon } => {
|
||||
blocks.push(ContentBlock::Text {
|
||||
text: format!("[User shared location: {lat}, {lon}]"),
|
||||
provider_metadata: None,
|
||||
});
|
||||
}
|
||||
ChannelContent::FileData { filename, .. } => {
|
||||
blocks.push(ContentBlock::Text {
|
||||
text: format!("[User sent a local file: {filename}]"),
|
||||
provider_metadata: None,
|
||||
});
|
||||
}
|
||||
// Commands aren't expected inside Multipart, but render as
|
||||
// text rather than drop the message if one slips through.
|
||||
ChannelContent::Command { name, args } => {
|
||||
blocks.push(ContentBlock::Text {
|
||||
text: format!("/{name} {}", args.join(" ")),
|
||||
provider_metadata: None,
|
||||
});
|
||||
}
|
||||
// Defensive: debug_assert above catches this in dev; ignore
|
||||
// gracefully in release.
|
||||
ChannelContent::Multipart(_) => {}
|
||||
}
|
||||
}
|
||||
|
||||
if blocks
|
||||
.iter()
|
||||
.any(|b| matches!(b, ContentBlock::Image { .. }))
|
||||
{
|
||||
let prefix_style = overrides
|
||||
.as_ref()
|
||||
.map(|o| o.prefix_agent_name)
|
||||
.unwrap_or(PrefixStyle::Off);
|
||||
dispatch_with_blocks(
|
||||
blocks,
|
||||
message,
|
||||
handle,
|
||||
router,
|
||||
adapter,
|
||||
adapter_arc,
|
||||
ct_str,
|
||||
thread_id,
|
||||
output_format,
|
||||
lifecycle_reactions,
|
||||
prefix_style,
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
// No image blocks — fall through to text path below.
|
||||
}
|
||||
|
||||
// For images: download, base64 encode, and send as multimodal content blocks
|
||||
if let ChannelContent::Image {
|
||||
ref url,
|
||||
@@ -853,6 +1036,7 @@ async fn dispatch_message(
|
||||
ChannelContent::File {
|
||||
ref url,
|
||||
ref filename,
|
||||
..
|
||||
} => {
|
||||
format!("[User sent a file ({filename}): {url}]")
|
||||
}
|
||||
@@ -868,6 +1052,37 @@ async fn dispatch_message(
|
||||
ChannelContent::FileData { ref filename, .. } => {
|
||||
format!("[User sent a local file: {filename}]")
|
||||
}
|
||||
ChannelContent::Multipart(parts) => parts
|
||||
.iter()
|
||||
.map(|p| match p {
|
||||
ChannelContent::Text(t) => t.clone(),
|
||||
ChannelContent::Image { url, caption } => match caption {
|
||||
Some(c) => format!("[User sent a photo: {url}]\nCaption: {c}"),
|
||||
None => format!("[User sent a photo: {url}]"),
|
||||
},
|
||||
ChannelContent::File { url, filename, .. } => {
|
||||
format!("[User sent a file ({filename}): {url}]")
|
||||
}
|
||||
ChannelContent::Voice {
|
||||
url,
|
||||
duration_seconds,
|
||||
} => format!("[User sent a voice message ({duration_seconds}s): {url}]"),
|
||||
ChannelContent::Location { lat, lon } => {
|
||||
format!("[User shared location: {lat}, {lon}]")
|
||||
}
|
||||
ChannelContent::FileData { filename, .. } => {
|
||||
format!("[User sent a local file: {filename}]")
|
||||
}
|
||||
ChannelContent::Command { name, args } => {
|
||||
format!("/{name} {}", args.join(" "))
|
||||
}
|
||||
// Nesting is rejected by adapters; emit empty so the join
|
||||
// doesn't insert spurious separators.
|
||||
ChannelContent::Multipart(_) => String::new(),
|
||||
})
|
||||
.filter(|s| !s.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
};
|
||||
|
||||
// Check if it's a slash command embedded in text (e.g. "/agents")
|
||||
@@ -881,7 +1096,15 @@ async fn dispatch_message(
|
||||
};
|
||||
|
||||
if is_channel_command(cmd) {
|
||||
let result = handle_command(cmd, &args, handle, router, &message.sender).await;
|
||||
let result = handle_command(
|
||||
cmd,
|
||||
&args,
|
||||
handle,
|
||||
router,
|
||||
&message.sender,
|
||||
sender_user_id(message),
|
||||
)
|
||||
.await;
|
||||
send_response(adapter, &message.sender, result, thread_id, output_format).await;
|
||||
return;
|
||||
}
|
||||
@@ -889,8 +1112,12 @@ async fn dispatch_message(
|
||||
}
|
||||
|
||||
// Check broadcast routing first
|
||||
if router.has_broadcast(&message.sender.platform_id) {
|
||||
let targets = router.resolve_broadcast(&message.sender.platform_id);
|
||||
// Broadcast lookup is keyed on the user, matching the read path's
|
||||
// sender_user_id() resolution. On Discord/Slack `sender.platform_id` is the
|
||||
// channel ID, so keying on it would collide with channel routing — see the
|
||||
// companion fix on `set_user_default` writes below.
|
||||
if router.has_broadcast(sender_user_id(message)) {
|
||||
let targets = router.resolve_broadcast(sender_user_id(message));
|
||||
if !targets.is_empty() {
|
||||
// RBAC check applies to broadcast too
|
||||
if let Err(denied) = handle
|
||||
@@ -958,12 +1185,39 @@ async fn dispatch_message(
|
||||
}
|
||||
}
|
||||
|
||||
// Route to agent (standard path)
|
||||
let agent_id = router.resolve(
|
||||
&message.channel,
|
||||
&message.sender.platform_id,
|
||||
message.sender.openfang_user.as_deref(),
|
||||
);
|
||||
// Route to agent (standard path).
|
||||
// Use sender_user_id() so user-keyed bindings (peer_id) match for adapters like
|
||||
// Discord/Slack where sender.platform_id is the channel ID, not the user ID.
|
||||
// Use resolve_with_context so channel_id-scoped (and guild_id-scoped)
|
||||
// bindings can route per channel — see binding_context_for() for the
|
||||
// single-source-of-truth allowlist.
|
||||
//
|
||||
// Issue #780: when the adapter stamped a per-thread target agent in
|
||||
// metadata (e.g. Telegram forum-topic routing via `thread_routes`), prefer
|
||||
// it over the standard router so operators can scope topics to specific
|
||||
// agents from config.toml.
|
||||
let target_agent_name = message
|
||||
.metadata
|
||||
.get("target_agent_name")
|
||||
.and_then(|v| v.as_str());
|
||||
let routed_by_name = if let Some(name) = target_agent_name {
|
||||
match handle.find_agent_by_name(name).await {
|
||||
Ok(Some(id)) => Some(id),
|
||||
_ => None,
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let binding_ctx = binding_context_for(message);
|
||||
let agent_id = routed_by_name.or_else(|| {
|
||||
router.resolve_with_context(
|
||||
&message.channel,
|
||||
sender_user_id(message),
|
||||
message.sender.openfang_user.as_deref(),
|
||||
&binding_ctx,
|
||||
)
|
||||
});
|
||||
|
||||
let agent_id = match agent_id {
|
||||
Some(id) => id,
|
||||
@@ -980,8 +1234,10 @@ async fn dispatch_message(
|
||||
};
|
||||
match fallback {
|
||||
Some(id) => {
|
||||
// Auto-set this as the user's default so future messages route directly
|
||||
router.set_user_default(message.sender.platform_id.clone(), id);
|
||||
// Auto-set this as the user's default so future messages route directly.
|
||||
// Key on sender_user_id() (not platform_id) so Discord/Slack — where
|
||||
// platform_id is the channel — store per-user, matching the read path.
|
||||
router.set_user_default(sender_user_id(message).to_string(), id);
|
||||
id
|
||||
}
|
||||
None => {
|
||||
@@ -1051,15 +1307,24 @@ async fn dispatch_message(
|
||||
|
||||
// Prepend sender context so the agent knows who is speaking.
|
||||
// In group spaces this is essential for multi-user conversations.
|
||||
//
|
||||
// For Telegram we also inject the numeric `tg_id` because display names are
|
||||
// not unique and can change — agents that key per-user state (RBAC, per-user
|
||||
// workspaces) need a stable identifier. See issue #915.
|
||||
let sender_name = &message.sender.display_name;
|
||||
let sender_email = message
|
||||
.metadata
|
||||
.get("sender_email")
|
||||
.and_then(|v| v.as_str());
|
||||
let telegram_user_id = message
|
||||
.metadata
|
||||
.get("telegram_user_id")
|
||||
.and_then(|v| v.as_str());
|
||||
let prefixed_text = if !sender_name.is_empty() {
|
||||
match sender_email {
|
||||
Some(email) => format!("[From: {sender_name} <{email}>] {text}"),
|
||||
None => format!("[From: {sender_name}] {text}"),
|
||||
match (sender_email, telegram_user_id) {
|
||||
(Some(email), _) => format!("[From: {sender_name} <{email}>] {text}"),
|
||||
(None, Some(tg_id)) => format!("[From: {sender_name} (tg_id:{tg_id})] {text}"),
|
||||
(None, None) => format!("[From: {sender_name}] {text}"),
|
||||
}
|
||||
} else {
|
||||
text.clone()
|
||||
@@ -1297,6 +1562,10 @@ fn media_type_from_url(url: &str) -> String {
|
||||
|
||||
/// Download an image from a URL and build content blocks for multimodal LLM input.
|
||||
///
|
||||
/// Accepts both `http(s)://` URLs (fetched via reqwest) and `file://` URLs
|
||||
/// (read from local disk — used by the channel inbox materialization path so
|
||||
/// agents see a stable local path even after a Discord CDN URL has expired).
|
||||
///
|
||||
/// Returns a `Vec<ContentBlock>` containing an image block (base64-encoded) and
|
||||
/// optionally a text block for the caption. If the download fails, returns a
|
||||
/// text-only block describing the failure.
|
||||
@@ -1306,38 +1575,79 @@ async fn download_image_to_blocks(url: &str, caption: Option<&str>) -> Vec<Conte
|
||||
// 5 MB limit to prevent memory abuse from oversized images
|
||||
const MAX_IMAGE_BYTES: usize = 5 * 1024 * 1024;
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let resp = match client.get(url).send().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
warn!("Failed to download image from channel: {e}");
|
||||
return vec![ContentBlock::Text {
|
||||
text: format!("[Image download failed: {e}]"),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
}
|
||||
};
|
||||
// Branch on URL scheme: file:// reads from local disk, everything else
|
||||
// goes through HTTP. We unify both paths into (bytes, header_type) before
|
||||
// the size/magic-byte logic below.
|
||||
let (bytes, header_type): (Vec<u8>, Option<String>) =
|
||||
if let Some(path) = url.strip_prefix("file://") {
|
||||
// file:// — local read. No content-type header to honor; magic-byte
|
||||
// sniffing and URL extension fallback do all the work. We don't
|
||||
// percent-decode: the inbox writer controls filenames and avoids
|
||||
// characters that would need encoding.
|
||||
match tokio::fs::read(path).await {
|
||||
Ok(b) => (b, None),
|
||||
Err(e) => {
|
||||
warn!("Failed to read image from local path {path}: {e}");
|
||||
return vec![ContentBlock::Text {
|
||||
text: format!("[Image read failed: {e}]"),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Build the client with transparent decompression DISABLED. Discord's
|
||||
// CDN edges occasionally advertise `content-encoding: gzip` (or br)
|
||||
// on PNG/JPEG passthroughs while the body is the raw, uncompressed
|
||||
// image bytes. With the default reqwest client (gzip/deflate/brotli
|
||||
// features enabled at the workspace level), this causes the
|
||||
// decompression layer to choke on the image header and reqwest
|
||||
// returns "error decoding response body" only on `bytes().await`,
|
||||
// not on `send()`. Forcing identity encoding sidesteps the whole
|
||||
// class of CDN content-encoding-flapping bugs. We also set a UA
|
||||
// (some CDNs 403 clients without one) and a 30s timeout aligned
|
||||
// with the upstream 5 MB cap.
|
||||
let client = reqwest::Client::builder()
|
||||
.no_gzip()
|
||||
.no_deflate()
|
||||
.no_brotli()
|
||||
.user_agent("openfang/0.1 (+https://openfang.ai)")
|
||||
.timeout(std::time::Duration::from_secs(30))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new());
|
||||
let resp = match client.get(url).send().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
warn!("Failed to download image from channel: {e}");
|
||||
return vec![ContentBlock::Text {
|
||||
text: format!("[Image download failed: {e}]"),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
}
|
||||
};
|
||||
|
||||
// Detect media type from Content-Type header — but only trust it if it's
|
||||
// actually an image/* type. Many APIs (Telegram, S3 pre-signed URLs) return
|
||||
// `application/octet-stream` for all files, which breaks vision.
|
||||
let header_type = resp
|
||||
.headers()
|
||||
.get("content-type")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.map(|ct| ct.split(';').next().unwrap_or(ct).trim().to_string())
|
||||
.filter(|ct| ct.starts_with("image/"));
|
||||
// Detect media type from Content-Type header — but only trust it if
|
||||
// it's actually an image/* type. Many APIs (Telegram, S3 pre-signed
|
||||
// URLs) return `application/octet-stream` for all files, which
|
||||
// breaks vision.
|
||||
let header_type = resp
|
||||
.headers()
|
||||
.get("content-type")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.map(|ct| ct.split(';').next().unwrap_or(ct).trim().to_string())
|
||||
.filter(|ct| ct.starts_with("image/"));
|
||||
|
||||
let bytes = match resp.bytes().await {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
warn!("Failed to read image bytes: {e}");
|
||||
return vec![ContentBlock::Text {
|
||||
text: format!("[Image read failed: {e}]"),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
}
|
||||
};
|
||||
let bytes = match resp.bytes().await {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
warn!("Failed to read image bytes: {e}");
|
||||
return vec![ContentBlock::Text {
|
||||
text: format!("[Image read failed: {e}]"),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
}
|
||||
};
|
||||
(bytes.to_vec(), header_type)
|
||||
};
|
||||
|
||||
// Three-tier media type detection:
|
||||
// 1. Trusted Content-Type header (only if image/*)
|
||||
@@ -1399,11 +1709,15 @@ async fn dispatch_with_blocks(
|
||||
lifecycle_reactions: bool,
|
||||
prefix_style: PrefixStyle,
|
||||
) {
|
||||
// Route to agent (same logic as text path)
|
||||
let agent_id = router.resolve(
|
||||
// Route to agent (same logic as text path).
|
||||
// Use sender_user_id() so user-keyed bindings match for Discord/Slack;
|
||||
// resolve_with_context lets channel_id-scoped bindings match per room.
|
||||
let binding_ctx = binding_context_for(message);
|
||||
let agent_id = router.resolve_with_context(
|
||||
&message.channel,
|
||||
&message.sender.platform_id,
|
||||
sender_user_id(message),
|
||||
message.sender.openfang_user.as_deref(),
|
||||
&binding_ctx,
|
||||
);
|
||||
|
||||
let agent_id = match agent_id {
|
||||
@@ -1420,7 +1734,9 @@ async fn dispatch_with_blocks(
|
||||
};
|
||||
match fallback {
|
||||
Some(id) => {
|
||||
router.set_user_default(message.sender.platform_id.clone(), id);
|
||||
// Key on sender_user_id() (not platform_id) so Discord/Slack — where
|
||||
// platform_id is the channel — store per-user, matching the read path.
|
||||
router.set_user_default(sender_user_id(message).to_string(), id);
|
||||
id
|
||||
}
|
||||
None => {
|
||||
@@ -1609,12 +1925,19 @@ async fn dispatch_with_blocks(
|
||||
}
|
||||
|
||||
/// Handle a bot command (returns the response text).
|
||||
///
|
||||
/// `user_id` is the platform user ID (e.g. Discord author ID, Slack user ID).
|
||||
/// For adapters that set `sender.platform_id` to the channel/conversation ID
|
||||
/// (Discord, Slack), callers must pass `sender_user_id(message)` here so that
|
||||
/// per-user agent routing works correctly. For adapters where platform_id is
|
||||
/// already the user (CLI, Telegram DM), the two are equivalent.
|
||||
async fn handle_command(
|
||||
name: &str,
|
||||
args: &[String],
|
||||
handle: &Arc<dyn ChannelBridgeHandle>,
|
||||
router: &Arc<AgentRouter>,
|
||||
sender: &ChannelUser,
|
||||
user_id: &str,
|
||||
) -> String {
|
||||
// Canonicalise through the unified command registry: aliases resolve to
|
||||
// their canonical name and matching is case-insensitive. If the command
|
||||
@@ -1668,14 +1991,17 @@ async fn handle_command(
|
||||
let agent_name = &args[0];
|
||||
match handle.find_agent_by_name(agent_name).await {
|
||||
Ok(Some(agent_id)) => {
|
||||
router.set_user_default(sender.platform_id.clone(), agent_id);
|
||||
// Key on user_id (the param wired in by the Discord/Slack call sites
|
||||
// via sender_user_id(message)) — not sender.platform_id, which is the
|
||||
// channel ID on those adapters. Matches the read-path resolution.
|
||||
router.set_user_default(user_id.to_string(), agent_id);
|
||||
format!("Now talking to agent: {agent_name}")
|
||||
}
|
||||
Ok(None) => {
|
||||
// Try to spawn it
|
||||
match handle.spawn_agent_by_name(agent_name).await {
|
||||
Ok(agent_id) => {
|
||||
router.set_user_default(sender.platform_id.clone(), agent_id);
|
||||
router.set_user_default(user_id.to_string(), agent_id);
|
||||
format!("Spawned and connected to agent: {agent_name}")
|
||||
}
|
||||
Err(e) => {
|
||||
@@ -1690,7 +2016,7 @@ async fn handle_command(
|
||||
// Need to resolve the user's current agent
|
||||
let agent_id = router.resolve(
|
||||
&crate::types::ChannelType::CLI,
|
||||
&sender.platform_id,
|
||||
user_id,
|
||||
sender.openfang_user.as_deref(),
|
||||
);
|
||||
match agent_id {
|
||||
@@ -1704,7 +2030,7 @@ async fn handle_command(
|
||||
"compact" => {
|
||||
let agent_id = router.resolve(
|
||||
&crate::types::ChannelType::CLI,
|
||||
&sender.platform_id,
|
||||
user_id,
|
||||
sender.openfang_user.as_deref(),
|
||||
);
|
||||
match agent_id {
|
||||
@@ -1718,7 +2044,7 @@ async fn handle_command(
|
||||
"model" => {
|
||||
let agent_id = router.resolve(
|
||||
&crate::types::ChannelType::CLI,
|
||||
&sender.platform_id,
|
||||
user_id,
|
||||
sender.openfang_user.as_deref(),
|
||||
);
|
||||
match agent_id {
|
||||
@@ -1742,7 +2068,7 @@ async fn handle_command(
|
||||
"stop" => {
|
||||
let agent_id = router.resolve(
|
||||
&crate::types::ChannelType::CLI,
|
||||
&sender.platform_id,
|
||||
user_id,
|
||||
sender.openfang_user.as_deref(),
|
||||
);
|
||||
match agent_id {
|
||||
@@ -1756,7 +2082,7 @@ async fn handle_command(
|
||||
"usage" => {
|
||||
let agent_id = router.resolve(
|
||||
&crate::types::ChannelType::CLI,
|
||||
&sender.platform_id,
|
||||
user_id,
|
||||
sender.openfang_user.as_deref(),
|
||||
);
|
||||
match agent_id {
|
||||
@@ -1770,7 +2096,7 @@ async fn handle_command(
|
||||
"think" => {
|
||||
let agent_id = router.resolve(
|
||||
&crate::types::ChannelType::CLI,
|
||||
&sender.platform_id,
|
||||
user_id,
|
||||
sender.openfang_user.as_deref(),
|
||||
);
|
||||
match agent_id {
|
||||
@@ -1939,10 +2265,10 @@ mod tests {
|
||||
openfang_user: None,
|
||||
};
|
||||
|
||||
let result = handle_command("agents", &[], &handle, &router, &sender).await;
|
||||
let result = handle_command("agents", &[], &handle, &router, &sender, "user1").await;
|
||||
assert!(result.contains("coder"));
|
||||
|
||||
let result = handle_command("help", &[], &handle, &router, &sender).await;
|
||||
let result = handle_command("help", &[], &handle, &router, &sender, "user1").await;
|
||||
assert!(result.contains("/agents"));
|
||||
}
|
||||
|
||||
@@ -1960,8 +2286,15 @@ mod tests {
|
||||
};
|
||||
|
||||
// Select existing agent
|
||||
let result =
|
||||
handle_command("agent", &["coder".to_string()], &handle, &router, &sender).await;
|
||||
let result = handle_command(
|
||||
"agent",
|
||||
&["coder".to_string()],
|
||||
&handle,
|
||||
&router,
|
||||
&sender,
|
||||
"user1",
|
||||
)
|
||||
.await;
|
||||
assert!(result.contains("Now talking to agent: coder"));
|
||||
|
||||
// Verify router was updated
|
||||
@@ -1969,6 +2302,48 @@ mod tests {
|
||||
assert_eq!(resolved, Some(agent_id));
|
||||
}
|
||||
|
||||
/// Discord/Slack-shaped: sender.platform_id is the *channel* id, user_id is
|
||||
/// the actual user. After /agent <name>, the default must be stored under
|
||||
/// user_id and resolvable by user_id — NOT by the channel id. This is the
|
||||
/// "split-keying" fix the read path has and the write path now matches.
|
||||
#[tokio::test]
|
||||
async fn test_handle_command_agent_select_keys_on_user_id_not_platform_id() {
|
||||
let agent_id = AgentId::new();
|
||||
let handle: Arc<dyn ChannelBridgeHandle> = Arc::new(MockHandle {
|
||||
agents: Mutex::new(vec![(agent_id, "coder".to_string())]),
|
||||
});
|
||||
let router = Arc::new(AgentRouter::new());
|
||||
// Discord-shape: platform_id is the channel, the real user is in user_id.
|
||||
let sender = ChannelUser {
|
||||
platform_id: "channel-123".to_string(),
|
||||
display_name: "Test".to_string(),
|
||||
openfang_user: None,
|
||||
};
|
||||
let user_id = "user-789";
|
||||
|
||||
let result = handle_command(
|
||||
"agent",
|
||||
&["coder".to_string()],
|
||||
&handle,
|
||||
&router,
|
||||
&sender,
|
||||
user_id,
|
||||
)
|
||||
.await;
|
||||
assert!(result.contains("Now talking to agent: coder"));
|
||||
|
||||
// Resolves under the user's id (correct).
|
||||
let by_user = router.resolve(&ChannelType::Discord, user_id, None);
|
||||
assert_eq!(by_user, Some(agent_id), "should resolve by user_id");
|
||||
|
||||
// Does NOT resolve under the channel id (the bug we just fixed).
|
||||
let by_channel = router.resolve(&ChannelType::Discord, "channel-123", None);
|
||||
assert_eq!(
|
||||
by_channel, None,
|
||||
"must NOT resolve by sender.platform_id (channel id)"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_handle_command_agent_without_args_lists_agents() {
|
||||
let agent_id = AgentId::new();
|
||||
@@ -1982,7 +2357,7 @@ mod tests {
|
||||
openfang_user: None,
|
||||
};
|
||||
|
||||
let result = handle_command("agent", &[], &handle, &router, &sender).await;
|
||||
let result = handle_command("agent", &[], &handle, &router, &sender, "user1").await;
|
||||
assert!(result.contains("Usage: /agent <name>"));
|
||||
assert!(result.contains("coder"));
|
||||
}
|
||||
@@ -2042,6 +2417,122 @@ mod tests {
|
||||
assert_eq!(GroupPolicy::default(), GroupPolicy::MentionOnly);
|
||||
}
|
||||
|
||||
// -- binding_context_for / ChannelMessage::channel_id() coverage --
|
||||
//
|
||||
// These tests pin the routing-time behavior so future adapter additions to
|
||||
// CHANNELS_WITH_PLATFORM_ID_AS_CHANNEL cannot silently regress the bridge.
|
||||
|
||||
fn make_msg_for_ctx(
|
||||
channel: ChannelType,
|
||||
platform_id: &str,
|
||||
metadata: Vec<(&str, serde_json::Value)>,
|
||||
) -> ChannelMessage {
|
||||
let mut md = std::collections::HashMap::new();
|
||||
for (k, v) in metadata {
|
||||
md.insert(k.to_string(), v);
|
||||
}
|
||||
ChannelMessage {
|
||||
channel,
|
||||
platform_message_id: "msg-1".to_string(),
|
||||
sender: crate::types::ChannelUser {
|
||||
platform_id: platform_id.to_string(),
|
||||
display_name: "Tester".to_string(),
|
||||
openfang_user: None,
|
||||
},
|
||||
content: ChannelContent::Text("hi".to_string()),
|
||||
target_agent: None,
|
||||
timestamp: chrono::Utc::now(),
|
||||
is_group: true,
|
||||
thread_id: None,
|
||||
metadata: md,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_context_for_discord_uses_platform_id_as_channel() {
|
||||
let msg = make_msg_for_ctx(ChannelType::Discord, "1234567890", vec![]);
|
||||
let ctx = binding_context_for(&msg);
|
||||
assert_eq!(ctx.channel, "discord");
|
||||
assert_eq!(ctx.channel_id.as_deref(), Some("1234567890"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_context_for_telegram_uses_platform_id_as_channel() {
|
||||
// Regression guard: Telegram is on the channel-ID allowlist.
|
||||
let msg = make_msg_for_ctx(ChannelType::Telegram, "-100123", vec![]);
|
||||
let ctx = binding_context_for(&msg);
|
||||
assert_eq!(ctx.channel, "telegram");
|
||||
assert_eq!(ctx.channel_id.as_deref(), Some("-100123"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_context_for_matrix_uses_room_id_from_platform_id() {
|
||||
let msg = make_msg_for_ctx(ChannelType::Matrix, "!room:server.tld", vec![]);
|
||||
let ctx = binding_context_for(&msg);
|
||||
assert_eq!(ctx.channel_id.as_deref(), Some("!room:server.tld"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_context_for_custom_supported_adapter() {
|
||||
// Custom("twitch") is on the allowlist.
|
||||
let msg = make_msg_for_ctx(
|
||||
ChannelType::Custom("twitch".to_string()),
|
||||
"channel-foo",
|
||||
vec![],
|
||||
);
|
||||
let ctx = binding_context_for(&msg);
|
||||
assert_eq!(ctx.channel, "twitch");
|
||||
assert_eq!(ctx.channel_id.as_deref(), Some("channel-foo"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_context_for_user_id_adapter_returns_none() {
|
||||
// Reddit's platform_id is the post author, not a subreddit/conversation.
|
||||
// The bridge must not surface that as `channel_id` (would silently match
|
||||
// user-scoped bindings against a user ID).
|
||||
let msg = make_msg_for_ctx(
|
||||
ChannelType::Custom("reddit".to_string()),
|
||||
"u/some-user",
|
||||
vec![],
|
||||
);
|
||||
let ctx = binding_context_for(&msg);
|
||||
assert_eq!(ctx.channel, "reddit");
|
||||
assert!(ctx.channel_id.is_none());
|
||||
// peer_id still falls through to platform_id (sender_user_id default).
|
||||
assert_eq!(ctx.peer_id, "u/some-user");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_context_for_metadata_fallback() {
|
||||
// For non-allowlisted adapters, metadata["channel_id"] is the
|
||||
// documented escape hatch — verify the bridge honors it.
|
||||
let msg = make_msg_for_ctx(
|
||||
ChannelType::Custom("reddit".to_string()),
|
||||
"u/some-user",
|
||||
vec![("channel_id", serde_json::json!("r/rust"))],
|
||||
);
|
||||
let ctx = binding_context_for(&msg);
|
||||
assert_eq!(ctx.channel_id.as_deref(), Some("r/rust"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_context_for_metadata_guild_id() {
|
||||
let msg = make_msg_for_ctx(
|
||||
ChannelType::Discord,
|
||||
"1234567890",
|
||||
vec![("guild_id", serde_json::json!("99999"))],
|
||||
);
|
||||
let ctx = binding_context_for(&msg);
|
||||
assert_eq!(ctx.guild_id.as_deref(), Some("99999"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_channel_message_channel_id_email_returns_none() {
|
||||
// Email's platform_id is the sender address — not a channel.
|
||||
let msg = make_msg_for_ctx(ChannelType::Email, "alice@example.com", vec![]);
|
||||
assert!(msg.channel_id().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_channel_type_str() {
|
||||
assert_eq!(channel_type_str(&ChannelType::Telegram), "telegram");
|
||||
|
||||
@@ -8,7 +8,7 @@ use crate::types::{
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use futures::{SinkExt, Stream, StreamExt};
|
||||
use std::collections::HashMap;
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::pin::Pin;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
@@ -22,6 +22,10 @@ const DISCORD_API_BASE: &str = "https://discord.com/api/v10";
|
||||
const MAX_BACKOFF: Duration = Duration::from_secs(60);
|
||||
const INITIAL_BACKOFF: Duration = Duration::from_secs(1);
|
||||
const DISCORD_MSG_LIMIT: usize = 2000;
|
||||
/// Maximum number of seen message IDs kept in the dedup set.
|
||||
/// MESSAGE_UPDATE (embed resolution) events arrive within seconds of the
|
||||
/// original CREATE; entries older than this cap are safe to discard.
|
||||
const MAX_DEDUP_MSG_IDS: usize = 2_000;
|
||||
|
||||
/// Discord Gateway opcodes.
|
||||
mod opcode {
|
||||
@@ -56,6 +60,8 @@ pub struct DiscordAdapter {
|
||||
allowed_users: Vec<String>,
|
||||
ignore_bots: bool,
|
||||
intents: u64,
|
||||
/// Auto-thread behavior: "true", "false", or "smart"
|
||||
auto_thread: String,
|
||||
shutdown_tx: Arc<watch::Sender<bool>>,
|
||||
shutdown_rx: watch::Receiver<bool>,
|
||||
/// Bot's own user ID (populated after READY event).
|
||||
@@ -64,6 +70,13 @@ pub struct DiscordAdapter {
|
||||
session_id: Arc<RwLock<Option<String>>>,
|
||||
/// Resume gateway URL.
|
||||
resume_gateway_url: Arc<RwLock<Option<String>>>,
|
||||
/// Thread channel IDs created by this bot (thread_id → parent_channel_id).
|
||||
/// Used to detect when incoming messages are inside a bot-created thread.
|
||||
created_thread_ids: Arc<RwLock<HashMap<String, String>>>,
|
||||
/// Message IDs seen via MESSAGE_CREATE (used to drop duplicate MESSAGE_UPDATE events).
|
||||
/// Populated immediately when MESSAGE_CREATE is forwarded — before bridge processing —
|
||||
/// to eliminate the race window where MESSAGE_UPDATE arrives before thread creation completes.
|
||||
threaded_message_ids: Arc<RwLock<HashSet<String>>>,
|
||||
}
|
||||
|
||||
impl DiscordAdapter {
|
||||
@@ -73,6 +86,7 @@ impl DiscordAdapter {
|
||||
allowed_users: Vec<String>,
|
||||
ignore_bots: bool,
|
||||
intents: u64,
|
||||
auto_thread: String,
|
||||
) -> Self {
|
||||
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||
Self {
|
||||
@@ -82,11 +96,14 @@ impl DiscordAdapter {
|
||||
allowed_users,
|
||||
ignore_bots,
|
||||
intents,
|
||||
auto_thread,
|
||||
shutdown_tx: Arc::new(shutdown_tx),
|
||||
shutdown_rx,
|
||||
bot_user_id: Arc::new(RwLock::new(None)),
|
||||
session_id: Arc::new(RwLock::new(None)),
|
||||
resume_gateway_url: Arc::new(RwLock::new(None)),
|
||||
created_thread_ids: Arc::new(RwLock::new(HashMap::new())),
|
||||
threaded_message_ids: Arc::new(RwLock::new(HashSet::new())),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -147,6 +164,79 @@ impl DiscordAdapter {
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Create a thread from a message in a Discord channel.
|
||||
async fn api_create_thread(
|
||||
&self,
|
||||
channel_id: &str,
|
||||
message_id: &str,
|
||||
name: &str,
|
||||
) -> Result<String, Box<dyn std::error::Error>> {
|
||||
let url = format!(
|
||||
"{DISCORD_API_BASE}/channels/{channel_id}/messages/{message_id}/threads",
|
||||
channel_id = channel_id,
|
||||
message_id = message_id
|
||||
);
|
||||
let body = serde_json::json!({
|
||||
"name": name,
|
||||
"auto_archive_duration": 1440 // 24 hours
|
||||
});
|
||||
let resp = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bot {}", self.token.as_str()))
|
||||
.json(&body)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let body_text = resp.text().await.unwrap_or_default();
|
||||
return Err(format!("Discord createThread failed: {}", body_text).into());
|
||||
}
|
||||
|
||||
let response: serde_json::Value = resp.json().await?;
|
||||
let thread_id = response["id"].as_str().unwrap_or("").to_string();
|
||||
|
||||
// Track thread_id → parent channel_id so we can recognise messages
|
||||
// that arrive inside this thread.
|
||||
if !thread_id.is_empty() {
|
||||
self.created_thread_ids
|
||||
.write()
|
||||
.await
|
||||
.insert(thread_id.clone(), channel_id.to_string());
|
||||
}
|
||||
|
||||
Ok(thread_id)
|
||||
}
|
||||
|
||||
/// Send a message to an existing thread.
|
||||
/// Discord threads are channels — post directly to channels/{thread_id}/messages.
|
||||
async fn api_send_thread_message(
|
||||
&self,
|
||||
_channel_id: &str,
|
||||
thread_id: &str,
|
||||
text: &str,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let url = format!("{DISCORD_API_BASE}/channels/{thread_id}/messages");
|
||||
let chunks = split_message(text, DISCORD_MSG_LIMIT);
|
||||
|
||||
for chunk in chunks {
|
||||
let body = serde_json::json!({ "content": chunk });
|
||||
let resp = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bot {}", self.token.as_str()))
|
||||
.json(&body)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let body_text = resp.text().await.unwrap_or_default();
|
||||
warn!("Discord sendThreadMessage failed: {body_text}");
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -159,6 +249,33 @@ impl ChannelAdapter for DiscordAdapter {
|
||||
ChannelType::Discord
|
||||
}
|
||||
|
||||
async fn should_auto_thread(&self, message: &ChannelMessage) -> Option<String> {
|
||||
// Only auto-thread in group channels (servers), not DMs
|
||||
if !message.is_group {
|
||||
return None;
|
||||
}
|
||||
|
||||
// Check auto_thread mode
|
||||
match self.auto_thread.as_str() {
|
||||
"true" => Some(thread_name_from_message(message)),
|
||||
"false" => None,
|
||||
"smart" => {
|
||||
// Only create thread if bot was @mentioned
|
||||
let was_mentioned = message
|
||||
.metadata
|
||||
.get("was_mentioned")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
if was_mentioned {
|
||||
Some(thread_name_from_message(message))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
async fn start(
|
||||
&self,
|
||||
) -> Result<Pin<Box<dyn Stream<Item = ChannelMessage> + Send>>, Box<dyn std::error::Error>>
|
||||
@@ -176,6 +293,8 @@ impl ChannelAdapter for DiscordAdapter {
|
||||
let bot_user_id = self.bot_user_id.clone();
|
||||
let session_id_store = self.session_id.clone();
|
||||
let resume_url_store = self.resume_gateway_url.clone();
|
||||
let created_thread_ids = self.created_thread_ids.clone();
|
||||
let threaded_message_ids = self.threaded_message_ids.clone();
|
||||
let mut shutdown = self.shutdown_rx.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
@@ -414,19 +533,66 @@ impl ChannelAdapter for DiscordAdapter {
|
||||
&allowed_guilds,
|
||||
&allowed_users,
|
||||
ignore_bots,
|
||||
&created_thread_ids,
|
||||
)
|
||||
.await
|
||||
{
|
||||
// MESSAGE_UPDATE must be suppressed if we already
|
||||
// forwarded a MESSAGE_CREATE for this message ID.
|
||||
// The check uses `seen_message_ids` (tracked below)
|
||||
// which is populated the moment MESSAGE_CREATE is
|
||||
// forwarded — before the bridge even processes it.
|
||||
// This closes the race window where MESSAGE_UPDATE
|
||||
// arrives before adapter.create_thread() completes.
|
||||
if event_name == "MESSAGE_UPDATE"
|
||||
&& threaded_message_ids
|
||||
.read()
|
||||
.await
|
||||
.contains(&msg.platform_message_id)
|
||||
{
|
||||
debug!(
|
||||
"Discord MESSAGE_UPDATE skipped (already seen {})",
|
||||
msg.platform_message_id
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
debug!(
|
||||
"Discord {event_name} from {}: {:?}",
|
||||
msg.sender.display_name, msg.content
|
||||
);
|
||||
|
||||
// Mark this message as seen immediately so any
|
||||
// concurrent or subsequent MESSAGE_UPDATE is dropped.
|
||||
if event_name == "MESSAGE_CREATE" {
|
||||
threaded_message_ids
|
||||
.write()
|
||||
.await
|
||||
.insert(msg.platform_message_id.clone());
|
||||
}
|
||||
|
||||
if tx.send(msg).await.is_err() {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
"THREAD_DELETE" | "CHANNEL_DELETE" => {
|
||||
// Clean up tracking when a thread is deleted so the
|
||||
// next message in the parent channel is treated fresh.
|
||||
if let Some(tid) = d["id"].as_str() {
|
||||
created_thread_ids.write().await.remove(tid);
|
||||
// Prune the dedup set to prevent unbounded growth.
|
||||
// Entries older than MAX_DEDUP_MSG_IDS are safe to
|
||||
// discard — embed UPDATE events arrive within seconds.
|
||||
let mut ids = threaded_message_ids.write().await;
|
||||
if ids.len() > MAX_DEDUP_MSG_IDS {
|
||||
ids.clear();
|
||||
}
|
||||
debug!("Discord thread/channel deleted: {tid}");
|
||||
}
|
||||
}
|
||||
|
||||
"RESUMED" => {
|
||||
info!("Discord session resumed successfully");
|
||||
}
|
||||
@@ -532,12 +698,123 @@ impl ChannelAdapter for DiscordAdapter {
|
||||
self.api_send_typing(&user.platform_id).await
|
||||
}
|
||||
|
||||
async fn send_in_thread(
|
||||
&self,
|
||||
user: &ChannelUser,
|
||||
content: ChannelContent,
|
||||
thread_id: &str,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let channel_id = &user.platform_id;
|
||||
match content {
|
||||
ChannelContent::Text(text) => {
|
||||
self.api_send_thread_message(channel_id, thread_id, &text)
|
||||
.await?;
|
||||
}
|
||||
_ => {
|
||||
self.api_send_thread_message(channel_id, thread_id, "(Unsupported content type)")
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn create_thread(
|
||||
&self,
|
||||
user: &ChannelUser,
|
||||
message_id: &str,
|
||||
thread_name: &str,
|
||||
) -> Result<String, Box<dyn std::error::Error>> {
|
||||
let channel_id = &user.platform_id;
|
||||
let thread_id = self
|
||||
.api_create_thread(channel_id, message_id, thread_name)
|
||||
.await?;
|
||||
// Also ensure the message_id is marked as seen (belt-and-suspenders:
|
||||
// the gateway loop already inserts on MESSAGE_CREATE, but keep this
|
||||
// in case create_thread is ever called from another path).
|
||||
self.threaded_message_ids
|
||||
.write()
|
||||
.await
|
||||
.insert(message_id.to_string());
|
||||
Ok(thread_id)
|
||||
}
|
||||
|
||||
async fn stop(&self) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let _ = self.shutdown_tx.send(true);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Maximum byte size for an attachment to be classified as a vision-eligible
|
||||
/// image. Anthropic's image content blocks are capped at 5 MB; oversize images
|
||||
/// fall through to `File` so the bridge passes the URL as text instead of
|
||||
/// attempting an inline image block.
|
||||
const VISION_IMAGE_MAX_BYTES: u64 = 5 * 1024 * 1024;
|
||||
|
||||
/// Best-effort MIME inference from a filename extension. Used as a fallback
|
||||
/// when Discord's `content_type` field is missing or empty (we've observed
|
||||
/// this on some bot-relayed attachments).
|
||||
fn mime_from_extension(filename: &str) -> Option<&'static str> {
|
||||
let ext = filename.rsplit('.').next()?.to_ascii_lowercase();
|
||||
match ext.as_str() {
|
||||
"jpg" | "jpeg" => Some("image/jpeg"),
|
||||
"png" => Some("image/png"),
|
||||
"gif" => Some("image/gif"),
|
||||
"webp" => Some("image/webp"),
|
||||
"heic" => Some("image/heic"),
|
||||
"heif" => Some("image/heif"),
|
||||
"pdf" => Some("application/pdf"),
|
||||
"txt" => Some("text/plain"),
|
||||
"md" => Some("text/markdown"),
|
||||
"json" => Some("application/json"),
|
||||
"mp4" => Some("video/mp4"),
|
||||
"mov" => Some("video/quicktime"),
|
||||
"mp3" => Some("audio/mpeg"),
|
||||
"wav" => Some("audio/wav"),
|
||||
"ogg" => Some("audio/ogg"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Classify a single Discord attachment JSON object into a `ChannelContent`
|
||||
/// block. Vision-eligible image MIME types (jpeg/png/gif/webp) under
|
||||
/// `VISION_IMAGE_MAX_BYTES` become `Image`; everything else becomes `File`
|
||||
/// (URL-pass-through; the bridge will surface it as a text descriptor in v1).
|
||||
///
|
||||
/// MIME resolution chain: `attachments[].content_type` (if non-empty) →
|
||||
/// extension lookup → `application/octet-stream`.
|
||||
fn classify_discord_attachment(att: &serde_json::Value) -> ChannelContent {
|
||||
let url = att["url"].as_str().unwrap_or("").to_string();
|
||||
let filename = att["filename"].as_str().unwrap_or("file").to_string();
|
||||
let size = att["size"].as_u64();
|
||||
|
||||
let resolved_mime: String = att["content_type"]
|
||||
.as_str()
|
||||
.filter(|s| !s.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| mime_from_extension(&filename).map(str::to_string))
|
||||
.unwrap_or_else(|| "application/octet-stream".to_string());
|
||||
|
||||
let is_vision_mime = matches!(
|
||||
resolved_mime.as_str(),
|
||||
"image/jpeg" | "image/png" | "image/gif" | "image/webp"
|
||||
);
|
||||
// If size is unknown, optimistically allow the image — the bridge will
|
||||
// surface a 4xx if Anthropic rejects it, which is better than silently
|
||||
// demoting to a text URL.
|
||||
let within_vision_limit = size.map(|s| s <= VISION_IMAGE_MAX_BYTES).unwrap_or(true);
|
||||
|
||||
if is_vision_mime && within_vision_limit {
|
||||
ChannelContent::Image { url, caption: None }
|
||||
} else {
|
||||
ChannelContent::File {
|
||||
url,
|
||||
filename,
|
||||
mime: Some(resolved_mime),
|
||||
size,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse a Discord MESSAGE_CREATE or MESSAGE_UPDATE payload into a `ChannelMessage`.
|
||||
async fn parse_discord_message(
|
||||
d: &serde_json::Value,
|
||||
@@ -545,7 +822,13 @@ async fn parse_discord_message(
|
||||
allowed_guilds: &[String],
|
||||
allowed_users: &[String],
|
||||
ignore_bots: bool,
|
||||
created_thread_ids: &Arc<RwLock<HashMap<String, String>>>,
|
||||
) -> Option<ChannelMessage> {
|
||||
// Diagnostic: dump the raw Discord payload so we can ground attachment
|
||||
// parsing in real JSON. Gated by RUST_LOG; silent at default `info` level.
|
||||
// Enable with: RUST_LOG=openfang_channels::discord=debug
|
||||
debug!(target: "openfang_channels::discord", payload = %d, "discord raw message payload");
|
||||
|
||||
let author = d.get("author")?;
|
||||
let author_id = author["id"].as_str()?;
|
||||
|
||||
@@ -577,12 +860,22 @@ async fn parse_discord_message(
|
||||
}
|
||||
|
||||
let content_text = d["content"].as_str().unwrap_or("");
|
||||
if content_text.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let channel_id = d["channel_id"].as_str()?;
|
||||
let message_id = d["id"].as_str().unwrap_or("0");
|
||||
|
||||
// Detect if this message is inside a bot-created thread.
|
||||
// In Discord, a thread is its own channel — channel_id will be the thread's ID.
|
||||
// If so, use the parent channel as platform_id and set thread_id so that:
|
||||
// (a) auto-thread logic is skipped (message.thread_id.is_some())
|
||||
// (b) responses are sent back into the same thread
|
||||
let (effective_channel_id, parsed_thread_id) = {
|
||||
let threads = created_thread_ids.read().await;
|
||||
if let Some(parent_channel_id) = threads.get(channel_id) {
|
||||
(parent_channel_id.clone(), Some(channel_id.to_string()))
|
||||
} else {
|
||||
(channel_id.to_string(), None)
|
||||
}
|
||||
};
|
||||
let username = author["username"].as_str().unwrap_or("Unknown");
|
||||
let discriminator = author["discriminator"].as_str().unwrap_or("0000");
|
||||
let display_name = if discriminator == "0" {
|
||||
@@ -597,7 +890,8 @@ async fn parse_discord_message(
|
||||
.map(|dt| dt.with_timezone(&chrono::Utc))
|
||||
.unwrap_or_else(chrono::Utc::now);
|
||||
|
||||
// Parse commands (messages starting with /)
|
||||
// Parse commands (messages starting with /). Commands do not carry
|
||||
// attachments in v1; attachment processing only runs in the non-command path.
|
||||
let content = if content_text.starts_with('/') {
|
||||
let parts: Vec<&str> = content_text.splitn(2, ' ').collect();
|
||||
let cmd_name = &parts[0][1..];
|
||||
@@ -611,7 +905,50 @@ async fn parse_discord_message(
|
||||
args,
|
||||
}
|
||||
} else {
|
||||
ChannelContent::Text(content_text.to_string())
|
||||
let attachment_blocks: Vec<ChannelContent> = d["attachments"]
|
||||
.as_array()
|
||||
.map(|arr| arr.iter().map(classify_discord_attachment).collect())
|
||||
.unwrap_or_default();
|
||||
|
||||
match (content_text.is_empty(), attachment_blocks.len()) {
|
||||
// No text, no attachments → nothing to ingest.
|
||||
(true, 0) => return None,
|
||||
// Text only.
|
||||
(false, 0) => ChannelContent::Text(content_text.to_string()),
|
||||
// Single attachment, no caption.
|
||||
(true, 1) => attachment_blocks.into_iter().next().unwrap(),
|
||||
// Single attachment + caption: emit Multipart with the caption as
|
||||
// a sibling Text block. This keeps the caption visible to providers
|
||||
// that flatten content to text only (e.g. claude-code/*, which
|
||||
// currently drops Image blocks) — the user gets a coherent
|
||||
// text-only response instead of a hallucination. Vision-capable
|
||||
// providers see the same blocks and dispatch multimodally.
|
||||
(false, 1) => {
|
||||
let block = attachment_blocks.into_iter().next().unwrap();
|
||||
let normalized = match block {
|
||||
// Drop any caption that classify_discord_attachment may have
|
||||
// attached; the sibling Text block is now the caption.
|
||||
ChannelContent::Image { url, caption: _ } => {
|
||||
ChannelContent::Image { url, caption: None }
|
||||
}
|
||||
other => other,
|
||||
};
|
||||
ChannelContent::Multipart(vec![
|
||||
ChannelContent::Text(content_text.to_string()),
|
||||
normalized,
|
||||
])
|
||||
}
|
||||
// Multiple attachments, no caption.
|
||||
(true, _) => ChannelContent::Multipart(attachment_blocks),
|
||||
// Multiple attachments + caption: text first, then attachments
|
||||
// (matches Discord's visual ordering: text above attachments).
|
||||
(false, _) => {
|
||||
let mut blocks = Vec::with_capacity(attachment_blocks.len() + 1);
|
||||
blocks.push(ChannelContent::Text(content_text.to_string()));
|
||||
blocks.extend(attachment_blocks);
|
||||
ChannelContent::Multipart(blocks)
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Determine if this is a group message (guild_id present = server channel)
|
||||
@@ -636,12 +973,15 @@ async fn parse_discord_message(
|
||||
if was_mentioned {
|
||||
metadata.insert("was_mentioned".to_string(), serde_json::json!(true));
|
||||
}
|
||||
// Stash the Discord author ID so the router can key bindings on user, not channel.
|
||||
// (`sender.platform_id` below is the channel ID, used for the send path.)
|
||||
metadata.insert("sender_user_id".to_string(), serde_json::json!(author_id));
|
||||
|
||||
Some(ChannelMessage {
|
||||
channel: ChannelType::Discord,
|
||||
platform_message_id: message_id.to_string(),
|
||||
sender: ChannelUser {
|
||||
platform_id: channel_id.to_string(),
|
||||
platform_id: effective_channel_id,
|
||||
display_name,
|
||||
openfang_user: None,
|
||||
},
|
||||
@@ -649,15 +989,50 @@ async fn parse_discord_message(
|
||||
target_agent: None,
|
||||
timestamp,
|
||||
is_group,
|
||||
thread_id: None,
|
||||
thread_id: parsed_thread_id,
|
||||
metadata,
|
||||
})
|
||||
}
|
||||
|
||||
/// Build a Discord thread name from the message content.
|
||||
/// Strips @mention prefixes (`<@...>`), trims whitespace, and truncates to
|
||||
/// Discord's 100-character thread name limit. Falls back to the sender's
|
||||
/// display name if the message has no usable text (e.g. image-only).
|
||||
fn thread_name_from_message(message: &ChannelMessage) -> String {
|
||||
let raw = match &message.content {
|
||||
ChannelContent::Text(t) => t.clone(),
|
||||
ChannelContent::Image { caption, .. } => caption.clone().unwrap_or_default(),
|
||||
_ => String::new(),
|
||||
};
|
||||
|
||||
// Strip leading Discord mention tokens (<@id> / <@!id>)
|
||||
let stripped = regex_lite::Regex::new(r"^(<@!?\d+>\s*)+")
|
||||
.map(|re| re.replace(&raw, "").into_owned())
|
||||
.unwrap_or(raw);
|
||||
|
||||
let trimmed = stripped.trim().to_string();
|
||||
|
||||
if trimmed.is_empty() {
|
||||
return message.sender.display_name.clone();
|
||||
}
|
||||
|
||||
// Truncate to Discord's 100-char limit
|
||||
if trimmed.chars().count() <= 100 {
|
||||
trimmed
|
||||
} else {
|
||||
trimmed.chars().take(97).collect::<String>() + "…"
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// Convenience helper: empty thread-tracking map for tests that don't exercise threading.
|
||||
fn empty_threads() -> Arc<RwLock<HashMap<String, String>>> {
|
||||
Arc::new(RwLock::new(HashMap::new()))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_discord_message_basic() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
@@ -674,7 +1049,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true)
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(msg.channel, ChannelType::Discord);
|
||||
@@ -698,7 +1073,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true).await;
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads()).await;
|
||||
assert!(msg.is_none());
|
||||
}
|
||||
|
||||
@@ -718,7 +1093,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true).await;
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads()).await;
|
||||
assert!(msg.is_none());
|
||||
}
|
||||
|
||||
@@ -739,7 +1114,7 @@ mod tests {
|
||||
});
|
||||
|
||||
// With ignore_bots=false, other bots' messages should be allowed
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], false).await;
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], false, &empty_threads()).await;
|
||||
assert!(msg.is_some());
|
||||
let msg = msg.unwrap();
|
||||
assert_eq!(msg.sender.display_name, "somebot");
|
||||
@@ -763,7 +1138,7 @@ mod tests {
|
||||
});
|
||||
|
||||
// Even with ignore_bots=false, the bot's own messages must still be filtered
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], false).await;
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], false, &empty_threads()).await;
|
||||
assert!(msg.is_none());
|
||||
}
|
||||
|
||||
@@ -784,12 +1159,20 @@ mod tests {
|
||||
});
|
||||
|
||||
// Not in allowed guilds
|
||||
let msg =
|
||||
parse_discord_message(&d, &bot_id, &["111".into(), "222".into()], &[], true).await;
|
||||
let msg = parse_discord_message(
|
||||
&d,
|
||||
&bot_id,
|
||||
&["111".into(), "222".into()],
|
||||
&[],
|
||||
true,
|
||||
&empty_threads(),
|
||||
)
|
||||
.await;
|
||||
assert!(msg.is_none());
|
||||
|
||||
// In allowed guilds
|
||||
let msg = parse_discord_message(&d, &bot_id, &["999".into()], &[], true).await;
|
||||
let msg =
|
||||
parse_discord_message(&d, &bot_id, &["999".into()], &[], true, &empty_threads()).await;
|
||||
assert!(msg.is_some());
|
||||
}
|
||||
|
||||
@@ -808,7 +1191,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true)
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match &msg.content {
|
||||
@@ -835,7 +1218,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true).await;
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads()).await;
|
||||
assert!(msg.is_none());
|
||||
}
|
||||
|
||||
@@ -854,7 +1237,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true)
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(msg.sender.display_name, "alice#1234");
|
||||
@@ -878,7 +1261,7 @@ mod tests {
|
||||
});
|
||||
|
||||
// MESSAGE_UPDATE uses the same parse function as MESSAGE_CREATE
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true)
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(msg.channel, ChannelType::Discord);
|
||||
@@ -909,16 +1292,25 @@ mod tests {
|
||||
&[],
|
||||
&["user111".into(), "user222".into()],
|
||||
true,
|
||||
&empty_threads(),
|
||||
)
|
||||
.await;
|
||||
assert!(msg.is_none());
|
||||
|
||||
// In allowed users
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &["user999".into()], true).await;
|
||||
let msg = parse_discord_message(
|
||||
&d,
|
||||
&bot_id,
|
||||
&[],
|
||||
&["user999".into()],
|
||||
true,
|
||||
&empty_threads(),
|
||||
)
|
||||
.await;
|
||||
assert!(msg.is_some());
|
||||
|
||||
// Empty allowed_users = allow all
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true).await;
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads()).await;
|
||||
assert!(msg.is_some());
|
||||
}
|
||||
|
||||
@@ -941,7 +1333,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true)
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(msg.is_group);
|
||||
@@ -964,7 +1356,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg2 = parse_discord_message(&d2, &bot_id, &[], &[], true)
|
||||
let msg2 = parse_discord_message(&d2, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(msg2.is_group);
|
||||
@@ -986,7 +1378,7 @@ mod tests {
|
||||
"timestamp": "2024-01-01T00:00:00+00:00"
|
||||
});
|
||||
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true)
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(!msg.is_group);
|
||||
@@ -1025,8 +1417,219 @@ mod tests {
|
||||
vec![],
|
||||
true,
|
||||
37376,
|
||||
"true".to_string(),
|
||||
);
|
||||
assert_eq!(adapter.name(), "discord");
|
||||
assert_eq!(adapter.channel_type(), ChannelType::Discord);
|
||||
}
|
||||
|
||||
// -- Multipart / attachment parsing tests (commit 4) ----------------------
|
||||
|
||||
fn att(filename: &str, content_type: Option<&str>, size: u64) -> serde_json::Value {
|
||||
let mut obj = serde_json::json!({
|
||||
"url": format!("https://cdn.discordapp.com/attachments/1/2/{filename}"),
|
||||
"filename": filename,
|
||||
"size": size,
|
||||
});
|
||||
if let Some(ct) = content_type {
|
||||
obj["content_type"] = serde_json::Value::String(ct.to_string());
|
||||
}
|
||||
obj
|
||||
}
|
||||
|
||||
fn payload_with(content: &str, attachments: Vec<serde_json::Value>) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"id": "msg1",
|
||||
"channel_id": "ch1",
|
||||
"content": content,
|
||||
"author": {
|
||||
"id": "user456",
|
||||
"username": "alice",
|
||||
"discriminator": "0",
|
||||
"bot": false
|
||||
},
|
||||
"timestamp": "2024-01-01T00:00:00+00:00",
|
||||
"attachments": attachments,
|
||||
})
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_image_only_no_caption() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with("", vec![att("photo.png", Some("image/png"), 100_000)]);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match msg.content {
|
||||
ChannelContent::Image { caption, url } => {
|
||||
assert!(caption.is_none());
|
||||
assert!(url.contains("photo.png"));
|
||||
}
|
||||
other => panic!("expected Image, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_image_with_caption() {
|
||||
// Single image + caption is emitted as Multipart([Text, Image]) so the
|
||||
// caption survives providers that flatten content blocks to text only
|
||||
// (e.g. claude-code/*). The Image carries no caption of its own; the
|
||||
// sibling Text block IS the caption.
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with(
|
||||
"look at this",
|
||||
vec![att("photo.jpg", Some("image/jpeg"), 50_000)],
|
||||
);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match msg.content {
|
||||
ChannelContent::Multipart(parts) => {
|
||||
assert_eq!(parts.len(), 2);
|
||||
assert!(matches!(&parts[0], ChannelContent::Text(t) if t == "look at this"));
|
||||
match &parts[1] {
|
||||
ChannelContent::Image { caption, url } => {
|
||||
assert!(
|
||||
caption.is_none(),
|
||||
"image caption should be None; the sibling Text block is the caption"
|
||||
);
|
||||
assert!(url.contains("photo.jpg"));
|
||||
}
|
||||
other => panic!("expected Image as second part, got {other:?}"),
|
||||
}
|
||||
}
|
||||
other => panic!("expected Multipart, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_multi_image_no_caption() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with(
|
||||
"",
|
||||
vec![
|
||||
att("a.png", Some("image/png"), 10_000),
|
||||
att("b.png", Some("image/png"), 20_000),
|
||||
],
|
||||
);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match msg.content {
|
||||
ChannelContent::Multipart(parts) => {
|
||||
assert_eq!(parts.len(), 2);
|
||||
assert!(parts
|
||||
.iter()
|
||||
.all(|p| matches!(p, ChannelContent::Image { .. })));
|
||||
}
|
||||
other => panic!("expected Multipart, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_multi_image_with_caption() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with(
|
||||
"two pics",
|
||||
vec![
|
||||
att("a.png", Some("image/png"), 10_000),
|
||||
att("b.png", Some("image/png"), 20_000),
|
||||
],
|
||||
);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match msg.content {
|
||||
ChannelContent::Multipart(parts) => {
|
||||
assert_eq!(parts.len(), 3);
|
||||
// Text first, then images.
|
||||
assert!(matches!(&parts[0], ChannelContent::Text(t) if t == "two pics"));
|
||||
assert!(matches!(&parts[1], ChannelContent::Image { .. }));
|
||||
assert!(matches!(&parts[2], ChannelContent::Image { .. }));
|
||||
}
|
||||
other => panic!("expected Multipart, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_heic_falls_to_file() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with("", vec![att("photo.heic", Some("image/heic"), 100_000)]);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match msg.content {
|
||||
ChannelContent::File { mime, filename, .. } => {
|
||||
assert_eq!(filename, "photo.heic");
|
||||
assert_eq!(mime.as_deref(), Some("image/heic"));
|
||||
}
|
||||
other => panic!("expected File, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_oversize_image_falls_to_file() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
// 6 MB exceeds VISION_IMAGE_MAX_BYTES (5 MB).
|
||||
let d = payload_with(
|
||||
"",
|
||||
vec![att("huge.png", Some("image/png"), 6 * 1024 * 1024)],
|
||||
);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match msg.content {
|
||||
ChannelContent::File {
|
||||
filename,
|
||||
mime,
|
||||
size,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(filename, "huge.png");
|
||||
assert_eq!(mime.as_deref(), Some("image/png"));
|
||||
assert_eq!(size, Some(6 * 1024 * 1024));
|
||||
}
|
||||
other => panic!("expected File, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_file_with_caption_yields_multipart() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with(
|
||||
"see attached",
|
||||
vec![att("doc.pdf", Some("application/pdf"), 200_000)],
|
||||
);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
match msg.content {
|
||||
ChannelContent::Multipart(parts) => {
|
||||
assert_eq!(parts.len(), 2);
|
||||
assert!(matches!(&parts[0], ChannelContent::Text(t) if t == "see attached"));
|
||||
assert!(matches!(&parts[1], ChannelContent::File { .. }));
|
||||
}
|
||||
other => panic!("expected Multipart, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_extension_fallback_when_content_type_missing() {
|
||||
// Discord occasionally omits content_type on bot-relayed attachments;
|
||||
// we should fall back to the filename extension.
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with("", vec![att("pic.png", None, 50_000)]);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(msg.content, ChannelContent::Image { .. }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_empty_message_with_no_attachments_returns_none() {
|
||||
let bot_id = Arc::new(RwLock::new(Some("bot123".to_string())));
|
||||
let d = payload_with("", vec![]);
|
||||
let msg = parse_discord_message(&d, &bot_id, &[], &[], true, &empty_threads()).await;
|
||||
assert!(msg.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,8 +42,8 @@ const MAX_MESSAGE_LEN: usize = 4000;
|
||||
/// Token refresh buffer — refresh 5 minutes before actual expiry.
|
||||
const TOKEN_REFRESH_BUFFER_SECS: u64 = 300;
|
||||
|
||||
/// Feishu websocket endpoint discovery API.
|
||||
const FEISHU_WS_ENDPOINT_URL: &str = "https://open.feishu.cn/callback/ws/endpoint";
|
||||
/// WebSocket endpoint path (appended to the region domain).
|
||||
const FEISHU_WS_ENDPOINT_PATH: &str = "/callback/ws/endpoint";
|
||||
|
||||
const INITIAL_BACKOFF: Duration = Duration::from_secs(1);
|
||||
const MAX_BACKOFF: Duration = Duration::from_secs(60);
|
||||
@@ -269,13 +269,24 @@ impl FeishuAdapter {
|
||||
///
|
||||
/// WebSocket mode does not require a public IP or webhook configuration.
|
||||
pub fn new_websocket(app_id: String, app_secret: String) -> Self {
|
||||
Self::new_websocket_with_region(app_id, app_secret, FeishuRegion::Cn)
|
||||
}
|
||||
|
||||
/// Create a new Feishu adapter in WebSocket mode with an explicit region.
|
||||
///
|
||||
/// Use this when the app is registered on Lark international (`open.larksuite.com`).
|
||||
pub fn new_websocket_with_region(
|
||||
app_id: String,
|
||||
app_secret: String,
|
||||
region: FeishuRegion,
|
||||
) -> Self {
|
||||
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||
Self {
|
||||
app_id,
|
||||
app_secret: Zeroizing::new(app_secret),
|
||||
connection_mode: FeishuConnectionMode::WebSocket,
|
||||
webhook_port: 0,
|
||||
region: FeishuRegion::Cn,
|
||||
region,
|
||||
webhook_path: String::new(),
|
||||
verification_token: None,
|
||||
encrypt_key: None,
|
||||
@@ -918,9 +929,10 @@ struct FeishuAdapterClone {
|
||||
impl FeishuAdapterClone {
|
||||
/// Get WebSocket endpoint from Feishu API.
|
||||
async fn get_websocket_endpoint(&self) -> Result<FeishuWsEndpoint, Box<dyn std::error::Error>> {
|
||||
let url = format!("{}{}", self.region.domain(), FEISHU_WS_ENDPOINT_PATH);
|
||||
let resp = self
|
||||
.client
|
||||
.post(FEISHU_WS_ENDPOINT_URL)
|
||||
.post(&url)
|
||||
.json(&serde_json::json!({
|
||||
"AppID": self.app_id,
|
||||
"AppSecret": self.app_secret.as_str(),
|
||||
|
||||
@@ -326,19 +326,17 @@ impl ChannelAdapter for IrcAdapter {
|
||||
}
|
||||
|
||||
// RPL_WELCOME (001) — registration complete, join channels
|
||||
"001" => {
|
||||
if !joined {
|
||||
info!("IRC registered as {nick_clone}");
|
||||
for ch in &channels_clone {
|
||||
let join_cmd = format!("JOIN {ch}\r\n");
|
||||
if let Err(e) = writer.write_all(join_cmd.as_bytes()).await {
|
||||
warn!("IRC JOIN send failed: {e}");
|
||||
break 'inner true;
|
||||
}
|
||||
info!("IRC joining {ch}");
|
||||
"001" if !joined => {
|
||||
info!("IRC registered as {nick_clone}");
|
||||
for ch in &channels_clone {
|
||||
let join_cmd = format!("JOIN {ch}\r\n");
|
||||
if let Err(e) = writer.write_all(join_cmd.as_bytes()).await {
|
||||
warn!("IRC JOIN send failed: {e}");
|
||||
break 'inner true;
|
||||
}
|
||||
joined = true;
|
||||
info!("IRC joining {ch}");
|
||||
}
|
||||
joined = true;
|
||||
}
|
||||
|
||||
// PRIVMSG — incoming message
|
||||
|
||||
@@ -18,14 +18,20 @@ use zeroize::Zeroizing;
|
||||
const SYNC_TIMEOUT_MS: u64 = 30000;
|
||||
const MAX_MESSAGE_LEN: usize = 4096;
|
||||
|
||||
/// Shared access + refresh token pair. Tokens are zeroized on drop and rotated
|
||||
/// in place when MSC2918 refresh succeeds.
|
||||
type TokenPair = Arc<RwLock<(Zeroizing<String>, Option<Zeroizing<String>>)>>;
|
||||
|
||||
/// Matrix channel adapter using the Client-Server API.
|
||||
pub struct MatrixAdapter {
|
||||
/// Matrix homeserver URL (e.g., `"https://matrix.org"`).
|
||||
homeserver_url: String,
|
||||
/// Bot's user ID (e.g., "@openfang:matrix.org").
|
||||
user_id: String,
|
||||
/// SECURITY: Access token is zeroized on drop.
|
||||
access_token: Zeroizing<String>,
|
||||
/// SECURITY: Access + refresh tokens are zeroized on drop. Stored behind
|
||||
/// an RwLock so the sync loop and send paths see rotated tokens after a
|
||||
/// MSC2918 /refresh call (matrix.org/MAS rotates both tokens every refresh).
|
||||
tokens: TokenPair,
|
||||
/// HTTP client.
|
||||
client: reqwest::Client,
|
||||
/// Allowed room IDs (empty = all joined rooms).
|
||||
@@ -40,19 +46,47 @@ pub struct MatrixAdapter {
|
||||
}
|
||||
|
||||
impl MatrixAdapter {
|
||||
/// Create a new Matrix adapter.
|
||||
/// Create a new Matrix adapter without a refresh token.
|
||||
pub fn new(
|
||||
homeserver_url: String,
|
||||
user_id: String,
|
||||
access_token: String,
|
||||
allowed_rooms: Vec<String>,
|
||||
auto_accept_invites: bool,
|
||||
) -> Self {
|
||||
Self::with_refresh_token(
|
||||
homeserver_url,
|
||||
user_id,
|
||||
access_token,
|
||||
None,
|
||||
allowed_rooms,
|
||||
auto_accept_invites,
|
||||
)
|
||||
}
|
||||
|
||||
/// Create a new Matrix adapter with an optional refresh token (MSC2918).
|
||||
///
|
||||
/// When `refresh_token` is `Some`, the adapter will automatically call
|
||||
/// `POST /_matrix/client/v3/refresh` on `401 M_UNKNOWN_TOKEN` responses
|
||||
/// and retry the failed request once. Both tokens rotate on each refresh
|
||||
/// under Matrix Authentication Service (MAS).
|
||||
pub fn with_refresh_token(
|
||||
homeserver_url: String,
|
||||
user_id: String,
|
||||
access_token: String,
|
||||
refresh_token: Option<String>,
|
||||
allowed_rooms: Vec<String>,
|
||||
auto_accept_invites: bool,
|
||||
) -> Self {
|
||||
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||
let tokens: TokenPair = Arc::new(RwLock::new((
|
||||
Zeroizing::new(access_token),
|
||||
refresh_token.map(Zeroizing::new),
|
||||
)));
|
||||
Self {
|
||||
homeserver_url,
|
||||
user_id,
|
||||
access_token: Zeroizing::new(access_token),
|
||||
tokens,
|
||||
client: reqwest::Client::new(),
|
||||
allowed_rooms,
|
||||
shutdown_tx: Arc::new(shutdown_tx),
|
||||
@@ -62,6 +96,11 @@ impl MatrixAdapter {
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the current access token (cloned).
|
||||
async fn current_access_token(&self) -> String {
|
||||
self.tokens.read().await.0.as_str().to_string()
|
||||
}
|
||||
|
||||
/// Send a text message to a Matrix room.
|
||||
async fn api_send_message(
|
||||
&self,
|
||||
@@ -81,18 +120,46 @@ impl MatrixAdapter {
|
||||
"body": chunk,
|
||||
});
|
||||
|
||||
let resp = self
|
||||
.client
|
||||
.put(&url)
|
||||
.bearer_auth(&*self.access_token)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await?;
|
||||
let mut attempt = 0;
|
||||
loop {
|
||||
attempt += 1;
|
||||
let token = self.current_access_token().await;
|
||||
let resp = self
|
||||
.client
|
||||
.put(&url)
|
||||
.bearer_auth(&token)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if resp.status().is_success() {
|
||||
break;
|
||||
}
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
return Err(format!("Matrix API error {status}: {body}").into());
|
||||
let body_text = resp.text().await.unwrap_or_default();
|
||||
|
||||
// Try a single refresh+retry on M_UNKNOWN_TOKEN (MSC2918).
|
||||
if attempt == 1
|
||||
&& status == reqwest::StatusCode::UNAUTHORIZED
|
||||
&& is_unknown_token_body(&body_text)
|
||||
{
|
||||
match try_refresh_tokens(&self.client, &self.homeserver_url, &self.tokens).await
|
||||
{
|
||||
Ok(()) => {
|
||||
info!("Matrix: access token refreshed via MSC2918, retrying send");
|
||||
continue;
|
||||
}
|
||||
Err(e) => {
|
||||
return Err(format!(
|
||||
"Matrix API error {status}: {body_text} (refresh failed: {e})"
|
||||
)
|
||||
.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return Err(format!("Matrix API error {status}: {body_text}").into());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -103,21 +170,32 @@ impl MatrixAdapter {
|
||||
async fn validate(&self) -> Result<String, Box<dyn std::error::Error>> {
|
||||
let url = format!("{}/_matrix/client/v3/account/whoami", self.homeserver_url);
|
||||
|
||||
let resp = self
|
||||
.client
|
||||
.get(&url)
|
||||
.bearer_auth(&*self.access_token)
|
||||
.send()
|
||||
.await?;
|
||||
let mut attempt = 0;
|
||||
loop {
|
||||
attempt += 1;
|
||||
let token = self.current_access_token().await;
|
||||
let resp = self.client.get(&url).bearer_auth(&token).send().await?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
if resp.status().is_success() {
|
||||
let body: serde_json::Value = resp.json().await?;
|
||||
let user_id = body["user_id"].as_str().unwrap_or("unknown").to_string();
|
||||
return Ok(user_id);
|
||||
}
|
||||
|
||||
let status = resp.status();
|
||||
let body_text = resp.text().await.unwrap_or_default();
|
||||
if attempt == 1
|
||||
&& status == reqwest::StatusCode::UNAUTHORIZED
|
||||
&& is_unknown_token_body(&body_text)
|
||||
&& try_refresh_tokens(&self.client, &self.homeserver_url, &self.tokens)
|
||||
.await
|
||||
.is_ok()
|
||||
{
|
||||
info!("Matrix: access token refreshed via MSC2918, retrying /whoami");
|
||||
continue;
|
||||
}
|
||||
return Err("Matrix authentication failed".into());
|
||||
}
|
||||
|
||||
let body: serde_json::Value = resp.json().await?;
|
||||
let user_id = body["user_id"].as_str().unwrap_or("unknown").to_string();
|
||||
|
||||
Ok(user_id)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -126,6 +204,87 @@ impl MatrixAdapter {
|
||||
}
|
||||
}
|
||||
|
||||
/// Detect `M_UNKNOWN_TOKEN` errors in a Matrix response body.
|
||||
///
|
||||
/// Matrix returns 401 for multiple reasons; we only want to refresh on
|
||||
/// `M_UNKNOWN_TOKEN` (the access token expired or was revoked). See
|
||||
/// <https://spec.matrix.org/latest/client-server-api/#soft-logout>.
|
||||
fn is_unknown_token_body(body: &str) -> bool {
|
||||
serde_json::from_str::<serde_json::Value>(body)
|
||||
.ok()
|
||||
.and_then(|v| v.get("errcode").and_then(|c| c.as_str()).map(String::from))
|
||||
.map(|c| c == "M_UNKNOWN_TOKEN")
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// Whether a Matrix 401 body indicates a hard logout (operator must re-login).
|
||||
///
|
||||
/// `soft_logout: true` (or absent — default per spec) means the device is still
|
||||
/// known to the server and a refresh-token grant is valid. `soft_logout: false`
|
||||
/// means the device was invalidated and the operator must perform a new
|
||||
/// `m.login.password` flow.
|
||||
fn is_hard_logout(body: &str) -> bool {
|
||||
serde_json::from_str::<serde_json::Value>(body)
|
||||
.ok()
|
||||
.and_then(|v| v.get("soft_logout").and_then(|s| s.as_bool()))
|
||||
.map(|soft| !soft)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// Call `POST /_matrix/client/v3/refresh` (MSC2918) and rotate the stored tokens.
|
||||
///
|
||||
/// On success, replaces the access token and (if the server returned one) the
|
||||
/// refresh token. MAS (matrix.org since 2025-04-07) rotates the refresh token
|
||||
/// on every call, so callers must use the new value next time.
|
||||
async fn try_refresh_tokens(
|
||||
client: &reqwest::Client,
|
||||
homeserver: &str,
|
||||
tokens: &TokenPair,
|
||||
) -> Result<(), String> {
|
||||
let refresh_token = {
|
||||
let guard = tokens.read().await;
|
||||
match guard.1.as_ref() {
|
||||
Some(rt) => rt.as_str().to_string(),
|
||||
None => return Err("no refresh token configured".to_string()),
|
||||
}
|
||||
};
|
||||
|
||||
let url = format!("{homeserver}/_matrix/client/v3/refresh");
|
||||
let resp = client
|
||||
.post(&url)
|
||||
.json(&serde_json::json!({ "refresh_token": refresh_token }))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("refresh request failed: {e}"))?;
|
||||
|
||||
let status = resp.status();
|
||||
if !status.is_success() {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
return Err(format!("refresh returned {status}: {body}"));
|
||||
}
|
||||
|
||||
let body: serde_json::Value = resp
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| format!("refresh response parse error: {e}"))?;
|
||||
|
||||
let new_access = body
|
||||
.get("access_token")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| "refresh response missing access_token".to_string())?;
|
||||
let new_refresh = body
|
||||
.get("refresh_token")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
|
||||
let mut guard = tokens.write().await;
|
||||
guard.0 = Zeroizing::new(new_access.to_string());
|
||||
if let Some(rt) = new_refresh {
|
||||
guard.1 = Some(Zeroizing::new(rt));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Accept a room invite by calling POST /_matrix/client/v3/rooms/{room_id}/join.
|
||||
async fn accept_invite(
|
||||
client: &reqwest::Client,
|
||||
@@ -218,7 +377,7 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
|
||||
let (tx, rx) = mpsc::channel::<ChannelMessage>(256);
|
||||
let homeserver = self.homeserver_url.clone();
|
||||
let access_token = self.access_token.clone();
|
||||
let tokens = Arc::clone(&self.tokens);
|
||||
// Use the validated user ID from /whoami instead of the config value.
|
||||
// Matrix server delegation or casing differences can cause self.user_id
|
||||
// to not match the sender field in timeline events, making the bot
|
||||
@@ -232,9 +391,10 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
|
||||
// FIX #4: Do an initial sync to get the since token, skipping old messages.
|
||||
if since_token.read().await.is_none() {
|
||||
if let Some(token) = initial_sync(&client, &homeserver, access_token.as_str()).await {
|
||||
let token = self.current_access_token().await;
|
||||
if let Some(next) = initial_sync(&client, &homeserver, &token).await {
|
||||
info!("Matrix: initial sync complete, skipping old messages");
|
||||
*since_token.write().await = Some(token);
|
||||
*since_token.write().await = Some(next);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -257,12 +417,13 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
url.push_str(&format!("&since={token}"));
|
||||
}
|
||||
|
||||
let current_token = tokens.read().await.0.as_str().to_string();
|
||||
let resp = tokio::select! {
|
||||
_ = shutdown_rx.changed() => {
|
||||
info!("Matrix adapter shutting down");
|
||||
break;
|
||||
}
|
||||
result = client.get(&url).bearer_auth(access_token.as_str()).send() => {
|
||||
result = client.get(&url).bearer_auth(¤t_token).send() => {
|
||||
match result {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
@@ -276,7 +437,38 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
};
|
||||
|
||||
if !resp.status().is_success() {
|
||||
warn!("Matrix sync returned {}", resp.status());
|
||||
let status = resp.status();
|
||||
// MSC2918: on 401 M_UNKNOWN_TOKEN with a refresh token configured,
|
||||
// try refreshing once and loop again immediately. Hard logout
|
||||
// (soft_logout:false) is unrecoverable here — the operator must
|
||||
// perform a fresh m.login.password.
|
||||
if status == reqwest::StatusCode::UNAUTHORIZED {
|
||||
let body_text = resp.text().await.unwrap_or_default();
|
||||
if is_unknown_token_body(&body_text) {
|
||||
if is_hard_logout(&body_text) {
|
||||
warn!(
|
||||
"Matrix: hard logout (soft_logout=false), operator must re-login"
|
||||
);
|
||||
} else {
|
||||
match try_refresh_tokens(&client, &homeserver, &tokens).await {
|
||||
Ok(()) => {
|
||||
info!(
|
||||
"Matrix: access token refreshed via MSC2918, resuming /sync"
|
||||
);
|
||||
backoff = Duration::from_secs(1);
|
||||
continue;
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("Matrix: token refresh failed: {e}");
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
warn!("Matrix sync returned {status}: {body_text}");
|
||||
}
|
||||
} else {
|
||||
warn!("Matrix sync returned {status}");
|
||||
}
|
||||
tokio::time::sleep(backoff).await;
|
||||
backoff = (backoff * 2).min(Duration::from_secs(60));
|
||||
continue;
|
||||
@@ -309,8 +501,8 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
);
|
||||
continue;
|
||||
}
|
||||
accept_invite(&client, &homeserver, access_token.as_str(), room_id)
|
||||
.await;
|
||||
let tok = tokens.read().await.0.as_str().to_string();
|
||||
accept_invite(&client, &homeserver, &tok, room_id).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -380,10 +572,11 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
}
|
||||
|
||||
// FIX #3: Determine if room is a DM (2 members) or group.
|
||||
let tok_for_count = tokens.read().await.0.as_str().to_string();
|
||||
let is_group = get_room_member_count(
|
||||
&client,
|
||||
&homeserver,
|
||||
access_token.as_str(),
|
||||
&tok_for_count,
|
||||
room_id,
|
||||
)
|
||||
.await
|
||||
@@ -409,10 +602,11 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
}
|
||||
|
||||
// FIX #3: Determine if room is a DM (2 members) or group.
|
||||
let tok_for_count = tokens.read().await.0.as_str().to_string();
|
||||
let is_group = get_room_member_count(
|
||||
&client,
|
||||
&homeserver,
|
||||
access_token.as_str(),
|
||||
&tok_for_count,
|
||||
room_id,
|
||||
)
|
||||
.await
|
||||
@@ -485,10 +679,11 @@ impl ChannelAdapter for MatrixAdapter {
|
||||
"timeout": 5000,
|
||||
});
|
||||
|
||||
let token = self.current_access_token().await;
|
||||
let _ = self
|
||||
.client
|
||||
.put(&url)
|
||||
.bearer_auth(&*self.access_token)
|
||||
.bearer_auth(&token)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await;
|
||||
@@ -518,6 +713,79 @@ mod tests {
|
||||
assert_eq!(adapter.name(), "matrix");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_unknown_token_body() {
|
||||
// Real matrix.org body for M_UNKNOWN_TOKEN under MAS.
|
||||
let body =
|
||||
r#"{"errcode":"M_UNKNOWN_TOKEN","error":"Token is not active","soft_logout":true}"#;
|
||||
assert!(is_unknown_token_body(body));
|
||||
assert!(!is_hard_logout(body));
|
||||
|
||||
let hard = r#"{"errcode":"M_UNKNOWN_TOKEN","error":"Invalidated","soft_logout":false}"#;
|
||||
assert!(is_unknown_token_body(hard));
|
||||
assert!(is_hard_logout(hard));
|
||||
|
||||
let other = r#"{"errcode":"M_FORBIDDEN","error":"You are not allowed"}"#;
|
||||
assert!(!is_unknown_token_body(other));
|
||||
assert!(!is_hard_logout(other));
|
||||
|
||||
// Empty / non-JSON must not trigger refresh.
|
||||
assert!(!is_unknown_token_body(""));
|
||||
assert!(!is_unknown_token_body("not json"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_tokens_rotates_pair() {
|
||||
// Spin up a tiny axum server that mimics MSC2918 /refresh: rotates both
|
||||
// access and refresh tokens and returns the new pair.
|
||||
use axum::{routing::post, Json, Router};
|
||||
|
||||
async fn refresh_handler(Json(body): Json<serde_json::Value>) -> Json<serde_json::Value> {
|
||||
let incoming = body
|
||||
.get("refresh_token")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("");
|
||||
assert_eq!(incoming, "old_refresh");
|
||||
Json(serde_json::json!({
|
||||
"access_token": "new_access",
|
||||
"refresh_token": "new_refresh",
|
||||
"expires_in_ms": 3_600_000u64,
|
||||
}))
|
||||
}
|
||||
|
||||
let app = Router::new().route("/_matrix/client/v3/refresh", post(refresh_handler));
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
let homeserver = format!("http://{addr}");
|
||||
let tokens: TokenPair = Arc::new(RwLock::new((
|
||||
Zeroizing::new("old_access".to_string()),
|
||||
Some(Zeroizing::new("old_refresh".to_string())),
|
||||
)));
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
try_refresh_tokens(&client, &homeserver, &tokens)
|
||||
.await
|
||||
.expect("refresh succeeds");
|
||||
|
||||
let guard = tokens.read().await;
|
||||
assert_eq!(guard.0.as_str(), "new_access");
|
||||
assert_eq!(guard.1.as_ref().map(|s| s.as_str()), Some("new_refresh"));
|
||||
drop(guard);
|
||||
|
||||
// Refresh with no refresh token configured must fail cleanly.
|
||||
let no_refresh: TokenPair = Arc::new(RwLock::new((Zeroizing::new("a".to_string()), None)));
|
||||
let err = try_refresh_tokens(&client, &homeserver, &no_refresh)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(err.contains("no refresh token"));
|
||||
|
||||
server.abort();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_matrix_allowed_rooms() {
|
||||
let adapter = MatrixAdapter::new(
|
||||
|
||||
@@ -18,6 +18,10 @@ pub struct BindingContext {
|
||||
pub peer_id: String,
|
||||
/// Guild/server ID.
|
||||
pub guild_id: Option<String>,
|
||||
/// Channel/conversation ID (e.g. Discord channel, Slack conversation,
|
||||
/// Telegram chat, IRC channel name). Populated by bridges so bindings can
|
||||
/// route by room independent of which user posted.
|
||||
pub channel_id: Option<String>,
|
||||
/// User's roles.
|
||||
pub roles: Vec<String>,
|
||||
}
|
||||
@@ -143,6 +147,19 @@ impl AgentRouter {
|
||||
channel_type: &ChannelType,
|
||||
platform_user_id: &str,
|
||||
user_key: Option<&str>,
|
||||
) -> Option<AgentId> {
|
||||
self.resolve_with_channel_id(channel_type, platform_user_id, user_key, None)
|
||||
}
|
||||
|
||||
/// Resolve with an explicit channel/conversation ID, so bindings whose
|
||||
/// `match_rule.channel_id` is set can match. Used by bridges that know the
|
||||
/// room/conversation the message arrived in (Discord/Slack/Telegram/IRC).
|
||||
pub fn resolve_with_channel_id(
|
||||
&self,
|
||||
channel_type: &ChannelType,
|
||||
platform_user_id: &str,
|
||||
user_key: Option<&str>,
|
||||
channel_id: Option<&str>,
|
||||
) -> Option<AgentId> {
|
||||
let channel_key = format!("{channel_type:?}");
|
||||
|
||||
@@ -152,6 +169,7 @@ impl AgentRouter {
|
||||
account_id: None,
|
||||
peer_id: platform_user_id.to_string(),
|
||||
guild_id: None,
|
||||
channel_id: channel_id.map(|s| s.to_string()),
|
||||
roles: Vec::new(),
|
||||
};
|
||||
if let Some(agent_id) = self.resolve_binding(&ctx) {
|
||||
@@ -329,6 +347,11 @@ impl AgentRouter {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
if let Some(ref cid) = rule.channel_id {
|
||||
if ctx.channel_id.as_ref() != Some(cid) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
if !rule.roles.is_empty() {
|
||||
// User must have at least one of the specified roles
|
||||
let has_role = rule.roles.iter().any(|r| ctx.roles.contains(r));
|
||||
@@ -639,7 +662,129 @@ mod tests {
|
||||
guild_id: Some("guild".to_string()),
|
||||
roles: vec!["admin".to_string()],
|
||||
account_id: Some("bot".to_string()),
|
||||
channel_id: Some("ch_42".to_string()),
|
||||
};
|
||||
assert_eq!(full.specificity(), 17); // 8+4+2+2+1
|
||||
assert_eq!(full.specificity(), 25); // 8+8+4+2+2+1
|
||||
|
||||
// peer_id alone vs channel_id alone — both worth 8.
|
||||
let peer_only = BindingMatchRule {
|
||||
peer_id: Some("u".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let channel_id_only = BindingMatchRule {
|
||||
channel_id: Some("c".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(peer_only.specificity(), 8);
|
||||
assert_eq!(channel_id_only.specificity(), 8);
|
||||
|
||||
// Combined peer_id + channel_id (16) outranks either alone (8).
|
||||
let peer_and_channel = BindingMatchRule {
|
||||
peer_id: Some("u".to_string()),
|
||||
channel_id: Some("c".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(peer_and_channel.specificity(), 16);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_channel_id_match() {
|
||||
// A binding scoped to a specific Discord channel should match messages
|
||||
// from that channel and reject messages from other channels.
|
||||
let router = AgentRouter::new();
|
||||
let agent_id = AgentId::new();
|
||||
router.register_agent("ops-bot".to_string(), agent_id);
|
||||
router.load_bindings(&[AgentBinding {
|
||||
agent: "ops-bot".to_string(),
|
||||
match_rule: openfang_types::config::BindingMatchRule {
|
||||
channel: Some("discord".to_string()),
|
||||
channel_id: Some("1477803840265781391".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
}]);
|
||||
|
||||
// Same channel, any user — matches.
|
||||
let resolved = router.resolve_with_channel_id(
|
||||
&ChannelType::Discord,
|
||||
"any-user",
|
||||
None,
|
||||
Some("1477803840265781391"),
|
||||
);
|
||||
assert_eq!(resolved, Some(agent_id));
|
||||
|
||||
// Different channel — no match.
|
||||
let resolved = router.resolve_with_channel_id(
|
||||
&ChannelType::Discord,
|
||||
"any-user",
|
||||
None,
|
||||
Some("9999999999999999999"),
|
||||
);
|
||||
assert_eq!(resolved, None);
|
||||
|
||||
// Missing channel_id on the wire — no match (the binding is restrictive).
|
||||
let resolved =
|
||||
router.resolve_with_channel_id(&ChannelType::Discord, "any-user", None, None);
|
||||
assert_eq!(resolved, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_channel_id_plus_peer_outranks_channel_id_alone() {
|
||||
// user A in #medical → researcher; anyone else in #medical → general.
|
||||
let router = AgentRouter::new();
|
||||
let researcher = AgentId::new();
|
||||
let general = AgentId::new();
|
||||
router.register_agent("researcher".to_string(), researcher);
|
||||
router.register_agent("general".to_string(), general);
|
||||
router.load_bindings(&[
|
||||
AgentBinding {
|
||||
agent: "general".to_string(),
|
||||
match_rule: openfang_types::config::BindingMatchRule {
|
||||
channel_id: Some("ch-medical".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
},
|
||||
AgentBinding {
|
||||
agent: "researcher".to_string(),
|
||||
match_rule: openfang_types::config::BindingMatchRule {
|
||||
channel_id: Some("ch-medical".to_string()),
|
||||
peer_id: Some("user-a".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
},
|
||||
]);
|
||||
|
||||
// user-a in #medical → researcher (more specific wins)
|
||||
let r = router.resolve_with_channel_id(
|
||||
&ChannelType::Discord,
|
||||
"user-a",
|
||||
None,
|
||||
Some("ch-medical"),
|
||||
);
|
||||
assert_eq!(r, Some(researcher));
|
||||
|
||||
// user-b in #medical → general (channel_id alone matches)
|
||||
let r = router.resolve_with_channel_id(
|
||||
&ChannelType::Discord,
|
||||
"user-b",
|
||||
None,
|
||||
Some("ch-medical"),
|
||||
);
|
||||
assert_eq!(r, Some(general));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_binding_match_rule_unknown_field_rejected() {
|
||||
// Typos like `channnel_id` must fail loudly at deserialization rather
|
||||
// than silently producing a wide-open binding. This is the highest-
|
||||
// leverage line in the patch from issue #1127.
|
||||
let bad = r#"{ "channnel_id": "ch-1" }"#;
|
||||
let r: Result<openfang_types::config::BindingMatchRule, _> = serde_json::from_str(bad);
|
||||
assert!(r.is_err(), "unknown field must be rejected by serde");
|
||||
|
||||
// Sanity: known fields still parse.
|
||||
let good = r#"{ "channel_id": "ch-1", "channel": "discord" }"#;
|
||||
let r: openfang_types::config::BindingMatchRule = serde_json::from_str(good).unwrap();
|
||||
assert_eq!(r.channel_id.as_deref(), Some("ch-1"));
|
||||
assert_eq!(r.channel.as_deref(), Some("discord"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,38 @@ const SLACK_API_BASE: &str = "https://slack.com/api";
|
||||
const MAX_BACKOFF: Duration = Duration::from_secs(60);
|
||||
const INITIAL_BACKOFF: Duration = Duration::from_secs(1);
|
||||
const SLACK_MSG_LIMIT: usize = 3000;
|
||||
/// TTL for envelope_id dedup entries. Well above the typical Slack
|
||||
/// connection-rotation overlap window (< 10s).
|
||||
const ENVELOPE_TTL: Duration = Duration::from_secs(60);
|
||||
/// Soft cap on the dedup cache size. When exceeded we GC expired entries.
|
||||
/// Recent envelope IDs are not reused by Slack, so 10k is more than enough.
|
||||
const ENVELOPE_CACHE_CAP: usize = 10_000;
|
||||
|
||||
/// Returns true if `envelope_id` was already seen within `ENVELOPE_TTL`.
|
||||
/// On first sight, records the timestamp and returns false. Performs
|
||||
/// opportunistic GC of expired entries when the cache grows large.
|
||||
///
|
||||
/// Slack Socket Mode delivers the same event to multiple active WebSocket
|
||||
/// connections during connection rotation. Apps must dedupe on `envelope_id`
|
||||
/// to avoid double-processing.
|
||||
fn is_duplicate_envelope(cache: &DashMap<String, Instant>, envelope_id: &str) -> bool {
|
||||
if envelope_id.is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Opportunistic GC: bound growth without per-call work.
|
||||
if cache.len() > ENVELOPE_CACHE_CAP {
|
||||
cache.retain(|_, ts| ts.elapsed() < ENVELOPE_TTL);
|
||||
}
|
||||
|
||||
if let Some(prev) = cache.get(envelope_id) {
|
||||
if prev.elapsed() < ENVELOPE_TTL {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
cache.insert(envelope_id.to_string(), Instant::now());
|
||||
false
|
||||
}
|
||||
|
||||
/// Slack Socket Mode adapter.
|
||||
pub struct SlackAdapter {
|
||||
@@ -41,6 +73,9 @@ pub struct SlackAdapter {
|
||||
auto_thread_reply: bool,
|
||||
/// Whether to unfurl (expand previews for) links in posted messages.
|
||||
unfurl_links: bool,
|
||||
/// Recently-seen envelope_ids. Slack Socket Mode redelivers the same event
|
||||
/// across rotated WebSocket connections; this prevents double-processing.
|
||||
seen_envelopes: Arc<DashMap<String, Instant>>,
|
||||
}
|
||||
|
||||
impl SlackAdapter {
|
||||
@@ -65,6 +100,7 @@ impl SlackAdapter {
|
||||
thread_ttl: Duration::from_secs(thread_ttl_hours * 3600),
|
||||
auto_thread_reply,
|
||||
unfurl_links,
|
||||
seen_envelopes: Arc::new(DashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -161,6 +197,7 @@ impl ChannelAdapter for SlackAdapter {
|
||||
let mut shutdown = self.shutdown_rx.clone();
|
||||
let active_threads = self.active_threads.clone();
|
||||
let auto_thread_reply = self.auto_thread_reply;
|
||||
let seen_envelopes = self.seen_envelopes.clone();
|
||||
|
||||
// Spawn periodic cleanup of expired thread entries.
|
||||
{
|
||||
@@ -288,6 +325,14 @@ impl ChannelAdapter for SlackAdapter {
|
||||
}
|
||||
}
|
||||
|
||||
// Dedup: Slack redelivers the same event on the new
|
||||
// connection during the rotation overlap. Ack on
|
||||
// both, but only forward to the agent once.
|
||||
if is_duplicate_envelope(&seen_envelopes, envelope_id) {
|
||||
debug!("Slack: skipping duplicate envelope_id {envelope_id}");
|
||||
continue;
|
||||
}
|
||||
|
||||
// Extract the event
|
||||
let event = &payload["payload"]["event"];
|
||||
if let Some(msg) = parse_slack_event(
|
||||
@@ -501,6 +546,9 @@ async fn parse_slack_event(
|
||||
|
||||
// Check if the bot was @-mentioned (for group_policy = "mention_only")
|
||||
let mut metadata = HashMap::new();
|
||||
// Stash the Slack user ID so the router can key bindings on user, not channel.
|
||||
// (`sender.platform_id` below is the channel ID, used for the send path.)
|
||||
metadata.insert("sender_user_id".to_string(), serde_json::json!(user_id));
|
||||
if event_type == "app_mention" {
|
||||
metadata.insert("was_mentioned".to_string(), serde_json::Value::Bool(true));
|
||||
}
|
||||
@@ -742,4 +790,58 @@ mod tests {
|
||||
);
|
||||
assert!(!adapter.unfurl_links);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_envelope_dedup_skips_second_delivery() {
|
||||
// Simulates Slack redelivering the same event across a connection
|
||||
// rotation: the envelope is acked on both connections but the agent
|
||||
// must only see it once.
|
||||
let cache: DashMap<String, Instant> = DashMap::new();
|
||||
let envelope_id = "8d2e1c5a-4f3b-49a1-b6e2-7c0a9f1234ab";
|
||||
|
||||
// First delivery on the old connection: not a duplicate, forward.
|
||||
assert!(
|
||||
!is_duplicate_envelope(&cache, envelope_id),
|
||||
"first sight of envelope must not be flagged as duplicate"
|
||||
);
|
||||
|
||||
// Second delivery on the new connection: duplicate, skip.
|
||||
assert!(
|
||||
is_duplicate_envelope(&cache, envelope_id),
|
||||
"second sight of same envelope must be flagged as duplicate"
|
||||
);
|
||||
|
||||
// Simulate the receive-loop pattern: count how many times the agent
|
||||
// would actually be invoked across two deliveries.
|
||||
let mut agent_invocations = 0;
|
||||
for _delivery in 0..2 {
|
||||
if !is_duplicate_envelope(&cache, envelope_id) {
|
||||
agent_invocations += 1;
|
||||
}
|
||||
}
|
||||
assert_eq!(
|
||||
agent_invocations, 0,
|
||||
"after initial double-delivery, no further invocations should occur within TTL"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_envelope_dedup_distinct_ids_pass_through() {
|
||||
let cache: DashMap<String, Instant> = DashMap::new();
|
||||
assert!(!is_duplicate_envelope(&cache, "envelope-a"));
|
||||
assert!(!is_duplicate_envelope(&cache, "envelope-b"));
|
||||
assert!(!is_duplicate_envelope(&cache, "envelope-c"));
|
||||
// Each unique envelope_id should be seen exactly once.
|
||||
assert_eq!(cache.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_envelope_dedup_empty_id_never_dedupes() {
|
||||
// Defensive: malformed payloads with no envelope_id should not poison
|
||||
// the cache or short-circuit forwarding.
|
||||
let cache: DashMap<String, Instant> = DashMap::new();
|
||||
assert!(!is_duplicate_envelope(&cache, ""));
|
||||
assert!(!is_duplicate_envelope(&cache, ""));
|
||||
assert_eq!(cache.len(), 0);
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -50,6 +50,16 @@ pub enum ChannelContent {
|
||||
File {
|
||||
url: String,
|
||||
filename: String,
|
||||
/// Best-effort MIME type from the source platform (e.g. Discord's
|
||||
/// `attachments[].content_type`). `None` if the platform did not
|
||||
/// provide one; downstream consumers may sniff bytes or fall back
|
||||
/// to extension-based detection.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
mime: Option<String>,
|
||||
/// Size in bytes, when known. Useful for capacity gating before
|
||||
/// the bridge attempts to materialize or transmit the file.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
size: Option<u64>,
|
||||
},
|
||||
/// Local file data (bytes read from disk). Used by the proactive `channel_send`
|
||||
/// tool when `file_path` is provided instead of `file_url`.
|
||||
@@ -70,6 +80,12 @@ pub enum ChannelContent {
|
||||
name: String,
|
||||
args: Vec<String>,
|
||||
},
|
||||
/// A composite message carrying multiple content blocks (e.g. a Discord
|
||||
/// message with several attachments, or an image with a separate file
|
||||
/// sibling). Blocks are flat-mapped by the bridge into multiple LLM
|
||||
/// content blocks. Implementations should not produce nested `Multipart`
|
||||
/// values; consumers may `debug_assert!` against nesting.
|
||||
Multipart(Vec<ChannelContent>),
|
||||
}
|
||||
|
||||
/// A unified message from any channel.
|
||||
@@ -97,6 +113,60 @@ pub struct ChannelMessage {
|
||||
pub metadata: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
// Re-export the adapter allowlist from openfang-types so config validation
|
||||
// and routing share a single source of truth (no drift between the two).
|
||||
pub use openfang_types::config::CHANNELS_WITH_PLATFORM_ID_AS_CHANNEL;
|
||||
|
||||
impl ChannelMessage {
|
||||
/// Return the platform-native channel/conversation ID for this message,
|
||||
/// suitable for matching against an `AgentBinding`'s `channel_id` field.
|
||||
///
|
||||
/// Resolution order:
|
||||
/// 1. For adapters in [`CHANNELS_WITH_PLATFORM_ID_AS_CHANNEL`],
|
||||
/// `sender.platform_id` already *is* the channel ID (these adapters
|
||||
/// overload the field because it doubles as the send target).
|
||||
/// 2. Otherwise, fall back to `metadata["channel_id"]` if present (any
|
||||
/// adapter can opt in by populating that key).
|
||||
/// 3. Otherwise, `None`.
|
||||
///
|
||||
/// This is the central routing-time accessor — config validation and the
|
||||
/// router both consult it (directly or via the same allowlist) so the two
|
||||
/// cannot drift.
|
||||
pub fn channel_id(&self) -> Option<String> {
|
||||
// For builtin variants the string is already lowercase by construction.
|
||||
// For `Custom(s)`, adapters _should_ register lowercase names but we
|
||||
// case-fold here so a stray `Custom("Twitch")` cannot silently slip
|
||||
// past the allowlist (and out of step with the validation path, which
|
||||
// already lowercases user input). Allocates only on the Custom arm.
|
||||
let channel_str: std::borrow::Cow<'_, str> = match &self.channel {
|
||||
ChannelType::Telegram => "telegram".into(),
|
||||
ChannelType::Discord => "discord".into(),
|
||||
ChannelType::Slack => "slack".into(),
|
||||
ChannelType::WhatsApp => "whatsapp".into(),
|
||||
ChannelType::Signal => "signal".into(),
|
||||
ChannelType::Matrix => "matrix".into(),
|
||||
ChannelType::Email => "email".into(),
|
||||
ChannelType::Teams => "teams".into(),
|
||||
ChannelType::Mattermost => "mattermost".into(),
|
||||
ChannelType::WebChat => "webchat".into(),
|
||||
ChannelType::CLI => "cli".into(),
|
||||
ChannelType::Mqtt => "mqtt".into(),
|
||||
ChannelType::Custom(s) => s.to_lowercase().into(),
|
||||
};
|
||||
if CHANNELS_WITH_PLATFORM_ID_AS_CHANNEL
|
||||
.iter()
|
||||
.any(|c| *c == channel_str.as_ref())
|
||||
{
|
||||
Some(self.sender.platform_id.clone())
|
||||
} else {
|
||||
self.metadata
|
||||
.get("channel_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Agent lifecycle phase for UX indicators.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
@@ -271,6 +341,24 @@ pub trait ChannelAdapter: Send + Sync {
|
||||
self.send(user, content).await
|
||||
}
|
||||
|
||||
/// Determine whether to auto-create a thread for an incoming message.
|
||||
/// Returns Some(thread_name) to create a thread, or None to reply directly.
|
||||
/// Default implementation returns None (no auto-threading).
|
||||
async fn should_auto_thread(&self, _message: &ChannelMessage) -> Option<String> {
|
||||
None
|
||||
}
|
||||
|
||||
/// Create a new thread (typically triggered after should_auto_thread returns Some).
|
||||
/// Returns the new thread ID on success.
|
||||
async fn create_thread(
|
||||
&self,
|
||||
_user: &ChannelUser,
|
||||
_message_id: &str,
|
||||
_thread_name: &str,
|
||||
) -> Result<String, Box<dyn std::error::Error>> {
|
||||
Err("Thread creation not supported for this adapter".into())
|
||||
}
|
||||
|
||||
/// Whether this adapter should suppress sending internal agent errors back to the user.
|
||||
///
|
||||
/// Returns `true` for public broadcast channels (e.g. Mastodon) where posting
|
||||
@@ -365,6 +453,34 @@ mod tests {
|
||||
assert_eq!(back, ChannelType::Email);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_channel_id_custom_arm_is_case_insensitive() {
|
||||
// A stray capitalized Custom variant must still resolve through the
|
||||
// allowlist. The validation path lowercases user input; the routing
|
||||
// path needs the same case-fold to stay in sync.
|
||||
let make = |name: &str| ChannelMessage {
|
||||
channel: ChannelType::Custom(name.to_string()),
|
||||
platform_message_id: "m".to_string(),
|
||||
sender: ChannelUser {
|
||||
platform_id: "C123".to_string(),
|
||||
display_name: "x".to_string(),
|
||||
openfang_user: None,
|
||||
},
|
||||
content: ChannelContent::Text("hi".to_string()),
|
||||
target_agent: None,
|
||||
timestamp: Utc::now(),
|
||||
is_group: false,
|
||||
thread_id: None,
|
||||
metadata: HashMap::new(),
|
||||
};
|
||||
assert_eq!(make("twitch").channel_id().as_deref(), Some("C123"));
|
||||
assert_eq!(make("Twitch").channel_id().as_deref(), Some("C123"));
|
||||
assert_eq!(make("TWITCH").channel_id().as_deref(), Some("C123"));
|
||||
// Lark spelling (Feishu Intl) must also match.
|
||||
assert_eq!(make("lark").channel_id().as_deref(), Some("C123"));
|
||||
assert_eq!(make("Lark").channel_id().as_deref(), Some("C123"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_channel_content_variants() {
|
||||
let text = ChannelContent::Text("hello".to_string());
|
||||
|
||||
@@ -271,7 +271,7 @@ impl ChannelAdapter for WhatsAppAdapter {
|
||||
return Err(format!("WhatsApp API error {status}: {body}").into());
|
||||
}
|
||||
}
|
||||
ChannelContent::File { url, filename } => {
|
||||
ChannelContent::File { url, filename, .. } => {
|
||||
let body = serde_json::json!({
|
||||
"messaging_product": "whatsapp",
|
||||
"to": user.platform_id,
|
||||
|
||||
@@ -1476,9 +1476,30 @@ fn provider_list() -> Vec<(&'static str, &'static str, &'static str, &'static st
|
||||
"openrouter/google/gemini-2.5-flash",
|
||||
"OpenRouter",
|
||||
),
|
||||
("minimax", "MINIMAX_API_KEY", "MiniMax-M2.7", "MiniMax"),
|
||||
]
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod provider_list_tests {
|
||||
use super::provider_list;
|
||||
|
||||
#[test]
|
||||
fn provider_list_includes_minimax() {
|
||||
let minimax = provider_list()
|
||||
.into_iter()
|
||||
.find(|(provider, _, _, _)| *provider == "minimax");
|
||||
assert!(
|
||||
minimax.is_some(),
|
||||
"MiniMax should be exposed by provider_list()"
|
||||
);
|
||||
let (_, env_var, model, display) = minimax.unwrap();
|
||||
assert_eq!(env_var, "MINIMAX_API_KEY");
|
||||
assert_eq!(model, "MiniMax-M2.7");
|
||||
assert_eq!(display, "MiniMax");
|
||||
}
|
||||
}
|
||||
|
||||
/// Quick probe to check if Ollama is running on localhost.
|
||||
fn check_ollama_available() -> bool {
|
||||
std::net::TcpStream::connect_timeout(
|
||||
|
||||
@@ -2023,10 +2023,8 @@ impl App {
|
||||
match canonical_head.as_str() {
|
||||
"/exit" => self.handle_chat_action(chat::ChatAction::Back),
|
||||
"/help" => {
|
||||
self.chat.push_message(
|
||||
chat::Role::System,
|
||||
commands::render_help(Surfaces::CLI),
|
||||
);
|
||||
self.chat
|
||||
.push_message(chat::Role::System, commands::render_help(Surfaces::CLI));
|
||||
}
|
||||
"/status" => {
|
||||
let mut s = Vec::new();
|
||||
|
||||
@@ -576,13 +576,11 @@ impl AgentSelectState {
|
||||
KeyCode::Esc => {
|
||||
self.sub = AgentSubScreen::CreateMethod;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if !self.custom_name.is_empty() {
|
||||
if self.custom_desc.is_empty() {
|
||||
self.custom_desc = format!("A custom {} agent", self.custom_name);
|
||||
}
|
||||
self.sub = AgentSubScreen::CustomDesc;
|
||||
KeyCode::Enter if !self.custom_name.is_empty() => {
|
||||
if self.custom_desc.is_empty() {
|
||||
self.custom_desc = format!("A custom {} agent", self.custom_name);
|
||||
}
|
||||
self.sub = AgentSubScreen::CustomDesc;
|
||||
}
|
||||
KeyCode::Char(c) => {
|
||||
self.custom_name.push(c);
|
||||
@@ -641,15 +639,11 @@ impl AgentSelectState {
|
||||
KeyCode::Esc => {
|
||||
self.sub = AgentSubScreen::CustomPrompt;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if self.tool_cursor > 0 {
|
||||
self.tool_cursor -= 1;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if self.tool_cursor > 0 => {
|
||||
self.tool_cursor -= 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if self.tool_cursor < TOOL_OPTIONS.len() - 1 {
|
||||
self.tool_cursor += 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if self.tool_cursor < TOOL_OPTIONS.len() - 1 => {
|
||||
self.tool_cursor += 1;
|
||||
}
|
||||
KeyCode::Char(' ') => {
|
||||
self.tool_checks[self.tool_cursor] = !self.tool_checks[self.tool_cursor];
|
||||
@@ -674,21 +668,15 @@ impl AgentSelectState {
|
||||
KeyCode::Esc => {
|
||||
self.sub = AgentSubScreen::CustomTools;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if self.skill_cursor > 0 {
|
||||
self.skill_cursor -= 1;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if self.skill_cursor > 0 => {
|
||||
self.skill_cursor -= 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if len > 0 && self.skill_cursor < len - 1 {
|
||||
self.skill_cursor += 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if len > 0 && self.skill_cursor < len - 1 => {
|
||||
self.skill_cursor += 1;
|
||||
}
|
||||
KeyCode::Char(' ') => {
|
||||
if len > 0 {
|
||||
let checked = &mut self.available_skills[self.skill_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Char(' ') if len > 0 => {
|
||||
let checked = &mut self.available_skills[self.skill_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
// Advance to MCP server selection
|
||||
@@ -706,21 +694,15 @@ impl AgentSelectState {
|
||||
KeyCode::Esc => {
|
||||
self.sub = AgentSubScreen::CustomSkills;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if self.mcp_cursor > 0 {
|
||||
self.mcp_cursor -= 1;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if self.mcp_cursor > 0 => {
|
||||
self.mcp_cursor -= 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if len > 0 && self.mcp_cursor < len - 1 {
|
||||
self.mcp_cursor += 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if len > 0 && self.mcp_cursor < len - 1 => {
|
||||
self.mcp_cursor += 1;
|
||||
}
|
||||
KeyCode::Char(' ') => {
|
||||
if len > 0 {
|
||||
let checked = &mut self.available_mcp[self.mcp_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Char(' ') if len > 0 => {
|
||||
let checked = &mut self.available_mcp[self.mcp_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
let toml = self.build_custom_toml();
|
||||
@@ -737,21 +719,15 @@ impl AgentSelectState {
|
||||
KeyCode::Esc => {
|
||||
self.sub = AgentSubScreen::AgentDetail;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if self.skill_cursor > 0 {
|
||||
self.skill_cursor -= 1;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if self.skill_cursor > 0 => {
|
||||
self.skill_cursor -= 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if len > 0 && self.skill_cursor < len - 1 {
|
||||
self.skill_cursor += 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if len > 0 && self.skill_cursor < len - 1 => {
|
||||
self.skill_cursor += 1;
|
||||
}
|
||||
KeyCode::Char(' ') => {
|
||||
if len > 0 {
|
||||
let checked = &mut self.available_skills[self.skill_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Char(' ') if len > 0 => {
|
||||
let checked = &mut self.available_skills[self.skill_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
// Save — collect checked skill names (none checked = "all")
|
||||
@@ -780,21 +756,15 @@ impl AgentSelectState {
|
||||
KeyCode::Esc => {
|
||||
self.sub = AgentSubScreen::AgentDetail;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if self.mcp_cursor > 0 {
|
||||
self.mcp_cursor -= 1;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if self.mcp_cursor > 0 => {
|
||||
self.mcp_cursor -= 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if len > 0 && self.mcp_cursor < len - 1 {
|
||||
self.mcp_cursor += 1;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if len > 0 && self.mcp_cursor < len - 1 => {
|
||||
self.mcp_cursor += 1;
|
||||
}
|
||||
KeyCode::Char(' ') => {
|
||||
if len > 0 {
|
||||
let checked = &mut self.available_mcp[self.mcp_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Char(' ') if len > 0 => {
|
||||
let checked = &mut self.available_mcp[self.mcp_cursor].1;
|
||||
*checked = !*checked;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
// Save — collect checked server names (none checked = "all")
|
||||
|
||||
@@ -164,19 +164,15 @@ impl AuditState {
|
||||
|
||||
let total = self.filtered.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('f') => {
|
||||
self.action_filter = self.action_filter.next();
|
||||
|
||||
@@ -155,19 +155,19 @@ impl CommsState {
|
||||
self.task_field = 0;
|
||||
}
|
||||
KeyCode::Char('r') => return CommsAction::Refresh,
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if self.focus == CommsFocus::EventList && !self.events.is_empty() {
|
||||
let i = self.event_list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { self.events.len() - 1 } else { i - 1 };
|
||||
self.event_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k')
|
||||
if self.focus == CommsFocus::EventList && !self.events.is_empty() =>
|
||||
{
|
||||
let i = self.event_list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { self.events.len() - 1 } else { i - 1 };
|
||||
self.event_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if self.focus == CommsFocus::EventList && !self.events.is_empty() {
|
||||
let i = self.event_list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % self.events.len();
|
||||
self.event_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j')
|
||||
if self.focus == CommsFocus::EventList && !self.events.is_empty() =>
|
||||
{
|
||||
let i = self.event_list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % self.events.len();
|
||||
self.event_list_state.select(Some(next));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
@@ -189,18 +189,17 @@ impl CommsState {
|
||||
self.send_field - 1
|
||||
};
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
KeyCode::Enter
|
||||
if !self.send_from.is_empty()
|
||||
&& !self.send_to.is_empty()
|
||||
&& !self.send_msg.is_empty()
|
||||
{
|
||||
self.show_send_modal = false;
|
||||
return CommsAction::SendMessage {
|
||||
from: self.send_from.clone(),
|
||||
to: self.send_to.clone(),
|
||||
msg: self.send_msg.clone(),
|
||||
};
|
||||
}
|
||||
&& !self.send_msg.is_empty() =>
|
||||
{
|
||||
self.show_send_modal = false;
|
||||
return CommsAction::SendMessage {
|
||||
from: self.send_from.clone(),
|
||||
to: self.send_to.clone(),
|
||||
msg: self.send_msg.clone(),
|
||||
};
|
||||
}
|
||||
KeyCode::Char(c) => match self.send_field {
|
||||
0 => self.send_from.push(c),
|
||||
@@ -238,15 +237,13 @@ impl CommsState {
|
||||
self.task_field - 1
|
||||
};
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if !self.task_title.is_empty() {
|
||||
self.show_task_modal = false;
|
||||
return CommsAction::PostTask {
|
||||
title: self.task_title.clone(),
|
||||
desc: self.task_desc.clone(),
|
||||
assign: self.task_assign.clone(),
|
||||
};
|
||||
}
|
||||
KeyCode::Enter if !self.task_title.is_empty() => {
|
||||
self.show_task_modal = false;
|
||||
return CommsAction::PostTask {
|
||||
title: self.task_title.clone(),
|
||||
desc: self.task_desc.clone(),
|
||||
assign: self.task_assign.clone(),
|
||||
};
|
||||
}
|
||||
KeyCode::Char(c) => match self.task_field {
|
||||
0 => self.task_title.push(c),
|
||||
|
||||
@@ -152,12 +152,10 @@ impl ExtensionsState {
|
||||
self.sub = ExtSub::Health;
|
||||
return ExtensionsAction::RefreshHealth;
|
||||
}
|
||||
KeyCode::Char('/') => {
|
||||
if self.sub == ExtSub::Browse {
|
||||
self.searching = true;
|
||||
self.search_query.clear();
|
||||
return ExtensionsAction::Continue;
|
||||
}
|
||||
KeyCode::Char('/') if self.sub == ExtSub::Browse => {
|
||||
self.searching = true;
|
||||
self.search_query.clear();
|
||||
return ExtensionsAction::Continue;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
@@ -172,19 +170,15 @@ impl ExtensionsState {
|
||||
fn handle_browse(&mut self, key: KeyEvent) -> ExtensionsAction {
|
||||
let total = self.filtered().len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.browse_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.browse_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.browse_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.browse_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.browse_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.browse_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.browse_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.browse_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
let filtered = self.filtered();
|
||||
@@ -222,24 +216,18 @@ impl ExtensionsState {
|
||||
|
||||
let total = self.installed_list_data().len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('d') | KeyCode::Delete => {
|
||||
if self.installed_list.selected().is_some() {
|
||||
self.confirm_remove = true;
|
||||
}
|
||||
KeyCode::Char('d') | KeyCode::Delete if self.installed_list.selected().is_some() => {
|
||||
self.confirm_remove = true;
|
||||
}
|
||||
KeyCode::Char('r') => return ExtensionsAction::RefreshAll,
|
||||
_ => {}
|
||||
@@ -250,19 +238,15 @@ impl ExtensionsState {
|
||||
fn handle_health(&mut self, key: KeyEvent) -> ExtensionsAction {
|
||||
let total = self.health_entries.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.health_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.health_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.health_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.health_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.health_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.health_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.health_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.health_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('r') | KeyCode::Enter => {
|
||||
if let Some(sel) = self.health_list.selected() {
|
||||
|
||||
@@ -109,19 +109,15 @@ impl HandsState {
|
||||
fn handle_marketplace(&mut self, key: KeyEvent) -> HandsAction {
|
||||
let total = self.definitions.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.marketplace_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.marketplace_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.marketplace_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.marketplace_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.marketplace_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.marketplace_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.marketplace_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.marketplace_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Enter | KeyCode::Char('a') => {
|
||||
if let Some(sel) = self.marketplace_list.selected() {
|
||||
@@ -157,24 +153,18 @@ impl HandsState {
|
||||
|
||||
let total = self.instances.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.active_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.active_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.active_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.active_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.active_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.active_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.active_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.active_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('d') | KeyCode::Delete => {
|
||||
if self.active_list.selected().is_some() {
|
||||
self.confirm_deactivate = true;
|
||||
}
|
||||
KeyCode::Char('d') | KeyCode::Delete if self.active_list.selected().is_some() => {
|
||||
self.confirm_deactivate = true;
|
||||
}
|
||||
KeyCode::Char('p') => {
|
||||
if let Some(sel) = self.active_list.selected() {
|
||||
|
||||
@@ -148,6 +148,14 @@ const PROVIDERS: &[ProviderInfo] = &[
|
||||
needs_key: true,
|
||||
hint: "",
|
||||
},
|
||||
ProviderInfo {
|
||||
name: "minimax",
|
||||
display: "MiniMax",
|
||||
env_var: "MINIMAX_API_KEY",
|
||||
default_model: "MiniMax-M2.7",
|
||||
needs_key: true,
|
||||
hint: "",
|
||||
},
|
||||
ProviderInfo {
|
||||
name: "huggingface",
|
||||
display: "Hugging Face",
|
||||
@@ -258,6 +266,23 @@ pub enum InitResult {
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::PROVIDERS;
|
||||
|
||||
#[test]
|
||||
fn init_wizard_lists_minimax_provider() {
|
||||
let minimax = PROVIDERS.iter().find(|provider| provider.name == "minimax");
|
||||
assert!(
|
||||
minimax.is_some(),
|
||||
"MiniMax should be selectable in openfang init"
|
||||
);
|
||||
let minimax = minimax.unwrap();
|
||||
assert_eq!(minimax.env_var, "MINIMAX_API_KEY");
|
||||
assert_eq!(minimax.default_model, "MiniMax-M2.7");
|
||||
}
|
||||
}
|
||||
|
||||
// ── Internal state ─────────────────────────────────────────────────────────
|
||||
|
||||
#[derive(Clone, Copy, PartialEq, Eq)]
|
||||
@@ -966,16 +991,15 @@ pub fn run() -> InitResult {
|
||||
state.step = Step::Provider;
|
||||
}
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
KeyCode::Enter
|
||||
if matches!(
|
||||
state.copilot_auth_status,
|
||||
CopilotAuthStatus::WaitingForUser
|
||||
) && !state.copilot_verification_uri.is_empty()
|
||||
{
|
||||
let _ = openfang_runtime::drivers::copilot::open_verification_url(
|
||||
&state.copilot_verification_uri,
|
||||
);
|
||||
}
|
||||
) && !state.copilot_verification_uri.is_empty() =>
|
||||
{
|
||||
let _ = openfang_runtime::drivers::copilot::open_verification_url(
|
||||
&state.copilot_verification_uri,
|
||||
);
|
||||
}
|
||||
_ => {}
|
||||
},
|
||||
@@ -990,41 +1014,36 @@ pub fn run() -> InitResult {
|
||||
state.key_test = KeyTestState::Idle;
|
||||
state.step = Step::Provider;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
KeyCode::Enter
|
||||
if !state.api_key_input.is_empty()
|
||||
&& state.key_test == KeyTestState::Idle
|
||||
{
|
||||
if let Some(p) = state.provider() {
|
||||
let _ = crate::dotenv::save_env_key(
|
||||
p.env_var,
|
||||
&state.api_key_input,
|
||||
);
|
||||
}
|
||||
state.key_test = KeyTestState::Testing;
|
||||
let provider_name = state
|
||||
.provider()
|
||||
.map(|p| p.name.to_string())
|
||||
.unwrap_or_default();
|
||||
let env_var = state
|
||||
.provider()
|
||||
.map(|p| p.env_var.to_string())
|
||||
.unwrap_or_default();
|
||||
let tx = test_tx.clone();
|
||||
std::thread::spawn(move || {
|
||||
let ok = crate::test_api_key(&provider_name, &env_var);
|
||||
let _ = tx.send(ok);
|
||||
});
|
||||
&& state.key_test == KeyTestState::Idle =>
|
||||
{
|
||||
if let Some(p) = state.provider() {
|
||||
let _ = crate::dotenv::save_env_key(
|
||||
p.env_var,
|
||||
&state.api_key_input,
|
||||
);
|
||||
}
|
||||
state.key_test = KeyTestState::Testing;
|
||||
let provider_name = state
|
||||
.provider()
|
||||
.map(|p| p.name.to_string())
|
||||
.unwrap_or_default();
|
||||
let env_var = state
|
||||
.provider()
|
||||
.map(|p| p.env_var.to_string())
|
||||
.unwrap_or_default();
|
||||
let tx = test_tx.clone();
|
||||
std::thread::spawn(move || {
|
||||
let ok = crate::test_api_key(&provider_name, &env_var);
|
||||
let _ = tx.send(ok);
|
||||
});
|
||||
}
|
||||
KeyCode::Char(c) => {
|
||||
if state.key_test == KeyTestState::Idle {
|
||||
state.api_key_input.push(c);
|
||||
}
|
||||
KeyCode::Char(c) if state.key_test == KeyTestState::Idle => {
|
||||
state.api_key_input.push(c);
|
||||
}
|
||||
KeyCode::Backspace => {
|
||||
if state.key_test == KeyTestState::Idle {
|
||||
state.api_key_input.pop();
|
||||
}
|
||||
KeyCode::Backspace if state.key_test == KeyTestState::Idle => {
|
||||
state.api_key_input.pop();
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
@@ -211,19 +211,15 @@ impl LogsState {
|
||||
|
||||
let total = self.filtered.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('f') => {
|
||||
self.level_filter = self.level_filter.next();
|
||||
@@ -237,15 +233,11 @@ impl LogsState {
|
||||
self.auto_refresh = !self.auto_refresh;
|
||||
}
|
||||
KeyCode::Char('r') => return LogsAction::Refresh,
|
||||
KeyCode::End => {
|
||||
if total > 0 {
|
||||
self.list_state.select(Some(total - 1));
|
||||
}
|
||||
KeyCode::End if total > 0 => {
|
||||
self.list_state.select(Some(total - 1));
|
||||
}
|
||||
KeyCode::Home => {
|
||||
if total > 0 {
|
||||
self.list_state.select(Some(0));
|
||||
}
|
||||
KeyCode::Home if total > 0 => {
|
||||
self.list_state.select(Some(0));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
@@ -106,19 +106,15 @@ impl MemoryState {
|
||||
fn handle_agent_select(&mut self, key: KeyEvent) -> MemoryAction {
|
||||
let total = self.agents.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.agent_list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.agent_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.agent_list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.agent_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.agent_list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.agent_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.agent_list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.agent_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if let Some(sel) = self.agent_list_state.selected() {
|
||||
@@ -166,19 +162,15 @@ impl MemoryState {
|
||||
self.kv_pairs.clear();
|
||||
self.selected_agent = None;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.kv_list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.kv_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.kv_list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.kv_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.kv_list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.kv_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.kv_list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.kv_list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('a') => {
|
||||
self.sub = MemorySub::AddKey;
|
||||
@@ -196,10 +188,8 @@ impl MemoryState {
|
||||
}
|
||||
}
|
||||
}
|
||||
KeyCode::Char('d') => {
|
||||
if self.kv_list_state.selected().is_some() {
|
||||
self.confirm_delete = true;
|
||||
}
|
||||
KeyCode::Char('d') if self.kv_list_state.selected().is_some() => {
|
||||
self.confirm_delete = true;
|
||||
}
|
||||
KeyCode::Char('r') => {
|
||||
if let Some(agent) = &self.selected_agent {
|
||||
|
||||
@@ -62,19 +62,15 @@ impl PeersState {
|
||||
}
|
||||
let total = self.peers.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('r') => return PeersAction::Refresh,
|
||||
_ => {}
|
||||
|
||||
@@ -130,19 +130,15 @@ impl SessionsState {
|
||||
|
||||
let total = self.filtered.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if let Some(sel) = self.list_state.selected() {
|
||||
@@ -155,10 +151,8 @@ impl SessionsState {
|
||||
}
|
||||
}
|
||||
}
|
||||
KeyCode::Char('d') => {
|
||||
if self.list_state.selected().is_some() {
|
||||
self.confirm_delete = true;
|
||||
}
|
||||
KeyCode::Char('d') if self.list_state.selected().is_some() => {
|
||||
self.confirm_delete = true;
|
||||
}
|
||||
KeyCode::Char('/') => {
|
||||
self.search_mode = true;
|
||||
|
||||
@@ -174,21 +174,17 @@ impl SettingsState {
|
||||
fn handle_providers(&mut self, key: KeyEvent) -> SettingsAction {
|
||||
let total = self.providers.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.provider_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.provider_list.select(Some(next));
|
||||
self.test_result = None;
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.provider_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.provider_list.select(Some(next));
|
||||
self.test_result = None;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.provider_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.provider_list.select(Some(next));
|
||||
self.test_result = None;
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.provider_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.provider_list.select(Some(next));
|
||||
self.test_result = None;
|
||||
}
|
||||
KeyCode::Char('e') => {
|
||||
if let Some(sel) = self.provider_list.selected() {
|
||||
@@ -223,19 +219,15 @@ impl SettingsState {
|
||||
fn handle_models(&mut self, key: KeyEvent) -> SettingsAction {
|
||||
let total = self.models.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('r') => return SettingsAction::RefreshModels,
|
||||
_ => {}
|
||||
@@ -246,19 +238,15 @@ impl SettingsState {
|
||||
fn handle_tools(&mut self, key: KeyEvent) -> SettingsAction {
|
||||
let total = self.tools.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.tool_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.tool_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.tool_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.tool_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.tool_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.tool_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.tool_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.tool_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('r') => return SettingsAction::RefreshTools,
|
||||
_ => {}
|
||||
|
||||
@@ -192,24 +192,18 @@ impl SkillsState {
|
||||
|
||||
let total = self.installed.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.installed_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.installed_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('u') => {
|
||||
if self.installed_list.selected().is_some() {
|
||||
self.confirm_uninstall = true;
|
||||
}
|
||||
KeyCode::Char('u') if self.installed_list.selected().is_some() => {
|
||||
self.confirm_uninstall = true;
|
||||
}
|
||||
KeyCode::Char('c') => {
|
||||
if let Some(sel) = self.installed_list.selected() {
|
||||
@@ -219,8 +213,7 @@ impl SkillsState {
|
||||
if self.installed[sel].config_declared > 0 {
|
||||
return SkillsAction::LoadSkillConfig(name);
|
||||
} else {
|
||||
self.status_msg =
|
||||
format!("'{}' declares no runtime config.", name);
|
||||
self.status_msg = format!("'{}' declares no runtime config.", name);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -256,19 +249,15 @@ impl SkillsState {
|
||||
|
||||
let total = self.clawhub_results.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.clawhub_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.clawhub_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.clawhub_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.clawhub_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.clawhub_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.clawhub_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.clawhub_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.clawhub_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('i') => {
|
||||
if let Some(sel) = self.clawhub_list.selected() {
|
||||
@@ -296,19 +285,15 @@ impl SkillsState {
|
||||
fn handle_mcp(&mut self, key: KeyEvent) -> SkillsAction {
|
||||
let total = self.mcp_servers.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.mcp_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.mcp_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.mcp_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.mcp_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.mcp_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.mcp_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.mcp_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.mcp_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('r') => return SkillsAction::RefreshMcp,
|
||||
_ => {}
|
||||
@@ -531,10 +516,7 @@ fn draw_skill_config_details(f: &mut Frame, area: Rect, state: &SkillsState) {
|
||||
|
||||
if rows.is_empty() {
|
||||
f.render_widget(
|
||||
Paragraph::new(Span::styled(
|
||||
"No config declared.",
|
||||
theme::dim_style(),
|
||||
)),
|
||||
Paragraph::new(Span::styled("No config declared.", theme::dim_style())),
|
||||
inner,
|
||||
);
|
||||
return;
|
||||
|
||||
@@ -194,19 +194,15 @@ impl TemplatesState {
|
||||
|
||||
let total = self.filtered.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.list_state.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.list_state.select(Some(next));
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if let Some(sel) = self.list_state.selected() {
|
||||
|
||||
@@ -156,19 +156,15 @@ impl TriggerState {
|
||||
self.create_step -= 1;
|
||||
}
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if self.create_step < 5 {
|
||||
self.create_step += 1;
|
||||
}
|
||||
KeyCode::Enter if self.create_step < 5 => {
|
||||
self.create_step += 1;
|
||||
}
|
||||
KeyCode::Char(c) => match self.create_step {
|
||||
0 => self.create_agent_id.push(c),
|
||||
2 => self.create_pattern_param.push(c),
|
||||
3 => self.create_prompt.push(c),
|
||||
4 => {
|
||||
if c.is_ascii_digit() {
|
||||
self.create_max_fires.push(c);
|
||||
}
|
||||
4 if c.is_ascii_digit() => {
|
||||
self.create_max_fires.push(c);
|
||||
}
|
||||
_ => {}
|
||||
},
|
||||
|
||||
@@ -111,19 +111,15 @@ impl UsageState {
|
||||
UsageSub::ByModel => {
|
||||
let total = self.by_model.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.model_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.model_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('r') => return UsageAction::Refresh,
|
||||
_ => {}
|
||||
@@ -132,19 +128,15 @@ impl UsageState {
|
||||
UsageSub::ByAgent => {
|
||||
let total = self.by_agent.len();
|
||||
match key.code {
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
if total > 0 {
|
||||
let i = self.agent_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.agent_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') if total > 0 => {
|
||||
let i = self.agent_list.selected().unwrap_or(0);
|
||||
let next = if i == 0 { total - 1 } else { i - 1 };
|
||||
self.agent_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if total > 0 {
|
||||
let i = self.agent_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.agent_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') if total > 0 => {
|
||||
let i = self.agent_list.selected().unwrap_or(0);
|
||||
let next = (i + 1) % total;
|
||||
self.agent_list.select(Some(next));
|
||||
}
|
||||
KeyCode::Char('r') => return UsageAction::Refresh,
|
||||
_ => {}
|
||||
|
||||
@@ -85,6 +85,12 @@ const PROVIDERS: &[ProviderInfo] = &[
|
||||
default_model: "qwen-plus",
|
||||
needs_key: true,
|
||||
},
|
||||
ProviderInfo {
|
||||
name: "minimax",
|
||||
env_var: "MINIMAX_API_KEY",
|
||||
default_model: "MiniMax-M2.7",
|
||||
needs_key: true,
|
||||
},
|
||||
ProviderInfo {
|
||||
name: "perplexity",
|
||||
env_var: "PERPLEXITY_API_KEY",
|
||||
@@ -322,13 +328,11 @@ impl WizardState {
|
||||
KeyCode::Esc => {
|
||||
self.step = WizardStep::Provider;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if !self.api_key_input.is_empty() {
|
||||
if let Some(p) = self.selected_provider_info() {
|
||||
self.model_input = p.default_model.to_string();
|
||||
}
|
||||
self.step = WizardStep::Model;
|
||||
KeyCode::Enter if !self.api_key_input.is_empty() => {
|
||||
if let Some(p) = self.selected_provider_info() {
|
||||
self.model_input = p.default_model.to_string();
|
||||
}
|
||||
self.step = WizardStep::Model;
|
||||
}
|
||||
KeyCode::Char(c) => {
|
||||
self.api_key_input.push(c);
|
||||
@@ -689,3 +693,20 @@ fn draw_done(f: &mut Frame, area: Rect, state: &WizardState) {
|
||||
f.render_widget(cont, chunks[1]);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::PROVIDERS;
|
||||
|
||||
#[test]
|
||||
fn wizard_lists_minimax_provider() {
|
||||
let minimax = PROVIDERS.iter().find(|provider| provider.name == "minimax");
|
||||
assert!(
|
||||
minimax.is_some(),
|
||||
"MiniMax should be selectable in wizard provider list"
|
||||
);
|
||||
let minimax = minimax.unwrap();
|
||||
assert_eq!(minimax.env_var, "MINIMAX_API_KEY");
|
||||
assert_eq!(minimax.default_model, "MiniMax-M2.7");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"$schema": "https://schema.tauri.app/config/2",
|
||||
"productName": "OpenFang",
|
||||
"version": "0.6.0",
|
||||
"version": "0.6.9",
|
||||
"identifier": "ai.openfang.desktop",
|
||||
"build": {},
|
||||
"app": {
|
||||
|
||||
@@ -16,6 +16,8 @@ uuid = { workspace = true }
|
||||
chrono = { workspace = true }
|
||||
dashmap = { workspace = true }
|
||||
dirs = { workspace = true }
|
||||
sha2 = { workspace = true }
|
||||
hex = { workspace = true }
|
||||
|
||||
[dev-dependencies]
|
||||
tokio-test = { workspace = true }
|
||||
|
||||
@@ -8,10 +8,27 @@ use crate::{
|
||||
use dashmap::DashMap;
|
||||
use openfang_types::agent::AgentId;
|
||||
use serde::Serialize;
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{Arc, RwLock};
|
||||
use tracing::{info, warn};
|
||||
use uuid::Uuid;
|
||||
|
||||
/// Callback signature invoked on every successful hand load / reload.
|
||||
///
|
||||
/// Arguments: `(hand_id, sha256_hex_of_hand_toml)`. The kernel wires this
|
||||
/// into the Merkle audit chain so reload events leave a tamper-evident
|
||||
/// record (issue #1172). The callback must be cheap and non-blocking; it
|
||||
/// runs inline on the loader thread.
|
||||
pub type HandAuditCallback = Arc<dyn Fn(&str, &str) + Send + Sync>;
|
||||
|
||||
/// Compute the SHA-256 hex digest of raw HAND.toml content.
|
||||
fn hand_toml_sha256(toml_content: &str) -> String {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(toml_content.as_bytes());
|
||||
hex::encode(hasher.finalize())
|
||||
}
|
||||
|
||||
// ─── Settings availability types ────────────────────────────────────────────
|
||||
|
||||
/// Availability status of a single setting option.
|
||||
@@ -41,6 +58,10 @@ pub struct HandRegistry {
|
||||
definitions: DashMap<String, HandDefinition>,
|
||||
/// Active hand instances, keyed by instance UUID.
|
||||
instances: DashMap<Uuid, HandInstance>,
|
||||
/// Optional callback invoked on every successful HAND.toml load with
|
||||
/// the computed SHA-256 of the file content. Wired by the kernel into
|
||||
/// the Merkle audit chain (issue #1172).
|
||||
audit_callback: RwLock<Option<HandAuditCallback>>,
|
||||
}
|
||||
|
||||
impl HandRegistry {
|
||||
@@ -49,9 +70,37 @@ impl HandRegistry {
|
||||
Self {
|
||||
definitions: DashMap::new(),
|
||||
instances: DashMap::new(),
|
||||
audit_callback: RwLock::new(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// Install a callback invoked on every successful HAND.toml load /
|
||||
/// reload with the file's SHA-256. The kernel wires this to
|
||||
/// `AuditLog::record(AuditAction::ConfigChange, ...)` so reload events
|
||||
/// leave a tamper-evident audit record (issue #1172).
|
||||
pub fn set_audit_callback(&self, callback: HandAuditCallback) {
|
||||
let mut guard = self
|
||||
.audit_callback
|
||||
.write()
|
||||
.unwrap_or_else(|e| e.into_inner());
|
||||
*guard = Some(callback);
|
||||
}
|
||||
|
||||
/// Compute SHA-256 of the given HAND.toml content and invoke the audit
|
||||
/// callback if one is registered. Returns the hex digest for callers
|
||||
/// that want to log or compare it.
|
||||
fn emit_hand_loaded_audit(&self, hand_id: &str, toml_content: &str) -> String {
|
||||
let hash = hand_toml_sha256(toml_content);
|
||||
let guard = self
|
||||
.audit_callback
|
||||
.read()
|
||||
.unwrap_or_else(|e| e.into_inner());
|
||||
if let Some(cb) = guard.as_ref() {
|
||||
cb(hand_id, &hash);
|
||||
}
|
||||
hash
|
||||
}
|
||||
|
||||
/// Persist active hand state to disk so it survives restarts.
|
||||
pub fn persist_state(&self, path: &std::path::Path) -> HandResult<()> {
|
||||
let entries: Vec<serde_json::Value> = self
|
||||
@@ -112,7 +161,8 @@ impl HandRegistry {
|
||||
for (id, toml_content, skill_content) in bundled {
|
||||
match bundled::parse_bundled(id, toml_content, skill_content) {
|
||||
Ok(def) => {
|
||||
info!(hand = %def.id, name = %def.name, "Loaded bundled hand");
|
||||
let hash = self.emit_hand_loaded_audit(&def.id, toml_content);
|
||||
info!(hand = %def.id, name = %def.name, sha256 = %hash, "Loaded bundled hand");
|
||||
self.definitions.insert(def.id.clone(), def);
|
||||
count += 1;
|
||||
}
|
||||
@@ -173,7 +223,8 @@ impl HandRegistry {
|
||||
match bundled::parse_bundled("custom", &contents, &skill_content) {
|
||||
Ok(def) => {
|
||||
let hand_id = def.id.clone();
|
||||
info!(hand = %hand_id, path = %path.display(), "Loaded workspace hand");
|
||||
let hash = self.emit_hand_loaded_audit(&hand_id, &contents);
|
||||
info!(hand = %hand_id, path = %path.display(), sha256 = %hash, "Loaded workspace hand");
|
||||
self.definitions.insert(hand_id, def);
|
||||
count += 1;
|
||||
}
|
||||
@@ -204,7 +255,8 @@ impl HandRegistry {
|
||||
)));
|
||||
}
|
||||
|
||||
info!(hand = %def.id, name = %def.name, path = %path.display(), "Installed hand from path");
|
||||
let hash = self.emit_hand_loaded_audit(&def.id, &toml_content);
|
||||
info!(hand = %def.id, name = %def.name, path = %path.display(), sha256 = %hash, "Installed hand from path");
|
||||
self.definitions.insert(def.id.clone(), def.clone());
|
||||
|
||||
// Persist the hand to the user's data dir so it survives daemon
|
||||
@@ -251,7 +303,8 @@ impl HandRegistry {
|
||||
)));
|
||||
}
|
||||
|
||||
info!(hand = %def.id, name = %def.name, "Installed hand from content");
|
||||
let hash = self.emit_hand_loaded_audit(&def.id, toml_content);
|
||||
info!(hand = %def.id, name = %def.name, sha256 = %hash, "Installed hand from content");
|
||||
self.definitions.insert(def.id.clone(), def.clone());
|
||||
Ok(def)
|
||||
}
|
||||
@@ -269,7 +322,8 @@ impl HandRegistry {
|
||||
let def = bundled::parse_bundled("custom", toml_content, skill_content)?;
|
||||
let existed = self.definitions.contains_key(&def.id);
|
||||
let verb = if existed { "Updated" } else { "Installed" };
|
||||
info!(hand = %def.id, name = %def.name, "{verb} hand from content");
|
||||
let hash = self.emit_hand_loaded_audit(&def.id, toml_content);
|
||||
info!(hand = %def.id, name = %def.name, sha256 = %hash, "{verb} hand from content");
|
||||
self.definitions.insert(def.id.clone(), def.clone());
|
||||
Ok(def)
|
||||
}
|
||||
@@ -1145,6 +1199,114 @@ metrics = []
|
||||
assert!(matches!(err, HandError::AlreadyActive(_)));
|
||||
}
|
||||
|
||||
/// Issue #1172: HAND.toml SHA-256 must be emitted to the audit
|
||||
/// callback on every successful load / reload.
|
||||
#[test]
|
||||
fn audit_callback_records_hand_toml_hash_on_load() {
|
||||
use std::sync::Mutex;
|
||||
|
||||
let captured: Arc<Mutex<Vec<(String, String)>>> = Arc::new(Mutex::new(Vec::new()));
|
||||
let sink = Arc::clone(&captured);
|
||||
|
||||
let reg = HandRegistry::new();
|
||||
reg.set_audit_callback(Arc::new(move |hand_id: &str, hash: &str| {
|
||||
sink.lock()
|
||||
.unwrap()
|
||||
.push((hand_id.to_string(), hash.to_string()));
|
||||
}));
|
||||
|
||||
let toml_str = r#"
|
||||
id = "audit-hand"
|
||||
name = "Audit Hand"
|
||||
description = "Used to verify audit-trail wiring"
|
||||
category = "other"
|
||||
tools = []
|
||||
|
||||
[agent]
|
||||
name = "audit-agent"
|
||||
description = "audit"
|
||||
system_prompt = "audit."
|
||||
"#;
|
||||
// Precompute the expected hash so the test fails loudly if the
|
||||
// registry ever changes how it digests the TOML content.
|
||||
let expected_hash = {
|
||||
let mut h = Sha256::new();
|
||||
h.update(toml_str.as_bytes());
|
||||
hex::encode(h.finalize())
|
||||
};
|
||||
|
||||
let def = reg.install_from_content(toml_str, "").unwrap();
|
||||
assert_eq!(def.id, "audit-hand");
|
||||
|
||||
let events = captured.lock().unwrap().clone();
|
||||
assert_eq!(
|
||||
events.len(),
|
||||
1,
|
||||
"exactly one audit event should be emitted per load"
|
||||
);
|
||||
assert_eq!(events[0].0, "audit-hand", "hand id propagated to callback");
|
||||
assert_eq!(
|
||||
events[0].1, expected_hash,
|
||||
"callback received SHA-256 of the HAND.toml content"
|
||||
);
|
||||
assert_eq!(events[0].1.len(), 64, "SHA-256 hex is 64 chars");
|
||||
}
|
||||
|
||||
/// Issue #1172: reloading the same HAND.toml via upsert must emit a
|
||||
/// fresh audit event so the chain records when the swap took effect.
|
||||
/// A content change must surface a different hash.
|
||||
#[test]
|
||||
fn audit_callback_fires_on_reload_with_new_hash() {
|
||||
use std::sync::Mutex;
|
||||
|
||||
let captured: Arc<Mutex<Vec<(String, String)>>> = Arc::new(Mutex::new(Vec::new()));
|
||||
let sink = Arc::clone(&captured);
|
||||
|
||||
let reg = HandRegistry::new();
|
||||
reg.set_audit_callback(Arc::new(move |hand_id: &str, hash: &str| {
|
||||
sink.lock()
|
||||
.unwrap()
|
||||
.push((hand_id.to_string(), hash.to_string()));
|
||||
}));
|
||||
|
||||
let v1 = r#"
|
||||
id = "reload-hand"
|
||||
name = "Reload Hand v1"
|
||||
description = "v1"
|
||||
category = "other"
|
||||
tools = []
|
||||
|
||||
[agent]
|
||||
name = "reload-agent"
|
||||
description = "reload"
|
||||
system_prompt = "v1."
|
||||
"#;
|
||||
let v2 = r#"
|
||||
id = "reload-hand"
|
||||
name = "Reload Hand v2"
|
||||
description = "v2"
|
||||
category = "other"
|
||||
tools = []
|
||||
|
||||
[agent]
|
||||
name = "reload-agent"
|
||||
description = "reload"
|
||||
system_prompt = "v2 — schedule changed."
|
||||
"#;
|
||||
|
||||
reg.upsert_from_content(v1, "").unwrap();
|
||||
reg.upsert_from_content(v2, "").unwrap();
|
||||
|
||||
let events = captured.lock().unwrap().clone();
|
||||
assert_eq!(events.len(), 2, "one event per upsert (load + reload)");
|
||||
assert_eq!(events[0].0, "reload-hand");
|
||||
assert_eq!(events[1].0, "reload-hand");
|
||||
assert_ne!(
|
||||
events[0].1, events[1].1,
|
||||
"different HAND.toml content must yield different SHA-256"
|
||||
);
|
||||
}
|
||||
|
||||
/// Integration test for issue #809: `hand config` round-trip.
|
||||
///
|
||||
/// Simulates what `openfang hand config <id> --set KEY=VAL` does against
|
||||
|
||||
@@ -32,6 +32,7 @@ futures = { workspace = true }
|
||||
subtle = { workspace = true }
|
||||
rand = { workspace = true }
|
||||
hex = { workspace = true }
|
||||
sha2 = { workspace = true }
|
||||
reqwest = { workspace = true }
|
||||
rustls = { workspace = true }
|
||||
cron = "0.16"
|
||||
|
||||
@@ -69,12 +69,26 @@ pub fn load_config(path: Option<&Path>) -> KernelConfig {
|
||||
}
|
||||
}
|
||||
|
||||
// GAP-012 (Tier 1): pre-validate the [[bindings]] array so a
|
||||
// single malformed entry doesn't poison the whole config and
|
||||
// force a fall-back to defaults (which would silently unbind
|
||||
// every agent). Bad entries are logged at ERROR and dropped;
|
||||
// survivors are passed through to typed deserialization.
|
||||
lenient_extract_bindings(&mut root_value);
|
||||
|
||||
match root_value.try_into::<KernelConfig>() {
|
||||
Ok(config) => {
|
||||
info!(path = %config_path.display(), "Loaded configuration");
|
||||
return config;
|
||||
}
|
||||
Err(e) => {
|
||||
// TODO(GAP-012-Tier-2): this fallback still silently
|
||||
// swaps the user's intent for `KernelConfig::default()`
|
||||
// on any non-binding deserialization failure. Tier 1
|
||||
// closes the binding-shape footgun; Tier 2 should
|
||||
// surface remaining failures via a health endpoint
|
||||
// and/or stderr banner so the silent-default path
|
||||
// can't hide a broken config.
|
||||
tracing::warn!(
|
||||
error = %e,
|
||||
path = %config_path.display(),
|
||||
@@ -242,6 +256,89 @@ pub fn deep_merge_toml(base: &mut toml::Value, overlay: &toml::Value) {
|
||||
}
|
||||
}
|
||||
|
||||
/// Lenient pre-pass over the `[[bindings]]` array (GAP-012 Tier 1).
|
||||
///
|
||||
/// Strict whole-config deserialization is fragile: any one malformed binding
|
||||
/// (e.g. a typo'd field that trips `deny_unknown_fields`) causes
|
||||
/// `try_into::<KernelConfig>()` to fail, which the caller then handles by
|
||||
/// falling back to `KernelConfig::default()` — silently unbinding *every*
|
||||
/// agent. That's the worst possible failure mode for a routing config: the
|
||||
/// user's intent is silently discarded, with only a single line in the logs.
|
||||
///
|
||||
/// This pass runs *before* typed deserialization. It walks the bindings
|
||||
/// array entry-by-entry, attempts to deserialize each into `AgentBinding`,
|
||||
/// logs malformed entries at ERROR with index + agent name + serde error,
|
||||
/// and replaces the array with the survivors. The downstream
|
||||
/// `try_into::<KernelConfig>()` then sees a clean array and succeeds.
|
||||
///
|
||||
/// `deny_unknown_fields` on `AgentBinding`/`BindingMatchRule` still applies
|
||||
/// per-entry — typos in surviving bindings would still produce errors here
|
||||
/// and be dropped. The strict-field guarantee is preserved at the entry
|
||||
/// level; only the all-or-nothing behavior is relaxed.
|
||||
///
|
||||
/// No-op if `root_value` is not a table or has no `bindings` array.
|
||||
fn lenient_extract_bindings(root_value: &mut toml::Value) {
|
||||
use openfang_types::config::AgentBinding;
|
||||
|
||||
let tbl = match root_value {
|
||||
toml::Value::Table(t) => t,
|
||||
_ => return,
|
||||
};
|
||||
|
||||
// Replace the array in place if (and only if) `bindings` is present
|
||||
// and is an array. Anything else (missing, wrong type) we leave alone
|
||||
// so the typed deserializer can produce its own targeted error.
|
||||
let original = match tbl.get("bindings") {
|
||||
Some(toml::Value::Array(arr)) => arr.clone(),
|
||||
_ => return,
|
||||
};
|
||||
|
||||
let mut survivors: Vec<toml::Value> = Vec::with_capacity(original.len());
|
||||
let mut dropped = 0usize;
|
||||
|
||||
for (idx, entry) in original.into_iter().enumerate() {
|
||||
match entry.clone().try_into::<AgentBinding>() {
|
||||
Ok(_) => survivors.push(entry),
|
||||
Err(e) => {
|
||||
dropped += 1;
|
||||
// Lazy: only allocate the agent-name fallback string when we
|
||||
// actually need it for an error log. The happy path skips this.
|
||||
let agent_name = entry
|
||||
.get("agent")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("<unknown>")
|
||||
.to_string();
|
||||
tracing::error!(
|
||||
binding_index = idx,
|
||||
agent = %agent_name,
|
||||
error = %e,
|
||||
"Skipping malformed binding #{} (agent='{}'): {}. \
|
||||
Other bindings will continue to load. \
|
||||
Fix the entry and reload to restore routing.",
|
||||
idx,
|
||||
agent_name,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if dropped > 0 {
|
||||
// Per-entry ERRORs above carry the root cause; this summary is a
|
||||
// grep-friendly one-liner, so WARN keeps ERROR == per-binding cause.
|
||||
tracing::warn!(
|
||||
dropped,
|
||||
survivors = survivors.len(),
|
||||
"Dropped {} malformed binding(s); {} binding(s) will load. \
|
||||
See preceding ERROR lines for per-binding details.",
|
||||
dropped,
|
||||
survivors.len()
|
||||
);
|
||||
}
|
||||
|
||||
tbl.insert("bindings".to_string(), toml::Value::Array(survivors));
|
||||
}
|
||||
|
||||
/// Get the default config file path.
|
||||
///
|
||||
/// Respects `OPENFANG_HOME` env var (e.g. `OPENFANG_HOME=/opt/openfang`).
|
||||
@@ -442,6 +539,224 @@ mod tests {
|
||||
assert_eq!(config.log_level, "info"); // defaults
|
||||
}
|
||||
|
||||
// ─── GAP-012 Tier 1: lenient bindings extraction ───────────────────
|
||||
|
||||
#[test]
|
||||
fn test_lenient_bindings_drops_typo_keeps_rest() {
|
||||
// Two bindings; the first has a typo'd field (`channnel_id`) that
|
||||
// `BindingMatchRule`'s `deny_unknown_fields` would reject. The second
|
||||
// is well-formed. Pre-fix behavior: whole config falls back to
|
||||
// defaults (zero bindings). Post-fix: bad one dropped, good one loads.
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("config.toml");
|
||||
let mut f = std::fs::File::create(&path).unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"
|
||||
log_level = "info"
|
||||
|
||||
[[bindings]]
|
||||
agent = "researcher-broken"
|
||||
match_rule = {{ channel = "discord", channnel_id = "123" }}
|
||||
|
||||
[[bindings]]
|
||||
agent = "researcher-good"
|
||||
match_rule = {{ channel = "discord", channel_id = "456" }}
|
||||
"#
|
||||
)
|
||||
.unwrap();
|
||||
drop(f);
|
||||
|
||||
let config = load_config(Some(&path));
|
||||
assert_eq!(
|
||||
config.bindings.len(),
|
||||
1,
|
||||
"expected exactly the well-formed binding to survive"
|
||||
);
|
||||
assert_eq!(config.bindings[0].agent, "researcher-good");
|
||||
assert_eq!(
|
||||
config.bindings[0].match_rule.channel_id.as_deref(),
|
||||
Some("456")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lenient_bindings_all_valid_unchanged() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("config.toml");
|
||||
let mut f = std::fs::File::create(&path).unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"
|
||||
log_level = "info"
|
||||
|
||||
[[bindings]]
|
||||
agent = "a"
|
||||
match_rule = {{ channel = "discord", channel_id = "1" }}
|
||||
|
||||
[[bindings]]
|
||||
agent = "b"
|
||||
match_rule = {{ channel = "telegram", channel_id = "2" }}
|
||||
"#
|
||||
)
|
||||
.unwrap();
|
||||
drop(f);
|
||||
|
||||
let config = load_config(Some(&path));
|
||||
assert_eq!(config.bindings.len(), 2);
|
||||
assert_eq!(config.bindings[0].agent, "a");
|
||||
assert_eq!(config.bindings[1].agent, "b");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lenient_bindings_all_malformed_yields_empty_but_keeps_rest_of_config() {
|
||||
// Every binding is broken, but the rest of the config (log_level,
|
||||
// api_listen) must still load. Pre-fix: total fallback to defaults.
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("config.toml");
|
||||
let mut f = std::fs::File::create(&path).unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"
|
||||
log_level = "trace"
|
||||
api_listen = "127.0.0.1:9999"
|
||||
|
||||
[[bindings]]
|
||||
agent = "broken-1"
|
||||
match_rule = {{ channnel_id = "1" }}
|
||||
|
||||
[[bindings]]
|
||||
agent = "broken-2"
|
||||
match_rule = {{ peer_idd = "u" }}
|
||||
"#
|
||||
)
|
||||
.unwrap();
|
||||
drop(f);
|
||||
|
||||
let config = load_config(Some(&path));
|
||||
assert!(config.bindings.is_empty(), "all bindings should be dropped");
|
||||
assert_eq!(
|
||||
config.log_level, "trace",
|
||||
"non-binding config must still load"
|
||||
);
|
||||
assert_eq!(config.api_listen, "127.0.0.1:9999");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lenient_bindings_no_bindings_section_is_noop() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("config.toml");
|
||||
let mut f = std::fs::File::create(&path).unwrap();
|
||||
writeln!(f, "log_level = \"info\"").unwrap();
|
||||
drop(f);
|
||||
|
||||
let config = load_config(Some(&path));
|
||||
assert!(config.bindings.is_empty());
|
||||
assert_eq!(config.log_level, "info");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lenient_bindings_missing_agent_field_dropped() {
|
||||
// A binding missing the required `agent` field can't deserialize at
|
||||
// all; it should be dropped (logged as agent='<unknown>') and the
|
||||
// good one should still load.
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("config.toml");
|
||||
let mut f = std::fs::File::create(&path).unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"
|
||||
[[bindings]]
|
||||
match_rule = {{ channel = "discord" }}
|
||||
|
||||
[[bindings]]
|
||||
agent = "good"
|
||||
match_rule = {{ channel = "discord", channel_id = "1" }}
|
||||
"#
|
||||
)
|
||||
.unwrap();
|
||||
drop(f);
|
||||
|
||||
let config = load_config(Some(&path));
|
||||
assert_eq!(config.bindings.len(), 1);
|
||||
assert_eq!(config.bindings[0].agent, "good");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lenient_bindings_preserves_survivor_order() {
|
||||
// Three bindings with the *middle* one malformed. Survivors must
|
||||
// retain their original relative order (1st, 3rd) — match-rule
|
||||
// routing can be order-sensitive (first-match-wins), so silently
|
||||
// reshuffling on a drop would be a subtle regression.
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("config.toml");
|
||||
let mut f = std::fs::File::create(&path).unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"
|
||||
[[bindings]]
|
||||
agent = "first"
|
||||
match_rule = {{ channel = "discord", channel_id = "1" }}
|
||||
|
||||
[[bindings]]
|
||||
agent = "middle-broken"
|
||||
match_rule = {{ channnel_id = "2" }}
|
||||
|
||||
[[bindings]]
|
||||
agent = "third"
|
||||
match_rule = {{ channel = "telegram", channel_id = "3" }}
|
||||
"#
|
||||
)
|
||||
.unwrap();
|
||||
drop(f);
|
||||
|
||||
let config = load_config(Some(&path));
|
||||
assert_eq!(config.bindings.len(), 2, "middle binding should be dropped");
|
||||
assert_eq!(
|
||||
config.bindings[0].agent, "first",
|
||||
"first survivor must remain first"
|
||||
);
|
||||
assert_eq!(
|
||||
config.bindings[1].agent, "third",
|
||||
"third must remain after first (order preserved)"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lenient_bindings_top_level_field_typo_dropped() {
|
||||
// Operator typos `agnt` instead of `agent` on the binding itself
|
||||
// (not inside `match_rule`). `AgentBinding`'s `deny_unknown_fields`
|
||||
// should reject the entry, the lenient pass should drop it, and
|
||||
// the well-formed sibling should still load. This is the more
|
||||
// common operator mistake than missing-field-entirely, so we lock
|
||||
// the behavior in explicitly.
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("config.toml");
|
||||
let mut f = std::fs::File::create(&path).unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"
|
||||
[[bindings]]
|
||||
agnt = "typo-at-top-level"
|
||||
match_rule = {{ channel = "discord", channel_id = "1" }}
|
||||
|
||||
[[bindings]]
|
||||
agent = "good"
|
||||
match_rule = {{ channel = "discord", channel_id = "2" }}
|
||||
"#
|
||||
)
|
||||
.unwrap();
|
||||
drop(f);
|
||||
|
||||
let config = load_config(Some(&path));
|
||||
assert_eq!(
|
||||
config.bindings.len(),
|
||||
1,
|
||||
"binding with top-level field typo should be dropped"
|
||||
);
|
||||
assert_eq!(config.bindings[0].agent, "good");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_no_includes_works() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
|
||||
@@ -483,6 +483,79 @@ mod tests {
|
||||
assert!(plan.hot_actions.contains(&HotAction::ReloadProviderUrls));
|
||||
}
|
||||
|
||||
/// #1129: editing `[default_model].subprocess_timeout_secs` must produce
|
||||
/// a hot-reload action so cross-message timeout retunes don't require
|
||||
/// a daemon bounce. The whole `default_model` block round-trips through
|
||||
/// `UpdateDefaultModel`, which carries the new timeout into the override
|
||||
/// slot read by `resolve_driver`.
|
||||
#[test]
|
||||
fn test_default_model_subprocess_timeout_hot_reload() {
|
||||
let a = default_cfg();
|
||||
let mut b = default_cfg();
|
||||
b.default_model.subprocess_timeout_secs = Some(900);
|
||||
let plan = build_reload_plan(&a, &b);
|
||||
assert!(
|
||||
!plan.restart_required,
|
||||
"subprocess_timeout_secs edits on default_model must be hot-reloadable"
|
||||
);
|
||||
assert!(plan.hot_actions.contains(&HotAction::UpdateDefaultModel));
|
||||
}
|
||||
|
||||
/// #1129: editing `[[fallback_providers]]` (including
|
||||
/// `subprocess_timeout_secs` on a non-default provider) must produce a
|
||||
/// `ReloadFallbackProviders` hot-action. Without this, mixed-fleet
|
||||
/// operators have no live tuning knob for their non-default driver.
|
||||
#[test]
|
||||
fn test_fallback_providers_subprocess_timeout_hot_reload() {
|
||||
use openfang_types::config::FallbackProviderConfig;
|
||||
let mut a = default_cfg();
|
||||
let mut b = default_cfg();
|
||||
a.fallback_providers.push(FallbackProviderConfig {
|
||||
provider: "codex".to_string(),
|
||||
model: "gpt-5-codex".to_string(),
|
||||
api_key_env: String::new(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: Some(120),
|
||||
});
|
||||
b.fallback_providers.push(FallbackProviderConfig {
|
||||
provider: "codex".to_string(),
|
||||
model: "gpt-5-codex".to_string(),
|
||||
api_key_env: String::new(),
|
||||
base_url: None,
|
||||
// Operator raises the ceiling for slow Codex turns.
|
||||
subprocess_timeout_secs: Some(900),
|
||||
});
|
||||
let plan = build_reload_plan(&a, &b);
|
||||
assert!(
|
||||
!plan.restart_required,
|
||||
"[[fallback_providers]] edits must be hot-reloadable"
|
||||
);
|
||||
assert!(plan
|
||||
.hot_actions
|
||||
.contains(&HotAction::ReloadFallbackProviders));
|
||||
}
|
||||
|
||||
/// #1129: adding a brand-new `[[fallback_providers]]` entry on reload also
|
||||
/// emits the hot-action so the new provider is picked up without bounce.
|
||||
#[test]
|
||||
fn test_fallback_providers_add_entry_hot_reload() {
|
||||
use openfang_types::config::FallbackProviderConfig;
|
||||
let a = default_cfg();
|
||||
let mut b = default_cfg();
|
||||
b.fallback_providers.push(FallbackProviderConfig {
|
||||
provider: "ollama".to_string(),
|
||||
model: "llama3.2:latest".to_string(),
|
||||
api_key_env: String::new(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: Some(300),
|
||||
});
|
||||
let plan = build_reload_plan(&a, &b);
|
||||
assert!(!plan.restart_required);
|
||||
assert!(plan
|
||||
.hot_actions
|
||||
.contains(&HotAction::ReloadFallbackProviders));
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Mixed changes
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
@@ -436,7 +436,9 @@ mod tests {
|
||||
auth_header: Some("Bearer test-token".to_string()),
|
||||
};
|
||||
let engine = test_engine(MockBridge::new());
|
||||
let results = engine.deliver(&[target], "daily-report", "result body").await;
|
||||
let results = engine
|
||||
.deliver(&[target], "daily-report", "result body")
|
||||
.await;
|
||||
|
||||
assert!(results[0].success, "error: {:?}", results[0].error);
|
||||
|
||||
@@ -732,8 +734,6 @@ mod tests {
|
||||
}
|
||||
|
||||
fn find_subsequence(haystack: &[u8], needle: &[u8]) -> Option<usize> {
|
||||
haystack
|
||||
.windows(needle.len())
|
||||
.position(|w| w == needle)
|
||||
haystack.windows(needle.len()).position(|w| w == needle)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
use crate::registry::AgentRegistry;
|
||||
use chrono::Utc;
|
||||
use dashmap::DashMap;
|
||||
use openfang_types::agent::{AgentId, AgentState};
|
||||
use openfang_types::agent::{AgentEntry, AgentId, AgentState, ScheduleMode};
|
||||
use tracing::{debug, warn};
|
||||
|
||||
/// Default heartbeat check interval (seconds).
|
||||
@@ -132,6 +132,14 @@ impl Default for RecoveryTracker {
|
||||
/// and the initial `set_state(Running)` call.
|
||||
const IDLE_GRACE_SECS: i64 = 10;
|
||||
|
||||
/// Reactive agents are healthy while idle between user messages.
|
||||
///
|
||||
/// They should only participate in heartbeat failure detection while a turn is
|
||||
/// actively running. Otherwise silence is the expected steady state.
|
||||
pub(crate) fn should_exempt_idle_reactive_agent(entry: &AgentEntry, is_running_task: bool) -> bool {
|
||||
matches!(entry.manifest.schedule, ScheduleMode::Reactive) && !is_running_task
|
||||
}
|
||||
|
||||
/// Check all running and crashed agents and return their heartbeat status.
|
||||
///
|
||||
/// This is a pure function — it doesn't start a background task.
|
||||
@@ -331,11 +339,13 @@ mod tests {
|
||||
autonomous: None,
|
||||
pinned_model: None,
|
||||
workspace: None,
|
||||
state_dir: None,
|
||||
generate_identity_files: true,
|
||||
exec_policy: None,
|
||||
tool_allowlist: vec![],
|
||||
tool_blocklist: vec![],
|
||||
cache_context: false,
|
||||
max_history_messages: None,
|
||||
},
|
||||
state,
|
||||
mode: AgentMode::default(),
|
||||
@@ -376,6 +386,35 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_idle_reactive_agent_is_exempt_when_not_processing() {
|
||||
let mut agent = make_entry(
|
||||
"reactive-idle",
|
||||
AgentState::Running,
|
||||
Utc::now() - Duration::seconds(600),
|
||||
Utc::now() - Duration::seconds(300),
|
||||
);
|
||||
agent.manifest.schedule = ScheduleMode::Reactive;
|
||||
|
||||
assert!(should_exempt_idle_reactive_agent(&agent, false));
|
||||
assert!(!should_exempt_idle_reactive_agent(&agent, true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_periodic_agent_is_not_exempt_when_idle() {
|
||||
let mut agent = make_entry(
|
||||
"periodic-idle",
|
||||
AgentState::Running,
|
||||
Utc::now() - Duration::seconds(600),
|
||||
Utc::now() - Duration::seconds(300),
|
||||
);
|
||||
agent.manifest.schedule = ScheduleMode::Periodic {
|
||||
cron: "0 * * * *".to_string(),
|
||||
};
|
||||
|
||||
assert!(!should_exempt_idle_reactive_agent(&agent, false));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_active_agent_detected_unresponsive() {
|
||||
// An agent that WAS active (last_active >> created_at) but has gone
|
||||
|
||||
+1836
-116
File diff suppressed because it is too large
Load Diff
@@ -234,6 +234,12 @@ pub struct BudgetStatus {
|
||||
/// Order matters: more specific patterns must come before generic ones
|
||||
/// (e.g. "gpt-4o-mini" before "gpt-4o", "gpt-4.1-mini" before "gpt-4.1").
|
||||
fn estimate_cost_rates(model: &str) -> (f64, f64) {
|
||||
// ── Requesty (issue #995) ──────────────────────────────────
|
||||
// Router-style gateway. IDs are `requesty/<upstream>/<model>` and
|
||||
// resolve via substring match on the upstream model name below
|
||||
// (e.g. "sonnet", "gpt-4o", "gemini", "deepseek", "llama").
|
||||
// No early-return here — fall through to upstream patterns.
|
||||
|
||||
// ── Anthropic ──────────────────────────────────────────────
|
||||
if model.contains("haiku") {
|
||||
return (0.25, 1.25);
|
||||
|
||||
@@ -134,6 +134,23 @@ impl AgentRegistry {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Update an agent's private state directory path. The state directory
|
||||
/// holds identity files, sessions, and per-agent memory and is always
|
||||
/// kept separate from the user-facing workspace. See issue #1097.
|
||||
pub fn update_state_dir(
|
||||
&self,
|
||||
id: AgentId,
|
||||
state_dir: Option<std::path::PathBuf>,
|
||||
) -> OpenFangResult<()> {
|
||||
let mut entry = self
|
||||
.agents
|
||||
.get_mut(&id)
|
||||
.ok_or_else(|| OpenFangError::AgentNotFound(id.to_string()))?;
|
||||
entry.manifest.state_dir = state_dir;
|
||||
entry.last_active = chrono::Utc::now();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Update an agent's visual identity (emoji, avatar, color).
|
||||
pub fn update_identity(
|
||||
&self,
|
||||
@@ -391,11 +408,13 @@ mod tests {
|
||||
autonomous: None,
|
||||
pinned_model: None,
|
||||
workspace: None,
|
||||
state_dir: None,
|
||||
generate_identity_files: true,
|
||||
exec_policy: None,
|
||||
tool_allowlist: vec![],
|
||||
tool_blocklist: vec![],
|
||||
cache_context: false,
|
||||
max_history_messages: None,
|
||||
},
|
||||
state: AgentState::Created,
|
||||
mode: AgentMode::default(),
|
||||
|
||||
@@ -449,6 +449,14 @@ fn describe_event(event: &Event) -> String {
|
||||
"Health check failed: agent {agent_id}, unresponsive for {unresponsive_secs}s"
|
||||
)
|
||||
}
|
||||
SystemEvent::CronJobExecuted {
|
||||
agent_id,
|
||||
job_id,
|
||||
job_name,
|
||||
..
|
||||
} => {
|
||||
format!("Cron job executed: {job_name} ({job_id}) for agent {agent_id}")
|
||||
}
|
||||
},
|
||||
EventPayload::Custom(data) => {
|
||||
format!("Custom event ({} bytes)", data.len())
|
||||
|
||||
@@ -176,6 +176,7 @@ impl SetupWizard {
|
||||
autonomous: None,
|
||||
pinned_model: None,
|
||||
workspace: None,
|
||||
state_dir: None,
|
||||
generate_identity_files: true,
|
||||
profile: None,
|
||||
fallback_models: vec![],
|
||||
@@ -183,6 +184,7 @@ impl SetupWizard {
|
||||
tool_allowlist: vec![],
|
||||
tool_blocklist: vec![],
|
||||
cache_context: false,
|
||||
max_history_messages: None,
|
||||
};
|
||||
|
||||
let skills_to_install: Vec<String> = intent
|
||||
|
||||
@@ -19,6 +19,7 @@ fn test_config() -> KernelConfig {
|
||||
model: "llama-3.3-70b-versatile".to_string(),
|
||||
api_key_env: "GROQ_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
}
|
||||
|
||||
@@ -19,6 +19,7 @@ fn test_config() -> KernelConfig {
|
||||
model: "llama-3.3-70b-versatile".to_string(),
|
||||
api_key_env: "GROQ_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
}
|
||||
|
||||
@@ -115,6 +115,7 @@ fn test_config(tmp: &tempfile::TempDir) -> KernelConfig {
|
||||
model: "test".to_string(),
|
||||
api_key_env: "OLLAMA_API_KEY".to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
}
|
||||
|
||||
@@ -24,6 +24,7 @@ fn test_config(provider: &str, model: &str, api_key_env: &str) -> KernelConfig {
|
||||
model: model.to_string(),
|
||||
api_key_env: api_key_env.to_string(),
|
||||
base_url: None,
|
||||
subprocess_timeout_secs: None,
|
||||
},
|
||||
..KernelConfig::default()
|
||||
}
|
||||
|
||||
@@ -496,7 +496,7 @@ impl SessionStore {
|
||||
.conn
|
||||
.lock()
|
||||
.map_err(|e| OpenFangError::Internal(e.to_string()))?;
|
||||
let messages_blob = rmp_serde::to_vec(&canonical.messages)
|
||||
let messages_blob = rmp_serde::to_vec_named(&canonical.messages)
|
||||
.map_err(|e| OpenFangError::Serialization(e.to_string()))?;
|
||||
conn.execute(
|
||||
"INSERT INTO canonical_sessions (agent_id, messages, compaction_cursor, compacted_summary, updated_at)
|
||||
@@ -586,12 +586,15 @@ impl SessionStore {
|
||||
ContentBlock::Image { media_type, .. } => {
|
||||
text_parts.push(format!("[image: {media_type}]"));
|
||||
}
|
||||
ContentBlock::Thinking { thinking } => {
|
||||
ContentBlock::Thinking { thinking, .. } => {
|
||||
text_parts.push(format!(
|
||||
"[thinking: {}]",
|
||||
openfang_types::truncate_str(thinking, 200)
|
||||
));
|
||||
}
|
||||
ContentBlock::RedactedThinking { .. } => {
|
||||
text_parts.push("[redacted_thinking]".to_string());
|
||||
}
|
||||
ContentBlock::Unknown => {}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -896,10 +896,8 @@ fn derive_capabilities(tools: &[String]) -> AgentCapabilities {
|
||||
"shell_exec" => {
|
||||
caps.shell = vec!["*".to_string()];
|
||||
}
|
||||
"web_fetch" | "web_search" | "browser_navigate" => {
|
||||
if caps.network.is_empty() {
|
||||
caps.network = vec!["*".to_string()];
|
||||
}
|
||||
"web_fetch" | "web_search" | "browser_navigate" if caps.network.is_empty() => {
|
||||
caps.network = vec!["*".to_string()];
|
||||
}
|
||||
"agent_send" | "agent_list" => {
|
||||
if caps.agent_message.is_empty() {
|
||||
|
||||
@@ -40,21 +40,43 @@ const MAX_RETRIES: u32 = 3;
|
||||
/// Base delay for exponential backoff (milliseconds).
|
||||
const BASE_RETRY_DELAY_MS: u64 = 1000;
|
||||
|
||||
/// Timeout for individual tool executions (seconds).
|
||||
/// Default timeout for individual tool executions (seconds).
|
||||
/// Raised from 60s to 120s for browser automation and long-running builds.
|
||||
/// Overridable via `OPENFANG_TOOL_TIMEOUT_SECS` env var. Set to `0` to disable
|
||||
/// the timeout entirely (useful for slow local inference like vLLM on old GPUs).
|
||||
const TOOL_TIMEOUT_SECS: u64 = 120;
|
||||
|
||||
/// Timeout for inter-agent tool calls (seconds).
|
||||
/// Default timeout for inter-agent tool calls (seconds).
|
||||
/// Agent delegation (agent_send, agent_spawn) can involve a full agent loop on the
|
||||
/// target, so these need a significantly longer timeout than regular tools.
|
||||
/// Overridable via `OPENFANG_AGENT_TOOL_TIMEOUT_SECS` env var. Set to `0` to
|
||||
/// disable (issue #1125: slow vLLM rigs running Hands need unbounded waits).
|
||||
const AGENT_TOOL_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// Parse a u64 env var, returning `None` when unset or unparseable so the
|
||||
/// caller falls back to the compiled-in default.
|
||||
fn env_timeout_secs(var: &str) -> Option<u64> {
|
||||
std::env::var(var).ok().and_then(|s| s.trim().parse().ok())
|
||||
}
|
||||
|
||||
/// Returns the appropriate timeout duration for a given tool name.
|
||||
/// Inter-agent calls get a longer timeout since they may trigger full agent loops.
|
||||
fn tool_timeout_for(tool_name: &str) -> Duration {
|
||||
match tool_name {
|
||||
"agent_send" | "agent_spawn" => Duration::from_secs(AGENT_TOOL_TIMEOUT_SECS),
|
||||
_ => Duration::from_secs(TOOL_TIMEOUT_SECS),
|
||||
///
|
||||
/// Returns `None` when the operator opted out by setting the relevant env var
|
||||
/// to `0`. In that case the tool runs with no upper bound, which is what users
|
||||
/// on slow local inference (vLLM on old GPUs) want for Hands and inter-agent
|
||||
/// delegation (issue #1125).
|
||||
fn tool_timeout_for(tool_name: &str) -> Option<Duration> {
|
||||
let secs = match tool_name {
|
||||
"agent_send" | "agent_spawn" => {
|
||||
env_timeout_secs("OPENFANG_AGENT_TOOL_TIMEOUT_SECS").unwrap_or(AGENT_TOOL_TIMEOUT_SECS)
|
||||
}
|
||||
_ => env_timeout_secs("OPENFANG_TOOL_TIMEOUT_SECS").unwrap_or(TOOL_TIMEOUT_SECS),
|
||||
};
|
||||
if secs == 0 {
|
||||
None
|
||||
} else {
|
||||
Some(Duration::from_secs(secs))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -62,8 +84,10 @@ fn tool_timeout_for(tool_name: &str) -> Duration {
|
||||
/// Raised from 3 to 5 to allow longer-form generation.
|
||||
const MAX_CONTINUATIONS: u32 = 5;
|
||||
|
||||
/// Maximum message history size before auto-trimming to prevent context overflow.
|
||||
const MAX_HISTORY_MESSAGES: usize = 20;
|
||||
/// Default maximum message history size before auto-trimming to prevent context overflow.
|
||||
/// Per-agent overrides come from `AgentManifest::max_history_messages` (issue #871).
|
||||
#[allow(dead_code)]
|
||||
const MAX_HISTORY_MESSAGES: usize = openfang_types::agent::DEFAULT_MAX_HISTORY_MESSAGES;
|
||||
|
||||
/// Detect when the LLM claims to have performed an action (sent, posted, emailed)
|
||||
/// without actually calling any tools. Prevents hallucinated completions.
|
||||
@@ -109,6 +133,78 @@ fn append_tool_error_guidance(tool_result_blocks: &mut Vec<ContentBlock>) {
|
||||
}
|
||||
}
|
||||
|
||||
/// Build an assistant message that preserves Thinking blocks alongside the
|
||||
/// final visible text.
|
||||
///
|
||||
/// Issue #1098 — thinking-model state preservation. When the LLM response
|
||||
/// contains `ContentBlock::Thinking` (Anthropic extended thinking with
|
||||
/// signatures, Gemini 2.5+ thoughts, OpenAI-compat reasoning_content,
|
||||
/// MiniMax/Qwen inline `<think>` blocks), the prior code stored only the
|
||||
/// final text via `Message::assistant(text)` — discarding all reasoning
|
||||
/// state. On the next turn the model re-derived its answer from scratch
|
||||
/// and quality degraded.
|
||||
///
|
||||
/// This helper preserves the full block list whenever any Thinking block is
|
||||
/// present, otherwise returns the legacy `Message::assistant(text)` form so
|
||||
/// downstream consumers (channel formatters, JSONL mirrors, embeddings) keep
|
||||
/// working without changes.
|
||||
///
|
||||
/// Note: we deliberately replace any visible Text blocks in `response_blocks`
|
||||
/// with `final_text` so that any post-processing the agent loop applied
|
||||
/// (phantom-action recovery, accumulated_text fallback, EmptyResponse guard
|
||||
/// stub) is reflected in the persisted message.
|
||||
fn build_assistant_message_preserving_thinking(
|
||||
response_blocks: &[ContentBlock],
|
||||
final_text: &str,
|
||||
) -> Message {
|
||||
// Key on either Thinking or RedactedThinking — Anthropic/Bedrock both
|
||||
// reject extended-thinking history that drops the redacted variant, so a
|
||||
// turn that contains only RedactedThinking must still be preserved.
|
||||
let has_reasoning = response_blocks.iter().any(|b| {
|
||||
matches!(
|
||||
b,
|
||||
ContentBlock::Thinking { .. } | ContentBlock::RedactedThinking { .. }
|
||||
)
|
||||
});
|
||||
if !has_reasoning {
|
||||
return Message::assistant(final_text.to_string());
|
||||
}
|
||||
|
||||
// Preserve order: Thinking / RedactedThinking blocks first (in original
|
||||
// order), then a single Text block carrying `final_text`. Tool blocks
|
||||
// aren't expected here (StopReason::EndTurn path), but copy them through
|
||||
// if present so we don't drop information.
|
||||
let mut blocks: Vec<ContentBlock> = Vec::with_capacity(response_blocks.len() + 1);
|
||||
let mut emitted_text = false;
|
||||
for b in response_blocks {
|
||||
match b {
|
||||
ContentBlock::Thinking { .. } | ContentBlock::RedactedThinking { .. } => {
|
||||
blocks.push(b.clone())
|
||||
}
|
||||
ContentBlock::Text { .. } if !emitted_text => {
|
||||
blocks.push(ContentBlock::Text {
|
||||
text: final_text.to_string(),
|
||||
provider_metadata: None,
|
||||
});
|
||||
emitted_text = true;
|
||||
}
|
||||
ContentBlock::Text { .. } => {
|
||||
// Drop additional text blocks — final_text already captures
|
||||
// the canonical visible message.
|
||||
}
|
||||
other => blocks.push(other.clone()),
|
||||
}
|
||||
}
|
||||
if !emitted_text && !final_text.is_empty() {
|
||||
blocks.push(ContentBlock::Text {
|
||||
text: final_text.to_string(),
|
||||
provider_metadata: None,
|
||||
});
|
||||
}
|
||||
|
||||
Message::assistant_with_blocks(blocks)
|
||||
}
|
||||
|
||||
/// Strip a provider prefix from a model ID before sending to the API.
|
||||
///
|
||||
/// Many models are stored as `provider/org/model` (e.g. `openrouter/google/gemini-2.5-flash`)
|
||||
@@ -362,12 +458,15 @@ pub async fn run_agent_loop(
|
||||
// Safety valve: trim excessively long message histories to prevent context overflow.
|
||||
// The full compaction system handles sophisticated summarization, but this prevents
|
||||
// the catastrophic case where 200+ messages cause instant context overflow.
|
||||
if messages.len() > MAX_HISTORY_MESSAGES {
|
||||
let trim_count = messages.len() - MAX_HISTORY_MESSAGES;
|
||||
// Per-agent cap: manifest override -> runtime default (issue #871).
|
||||
let max_history = manifest.effective_max_history_messages();
|
||||
if messages.len() > max_history {
|
||||
let trim_count = messages.len() - max_history;
|
||||
warn!(
|
||||
agent = %manifest.name,
|
||||
total_messages = messages.len(),
|
||||
trimming = trim_count,
|
||||
max_history = max_history,
|
||||
"Trimming old messages to prevent context overflow"
|
||||
);
|
||||
messages.drain(..trim_count);
|
||||
@@ -605,7 +704,16 @@ pub async fn run_agent_loop(
|
||||
};
|
||||
|
||||
final_response = text.clone();
|
||||
session.messages.push(Message::assistant(text));
|
||||
// Issue #1098: persist Thinking blocks alongside the text so
|
||||
// reasoning models retain state across turns. When the
|
||||
// response carries any Thinking content (Anthropic extended
|
||||
// thinking, Gemini 2.5 thought signatures, DeepSeek-R1/Qwen3
|
||||
// `reasoning_content`, MiniMax inline `<think>`), save the
|
||||
// full content blocks; otherwise fall back to the legacy
|
||||
// Text shape so existing sessions/snapshots stay readable.
|
||||
let assistant_msg =
|
||||
build_assistant_message_preserving_thinking(&response.content, &text);
|
||||
session.messages.push(assistant_msg);
|
||||
|
||||
// Prune NO_REPLY heartbeat turns to save context budget
|
||||
crate::session_repair::prune_heartbeat_turns(&mut session.messages, 10);
|
||||
@@ -719,10 +827,12 @@ pub async fn run_agent_loop(
|
||||
session.messages.push(Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(assistant_blocks.clone()),
|
||||
..Default::default()
|
||||
});
|
||||
messages.push(Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(assistant_blocks),
|
||||
..Default::default()
|
||||
});
|
||||
|
||||
// Build allowed tool names list for capability enforcement
|
||||
@@ -813,49 +923,51 @@ pub async fn run_agent_loop(
|
||||
// Resolve effective exec policy (per-agent override or global)
|
||||
let effective_exec_policy = manifest.exec_policy.as_ref();
|
||||
|
||||
// Timeout-wrapped execution
|
||||
let timeout = tool_timeout_for(&tool_call.name);
|
||||
let timeout_secs = timeout.as_secs();
|
||||
let result = match tokio::time::timeout(
|
||||
timeout,
|
||||
tool_runner::execute_tool(
|
||||
&tool_call.id,
|
||||
&tool_call.name,
|
||||
&tool_call.input,
|
||||
kernel.as_ref(),
|
||||
Some(&allowed_tool_names),
|
||||
Some(&caller_id_str),
|
||||
skill_registry,
|
||||
mcp_connections,
|
||||
web_ctx,
|
||||
browser_ctx,
|
||||
if hand_allowed_env.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(&hand_allowed_env)
|
||||
},
|
||||
workspace_root,
|
||||
media_engine,
|
||||
effective_exec_policy,
|
||||
tts_engine,
|
||||
docker_config,
|
||||
process_manager,
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
warn!(tool = %tool_call.name, "Tool execution timed out after {}s", timeout_secs);
|
||||
openfang_types::tool::ToolResult {
|
||||
tool_use_id: tool_call.id.clone(),
|
||||
content: format!(
|
||||
"Tool '{}' timed out after {}s.",
|
||||
tool_call.name, timeout_secs
|
||||
),
|
||||
is_error: true,
|
||||
// Timeout-wrapped execution. `tool_timeout_for` returns None
|
||||
// when the operator disabled the timeout (issue #1125).
|
||||
let timeout_opt = tool_timeout_for(&tool_call.name);
|
||||
let exec_fut = tool_runner::execute_tool(
|
||||
&tool_call.id,
|
||||
&tool_call.name,
|
||||
&tool_call.input,
|
||||
kernel.as_ref(),
|
||||
Some(&allowed_tool_names),
|
||||
Some(&caller_id_str),
|
||||
skill_registry,
|
||||
mcp_connections,
|
||||
web_ctx,
|
||||
browser_ctx,
|
||||
if hand_allowed_env.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(&hand_allowed_env)
|
||||
},
|
||||
workspace_root,
|
||||
media_engine,
|
||||
effective_exec_policy,
|
||||
tts_engine,
|
||||
docker_config,
|
||||
process_manager,
|
||||
);
|
||||
let result = match timeout_opt {
|
||||
Some(timeout) => {
|
||||
let timeout_secs = timeout.as_secs();
|
||||
match tokio::time::timeout(timeout, exec_fut).await {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
warn!(tool = %tool_call.name, "Tool execution timed out after {}s", timeout_secs);
|
||||
openfang_types::tool::ToolResult {
|
||||
tool_use_id: tool_call.id.clone(),
|
||||
content: format!(
|
||||
"Tool '{}' timed out after {}s.",
|
||||
tool_call.name, timeout_secs
|
||||
),
|
||||
is_error: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
None => exec_fut.await,
|
||||
};
|
||||
|
||||
// Fire AfterToolCall hook
|
||||
@@ -939,6 +1051,7 @@ pub async fn run_agent_loop(
|
||||
let tool_results_msg = Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Blocks(tool_result_blocks.clone()),
|
||||
..Default::default()
|
||||
};
|
||||
session.messages.push(tool_results_msg.clone());
|
||||
messages.push(tool_results_msg);
|
||||
@@ -958,7 +1071,12 @@ pub async fn run_agent_loop(
|
||||
} else {
|
||||
text
|
||||
};
|
||||
session.messages.push(Message::assistant(&text));
|
||||
// Issue #1148: preserve Thinking / RedactedThinking blocks
|
||||
// present in the response so reasoning state survives
|
||||
// MaxTokens truncation — same as the EndTurn branch.
|
||||
let assistant_msg =
|
||||
build_assistant_message_preserving_thinking(&response.content, &text);
|
||||
session.messages.push(assistant_msg);
|
||||
if let Err(e) = memory.save_session_async(session).await {
|
||||
warn!("Failed to save session on max continuations: {e}");
|
||||
}
|
||||
@@ -989,10 +1107,15 @@ pub async fn run_agent_loop(
|
||||
directives: Default::default(),
|
||||
});
|
||||
}
|
||||
// Model hit token limit — add partial response and continue
|
||||
// Model hit token limit — add partial response and continue.
|
||||
// Issue #1148: preserve full response content (Thinking,
|
||||
// RedactedThinking, etc.) so reasoning state is not dropped
|
||||
// when continuing across the token-limit boundary.
|
||||
let text = response.text();
|
||||
session.messages.push(Message::assistant(&text));
|
||||
messages.push(Message::assistant(&text));
|
||||
let assistant_msg =
|
||||
build_assistant_message_preserving_thinking(&response.content, &text);
|
||||
session.messages.push(assistant_msg.clone());
|
||||
messages.push(assistant_msg);
|
||||
session.messages.push(Message::user("Please continue."));
|
||||
messages.push(Message::user("Please continue."));
|
||||
warn!(iteration, "Max tokens hit, continuing");
|
||||
@@ -1143,6 +1266,7 @@ async fn call_with_retry(
|
||||
api_key,
|
||||
base_url: fb.base_url.clone(),
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let fb_driver = match crate::drivers::create_driver(&fb_config) {
|
||||
Ok(d) => d,
|
||||
@@ -1326,6 +1450,7 @@ async fn stream_with_retry(
|
||||
api_key,
|
||||
base_url: fb.base_url.clone(),
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let fb_driver = match crate::drivers::create_driver(&fb_config) {
|
||||
Ok(d) => d,
|
||||
@@ -1552,12 +1677,15 @@ pub async fn run_agent_loop_streaming(
|
||||
let mut accumulated_text = String::new();
|
||||
|
||||
// Safety valve: trim excessively long message histories to prevent context overflow.
|
||||
if messages.len() > MAX_HISTORY_MESSAGES {
|
||||
let trim_count = messages.len() - MAX_HISTORY_MESSAGES;
|
||||
// Per-agent cap: manifest override -> runtime default (issue #871).
|
||||
let max_history = manifest.effective_max_history_messages();
|
||||
if messages.len() > max_history {
|
||||
let trim_count = messages.len() - max_history;
|
||||
warn!(
|
||||
agent = %manifest.name,
|
||||
total_messages = messages.len(),
|
||||
trimming = trim_count,
|
||||
max_history = max_history,
|
||||
"Trimming old messages to prevent context overflow (streaming)"
|
||||
);
|
||||
messages.drain(..trim_count);
|
||||
@@ -1657,6 +1785,12 @@ pub async fn run_agent_loop_streaming(
|
||||
}
|
||||
}
|
||||
|
||||
// Stamp last_active before the (potentially long) LLM call so the
|
||||
// heartbeat monitor doesn't flag us as unresponsive mid-iteration.
|
||||
if let Some(k) = &kernel {
|
||||
k.touch_agent(&agent_id_str);
|
||||
}
|
||||
|
||||
// Stream LLM call with retry, error classification, and circuit breaker
|
||||
let provider_name = manifest.model.provider.as_str();
|
||||
let mut response = stream_with_retry(
|
||||
@@ -1790,7 +1924,13 @@ pub async fn run_agent_loop_streaming(
|
||||
text
|
||||
};
|
||||
final_response = text.clone();
|
||||
session.messages.push(Message::assistant(text));
|
||||
// Issue #1098: preserve Thinking blocks (with Anthropic
|
||||
// signatures / Gemini thought signatures / inline-think /
|
||||
// reasoning_content) on the persisted assistant turn. See
|
||||
// build_assistant_message_preserving_thinking for details.
|
||||
let assistant_msg =
|
||||
build_assistant_message_preserving_thinking(&response.content, &text);
|
||||
session.messages.push(assistant_msg);
|
||||
|
||||
// Prune NO_REPLY heartbeat turns to save context budget
|
||||
crate::session_repair::prune_heartbeat_turns(&mut session.messages, 10);
|
||||
@@ -1898,10 +2038,12 @@ pub async fn run_agent_loop_streaming(
|
||||
session.messages.push(Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(assistant_blocks.clone()),
|
||||
..Default::default()
|
||||
});
|
||||
messages.push(Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(assistant_blocks),
|
||||
..Default::default()
|
||||
});
|
||||
|
||||
let allowed_tool_names: Vec<String> =
|
||||
@@ -1990,49 +2132,51 @@ pub async fn run_agent_loop_streaming(
|
||||
// Resolve effective exec policy (per-agent override or global)
|
||||
let effective_exec_policy = manifest.exec_policy.as_ref();
|
||||
|
||||
// Timeout-wrapped execution
|
||||
let timeout = tool_timeout_for(&tool_call.name);
|
||||
let timeout_secs = timeout.as_secs();
|
||||
let result = match tokio::time::timeout(
|
||||
timeout,
|
||||
tool_runner::execute_tool(
|
||||
&tool_call.id,
|
||||
&tool_call.name,
|
||||
&tool_call.input,
|
||||
kernel.as_ref(),
|
||||
Some(&allowed_tool_names),
|
||||
Some(&caller_id_str),
|
||||
skill_registry,
|
||||
mcp_connections,
|
||||
web_ctx,
|
||||
browser_ctx,
|
||||
if hand_allowed_env.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(&hand_allowed_env)
|
||||
},
|
||||
workspace_root,
|
||||
media_engine,
|
||||
effective_exec_policy,
|
||||
tts_engine,
|
||||
docker_config,
|
||||
process_manager,
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
warn!(tool = %tool_call.name, "Tool execution timed out after {}s (streaming)", timeout_secs);
|
||||
openfang_types::tool::ToolResult {
|
||||
tool_use_id: tool_call.id.clone(),
|
||||
content: format!(
|
||||
"Tool '{}' timed out after {}s.",
|
||||
tool_call.name, timeout_secs
|
||||
),
|
||||
is_error: true,
|
||||
// Timeout-wrapped execution. `tool_timeout_for` returns None
|
||||
// when the operator disabled the timeout (issue #1125).
|
||||
let timeout_opt = tool_timeout_for(&tool_call.name);
|
||||
let exec_fut = tool_runner::execute_tool(
|
||||
&tool_call.id,
|
||||
&tool_call.name,
|
||||
&tool_call.input,
|
||||
kernel.as_ref(),
|
||||
Some(&allowed_tool_names),
|
||||
Some(&caller_id_str),
|
||||
skill_registry,
|
||||
mcp_connections,
|
||||
web_ctx,
|
||||
browser_ctx,
|
||||
if hand_allowed_env.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(&hand_allowed_env)
|
||||
},
|
||||
workspace_root,
|
||||
media_engine,
|
||||
effective_exec_policy,
|
||||
tts_engine,
|
||||
docker_config,
|
||||
process_manager,
|
||||
);
|
||||
let result = match timeout_opt {
|
||||
Some(timeout) => {
|
||||
let timeout_secs = timeout.as_secs();
|
||||
match tokio::time::timeout(timeout, exec_fut).await {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
warn!(tool = %tool_call.name, "Tool execution timed out after {}s (streaming)", timeout_secs);
|
||||
openfang_types::tool::ToolResult {
|
||||
tool_use_id: tool_call.id.clone(),
|
||||
content: format!(
|
||||
"Tool '{}' timed out after {}s.",
|
||||
tool_call.name, timeout_secs
|
||||
),
|
||||
is_error: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
None => exec_fut.await,
|
||||
};
|
||||
|
||||
// Fire AfterToolCall hook
|
||||
@@ -2130,6 +2274,7 @@ pub async fn run_agent_loop_streaming(
|
||||
let tool_results_msg = Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Blocks(tool_result_blocks.clone()),
|
||||
..Default::default()
|
||||
};
|
||||
session.messages.push(tool_results_msg.clone());
|
||||
messages.push(tool_results_msg);
|
||||
@@ -2147,7 +2292,12 @@ pub async fn run_agent_loop_streaming(
|
||||
} else {
|
||||
text
|
||||
};
|
||||
session.messages.push(Message::assistant(&text));
|
||||
// Issue #1148: preserve Thinking / RedactedThinking blocks
|
||||
// present in the response so reasoning state survives
|
||||
// MaxTokens truncation — same as the EndTurn branch.
|
||||
let assistant_msg =
|
||||
build_assistant_message_preserving_thinking(&response.content, &text);
|
||||
session.messages.push(assistant_msg);
|
||||
if let Err(e) = memory.save_session_async(session).await {
|
||||
warn!("Failed to save session on max continuations: {e}");
|
||||
}
|
||||
@@ -2178,9 +2328,14 @@ pub async fn run_agent_loop_streaming(
|
||||
directives: Default::default(),
|
||||
});
|
||||
}
|
||||
// Issue #1148: preserve full response content (Thinking,
|
||||
// RedactedThinking, etc.) so reasoning state is not dropped
|
||||
// when continuing across the token-limit boundary.
|
||||
let text = response.text();
|
||||
session.messages.push(Message::assistant(&text));
|
||||
messages.push(Message::assistant(&text));
|
||||
let assistant_msg =
|
||||
build_assistant_message_preserving_thinking(&response.content, &text);
|
||||
session.messages.push(assistant_msg.clone());
|
||||
messages.push(assistant_msg);
|
||||
session.messages.push(Message::user("Please continue."));
|
||||
messages.push(Message::user("Please continue."));
|
||||
warn!(iteration, "Max tokens hit (streaming), continuing");
|
||||
@@ -3084,6 +3239,189 @@ mod tests {
|
||||
assert_eq!(MAX_ITERATIONS, 50);
|
||||
}
|
||||
|
||||
/// Issue #1098: when a response carries Thinking blocks, the persisted
|
||||
/// assistant turn must keep them so the next turn round-trips reasoning
|
||||
/// state to the model.
|
||||
#[test]
|
||||
fn test_build_assistant_message_preserves_thinking() {
|
||||
let response_blocks = vec![
|
||||
ContentBlock::Thinking {
|
||||
thinking: "Let me reason carefully...".to_string(),
|
||||
signature: Some("sig_anthropic_xyz".to_string()),
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "anthropic_extended_thinking"
|
||||
})),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "Initial response text".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
];
|
||||
// Final text might differ from the original Text block (phantom-action
|
||||
// recovery / synthesis fallback rewrites it). The helper should adopt
|
||||
// final_text into the persisted Text block.
|
||||
let final_text = "Initial response text";
|
||||
let msg = build_assistant_message_preserving_thinking(&response_blocks, final_text);
|
||||
assert_eq!(msg.role, Role::Assistant);
|
||||
let blocks = match &msg.content {
|
||||
MessageContent::Blocks(b) => b,
|
||||
other => panic!("expected blocks, got {other:?}"),
|
||||
};
|
||||
assert_eq!(blocks.len(), 2, "must preserve thinking + text");
|
||||
match &blocks[0] {
|
||||
ContentBlock::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(thinking, "Let me reason carefully...");
|
||||
assert_eq!(signature.as_deref(), Some("sig_anthropic_xyz"));
|
||||
}
|
||||
_ => panic!("expected Thinking first"),
|
||||
}
|
||||
match &blocks[1] {
|
||||
ContentBlock::Text { text, .. } => assert_eq!(text, "Initial response text"),
|
||||
_ => panic!("expected Text second"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Without thinking, fall back to the legacy `Message::assistant(text)`
|
||||
/// shape so existing JSONL mirrors and embeddings keep working.
|
||||
#[test]
|
||||
fn test_build_assistant_message_no_thinking_is_plain_text() {
|
||||
let response_blocks = vec![ContentBlock::Text {
|
||||
text: "Hi.".to_string(),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
let msg = build_assistant_message_preserving_thinking(&response_blocks, "Hi.");
|
||||
match msg.content {
|
||||
MessageContent::Text(t) => assert_eq!(t, "Hi."),
|
||||
_ => panic!("expected plain text content for non-thinking responses"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Final text supplied by the loop (e.g. recovery stub) must replace
|
||||
/// the original text part — the persisted message reflects what was
|
||||
/// actually returned to the user, not the raw LLM output.
|
||||
#[test]
|
||||
fn test_build_assistant_message_final_text_replaces_original_text() {
|
||||
let response_blocks = vec![
|
||||
ContentBlock::Thinking {
|
||||
thinking: "deliberation".to_string(),
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({"format": "inline_think"})),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "raw LLM output".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
];
|
||||
let final_text = "[Task completed — recovered after empty response.]";
|
||||
let msg = build_assistant_message_preserving_thinking(&response_blocks, final_text);
|
||||
let blocks = match &msg.content {
|
||||
MessageContent::Blocks(b) => b,
|
||||
_ => panic!("expected blocks"),
|
||||
};
|
||||
let saved_text = blocks.iter().find_map(|b| match b {
|
||||
ContentBlock::Text { text, .. } => Some(text.as_str()),
|
||||
_ => None,
|
||||
});
|
||||
assert_eq!(saved_text, Some(final_text));
|
||||
}
|
||||
|
||||
/// Issue #1148 — when the LLM hits MaxTokens, the persisted assistant
|
||||
/// turn must keep `Thinking` and `RedactedThinking` blocks so reasoning
|
||||
/// state survives across the token-limit boundary. The helper used by
|
||||
/// the MaxTokens branches is the same `build_assistant_message_preserving_thinking`
|
||||
/// that EndTurn uses; this test pins that contract for both block types
|
||||
/// so the four MaxTokens persistence sites stay correct.
|
||||
#[test]
|
||||
fn test_build_assistant_message_preserves_redacted_thinking_for_max_tokens() {
|
||||
let response_blocks = vec![
|
||||
ContentBlock::Thinking {
|
||||
thinking: "Mid-stream reasoning".to_string(),
|
||||
signature: Some("sig_xyz".to_string()),
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "anthropic_extended_thinking"
|
||||
})),
|
||||
},
|
||||
ContentBlock::RedactedThinking {
|
||||
data: "encrypted_blob_abc".to_string(),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "Partial answer before token limit".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
];
|
||||
let final_text = "Partial answer before token limit";
|
||||
let msg = build_assistant_message_preserving_thinking(&response_blocks, final_text);
|
||||
let blocks = match &msg.content {
|
||||
MessageContent::Blocks(b) => b,
|
||||
other => panic!("expected Blocks content for MaxTokens persistence, got {other:?}"),
|
||||
};
|
||||
|
||||
// All reasoning blocks must survive the persistence step so the
|
||||
// follow-up "Please continue." turn carries them back to the model.
|
||||
let has_thinking = blocks
|
||||
.iter()
|
||||
.any(|b| matches!(b, ContentBlock::Thinking { .. }));
|
||||
let has_redacted = blocks
|
||||
.iter()
|
||||
.any(|b| matches!(b, ContentBlock::RedactedThinking { .. }));
|
||||
assert!(
|
||||
has_thinking,
|
||||
"Thinking block must be preserved on MaxTokens"
|
||||
);
|
||||
assert!(
|
||||
has_redacted,
|
||||
"RedactedThinking block must be preserved on MaxTokens"
|
||||
);
|
||||
|
||||
// Verify the opaque blob is byte-identical (Anthropic rejects altered data).
|
||||
for b in blocks {
|
||||
if let ContentBlock::RedactedThinking { data } = b {
|
||||
assert_eq!(data, "encrypted_blob_abc");
|
||||
}
|
||||
}
|
||||
|
||||
// Final text reflects what the user will see.
|
||||
let saved_text = blocks.iter().find_map(|b| match b {
|
||||
ContentBlock::Text { text, .. } => Some(text.as_str()),
|
||||
_ => None,
|
||||
});
|
||||
assert_eq!(saved_text, Some(final_text));
|
||||
}
|
||||
|
||||
/// Issue #1187 — a turn that contains only `RedactedThinking` (no
|
||||
/// `Thinking` block) must still trigger the block-preserving path. The
|
||||
/// previous gate keyed solely on `Thinking`, so redacted-only turns were
|
||||
/// downgraded to plain text and the encrypted blob was lost on the next
|
||||
/// request, which Anthropic/Bedrock reject.
|
||||
#[test]
|
||||
fn test_build_assistant_message_preserves_redacted_only() {
|
||||
let response_blocks = vec![
|
||||
ContentBlock::RedactedThinking {
|
||||
data: "encrypted_only".to_string(),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "Answer".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
];
|
||||
let msg = build_assistant_message_preserving_thinking(&response_blocks, "Answer");
|
||||
let blocks = match &msg.content {
|
||||
MessageContent::Blocks(b) => b,
|
||||
other => panic!("expected Blocks content for redacted-only turn, got {other:?}"),
|
||||
};
|
||||
let has_redacted = blocks.iter().any(
|
||||
|b| matches!(b, ContentBlock::RedactedThinking { data } if data == "encrypted_only"),
|
||||
);
|
||||
assert!(
|
||||
has_redacted,
|
||||
"RedactedThinking-only turn must be preserved as Blocks"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_constants() {
|
||||
assert_eq!(MAX_RETRIES, 3);
|
||||
@@ -3135,17 +3473,129 @@ mod tests {
|
||||
assert_eq!(AGENT_TOOL_TIMEOUT_SECS, 600);
|
||||
}
|
||||
|
||||
/// All `tool_timeout_for` cases live in one test (defaults plus env
|
||||
/// overrides) to avoid env-var races between parallel test threads.
|
||||
/// Issue #1125: operators on slow local inference (vLLM on old GPUs) need
|
||||
/// to disable or extend the inter-agent timeout via env var.
|
||||
#[test]
|
||||
fn test_tool_timeout_for_agent_tools() {
|
||||
assert_eq!(tool_timeout_for("agent_send"), Duration::from_secs(600));
|
||||
assert_eq!(tool_timeout_for("agent_spawn"), Duration::from_secs(600));
|
||||
assert_eq!(tool_timeout_for("file_read"), Duration::from_secs(120));
|
||||
assert_eq!(tool_timeout_for("shell_exec"), Duration::from_secs(120));
|
||||
// Baseline: no env overrides → compiled-in defaults.
|
||||
std::env::remove_var("OPENFANG_AGENT_TOOL_TIMEOUT_SECS");
|
||||
std::env::remove_var("OPENFANG_TOOL_TIMEOUT_SECS");
|
||||
assert_eq!(
|
||||
tool_timeout_for("agent_send"),
|
||||
Some(Duration::from_secs(600))
|
||||
);
|
||||
assert_eq!(
|
||||
tool_timeout_for("agent_spawn"),
|
||||
Some(Duration::from_secs(600))
|
||||
);
|
||||
assert_eq!(
|
||||
tool_timeout_for("file_read"),
|
||||
Some(Duration::from_secs(120))
|
||||
);
|
||||
assert_eq!(
|
||||
tool_timeout_for("shell_exec"),
|
||||
Some(Duration::from_secs(120))
|
||||
);
|
||||
|
||||
// Override: set to 0 → timeout disabled.
|
||||
std::env::set_var("OPENFANG_AGENT_TOOL_TIMEOUT_SECS", "0");
|
||||
std::env::set_var("OPENFANG_TOOL_TIMEOUT_SECS", "0");
|
||||
assert_eq!(tool_timeout_for("agent_send"), None);
|
||||
assert_eq!(tool_timeout_for("agent_spawn"), None);
|
||||
assert_eq!(tool_timeout_for("file_read"), None);
|
||||
|
||||
// Override: custom positive values are honored verbatim.
|
||||
std::env::set_var("OPENFANG_AGENT_TOOL_TIMEOUT_SECS", "1800");
|
||||
std::env::set_var("OPENFANG_TOOL_TIMEOUT_SECS", "300");
|
||||
assert_eq!(
|
||||
tool_timeout_for("agent_send"),
|
||||
Some(Duration::from_secs(1800))
|
||||
);
|
||||
assert_eq!(
|
||||
tool_timeout_for("file_read"),
|
||||
Some(Duration::from_secs(300))
|
||||
);
|
||||
|
||||
// Override: unparseable values fall back to compiled-in defaults.
|
||||
std::env::set_var("OPENFANG_AGENT_TOOL_TIMEOUT_SECS", "not-a-number");
|
||||
std::env::set_var("OPENFANG_TOOL_TIMEOUT_SECS", "");
|
||||
assert_eq!(
|
||||
tool_timeout_for("agent_send"),
|
||||
Some(Duration::from_secs(600))
|
||||
);
|
||||
assert_eq!(
|
||||
tool_timeout_for("file_read"),
|
||||
Some(Duration::from_secs(120))
|
||||
);
|
||||
|
||||
std::env::remove_var("OPENFANG_AGENT_TOOL_TIMEOUT_SECS");
|
||||
std::env::remove_var("OPENFANG_TOOL_TIMEOUT_SECS");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_max_history_messages() {
|
||||
assert_eq!(MAX_HISTORY_MESSAGES, 20);
|
||||
assert_eq!(
|
||||
openfang_types::agent::DEFAULT_MAX_HISTORY_MESSAGES,
|
||||
MAX_HISTORY_MESSAGES
|
||||
);
|
||||
}
|
||||
|
||||
/// Issue #871: an agent with a manifest override uses that value.
|
||||
#[test]
|
||||
fn test_effective_max_history_uses_manifest_override() {
|
||||
let mut manifest = openfang_types::agent::AgentManifest {
|
||||
max_history_messages: Some(40),
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(manifest.effective_max_history_messages(), 40);
|
||||
|
||||
manifest.max_history_messages = Some(6);
|
||||
assert_eq!(manifest.effective_max_history_messages(), 6);
|
||||
}
|
||||
|
||||
/// Issue #871: an agent without an override falls back to the runtime
|
||||
/// default. `Some(0)` is also treated as the default to avoid an agent
|
||||
/// accidentally disabling history entirely.
|
||||
#[test]
|
||||
fn test_effective_max_history_falls_back_to_default() {
|
||||
let mut manifest = openfang_types::agent::AgentManifest {
|
||||
max_history_messages: None,
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
manifest.effective_max_history_messages(),
|
||||
MAX_HISTORY_MESSAGES
|
||||
);
|
||||
|
||||
manifest.max_history_messages = Some(0);
|
||||
assert_eq!(
|
||||
manifest.effective_max_history_messages(),
|
||||
MAX_HISTORY_MESSAGES
|
||||
);
|
||||
}
|
||||
|
||||
/// Issue #871: `max_history_messages` round-trips through serde with
|
||||
/// `#[serde(default)]`, so manifests without the field still deserialize.
|
||||
#[test]
|
||||
fn test_manifest_max_history_round_trip_json() {
|
||||
let json_no_override = r#"{"name":"worker","module":"builtin:chat"}"#;
|
||||
let manifest: openfang_types::agent::AgentManifest =
|
||||
serde_json::from_str(json_no_override).unwrap();
|
||||
assert_eq!(manifest.max_history_messages, None);
|
||||
assert_eq!(
|
||||
manifest.effective_max_history_messages(),
|
||||
MAX_HISTORY_MESSAGES
|
||||
);
|
||||
|
||||
let json_with_override =
|
||||
r#"{"name":"orchestrator","module":"builtin:chat","max_history_messages":40}"#;
|
||||
let manifest: openfang_types::agent::AgentManifest =
|
||||
serde_json::from_str(json_with_override).unwrap();
|
||||
assert_eq!(manifest.max_history_messages, Some(40));
|
||||
assert_eq!(manifest.effective_max_history_messages(), 40);
|
||||
}
|
||||
|
||||
fn sample_image_block() -> ContentBlock {
|
||||
|
||||
@@ -404,6 +404,7 @@ fn build_conversation_text(messages: &[Message], config: &CompactionConfig) -> S
|
||||
conversation_text.push_str(&format!("[Image: {media_type}]\n\n"));
|
||||
}
|
||||
ContentBlock::Thinking { .. } => {}
|
||||
ContentBlock::RedactedThinking { .. } => {}
|
||||
ContentBlock::Unknown => {}
|
||||
}
|
||||
}
|
||||
@@ -457,6 +458,7 @@ async fn summarize_messages(
|
||||
text: summarize_prompt,
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
}],
|
||||
tools: vec![],
|
||||
max_tokens: config.max_summary_tokens,
|
||||
@@ -575,6 +577,7 @@ async fn summarize_in_chunks(
|
||||
text: merge_prompt,
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
}],
|
||||
tools: vec![],
|
||||
max_tokens: config.max_summary_tokens,
|
||||
@@ -912,6 +915,7 @@ mod tests {
|
||||
input: serde_json::json!({"query": "test"}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
};
|
||||
messages[2] = Message {
|
||||
role: Role::User,
|
||||
@@ -921,6 +925,7 @@ mod tests {
|
||||
content: "Search results here".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let session = Session {
|
||||
@@ -1251,6 +1256,7 @@ mod tests {
|
||||
provider_metadata: None,
|
||||
},
|
||||
]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -1260,6 +1266,7 @@ mod tests {
|
||||
content: "Results found".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -1267,6 +1274,7 @@ mod tests {
|
||||
media_type: "image/png".to_string(),
|
||||
data: "base64data".to_string(),
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
|
||||
@@ -1401,6 +1409,7 @@ mod tests {
|
||||
content: tool_content,
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
}];
|
||||
let text = build_conversation_text(&messages, &config);
|
||||
// The base64 blob should be stripped/replaced by session_repair
|
||||
@@ -1421,6 +1430,7 @@ mod tests {
|
||||
content: large_result,
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
}];
|
||||
let text = build_conversation_text(&messages, &config);
|
||||
// Should be capped at ~2000 chars (plus the "..." suffix)
|
||||
@@ -1445,6 +1455,7 @@ mod tests {
|
||||
content: short_result.to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
}];
|
||||
let text = build_conversation_text(&messages, &config);
|
||||
assert!(text.contains(short_result));
|
||||
@@ -1464,6 +1475,7 @@ mod tests {
|
||||
input: serde_json::json!({}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -1473,6 +1485,7 @@ mod tests {
|
||||
content: "file contents".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Done reading."),
|
||||
];
|
||||
|
||||
@@ -290,6 +290,7 @@ mod tests {
|
||||
content: big_result.clone(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: openfang_types::message::Role::User,
|
||||
@@ -299,6 +300,7 @@ mod tests {
|
||||
content: big_result,
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
|
||||
@@ -350,6 +352,7 @@ mod tests {
|
||||
content: big_chinese,
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
}];
|
||||
// Must not panic on multi-byte content
|
||||
let compacted = apply_context_guard(&mut messages, &budget, &[]);
|
||||
|
||||
@@ -237,6 +237,7 @@ mod tests {
|
||||
Role::Assistant
|
||||
},
|
||||
content: MessageContent::Text(text),
|
||||
..Default::default()
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
@@ -295,6 +296,7 @@ mod tests {
|
||||
content: big_result.clone(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -304,6 +306,7 @@ mod tests {
|
||||
content: big_result,
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
// Tiny context window to force all stages
|
||||
@@ -342,6 +345,7 @@ mod tests {
|
||||
content: chinese_result,
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
// Tiny context window to force stage 3 tool truncation
|
||||
@@ -365,6 +369,7 @@ mod tests {
|
||||
input: serde_json::json!({}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -374,6 +379,7 @@ mod tests {
|
||||
content: "file contents".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::user("thanks"),
|
||||
];
|
||||
|
||||
@@ -84,6 +84,22 @@ enum ApiContentBlock {
|
||||
#[serde(skip_serializing_if = "std::ops::Not::not")]
|
||||
is_error: bool,
|
||||
},
|
||||
/// Extended-thinking block echoed back to the API.
|
||||
///
|
||||
/// Anthropic requires the original `signature` to be returned verbatim
|
||||
/// alongside the `thinking` text on subsequent turns; otherwise the
|
||||
/// model loses its prior reasoning state. Without `signature` the API
|
||||
/// rejects the block, so we omit thinking blocks that arrive without
|
||||
/// one (e.g. legacy sessions saved before this field was tracked).
|
||||
#[serde(rename = "thinking")]
|
||||
Thinking { thinking: String, signature: String },
|
||||
/// Redacted (encrypted) thinking block echoed back to the API.
|
||||
///
|
||||
/// Anthropic returns these when the model decides to hide reasoning;
|
||||
/// the `data` blob is opaque and MUST be echoed verbatim on the next
|
||||
/// turn or the API rejects the resubmitted history.
|
||||
#[serde(rename = "redacted_thinking")]
|
||||
RedactedThinking { data: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -120,8 +136,21 @@ enum ResponseContentBlock {
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
/// Extended-thinking block from Anthropic. The `signature` is opaque
|
||||
/// to us but MUST be persisted and echoed back on the next request,
|
||||
/// otherwise the API rejects the resubmitted thinking block and the
|
||||
/// model loses its reasoning state.
|
||||
#[serde(rename = "thinking")]
|
||||
Thinking { thinking: String },
|
||||
Thinking {
|
||||
thinking: String,
|
||||
#[serde(default)]
|
||||
signature: Option<String>,
|
||||
},
|
||||
/// Redacted (encrypted) thinking block. The `data` blob is opaque to
|
||||
/// us and must be persisted as-is so we can echo it back on the next
|
||||
/// request — Anthropic rejects history that strips these blocks.
|
||||
#[serde(rename = "redacted_thinking")]
|
||||
RedactedThinking { data: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -144,12 +173,25 @@ struct ApiErrorDetail {
|
||||
/// Accumulator for content blocks during streaming.
|
||||
enum ContentBlockAccum {
|
||||
Text(String),
|
||||
Thinking(String),
|
||||
/// Extended thinking — text plus an opaque signature delivered as
|
||||
/// `signature_delta` events (or as a single field on `content_block_stop`
|
||||
/// for older API versions). The signature is required to round-trip
|
||||
/// thinking blocks on subsequent turns.
|
||||
Thinking {
|
||||
thinking: String,
|
||||
signature: String,
|
||||
},
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input_json: String,
|
||||
},
|
||||
/// Redacted (encrypted) thinking block streamed from Anthropic.
|
||||
/// The opaque `data` blob arrives on `content_block_start` and must be
|
||||
/// persisted so the next turn can echo it back verbatim.
|
||||
RedactedThinking {
|
||||
data: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -412,7 +454,24 @@ impl LlmDriver for AnthropicDriver {
|
||||
});
|
||||
}
|
||||
"thinking" => {
|
||||
blocks.push(ContentBlockAccum::Thinking(String::new()));
|
||||
// Some API versions ship the signature on
|
||||
// content_block_start instead of as a delta.
|
||||
let initial_sig =
|
||||
block["signature"].as_str().unwrap_or("").to_string();
|
||||
blocks.push(ContentBlockAccum::Thinking {
|
||||
thinking: String::new(),
|
||||
signature: initial_sig,
|
||||
});
|
||||
}
|
||||
"redacted_thinking" => {
|
||||
// Anthropic delivers redacted_thinking
|
||||
// as a single block_start with the opaque
|
||||
// `data` blob (no delta events). Store it
|
||||
// verbatim so we can echo it back on the
|
||||
// next request — API rejects history
|
||||
// that strips redacted_thinking blocks.
|
||||
let data = block["data"].as_str().unwrap_or("").to_string();
|
||||
blocks.push(ContentBlockAccum::RedactedThinking { data });
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
@@ -452,11 +511,33 @@ impl LlmDriver for AnthropicDriver {
|
||||
}
|
||||
}
|
||||
"thinking_delta" => {
|
||||
if let Some(thinking) = delta["thinking"].as_str() {
|
||||
if let Some(ContentBlockAccum::Thinking(ref mut t)) =
|
||||
blocks.get_mut(block_idx)
|
||||
if let Some(t) = delta["thinking"].as_str() {
|
||||
if let Some(ContentBlockAccum::Thinking {
|
||||
thinking: ref mut buf,
|
||||
..
|
||||
}) = blocks.get_mut(block_idx)
|
||||
{
|
||||
t.push_str(thinking);
|
||||
buf.push_str(t);
|
||||
}
|
||||
// Forward to UI as ThinkingDelta event so dashboards can show reasoning.
|
||||
let _ = tx
|
||||
.send(StreamEvent::ThinkingDelta {
|
||||
text: t.to_string(),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
}
|
||||
"signature_delta" => {
|
||||
// Anthropic streams the thinking signature
|
||||
// as its own delta type; concatenate any
|
||||
// partial pieces into the accumulator.
|
||||
if let Some(sig) = delta["signature"].as_str() {
|
||||
if let Some(ContentBlockAccum::Thinking {
|
||||
ref mut signature,
|
||||
..
|
||||
}) = blocks.get_mut(block_idx)
|
||||
{
|
||||
signature.push_str(sig);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -512,8 +593,26 @@ impl LlmDriver for AnthropicDriver {
|
||||
provider_metadata: None,
|
||||
});
|
||||
}
|
||||
ContentBlockAccum::Thinking(thinking) => {
|
||||
content.push(ContentBlock::Thinking { thinking });
|
||||
ContentBlockAccum::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
} => {
|
||||
// Drop empty thinking blocks (rare, but happens if the
|
||||
// stream is interrupted mid-block). Always keep the
|
||||
// signature when present — it's required to round-trip.
|
||||
if !thinking.is_empty() || !signature.is_empty() {
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking,
|
||||
signature: if signature.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(signature)
|
||||
},
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "anthropic_extended_thinking"
|
||||
})),
|
||||
});
|
||||
}
|
||||
}
|
||||
ContentBlockAccum::ToolUse {
|
||||
id,
|
||||
@@ -530,6 +629,11 @@ impl LlmDriver for AnthropicDriver {
|
||||
});
|
||||
tool_calls.push(ToolCall { id, name, input });
|
||||
}
|
||||
ContentBlockAccum::RedactedThinking { data } => {
|
||||
if !data.is_empty() {
|
||||
content.push(ContentBlock::RedactedThinking { data });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -615,7 +719,38 @@ fn convert_message(msg: &Message) -> ApiMessage {
|
||||
content: content.clone(),
|
||||
is_error: *is_error,
|
||||
}),
|
||||
ContentBlock::Thinking { .. } => None,
|
||||
ContentBlock::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
..
|
||||
} => {
|
||||
// Anthropic's extended-thinking spec requires the
|
||||
// verbatim `signature` to accompany any thinking block
|
||||
// resubmitted in conversation history. Without one,
|
||||
// the API rejects the request, so we silently drop
|
||||
// legacy thinking blocks (saved before signature
|
||||
// tracking) instead of round-tripping them.
|
||||
signature.as_ref().and_then(|sig| {
|
||||
if sig.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(ApiContentBlock::Thinking {
|
||||
thinking: thinking.clone(),
|
||||
signature: sig.clone(),
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
ContentBlock::RedactedThinking { data } => {
|
||||
// Echo the encrypted blob verbatim. Anthropic
|
||||
// rejects history that drops redacted_thinking
|
||||
// blocks, so always include them on resubmission.
|
||||
if data.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(ApiContentBlock::RedactedThinking { data: data.clone() })
|
||||
}
|
||||
}
|
||||
ContentBlock::Unknown => None,
|
||||
})
|
||||
.collect();
|
||||
@@ -651,8 +786,20 @@ fn convert_response(api: ApiResponse) -> CompletionResponse {
|
||||
});
|
||||
tool_calls.push(ToolCall { id, name, input });
|
||||
}
|
||||
ResponseContentBlock::Thinking { thinking } => {
|
||||
content.push(ContentBlock::Thinking { thinking });
|
||||
ResponseContentBlock::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
} => {
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "anthropic_extended_thinking"
|
||||
})),
|
||||
});
|
||||
}
|
||||
ResponseContentBlock::RedactedThinking { data } => {
|
||||
content.push(ContentBlock::RedactedThinking { data });
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -758,6 +905,7 @@ mod tests {
|
||||
input: serde_json::Value::String(r#"{"query": "test"}"#.to_string()),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
};
|
||||
let api_msg = convert_message(&msg);
|
||||
if let ApiContent::Blocks(blocks) = api_msg.content {
|
||||
@@ -772,4 +920,267 @@ mod tests {
|
||||
panic!("Expected Blocks content");
|
||||
}
|
||||
}
|
||||
|
||||
/// Issue #1098: Anthropic extended-thinking blocks must round-trip
|
||||
/// through the driver — the inbound response carries a `signature` that
|
||||
/// MUST be echoed verbatim on the next request, otherwise the API
|
||||
/// rejects the resubmitted thinking block and the model loses prior
|
||||
/// reasoning state.
|
||||
#[test]
|
||||
fn test_thinking_block_signature_round_trip() {
|
||||
// Step 1: API delivers a thinking block with signature
|
||||
let api_response = ApiResponse {
|
||||
content: vec![
|
||||
ResponseContentBlock::Thinking {
|
||||
thinking: "Let me carefully consider this problem...".to_string(),
|
||||
signature: Some("WaUjzkypQ2mUEVM36O2TxuC".to_string()),
|
||||
},
|
||||
ResponseContentBlock::Text {
|
||||
text: "The answer is 42.".to_string(),
|
||||
},
|
||||
],
|
||||
stop_reason: "end_turn".to_string(),
|
||||
usage: ApiUsage {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
},
|
||||
};
|
||||
let response = convert_response(api_response);
|
||||
assert_eq!(response.content.len(), 2);
|
||||
|
||||
// Step 2: Verify the signature reached the ContentBlock
|
||||
let thinking_block = &response.content[0];
|
||||
match thinking_block {
|
||||
ContentBlock::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(thinking, "Let me carefully consider this problem...");
|
||||
assert_eq!(signature.as_deref(), Some("WaUjzkypQ2mUEVM36O2TxuC"));
|
||||
}
|
||||
_ => panic!("expected Thinking content block"),
|
||||
}
|
||||
|
||||
// Step 3: Now feed the assistant turn back into the driver as if
|
||||
// it were prior conversation history (next user turn). The signature
|
||||
// must survive into the outbound API request.
|
||||
let assistant_msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(response.content.clone()),
|
||||
..Default::default()
|
||||
};
|
||||
let api_msg = convert_message(&assistant_msg);
|
||||
let blocks = match api_msg.content {
|
||||
ApiContent::Blocks(b) => b,
|
||||
_ => panic!("expected Blocks content"),
|
||||
};
|
||||
|
||||
// The Thinking block must appear in the outbound payload with its signature.
|
||||
let mut found_thinking = false;
|
||||
for block in &blocks {
|
||||
if let ApiContentBlock::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
} = block
|
||||
{
|
||||
assert_eq!(thinking, "Let me carefully consider this problem...");
|
||||
assert_eq!(signature, "WaUjzkypQ2mUEVM36O2TxuC");
|
||||
found_thinking = true;
|
||||
}
|
||||
}
|
||||
assert!(
|
||||
found_thinking,
|
||||
"outbound API request must include the thinking block with signature"
|
||||
);
|
||||
|
||||
// Step 4: Verify on-the-wire JSON shape (`type=thinking`, `signature` present).
|
||||
let outbound_json = serde_json::to_value(&blocks).unwrap();
|
||||
let arr = outbound_json.as_array().unwrap();
|
||||
let thinking_json = arr
|
||||
.iter()
|
||||
.find(|v| v["type"] == "thinking")
|
||||
.expect("thinking block in JSON");
|
||||
assert_eq!(thinking_json["signature"], "WaUjzkypQ2mUEVM36O2TxuC");
|
||||
assert_eq!(
|
||||
thinking_json["thinking"],
|
||||
"Let me carefully consider this problem..."
|
||||
);
|
||||
}
|
||||
|
||||
/// Legacy thinking blocks saved before signature tracking should NOT
|
||||
/// be replayed — Anthropic rejects thinking blocks without signatures.
|
||||
#[test]
|
||||
fn test_thinking_block_without_signature_dropped_outbound() {
|
||||
let assistant_msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(vec![
|
||||
ContentBlock::Thinking {
|
||||
thinking: "old reasoning from before sig tracking".to_string(),
|
||||
signature: None,
|
||||
provider_metadata: None,
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "Hello.".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
]),
|
||||
..Default::default()
|
||||
};
|
||||
let api_msg = convert_message(&assistant_msg);
|
||||
let blocks = match api_msg.content {
|
||||
ApiContent::Blocks(b) => b,
|
||||
_ => panic!("expected Blocks content"),
|
||||
};
|
||||
// The legacy thinking block must be dropped (no sig = API would 400).
|
||||
for block in &blocks {
|
||||
assert!(
|
||||
!matches!(block, ApiContentBlock::Thinking { .. }),
|
||||
"thinking block without signature must be dropped"
|
||||
);
|
||||
}
|
||||
// The text part is still preserved.
|
||||
assert!(blocks
|
||||
.iter()
|
||||
.any(|b| matches!(b, ApiContentBlock::Text { .. })));
|
||||
}
|
||||
|
||||
/// Streaming path: signature_delta events accumulate into the final block.
|
||||
#[test]
|
||||
fn test_thinking_block_serde_with_signature_field() {
|
||||
// Verify the API response wire format is parsed correctly.
|
||||
let json = serde_json::json!({
|
||||
"type": "thinking",
|
||||
"thinking": "step 1, step 2",
|
||||
"signature": "abc123"
|
||||
});
|
||||
let block: ResponseContentBlock = serde_json::from_value(json).unwrap();
|
||||
match block {
|
||||
ResponseContentBlock::Thinking {
|
||||
thinking,
|
||||
signature,
|
||||
} => {
|
||||
assert_eq!(thinking, "step 1, step 2");
|
||||
assert_eq!(signature.as_deref(), Some("abc123"));
|
||||
}
|
||||
_ => panic!("expected Thinking response block"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Issue #1148 — Anthropic `redacted_thinking` blocks must survive the
|
||||
/// full driver round-trip. The opaque `data` blob is required verbatim
|
||||
/// on resubmission; dropping or mutating it causes the API to reject
|
||||
/// the assistant turn on the next request.
|
||||
#[test]
|
||||
fn test_redacted_thinking_round_trip() {
|
||||
// Step 1: API delivers a response with a redacted_thinking block.
|
||||
let api_response = ApiResponse {
|
||||
content: vec![
|
||||
ResponseContentBlock::RedactedThinking {
|
||||
data: "EncRyPt3D_BLO8".to_string(),
|
||||
},
|
||||
ResponseContentBlock::Text {
|
||||
text: "The answer is 42.".to_string(),
|
||||
},
|
||||
],
|
||||
stop_reason: "end_turn".to_string(),
|
||||
usage: ApiUsage {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
},
|
||||
};
|
||||
let response = convert_response(api_response);
|
||||
assert_eq!(response.content.len(), 2);
|
||||
|
||||
// Step 2: The opaque blob must reach the ContentBlock layer.
|
||||
match &response.content[0] {
|
||||
ContentBlock::RedactedThinking { data } => {
|
||||
assert_eq!(data, "EncRyPt3D_BLO8");
|
||||
}
|
||||
other => panic!("expected RedactedThinking content block, got {other:?}"),
|
||||
}
|
||||
|
||||
// Step 3: Resubmit the assistant turn as conversation history.
|
||||
let assistant_msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(response.content.clone()),
|
||||
..Default::default()
|
||||
};
|
||||
let api_msg = convert_message(&assistant_msg);
|
||||
let blocks = match api_msg.content {
|
||||
ApiContent::Blocks(b) => b,
|
||||
_ => panic!("expected Blocks content"),
|
||||
};
|
||||
|
||||
// The redacted_thinking block must appear in the outbound payload.
|
||||
let mut found_redacted = false;
|
||||
for block in &blocks {
|
||||
if let ApiContentBlock::RedactedThinking { data } = block {
|
||||
assert_eq!(data, "EncRyPt3D_BLO8");
|
||||
found_redacted = true;
|
||||
}
|
||||
}
|
||||
assert!(
|
||||
found_redacted,
|
||||
"outbound API request must include the redacted_thinking block"
|
||||
);
|
||||
|
||||
// Step 4: On-the-wire JSON shape (`type=redacted_thinking`, `data` present).
|
||||
let outbound_json = serde_json::to_value(&blocks).unwrap();
|
||||
let arr = outbound_json.as_array().unwrap();
|
||||
let redacted_json = arr
|
||||
.iter()
|
||||
.find(|v| v["type"] == "redacted_thinking")
|
||||
.expect("redacted_thinking block in JSON");
|
||||
assert_eq!(redacted_json["data"], "EncRyPt3D_BLO8");
|
||||
}
|
||||
|
||||
/// API response wire format for `redacted_thinking` is parsed correctly.
|
||||
#[test]
|
||||
fn test_redacted_thinking_serde() {
|
||||
let json = serde_json::json!({
|
||||
"type": "redacted_thinking",
|
||||
"data": "opaque_blob_xyz"
|
||||
});
|
||||
let block: ResponseContentBlock = serde_json::from_value(json).unwrap();
|
||||
match block {
|
||||
ResponseContentBlock::RedactedThinking { data } => {
|
||||
assert_eq!(data, "opaque_blob_xyz");
|
||||
}
|
||||
_ => panic!("expected RedactedThinking response block"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Empty redacted_thinking blocks (e.g. interrupted stream) must be
|
||||
/// dropped on outbound to avoid sending malformed history.
|
||||
#[test]
|
||||
fn test_redacted_thinking_empty_dropped_outbound() {
|
||||
let assistant_msg = Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(vec![
|
||||
ContentBlock::RedactedThinking {
|
||||
data: String::new(),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "Hello.".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
]),
|
||||
..Default::default()
|
||||
};
|
||||
let api_msg = convert_message(&assistant_msg);
|
||||
let blocks = match api_msg.content {
|
||||
ApiContent::Blocks(b) => b,
|
||||
_ => panic!("expected Blocks content"),
|
||||
};
|
||||
for block in &blocks {
|
||||
assert!(
|
||||
!matches!(block, ApiContentBlock::RedactedThinking { .. }),
|
||||
"empty redacted_thinking block must be dropped"
|
||||
);
|
||||
}
|
||||
assert!(blocks
|
||||
.iter()
|
||||
.any(|b| matches!(b, ApiContentBlock::Text { .. })));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -92,6 +92,20 @@ enum BedrockContentBlock {
|
||||
#[serde(rename = "toolResult")]
|
||||
tool_result: BedrockToolResult,
|
||||
},
|
||||
// Bedrock Converse representation of Anthropic's `redacted_thinking`.
|
||||
// The encrypted blob is echoed back verbatim under
|
||||
// reasoningContent.redactedContent so Claude extended-thinking history
|
||||
// is not rejected on resubmission.
|
||||
ReasoningContent {
|
||||
#[serde(rename = "reasoningContent")]
|
||||
reasoning_content: BedrockReasoningContent,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct BedrockReasoningContent {
|
||||
#[serde(rename = "redactedContent")]
|
||||
redacted_content: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -299,6 +313,20 @@ fn convert_content_block(block: &ContentBlock) -> Option<BedrockContentBlock> {
|
||||
},
|
||||
},
|
||||
}),
|
||||
// Echo redacted_thinking verbatim. Bedrock Converse rejects history
|
||||
// that drops these blocks on Claude extended-thinking models, mirroring
|
||||
// the anthropic.rs path. Drop empty blobs (e.g. interrupted stream).
|
||||
ContentBlock::RedactedThinking { data } => {
|
||||
if data.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(BedrockContentBlock::ReasoningContent {
|
||||
reasoning_content: BedrockReasoningContent {
|
||||
redacted_content: data.clone(),
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
// Image, Thinking, and Unknown are not supported — silently drop
|
||||
ContentBlock::Image { .. } | ContentBlock::Thinking { .. } | ContentBlock::Unknown => None,
|
||||
}
|
||||
@@ -782,6 +810,7 @@ mod tests {
|
||||
let messages = vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Text("Hello".to_string()),
|
||||
..Default::default()
|
||||
}];
|
||||
let (bedrock_msgs, system) = convert_messages(&messages, &None);
|
||||
assert_eq!(bedrock_msgs.len(), 1);
|
||||
@@ -794,6 +823,7 @@ mod tests {
|
||||
let messages = vec![Message {
|
||||
role: Role::System,
|
||||
content: MessageContent::Text("Be helpful".to_string()),
|
||||
..Default::default()
|
||||
}];
|
||||
let (bedrock_msgs, system) = convert_messages(&messages, &None);
|
||||
assert!(bedrock_msgs.is_empty());
|
||||
@@ -806,6 +836,7 @@ mod tests {
|
||||
let messages = vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Text("Hi".to_string()),
|
||||
..Default::default()
|
||||
}];
|
||||
let (_, system) = convert_messages(&messages, &Some("You are an AI".to_string()));
|
||||
assert!(system.is_some());
|
||||
@@ -1120,6 +1151,53 @@ mod tests {
|
||||
assert!(text_at_3 >= 1);
|
||||
}
|
||||
|
||||
/// Issue #1187 — Bedrock Converse history must preserve
|
||||
/// `redacted_thinking` blocks on Claude extended-thinking models.
|
||||
/// A message containing only RedactedThinking must round-trip through
|
||||
/// `convert_content_block` without being silently dropped, and the wire
|
||||
/// format must use `reasoningContent.redactedContent`.
|
||||
#[test]
|
||||
fn test_bedrock_redacted_thinking_round_trip() {
|
||||
let msg = Message::assistant_with_blocks(vec![ContentBlock::RedactedThinking {
|
||||
data: "encrypted-blob-abc123".to_string(),
|
||||
}]);
|
||||
|
||||
let bedrock_blocks = convert_message_content(&msg.content);
|
||||
assert_eq!(
|
||||
bedrock_blocks.len(),
|
||||
1,
|
||||
"RedactedThinking must survive convert_content_block"
|
||||
);
|
||||
match &bedrock_blocks[0] {
|
||||
BedrockContentBlock::ReasoningContent { reasoning_content } => {
|
||||
assert_eq!(reasoning_content.redacted_content, "encrypted-blob-abc123");
|
||||
}
|
||||
other => panic!("expected ReasoningContent block, got {other:?}"),
|
||||
}
|
||||
|
||||
// Wire format check: serialized JSON must carry
|
||||
// reasoningContent.redactedContent so Bedrock accepts the history.
|
||||
let json = serde_json::to_value(&bedrock_blocks[0]).unwrap();
|
||||
assert_eq!(
|
||||
json["reasoningContent"]["redactedContent"],
|
||||
"encrypted-blob-abc123"
|
||||
);
|
||||
}
|
||||
|
||||
/// Empty RedactedThinking blobs (interrupted stream) must be dropped on
|
||||
/// outbound, matching the anthropic.rs behavior.
|
||||
#[test]
|
||||
fn test_bedrock_redacted_thinking_empty_dropped() {
|
||||
let msg = Message::assistant_with_blocks(vec![ContentBlock::RedactedThinking {
|
||||
data: String::new(),
|
||||
}]);
|
||||
let bedrock_blocks = convert_message_content(&msg.content);
|
||||
assert!(
|
||||
bedrock_blocks.is_empty(),
|
||||
"empty redacted_thinking must be dropped"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_tool_pairing_noop_on_correct() {
|
||||
// already correct 2-for-2 → no change
|
||||
|
||||
@@ -11,7 +11,7 @@
|
||||
use crate::llm_driver::{CompletionRequest, CompletionResponse, LlmDriver, LlmError, StreamEvent};
|
||||
use async_trait::async_trait;
|
||||
use dashmap::DashMap;
|
||||
use openfang_types::message::{ContentBlock, Role, StopReason, TokenUsage};
|
||||
use openfang_types::message::{ContentBlock, MessageContent, Role, StopReason, TokenUsage};
|
||||
use serde::Deserialize;
|
||||
use std::sync::Arc;
|
||||
use tokio::io::{AsyncBufReadExt, AsyncReadExt};
|
||||
@@ -130,6 +130,14 @@ impl ClaudeCodeDriver {
|
||||
}
|
||||
|
||||
/// Build a text prompt from the completion request messages.
|
||||
///
|
||||
/// The Claude Code CLI is text-only (`-p <prompt>`), so non-text content
|
||||
/// blocks (images, etc.) cannot be sent natively. Rather than dropping
|
||||
/// them silently — which causes the model to hallucinate about content
|
||||
/// it can't see — we render each non-text block as a synthetic
|
||||
/// `[attachment: ...]` marker. The model still can't *view* the
|
||||
/// attachment, but it knows the attachment exists and can acknowledge
|
||||
/// it coherently instead of confabulating.
|
||||
fn build_prompt(request: &CompletionRequest) -> String {
|
||||
let mut parts = Vec::new();
|
||||
|
||||
@@ -139,15 +147,53 @@ impl ClaudeCodeDriver {
|
||||
Role::Assistant => "Assistant",
|
||||
Role::System => "System",
|
||||
};
|
||||
let text = msg.content.text_content();
|
||||
if !text.is_empty() {
|
||||
parts.push(format!("[{role_label}]\n{text}"));
|
||||
let rendered = Self::render_content(&msg.content);
|
||||
if !rendered.is_empty() {
|
||||
parts.push(format!("[{role_label}]\n{rendered}"));
|
||||
}
|
||||
}
|
||||
|
||||
parts.join("\n\n")
|
||||
}
|
||||
|
||||
/// Render message content for the text-only CLI prompt.
|
||||
///
|
||||
/// Text blocks pass through verbatim. Image blocks are rendered as
|
||||
/// `[attachment: <media_type> image, ~N KB — not viewable on this
|
||||
/// provider]` so the model receives a positive signal that an
|
||||
/// attachment arrived. ToolUse/ToolResult/Thinking are omitted —
|
||||
/// the CLI manages its own tool loop.
|
||||
fn render_content(content: &MessageContent) -> String {
|
||||
match content {
|
||||
MessageContent::Text(s) => s.clone(),
|
||||
MessageContent::Blocks(blocks) => blocks
|
||||
.iter()
|
||||
.filter_map(|b| match b {
|
||||
ContentBlock::Text { text, .. } => {
|
||||
if text.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(text.clone())
|
||||
}
|
||||
}
|
||||
ContentBlock::Image { media_type, data } => {
|
||||
// base64 → ~3/4 the length in decoded bytes.
|
||||
let approx_kb = (data.len().saturating_mul(3) / 4) / 1024;
|
||||
Some(format!(
|
||||
"[attachment: {media_type} image, ~{approx_kb} KB — not viewable on this provider]"
|
||||
))
|
||||
}
|
||||
ContentBlock::ToolUse { .. }
|
||||
| ContentBlock::ToolResult { .. }
|
||||
| ContentBlock::Thinking { .. }
|
||||
| ContentBlock::RedactedThinking { .. }
|
||||
| ContentBlock::Unknown => None,
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Map a model ID like "claude-code/opus" to CLI --model flag value.
|
||||
fn model_flag(model: &str) -> Option<String> {
|
||||
let stripped = model.strip_prefix("claude-code/").unwrap_or(model);
|
||||
@@ -711,6 +757,7 @@ mod tests {
|
||||
messages: vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::text("Hello"),
|
||||
..Default::default()
|
||||
}],
|
||||
tools: vec![],
|
||||
max_tokens: 1024,
|
||||
@@ -726,6 +773,79 @@ mod tests {
|
||||
assert!(prompt.contains("Hello"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_prompt_renders_image_attachment_marker() {
|
||||
use openfang_types::message::{ContentBlock, Message, MessageContent};
|
||||
|
||||
// ~12 KB of base64 — decoded ~9 KB.
|
||||
let fake_b64 = "A".repeat(12 * 1024);
|
||||
let request = CompletionRequest {
|
||||
model: "claude-code/sonnet".to_string(),
|
||||
messages: vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Blocks(vec![
|
||||
ContentBlock::Text {
|
||||
text: "what's in this?".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
ContentBlock::Image {
|
||||
media_type: "image/png".to_string(),
|
||||
data: fake_b64,
|
||||
},
|
||||
]),
|
||||
..Default::default()
|
||||
}],
|
||||
tools: vec![],
|
||||
max_tokens: 1024,
|
||||
temperature: 0.7,
|
||||
system: None,
|
||||
thinking: None,
|
||||
};
|
||||
|
||||
let prompt = ClaudeCodeDriver::build_prompt(&request);
|
||||
assert!(prompt.contains("what's in this?"), "text preserved");
|
||||
assert!(
|
||||
prompt.contains("[attachment: image/png image"),
|
||||
"image rendered as synthetic attachment marker, got: {prompt}"
|
||||
);
|
||||
assert!(
|
||||
prompt.contains("not viewable on this provider"),
|
||||
"marker explains the limitation, got: {prompt}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_prompt_image_only_still_emits_marker() {
|
||||
use openfang_types::message::{ContentBlock, Message, MessageContent};
|
||||
|
||||
let request = CompletionRequest {
|
||||
model: "claude-code/sonnet".to_string(),
|
||||
messages: vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Blocks(vec![ContentBlock::Image {
|
||||
media_type: "image/jpeg".to_string(),
|
||||
data: "Zm9v".to_string(),
|
||||
}]),
|
||||
..Default::default()
|
||||
}],
|
||||
tools: vec![],
|
||||
max_tokens: 1024,
|
||||
temperature: 0.7,
|
||||
system: None,
|
||||
thinking: None,
|
||||
};
|
||||
|
||||
let prompt = ClaudeCodeDriver::build_prompt(&request);
|
||||
assert!(
|
||||
prompt.contains("[User]"),
|
||||
"user role label emitted even with image-only content, got: {prompt}"
|
||||
);
|
||||
assert!(
|
||||
prompt.contains("[attachment: image/jpeg image"),
|
||||
"bare image renders marker, got: {prompt}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_model_flag_mapping() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -321,7 +321,39 @@ fn convert_messages(
|
||||
},
|
||||
});
|
||||
}
|
||||
ContentBlock::Thinking { .. } => {}
|
||||
ContentBlock::Thinking {
|
||||
thinking,
|
||||
provider_metadata,
|
||||
..
|
||||
} => {
|
||||
// Issue #1098: preserve Gemini 2.5+ thought parts
|
||||
// when the upstream model originally emitted them.
|
||||
// Most Gemini state actually rides on the
|
||||
// thoughtSignature attached to text/tool_use
|
||||
// parts above, but we round-trip the visible
|
||||
// thinking text + sig as a `Thought` part too
|
||||
// so the model's internal state is fully
|
||||
// preserved. Other providers' thinking blocks
|
||||
// are dropped here (they have their own
|
||||
// outbound paths in the OpenAI/Anthropic
|
||||
// drivers).
|
||||
let format = provider_metadata
|
||||
.as_ref()
|
||||
.and_then(|m| m.get("format"))
|
||||
.and_then(|v| v.as_str());
|
||||
if format == Some("gemini_thought") && !thinking.is_empty() {
|
||||
let sig = provider_metadata
|
||||
.as_ref()
|
||||
.and_then(|m| m.get("thought_signature"))
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
parts.push(GeminiPart::Thought {
|
||||
text: thinking.clone(),
|
||||
thought: true,
|
||||
thought_signature: sig,
|
||||
});
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
@@ -561,12 +593,29 @@ fn convert_response(resp: GeminiResponse) -> Result<CompletionResponse, LlmError
|
||||
input: function_call.args,
|
||||
});
|
||||
}
|
||||
GeminiPart::Thought { text, .. } => {
|
||||
GeminiPart::Thought {
|
||||
text,
|
||||
thought_signature,
|
||||
..
|
||||
} => {
|
||||
// Gemini 2.5+ thinking parts — internal reasoning.
|
||||
// Store as Thinking content block so the UI can
|
||||
// optionally display it (like <think> blocks).
|
||||
// optionally display it. Issue #1098: preserve the
|
||||
// part-level `thoughtSignature` in `provider_metadata`
|
||||
// (and on subsequent text/tool_use parts) so the
|
||||
// model retains state across turns.
|
||||
if !text.is_empty() {
|
||||
content.push(ContentBlock::Thinking { thinking: text });
|
||||
let provider_metadata = thought_signature.map(|sig| {
|
||||
serde_json::json!({
|
||||
"format": "gemini_thought",
|
||||
"thought_signature": sig,
|
||||
})
|
||||
});
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking: text,
|
||||
signature: None,
|
||||
provider_metadata,
|
||||
});
|
||||
}
|
||||
}
|
||||
GeminiPart::InlineData { .. } | GeminiPart::FunctionResponse { .. } => {
|
||||
@@ -788,6 +837,10 @@ impl LlmDriver for GeminiDriver {
|
||||
let mut text_content = String::new();
|
||||
// Thought signature for accumulated text content (last one wins)
|
||||
let mut text_thought_sig: Option<String> = None;
|
||||
// Accumulated thought (Gemini 2.5+) text + signature, for
|
||||
// round-tripping reasoning state across turns (issue #1098).
|
||||
let mut thought_text = String::new();
|
||||
let mut thought_sig: Option<String> = None;
|
||||
// Track function calls: (name, args_json, thought_signature)
|
||||
let mut fn_calls: Vec<(String, serde_json::Value, Option<String>)> = Vec::new();
|
||||
let mut finish_reason: Option<String> = None;
|
||||
@@ -894,17 +947,26 @@ impl LlmDriver for GeminiDriver {
|
||||
thought_signature.clone(),
|
||||
));
|
||||
}
|
||||
GeminiPart::Thought { ref text, .. } => {
|
||||
GeminiPart::Thought {
|
||||
ref text,
|
||||
ref thought_signature,
|
||||
..
|
||||
} => {
|
||||
// Gemini 2.5+ thinking chunk — emit as
|
||||
// thinking delta so UIs can optionally
|
||||
// show it; do NOT mix into text_content.
|
||||
// show it; accumulate the text + sig
|
||||
// for later persistence (issue #1098).
|
||||
if !text.is_empty() {
|
||||
thought_text.push_str(text);
|
||||
let _ = tx
|
||||
.send(StreamEvent::ThinkingDelta {
|
||||
text: text.clone(),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
if thought_signature.is_some() {
|
||||
thought_sig = thought_signature.clone();
|
||||
}
|
||||
}
|
||||
GeminiPart::InlineData { .. }
|
||||
| GeminiPart::FunctionResponse { .. } => {}
|
||||
@@ -985,14 +1047,24 @@ impl LlmDriver for GeminiDriver {
|
||||
thought_signature.clone(),
|
||||
));
|
||||
}
|
||||
GeminiPart::Thought { ref text, .. } => {
|
||||
GeminiPart::Thought {
|
||||
ref text,
|
||||
ref thought_signature,
|
||||
..
|
||||
} if !text.is_empty()
|
||||
|| thought_signature.is_some() =>
|
||||
{
|
||||
if !text.is_empty() {
|
||||
thought_text.push_str(text);
|
||||
let _ = tx
|
||||
.send(StreamEvent::ThinkingDelta {
|
||||
text: text.clone(),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
if thought_signature.is_some() {
|
||||
thought_sig = thought_signature.clone();
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
@@ -1034,6 +1106,25 @@ impl LlmDriver for GeminiDriver {
|
||||
let mut content = Vec::new();
|
||||
let mut tool_calls = Vec::new();
|
||||
|
||||
// Issue #1098: persist any accumulated Thought parts (Gemini
|
||||
// 2.5+ thinking) so reasoning state round-trips on the next
|
||||
// turn. The thoughtSignature also rides on text/tool_use
|
||||
// parts below; this Thinking block carries the human-readable
|
||||
// reasoning text for UI display + audit.
|
||||
if !thought_text.is_empty() || thought_sig.is_some() {
|
||||
let provider_metadata = thought_sig.as_ref().map(|sig| {
|
||||
serde_json::json!({
|
||||
"format": "gemini_thought",
|
||||
"thought_signature": sig,
|
||||
})
|
||||
});
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking: thought_text,
|
||||
signature: None,
|
||||
provider_metadata,
|
||||
});
|
||||
}
|
||||
|
||||
if !text_content.is_empty() {
|
||||
let provider_metadata =
|
||||
text_thought_sig.map(|sig| serde_json::json!({ "thought_signature": sig }));
|
||||
@@ -1362,6 +1453,7 @@ mod tests {
|
||||
Message {
|
||||
role: Role::System,
|
||||
content: MessageContent::Text("System prompt here.".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
Message::user("Hi"),
|
||||
];
|
||||
@@ -1495,6 +1587,7 @@ mod tests {
|
||||
"thought_signature": "sig_xyz789"
|
||||
})),
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -1504,6 +1597,7 @@ mod tests {
|
||||
content: "Results about Rust programming".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
|
||||
@@ -1541,6 +1635,7 @@ mod tests {
|
||||
"thought_signature": "text_sig_abc"
|
||||
})),
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
|
||||
@@ -1612,6 +1707,7 @@ mod tests {
|
||||
input: serde_json::json!({"path": "/tmp/test"}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -1621,6 +1717,7 @@ mod tests {
|
||||
content: "file contents".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
|
||||
@@ -1864,6 +1961,7 @@ mod tests {
|
||||
Message {
|
||||
role: Role::Assistant,
|
||||
content: MessageContent::Blocks(completion.content),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -1873,6 +1971,7 @@ mod tests {
|
||||
content: "search results".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
let (contents, _) = convert_messages(&messages, &None);
|
||||
@@ -1994,7 +2093,7 @@ mod tests {
|
||||
// Should have a Thinking block and a Text block
|
||||
assert_eq!(completion.content.len(), 2);
|
||||
match &completion.content[0] {
|
||||
ContentBlock::Thinking { thinking } => {
|
||||
ContentBlock::Thinking { thinking, .. } => {
|
||||
assert_eq!(thinking, "Let me reason...");
|
||||
}
|
||||
_ => panic!("Expected Thinking block, got {:?}", completion.content[0]),
|
||||
|
||||
@@ -21,9 +21,9 @@ use openfang_types::model_catalog::{
|
||||
HUGGINGFACE_BASE_URL, KIMI_CODING_BASE_URL, LEMONADE_BASE_URL, LMSTUDIO_BASE_URL,
|
||||
MINIMAX_BASE_URL, MISTRAL_BASE_URL, MOONSHOT_BASE_URL, NOVITA_BASE_URL, NVIDIA_NIM_BASE_URL,
|
||||
OLLAMA_BASE_URL, OPENAI_BASE_URL, OPENROUTER_BASE_URL, PERPLEXITY_BASE_URL, QIANFAN_BASE_URL,
|
||||
QWEN_BASE_URL, REPLICATE_BASE_URL, SAMBANOVA_BASE_URL, TOGETHER_BASE_URL, VENICE_BASE_URL,
|
||||
VLLM_BASE_URL, VOLCENGINE_BASE_URL, VOLCENGINE_CODING_BASE_URL, XAI_BASE_URL, ZAI_BASE_URL,
|
||||
ZAI_CODING_BASE_URL, ZHIPU_BASE_URL, ZHIPU_CODING_BASE_URL,
|
||||
QWEN_BASE_URL, REPLICATE_BASE_URL, REQUESTY_BASE_URL, SAMBANOVA_BASE_URL, TOGETHER_BASE_URL,
|
||||
VENICE_BASE_URL, VLLM_BASE_URL, VOLCENGINE_BASE_URL, VOLCENGINE_CODING_BASE_URL, XAI_BASE_URL,
|
||||
ZAI_BASE_URL, ZAI_CODING_BASE_URL, ZHIPU_BASE_URL, ZHIPU_CODING_BASE_URL,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -35,6 +35,64 @@ struct ProviderDefaults {
|
||||
key_required: bool,
|
||||
}
|
||||
|
||||
/// Resolve an OpenAI-compatible base URL for a local/self-hosted provider from
|
||||
/// well-known environment variables. Returns `None` if no override is set.
|
||||
///
|
||||
/// This lets users point Ollama / LM Studio / vLLM / Lemonade at a remote host
|
||||
/// (VPS, LXC, another box on the LAN) without editing `~/.openfang/config.toml`.
|
||||
///
|
||||
/// Recognised variables:
|
||||
/// - `ollama` → `OLLAMA_BASE_URL`, then `OLLAMA_HOST` (Ollama CLI convention)
|
||||
/// - `lmstudio` → `LMSTUDIO_BASE_URL`, then `LMSTUDIO_HOST`
|
||||
/// - `vllm` → `VLLM_BASE_URL`, then `VLLM_HOST`
|
||||
/// - `lemonade` → `LEMONADE_BASE_URL`, then `LEMONADE_HOST`
|
||||
///
|
||||
/// `*_HOST` values may omit the scheme and the `/v1` suffix
|
||||
/// (e.g. `OLLAMA_HOST=192.168.1.50:11434`); both are normalised.
|
||||
pub fn local_provider_url_from_env(provider: &str) -> Option<String> {
|
||||
fn read(var: &str) -> Option<String> {
|
||||
std::env::var(var)
|
||||
.ok()
|
||||
.map(|v| v.trim().to_string())
|
||||
.filter(|v| !v.is_empty())
|
||||
}
|
||||
|
||||
/// Normalise a host-style value into a full OpenAI-compatible base URL.
|
||||
/// - Adds `http://` if no scheme is present.
|
||||
/// - Appends `/v1` if not already present in the path.
|
||||
fn normalize(raw: &str) -> String {
|
||||
let mut url = if raw.contains("://") {
|
||||
raw.trim_end_matches('/').to_string()
|
||||
} else {
|
||||
format!("http://{}", raw.trim_end_matches('/'))
|
||||
};
|
||||
// Add /v1 suffix if missing (OpenAI-compatible endpoints expect it).
|
||||
// Be lenient: accept either `/v1` or `/v1/` already in place, and also
|
||||
// `/openai/v1` style proxies.
|
||||
let lower = url.to_lowercase();
|
||||
if !lower.ends_with("/v1") && !lower.contains("/v1/") {
|
||||
url.push_str("/v1");
|
||||
}
|
||||
url
|
||||
}
|
||||
|
||||
let (primary, host_fallback) = match provider {
|
||||
"ollama" => ("OLLAMA_BASE_URL", "OLLAMA_HOST"),
|
||||
"lmstudio" => ("LMSTUDIO_BASE_URL", "LMSTUDIO_HOST"),
|
||||
"vllm" => ("VLLM_BASE_URL", "VLLM_HOST"),
|
||||
"lemonade" => ("LEMONADE_BASE_URL", "LEMONADE_HOST"),
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
if let Some(v) = read(primary) {
|
||||
return Some(normalize(&v));
|
||||
}
|
||||
if let Some(v) = read(host_fallback) {
|
||||
return Some(normalize(&v));
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
/// Get defaults for known providers.
|
||||
fn provider_defaults(provider: &str) -> Option<ProviderDefaults> {
|
||||
match provider {
|
||||
@@ -48,6 +106,11 @@ fn provider_defaults(provider: &str) -> Option<ProviderDefaults> {
|
||||
api_key_env: "OPENROUTER_API_KEY",
|
||||
key_required: true,
|
||||
}),
|
||||
"requesty" => Some(ProviderDefaults {
|
||||
base_url: REQUESTY_BASE_URL,
|
||||
api_key_env: "REQUESTY_API_KEY",
|
||||
key_required: true,
|
||||
}),
|
||||
"deepseek" => Some(ProviderDefaults {
|
||||
base_url: DEEPSEEK_BASE_URL,
|
||||
api_key_env: "DEEPSEEK_API_KEY",
|
||||
@@ -325,10 +388,28 @@ pub fn create_driver(config: &DriverConfig) -> Result<Arc<dyn LlmDriver>, LlmErr
|
||||
// Claude Code CLI — subprocess-based, no API key needed
|
||||
if provider == "claude-code" {
|
||||
let cli_path = config.base_url.clone();
|
||||
return Ok(Arc::new(claude_code::ClaudeCodeDriver::new(
|
||||
cli_path,
|
||||
config.skip_permissions,
|
||||
)));
|
||||
// Timeout precedence (highest wins):
|
||||
// 1. OPENFANG_SUBPROCESS_TIMEOUT_SECS env var (no-rebuild override for emergencies)
|
||||
// 2. DriverConfig.subprocess_timeout_secs, populated upstream from
|
||||
// config.toml — `default_model.subprocess_timeout_secs` for the
|
||||
// primary driver, `[[fallback_providers]].subprocess_timeout_secs`
|
||||
// for global fallbacks. See kernel.rs::resolve_driver and
|
||||
// kernel.rs::create_drivers for the wiring.
|
||||
// 3. Driver default (currently 300s, set inside ClaudeCodeDriver::new)
|
||||
// NOTE: The field and env var are scope-named to apply to any subprocess
|
||||
// driver, but today only `provider = "claude-code"` reads them. Other
|
||||
// drivers accept the field silently (forward-compat); future subprocess
|
||||
// drivers (qwen-code, etc.) will opt in here individually.
|
||||
let timeout = std::env::var("OPENFANG_SUBPROCESS_TIMEOUT_SECS")
|
||||
.ok()
|
||||
.and_then(|s| s.parse::<u64>().ok())
|
||||
.or(config.subprocess_timeout_secs);
|
||||
return Ok(Arc::new(match timeout {
|
||||
Some(secs) => {
|
||||
claude_code::ClaudeCodeDriver::with_timeout(cli_path, config.skip_permissions, secs)
|
||||
}
|
||||
None => claude_code::ClaudeCodeDriver::new(cli_path, config.skip_permissions),
|
||||
}));
|
||||
}
|
||||
|
||||
// Qwen Code CLI — subprocess-based, uses Qwen OAuth (free tier)
|
||||
@@ -455,9 +536,14 @@ pub fn create_driver(config: &DriverConfig) -> Result<Arc<dyn LlmDriver>, LlmErr
|
||||
)));
|
||||
}
|
||||
|
||||
// Precedence for the base URL:
|
||||
// 1. Explicit `DriverConfig.base_url` (from config.toml or `[provider_urls]`)
|
||||
// 2. Well-known env vars for local providers (`OLLAMA_HOST`, etc.) — issue #1154
|
||||
// 3. Hard-coded provider default (localhost for ollama/lmstudio/vllm/lemonade)
|
||||
let base_url = config
|
||||
.base_url
|
||||
.clone()
|
||||
.or_else(|| local_provider_url_from_env(provider))
|
||||
.unwrap_or_else(|| defaults.base_url.to_string());
|
||||
|
||||
return Ok(Arc::new(openai::OpenAIDriver::new(api_key, base_url)));
|
||||
@@ -611,9 +697,53 @@ pub fn known_providers() -> &'static [&'static str] {
|
||||
]
|
||||
}
|
||||
|
||||
/// Cross-module env-var serialisation lock for tests that mutate process env.
|
||||
///
|
||||
/// Several tests in this crate (drivers, model_catalog) set/unset the same
|
||||
/// `OLLAMA_*` / `LMSTUDIO_*` env vars and would race under cargo's parallel
|
||||
/// test runner. Anything that mutates those vars must hold this lock.
|
||||
#[cfg(test)]
|
||||
pub(crate) fn env_lock_for_tests() -> &'static std::sync::Mutex<()> {
|
||||
use std::ops::Deref;
|
||||
tests::ENV_LOCK.deref()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::ffi::OsString;
|
||||
use std::sync::{LazyLock, Mutex};
|
||||
|
||||
pub(super) static ENV_LOCK: LazyLock<Mutex<()>> = LazyLock::new(|| Mutex::new(()));
|
||||
|
||||
struct EnvVarGuard {
|
||||
key: &'static str,
|
||||
original: Option<OsString>,
|
||||
}
|
||||
|
||||
impl EnvVarGuard {
|
||||
fn set(key: &'static str, value: &str) -> Self {
|
||||
let original = std::env::var_os(key);
|
||||
std::env::set_var(key, value);
|
||||
Self { key, original }
|
||||
}
|
||||
|
||||
fn remove(key: &'static str) -> Self {
|
||||
let original = std::env::var_os(key);
|
||||
std::env::remove_var(key);
|
||||
Self { key, original }
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for EnvVarGuard {
|
||||
fn drop(&mut self) {
|
||||
if let Some(value) = &self.original {
|
||||
std::env::set_var(self.key, value);
|
||||
} else {
|
||||
std::env::remove_var(self.key);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_defaults_groq() {
|
||||
@@ -648,6 +778,7 @@ mod tests {
|
||||
api_key: Some("test".to_string()),
|
||||
base_url: Some("http://localhost:9999/v1".to_string()),
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_ok());
|
||||
@@ -660,6 +791,7 @@ mod tests {
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_err());
|
||||
@@ -772,29 +904,33 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn test_novita_provider_with_env_key() {
|
||||
let _env_lock = ENV_LOCK.lock().unwrap();
|
||||
let unique_key = "test-novita-key-12345";
|
||||
std::env::set_var("NOVITA_API_KEY", unique_key);
|
||||
let _env = EnvVarGuard::set("NOVITA_API_KEY", unique_key);
|
||||
let config = DriverConfig {
|
||||
provider: "novita".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(
|
||||
driver.is_ok(),
|
||||
"Novita provider with env var should succeed"
|
||||
);
|
||||
std::env::remove_var("NOVITA_API_KEY");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_novita_provider_no_key_errors() {
|
||||
let _env_lock = ENV_LOCK.lock().unwrap();
|
||||
let _env = EnvVarGuard::remove("NOVITA_API_KEY");
|
||||
let config = DriverConfig {
|
||||
provider: "novita".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_err());
|
||||
@@ -803,30 +939,34 @@ mod tests {
|
||||
#[test]
|
||||
fn test_nvidia_provider_with_env_key() {
|
||||
// NVIDIA NIM is a known provider — set API key and verify driver creation succeeds.
|
||||
let _env_lock = ENV_LOCK.lock().unwrap();
|
||||
let unique_key = "test-nvidia-key-12345";
|
||||
std::env::set_var("NVIDIA_API_KEY", unique_key);
|
||||
let _env = EnvVarGuard::set("NVIDIA_API_KEY", unique_key);
|
||||
let config = DriverConfig {
|
||||
provider: "nvidia".to_string(),
|
||||
api_key: None, // picked up from env via provider_defaults
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(
|
||||
driver.is_ok(),
|
||||
"NVIDIA provider with env var should succeed"
|
||||
);
|
||||
std::env::remove_var("NVIDIA_API_KEY");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_nvidia_provider_no_key_errors() {
|
||||
// NVIDIA NIM provider with no API key should error.
|
||||
let _env_lock = ENV_LOCK.lock().unwrap();
|
||||
let _env = EnvVarGuard::remove("NVIDIA_API_KEY");
|
||||
let config = DriverConfig {
|
||||
provider: "nvidia".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_err());
|
||||
@@ -835,13 +975,15 @@ mod tests {
|
||||
#[test]
|
||||
fn test_custom_provider_key_no_url_helpful_error() {
|
||||
// Custom provider with key set (via env) but no base_url should give helpful error.
|
||||
let _env_lock = ENV_LOCK.lock().unwrap();
|
||||
let unique_key = "test-custom-key-67890";
|
||||
std::env::set_var("MYCUSTOM_API_KEY", unique_key);
|
||||
let _env = EnvVarGuard::set("MYCUSTOM_API_KEY", unique_key);
|
||||
let config = DriverConfig {
|
||||
provider: "mycustom".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let result = create_driver(&config);
|
||||
assert!(result.is_err());
|
||||
@@ -851,7 +993,6 @@ mod tests {
|
||||
"Error should mention base_url: {}",
|
||||
err
|
||||
);
|
||||
std::env::remove_var("MYCUSTOM_API_KEY");
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -870,6 +1011,7 @@ mod tests {
|
||||
api_key: Some("explicit-key".to_string()),
|
||||
base_url: Some("https://api.example.com/v1".to_string()),
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_ok());
|
||||
@@ -897,6 +1039,7 @@ mod tests {
|
||||
api_key: Some("test-azure-key".to_string()),
|
||||
base_url: Some("https://myresource.openai.azure.com/openai/deployments".to_string()),
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_ok(), "Azure driver with key + URL should succeed");
|
||||
@@ -909,6 +1052,7 @@ mod tests {
|
||||
api_key: None,
|
||||
base_url: Some("https://myresource.openai.azure.com/openai/deployments".to_string()),
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let result = create_driver(&config);
|
||||
assert!(result.is_err(), "Azure driver without key should error");
|
||||
@@ -927,6 +1071,7 @@ mod tests {
|
||||
api_key: Some("test-azure-key".to_string()),
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let result = create_driver(&config);
|
||||
assert!(result.is_err(), "Azure driver without URL should error");
|
||||
@@ -945,6 +1090,7 @@ mod tests {
|
||||
api_key: Some("test-azure-key".to_string()),
|
||||
base_url: Some("https://myresource.openai.azure.com/openai/deployments".to_string()),
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(
|
||||
@@ -969,6 +1115,7 @@ mod tests {
|
||||
api_key: Some("test-bedrock-api-key".to_string()),
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
// Should succeed because api_key is provided
|
||||
let driver = create_driver(&config);
|
||||
@@ -977,4 +1124,187 @@ mod tests {
|
||||
"Bedrock with explicit api_key should construct successfully"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_claude_code_driver_constructs_with_default_timeout() {
|
||||
// No timeout in config and no env override → driver uses its built-in default.
|
||||
std::env::remove_var("OPENFANG_SUBPROCESS_TIMEOUT_SECS");
|
||||
let config = DriverConfig {
|
||||
provider: "claude-code".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_ok(), "claude-code driver should construct");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_claude_code_driver_constructs_with_config_timeout() {
|
||||
// Timeout set via config field → with_timeout path is exercised.
|
||||
std::env::remove_var("OPENFANG_SUBPROCESS_TIMEOUT_SECS");
|
||||
let config = DriverConfig {
|
||||
provider: "claude-code".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: Some(480),
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(
|
||||
driver.is_ok(),
|
||||
"claude-code driver should construct with custom timeout"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_claude_code_driver_constructs_with_env_timeout_override() {
|
||||
// Env var present → wins over config field. We can't read the timeout off the
|
||||
// trait object here, but at minimum the construction path must not panic
|
||||
// when both are set and the env var parses cleanly.
|
||||
std::env::set_var("OPENFANG_SUBPROCESS_TIMEOUT_SECS", "600");
|
||||
let config = DriverConfig {
|
||||
provider: "claude-code".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: Some(120),
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
std::env::remove_var("OPENFANG_SUBPROCESS_TIMEOUT_SECS");
|
||||
assert!(
|
||||
driver.is_ok(),
|
||||
"claude-code driver should construct when env override is set"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_claude_code_driver_ignores_unparseable_env_timeout() {
|
||||
// Garbage env var → falls through to config field, doesn't error.
|
||||
std::env::set_var("OPENFANG_SUBPROCESS_TIMEOUT_SECS", "not-a-number");
|
||||
let config = DriverConfig {
|
||||
provider: "claude-code".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: Some(420),
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
std::env::remove_var("OPENFANG_SUBPROCESS_TIMEOUT_SECS");
|
||||
assert!(
|
||||
driver.is_ok(),
|
||||
"unparseable env override should fall through to config field"
|
||||
);
|
||||
}
|
||||
|
||||
// ── Issue #1154: env-var URL overrides for local providers ──
|
||||
|
||||
#[test]
|
||||
fn test_local_url_env_ollama_host_normalised() {
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::remove("OLLAMA_BASE_URL");
|
||||
let _g2 = EnvVarGuard::set("OLLAMA_HOST", "192.168.1.50:11434");
|
||||
let url = local_provider_url_from_env("ollama").expect("env should resolve");
|
||||
assert_eq!(url, "http://192.168.1.50:11434/v1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_local_url_env_ollama_base_url_wins() {
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::set("OLLAMA_BASE_URL", "https://llm.example.com/v1");
|
||||
let _g2 = EnvVarGuard::set("OLLAMA_HOST", "should-be-ignored:11434");
|
||||
let url = local_provider_url_from_env("ollama").expect("env should resolve");
|
||||
assert_eq!(url, "https://llm.example.com/v1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_local_url_env_lmstudio() {
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::remove("LMSTUDIO_BASE_URL");
|
||||
let _g2 = EnvVarGuard::set("LMSTUDIO_HOST", "http://10.0.0.5:1234");
|
||||
let url = local_provider_url_from_env("lmstudio").expect("env should resolve");
|
||||
assert_eq!(url, "http://10.0.0.5:1234/v1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_local_url_env_vllm() {
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::remove("VLLM_BASE_URL");
|
||||
let _g2 = EnvVarGuard::set("VLLM_HOST", "vps.internal:8000");
|
||||
let url = local_provider_url_from_env("vllm").expect("env should resolve");
|
||||
assert_eq!(url, "http://vps.internal:8000/v1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_local_url_env_unset_returns_none() {
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::remove("OLLAMA_BASE_URL");
|
||||
let _g2 = EnvVarGuard::remove("OLLAMA_HOST");
|
||||
assert!(local_provider_url_from_env("ollama").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_local_url_env_only_for_local_providers() {
|
||||
// Cloud providers should never resolve via these helpers — they have
|
||||
// their own *_API_KEY conventions and a fixed cloud base URL.
|
||||
assert!(local_provider_url_from_env("openai").is_none());
|
||||
assert!(local_provider_url_from_env("anthropic").is_none());
|
||||
assert!(local_provider_url_from_env("groq").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_local_url_env_preserves_existing_v1_suffix() {
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::set("OLLAMA_BASE_URL", "http://1.2.3.4:11434/v1");
|
||||
let _g2 = EnvVarGuard::remove("OLLAMA_HOST");
|
||||
let url = local_provider_url_from_env("ollama").expect("env should resolve");
|
||||
assert_eq!(url, "http://1.2.3.4:11434/v1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_driver_ollama_uses_env_host() {
|
||||
// End-to-end: when no explicit base_url and no OLLAMA_API_KEY, the
|
||||
// driver should be constructed pointed at the env-supplied host.
|
||||
// (We can't introspect the OpenAIDriver's base_url directly, but
|
||||
// construction succeeds — separate unit covers URL resolution.)
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::remove("OLLAMA_BASE_URL");
|
||||
let _g2 = EnvVarGuard::set("OLLAMA_HOST", "10.20.30.40:11434");
|
||||
let _g3 = EnvVarGuard::remove("OLLAMA_API_KEY");
|
||||
|
||||
let config = DriverConfig {
|
||||
provider: "ollama".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(
|
||||
driver.is_ok(),
|
||||
"ollama with OLLAMA_HOST set and no API key should construct: {:?}",
|
||||
driver.err()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_driver_lmstudio_no_key_no_env_still_works() {
|
||||
// Pre-#1154 regression guard: lmstudio with no env vars and no API key
|
||||
// should still construct (falls back to localhost default).
|
||||
let _lock = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let _g1 = EnvVarGuard::remove("LMSTUDIO_BASE_URL");
|
||||
let _g2 = EnvVarGuard::remove("LMSTUDIO_HOST");
|
||||
let _g3 = EnvVarGuard::remove("LMSTUDIO_API_KEY");
|
||||
|
||||
let config = DriverConfig {
|
||||
provider: "lmstudio".to_string(),
|
||||
api_key: None,
|
||||
base_url: None,
|
||||
skip_permissions: true,
|
||||
subprocess_timeout_secs: None,
|
||||
};
|
||||
let driver = create_driver(&config);
|
||||
assert!(driver.is_ok(), "lmstudio default should construct");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -190,9 +190,14 @@ struct OaiMessage {
|
||||
tool_calls: Option<Vec<OaiToolCall>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_call_id: Option<String>,
|
||||
/// Moonshot Kimi: sent as empty string on assistant messages with tool_calls when using Kimi (thinking is disabled for multi-turn compatibility).
|
||||
/// Legacy reasoning field. Pre-vLLM 0.19, DeepSeek, Moonshot/Kimi (empty string when thinking is disabled for tool_calls multi-turn).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
reasoning_content: Option<String>,
|
||||
/// New reasoning field per OpenAI GPT-OSS Responses-API convention.
|
||||
/// vLLM 0.19+ (PR #33402) renamed `reasoning_content` to `reasoning`.
|
||||
/// Issue #1157: emit both for backward compat across servers.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
reasoning: Option<String>,
|
||||
}
|
||||
|
||||
/// Content can be a plain string or an array of content parts (for images).
|
||||
@@ -263,8 +268,23 @@ struct OaiResponseMessage {
|
||||
content: Option<String>,
|
||||
tool_calls: Option<Vec<OaiToolCall>>,
|
||||
/// Reasoning/thinking content returned by some models (DeepSeek-R1, Qwen3, etc.)
|
||||
/// via LM Studio, Ollama, and other local inference servers.
|
||||
/// via LM Studio, Ollama, and pre-0.19 vLLM.
|
||||
reasoning_content: Option<String>,
|
||||
/// New reasoning field per OpenAI GPT-OSS Responses-API convention.
|
||||
/// vLLM 0.19+ (PR #33402) emits this name instead of `reasoning_content`.
|
||||
/// Issue #1157.
|
||||
reasoning: Option<String>,
|
||||
}
|
||||
|
||||
impl OaiResponseMessage {
|
||||
/// Return whichever reasoning field the server populated.
|
||||
/// vLLM ≥ 0.19 → `reasoning`. Older servers / DeepSeek / Qwen → `reasoning_content`.
|
||||
fn reasoning_text(&self) -> Option<&str> {
|
||||
self.reasoning
|
||||
.as_deref()
|
||||
.filter(|s| !s.is_empty())
|
||||
.or_else(|| self.reasoning_content.as_deref().filter(|s| !s.is_empty()))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -289,6 +309,144 @@ fn strip_trailing_empty_assistant(messages: &mut Vec<OaiMessage>) {
|
||||
}
|
||||
}
|
||||
|
||||
/// Assemble an outbound assistant `OaiMessage` from `ContentBlock`s, replaying
|
||||
/// any `Thinking` blocks in the format the upstream model originally emitted.
|
||||
///
|
||||
/// This is the fix for issue #1098 — thinking-model state preservation.
|
||||
/// Without this, `<think>...</think>` and `reasoning_content` are stripped on
|
||||
/// the next turn so the model loses its prior reasoning trace and re-derives
|
||||
/// the answer (degrading quality). We honour `provider_metadata.format`:
|
||||
///
|
||||
/// - `"reasoning_content"` → emitted on the OpenAI `reasoning_content` field
|
||||
/// (DeepSeek-R1, Qwen3, MiniMax M2 via LM Studio/Ollama)
|
||||
/// - `"inline_think"` → wrapped in `<think>...</think>` and prepended to
|
||||
/// the visible content (MiniMax M2.5, Llama-3.3-think variants)
|
||||
/// - missing/other → fall back to the legacy Moonshot/Kimi behaviour
|
||||
/// (only emit `reasoning_content` when `needs_reasoning_content()` is true)
|
||||
fn assemble_assistant_message(
|
||||
blocks: &[ContentBlock],
|
||||
model: &str,
|
||||
driver: &OpenAIDriver,
|
||||
) -> OaiMessage {
|
||||
let mut text_parts: Vec<String> = Vec::new();
|
||||
let mut tool_calls: Vec<OaiToolCall> = Vec::new();
|
||||
let mut reasoning_field: Option<String> = None;
|
||||
let mut inline_think: Option<String> = None;
|
||||
|
||||
for block in blocks {
|
||||
match block {
|
||||
ContentBlock::Text { text, .. } => text_parts.push(text.clone()),
|
||||
ContentBlock::ToolUse {
|
||||
id, name, input, ..
|
||||
} => {
|
||||
tool_calls.push(OaiToolCall {
|
||||
id: id.clone(),
|
||||
call_type: "function".to_string(),
|
||||
function: OaiFunction {
|
||||
name: name.clone(),
|
||||
arguments: serde_json::to_string(input).unwrap_or_default(),
|
||||
},
|
||||
});
|
||||
}
|
||||
ContentBlock::Thinking {
|
||||
thinking,
|
||||
provider_metadata,
|
||||
..
|
||||
} => {
|
||||
if thinking.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let format = provider_metadata
|
||||
.as_ref()
|
||||
.and_then(|m| m.get("format"))
|
||||
.and_then(|v| v.as_str());
|
||||
match format {
|
||||
Some("inline_think") => {
|
||||
// MiniMax / models trained to expect `<think>` in
|
||||
// historical assistant messages. Concatenate
|
||||
// multiple thinking blocks if present.
|
||||
let entry = format!("<think>{thinking}</think>");
|
||||
match &mut inline_think {
|
||||
Some(existing) => existing.push_str(&entry),
|
||||
None => inline_think = Some(entry),
|
||||
}
|
||||
}
|
||||
Some("reasoning_content") => {
|
||||
// DeepSeek-R1 / Qwen3 / OpenAI-compat servers that
|
||||
// expose a separate `reasoning_content` field.
|
||||
match &mut reasoning_field {
|
||||
Some(existing) => existing.push_str(thinking),
|
||||
None => reasoning_field = Some(thinking.clone()),
|
||||
}
|
||||
}
|
||||
_ => {
|
||||
// Unknown format — preserve as inline_think since it's
|
||||
// safe (visible to the model as ordinary text). The
|
||||
// legacy Moonshot path overrides this below.
|
||||
let entry = format!("<think>{thinking}</think>");
|
||||
match &mut inline_think {
|
||||
Some(existing) => existing.push_str(&entry),
|
||||
None => inline_think = Some(entry),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
// Build the visible content by prepending inline_think (if any).
|
||||
let mut visible = String::new();
|
||||
if let Some(it) = inline_think.as_ref() {
|
||||
visible.push_str(it);
|
||||
}
|
||||
if !text_parts.is_empty() {
|
||||
visible.push_str(&text_parts.join(""));
|
||||
}
|
||||
|
||||
let has_tool_calls = !tool_calls.is_empty();
|
||||
let needs_reasoning = driver.needs_reasoning_content(model);
|
||||
|
||||
// Final reasoning fields: the per-block format hint wins; otherwise
|
||||
// fall back to legacy Moonshot/Kimi behaviour (empty string when needed).
|
||||
//
|
||||
// Issue #1157: vLLM ≥ 0.19 renamed `reasoning_content` to `reasoning`.
|
||||
// Emit BOTH fields so the persisted thinking trace reaches the model
|
||||
// regardless of which server version we're talking to. Old servers
|
||||
// ignore `reasoning`; new vLLM ignores `reasoning_content` (and would
|
||||
// otherwise silently strip our thinking, see PR vllm#33402).
|
||||
let (reasoning_content, reasoning) = if let Some(text) = reasoning_field {
|
||||
(Some(text.clone()), Some(text))
|
||||
} else if needs_reasoning {
|
||||
// Moonshot/Kimi legacy contract: empty `reasoning_content` to disable
|
||||
// thinking on tool-call multi-turn. The `reasoning` field stays unset.
|
||||
(Some(String::new()), None)
|
||||
} else {
|
||||
(None, None)
|
||||
};
|
||||
|
||||
OaiMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: if visible.is_empty() {
|
||||
if has_tool_calls {
|
||||
Some(OaiMessageContent::Text(String::new()))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
} else {
|
||||
Some(OaiMessageContent::Text(visible))
|
||||
},
|
||||
tool_calls: if tool_calls.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(tool_calls)
|
||||
},
|
||||
tool_call_id: None,
|
||||
reasoning_content,
|
||||
reasoning,
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmDriver for OpenAIDriver {
|
||||
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
@@ -302,22 +460,22 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
|
||||
// Convert messages
|
||||
for msg in &request.messages {
|
||||
match (&msg.role, &msg.content) {
|
||||
(Role::System, MessageContent::Text(text)) => {
|
||||
if request.system.is_none() {
|
||||
oai_messages.push(OaiMessage {
|
||||
role: "system".to_string(),
|
||||
content: Some(OaiMessageContent::Text(text.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
(Role::System, MessageContent::Text(text)) if request.system.is_none() => {
|
||||
oai_messages.push(OaiMessage {
|
||||
role: "system".to_string(),
|
||||
content: Some(OaiMessageContent::Text(text.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
(Role::User, MessageContent::Text(text)) => {
|
||||
oai_messages.push(OaiMessage {
|
||||
@@ -326,6 +484,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
(Role::Assistant, MessageContent::Text(text)) => {
|
||||
@@ -335,6 +494,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
(Role::User, MessageContent::Blocks(blocks)) => {
|
||||
@@ -359,6 +519,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: Some(tool_use_id.clone()),
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
ContentBlock::Text { text, .. } => {
|
||||
@@ -382,63 +543,13 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
(Role::Assistant, MessageContent::Blocks(blocks)) => {
|
||||
let mut text_parts = Vec::new();
|
||||
let mut tool_calls = Vec::new();
|
||||
let mut reasoning_text = String::new();
|
||||
for block in blocks {
|
||||
match block {
|
||||
ContentBlock::Text { text, .. } => text_parts.push(text.clone()),
|
||||
ContentBlock::ToolUse {
|
||||
id, name, input, ..
|
||||
} => {
|
||||
tool_calls.push(OaiToolCall {
|
||||
id: id.clone(),
|
||||
call_type: "function".to_string(),
|
||||
function: OaiFunction {
|
||||
name: name.clone(),
|
||||
arguments: serde_json::to_string(input).unwrap_or_default(),
|
||||
},
|
||||
});
|
||||
}
|
||||
ContentBlock::Thinking { thinking, .. } => {
|
||||
reasoning_text = thinking.clone();
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
let has_tool_calls = !tool_calls.is_empty();
|
||||
let needs_reasoning = self.needs_reasoning_content(&request.model);
|
||||
oai_messages.push(OaiMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: if text_parts.is_empty() {
|
||||
if has_tool_calls {
|
||||
Some(OaiMessageContent::Text(String::new()))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
} else {
|
||||
Some(OaiMessageContent::Text(text_parts.join("")))
|
||||
},
|
||||
tool_calls: if tool_calls.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(tool_calls)
|
||||
},
|
||||
tool_call_id: None,
|
||||
reasoning_content: if needs_reasoning {
|
||||
Some(if reasoning_text.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
reasoning_text
|
||||
})
|
||||
} else {
|
||||
None
|
||||
},
|
||||
});
|
||||
let assembled = assemble_assistant_message(blocks, &request.model, self);
|
||||
oai_messages.push(assembled);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
@@ -600,7 +711,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Model doesn't support function calling — retry without tools
|
||||
// Model doesn't support function calling — retry without tools
|
||||
// (e.g. GLM-5 on DashScope returns 500 "internal error" when tools are sent)
|
||||
let body_lower = body.to_lowercase();
|
||||
if !oai_request.tools.is_empty()
|
||||
@@ -644,30 +755,46 @@ impl LlmDriver for OpenAIDriver {
|
||||
let mut content = Vec::new();
|
||||
let mut tool_calls = Vec::new();
|
||||
|
||||
// Capture reasoning_content from models that use a separate field
|
||||
// (DeepSeek-R1, Qwen3, etc. via LM Studio/Ollama)
|
||||
if let Some(ref reasoning) = choice.message.reasoning_content {
|
||||
// Capture reasoning text from models that use a separate field.
|
||||
// Issue #1098 (legacy `reasoning_content`) + #1157 (vLLM ≥ 0.19
|
||||
// renamed it to `reasoning`). Accept either.
|
||||
if let Some(reasoning) = choice.message.reasoning_text() {
|
||||
if !reasoning.is_empty() {
|
||||
debug!(
|
||||
len = reasoning.len(),
|
||||
"Captured reasoning_content from response"
|
||||
);
|
||||
debug!(len = reasoning.len(), "Captured reasoning from response");
|
||||
// Mark the format so the outbound path knows to re-emit
|
||||
// this on the reasoning field rather than as inline
|
||||
// `<think>` tags. The outbound assembler writes BOTH
|
||||
// `reasoning` and `reasoning_content` for cross-server
|
||||
// compat.
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking: reasoning.clone(),
|
||||
thinking: reasoning.to_string(),
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "reasoning_content"
|
||||
})),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let already_has_reasoning = choice.message.reasoning_text().is_some();
|
||||
if let Some(text) = choice.message.content {
|
||||
if !text.is_empty() {
|
||||
// Extract <think>...</think> blocks that some local models
|
||||
// embed directly in the content field.
|
||||
let (cleaned, thinking) = extract_think_tags(&text);
|
||||
if let Some(think_text) = thinking {
|
||||
// Only add if we didn't already get reasoning_content
|
||||
if choice.message.reasoning_content.is_none() {
|
||||
// Only add if we didn't already get a reasoning field
|
||||
// (either legacy `reasoning_content` or new vLLM 0.19+
|
||||
// `reasoning`). Issue #1157.
|
||||
if !already_has_reasoning {
|
||||
// Mark the format so we re-emit as inline `<think>`
|
||||
// tags on the next turn (MiniMax/M2.5 style).
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking: think_text,
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "inline_think"
|
||||
})),
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -694,7 +821,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
let thinking_text = content
|
||||
.iter()
|
||||
.find_map(|b| match b {
|
||||
ContentBlock::Thinking { thinking } => Some(thinking.as_str()),
|
||||
ContentBlock::Thinking { thinking, .. } => Some(thinking.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or("");
|
||||
@@ -754,7 +881,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
// this as a "silent failure" and loop unnecessarily.
|
||||
if !content.is_empty() && usage.input_tokens == 0 && usage.output_tokens == 0 {
|
||||
debug!(
|
||||
"Response has content but no usage stats — setting synthetic output_tokens=1"
|
||||
"Response has content but no usage stats — setting synthetic output_tokens=1"
|
||||
);
|
||||
usage.output_tokens = 1;
|
||||
}
|
||||
@@ -788,21 +915,21 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
|
||||
for msg in &request.messages {
|
||||
match (&msg.role, &msg.content) {
|
||||
(Role::System, MessageContent::Text(text)) => {
|
||||
if request.system.is_none() {
|
||||
oai_messages.push(OaiMessage {
|
||||
role: "system".to_string(),
|
||||
content: Some(OaiMessageContent::Text(text.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
});
|
||||
}
|
||||
(Role::System, MessageContent::Text(text)) if request.system.is_none() => {
|
||||
oai_messages.push(OaiMessage {
|
||||
role: "system".to_string(),
|
||||
content: Some(OaiMessageContent::Text(text.clone())),
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
(Role::User, MessageContent::Text(text)) => {
|
||||
oai_messages.push(OaiMessage {
|
||||
@@ -811,6 +938,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
(Role::Assistant, MessageContent::Text(text)) => {
|
||||
@@ -820,6 +948,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: None,
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
(Role::User, MessageContent::Blocks(blocks)) => {
|
||||
@@ -840,64 +969,14 @@ impl LlmDriver for OpenAIDriver {
|
||||
tool_calls: None,
|
||||
tool_call_id: Some(tool_use_id.clone()),
|
||||
reasoning_content: None,
|
||||
reasoning: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
(Role::Assistant, MessageContent::Blocks(blocks)) => {
|
||||
let mut text_parts = Vec::new();
|
||||
let mut tool_calls_out = Vec::new();
|
||||
let mut reasoning_text = String::new();
|
||||
for block in blocks {
|
||||
match block {
|
||||
ContentBlock::Text { text, .. } => text_parts.push(text.clone()),
|
||||
ContentBlock::ToolUse {
|
||||
id, name, input, ..
|
||||
} => {
|
||||
tool_calls_out.push(OaiToolCall {
|
||||
id: id.clone(),
|
||||
call_type: "function".to_string(),
|
||||
function: OaiFunction {
|
||||
name: name.clone(),
|
||||
arguments: serde_json::to_string(input).unwrap_or_default(),
|
||||
},
|
||||
});
|
||||
}
|
||||
ContentBlock::Thinking { thinking, .. } => {
|
||||
reasoning_text = thinking.clone();
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
let has_tool_calls = !tool_calls_out.is_empty();
|
||||
let needs_reasoning = self.needs_reasoning_content(&request.model);
|
||||
oai_messages.push(OaiMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: if text_parts.is_empty() {
|
||||
if has_tool_calls {
|
||||
Some(OaiMessageContent::Text(String::new()))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
} else {
|
||||
Some(OaiMessageContent::Text(text_parts.join("")))
|
||||
},
|
||||
tool_calls: if tool_calls_out.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(tool_calls_out)
|
||||
},
|
||||
tool_call_id: None,
|
||||
reasoning_content: if needs_reasoning {
|
||||
Some(if reasoning_text.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
reasoning_text
|
||||
})
|
||||
} else {
|
||||
None
|
||||
},
|
||||
});
|
||||
let assembled = assemble_assistant_message(blocks, &request.model, self);
|
||||
oai_messages.push(assembled);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
@@ -1055,7 +1134,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Provider doesn't support stream_options — retry without it
|
||||
// Provider doesn't support stream_options — retry without it
|
||||
if status == 400
|
||||
&& oai_request.stream_options.is_some()
|
||||
&& attempt < max_retries
|
||||
@@ -1068,7 +1147,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Model doesn't support function calling — retry without tools
|
||||
// Model doesn't support function calling — retry without tools
|
||||
let body_lower = body.to_lowercase();
|
||||
if !oai_request.tools.is_empty()
|
||||
&& attempt < max_retries
|
||||
@@ -1157,7 +1236,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
for choice in choices {
|
||||
let delta = &choice["delta"];
|
||||
|
||||
// Text content delta — route through think filter to
|
||||
// Text content delta — route through think filter to
|
||||
// strip <think>...</think> tags before they reach the client.
|
||||
if let Some(text) = delta["content"].as_str() {
|
||||
if !text.is_empty() {
|
||||
@@ -1273,7 +1352,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
sse_lines = sse_line_count,
|
||||
finish = ?finish_reason,
|
||||
buffer_remaining = buffer.len(),
|
||||
"SSE stream returned empty: 0 content, 0 tokens — likely a silently failed request"
|
||||
"SSE stream returned empty: 0 content, 0 tokens — likely a silently failed request"
|
||||
);
|
||||
} else {
|
||||
debug!(
|
||||
@@ -1296,8 +1375,15 @@ impl LlmDriver for OpenAIDriver {
|
||||
|
||||
// Add reasoning/thinking content if present
|
||||
if !reasoning_content.is_empty() {
|
||||
// Mark format so outbound path replays this as
|
||||
// `reasoning_content` (DeepSeek-R1, Qwen3, MiniMax via
|
||||
// LM Studio/Ollama). Issue #1098.
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking: reasoning_content.clone(),
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "reasoning_content"
|
||||
})),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -1307,8 +1393,14 @@ impl LlmDriver for OpenAIDriver {
|
||||
if let Some(think_text) = thinking {
|
||||
// Only add if we didn't already get reasoning_content
|
||||
if reasoning_content.is_empty() {
|
||||
// Mark as inline-think so the next outbound turn
|
||||
// re-emits the content wrapped in `<think>...</think>`.
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking: think_text,
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({
|
||||
"format": "inline_think"
|
||||
})),
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -1333,7 +1425,7 @@ impl LlmDriver for OpenAIDriver {
|
||||
let thinking_text = content
|
||||
.iter()
|
||||
.find_map(|b| match b {
|
||||
ContentBlock::Thinking { thinking } => Some(thinking.as_str()),
|
||||
ContentBlock::Thinking { thinking, .. } => Some(thinking.as_str()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or("");
|
||||
@@ -1400,7 +1492,9 @@ impl LlmDriver for OpenAIDriver {
|
||||
// non-zero output_tokens so the agent loop doesn't misclassify
|
||||
// this as a "silent failure" and loop unnecessarily.
|
||||
if !content.is_empty() && usage.input_tokens == 0 && usage.output_tokens == 0 {
|
||||
debug!("Stream has content but no usage stats — setting synthetic output_tokens=1");
|
||||
debug!(
|
||||
"Stream has content but no usage stats — setting synthetic output_tokens=1"
|
||||
);
|
||||
usage.output_tokens = 1;
|
||||
}
|
||||
|
||||
@@ -1453,7 +1547,7 @@ fn extract_think_tags(text: &str) -> (String, Option<String>) {
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
// Unclosed <think> tag — treat everything after as thinking
|
||||
// Unclosed <think> tag — treat everything after as thinking
|
||||
let thought = cleaned[start + "<think>".len()..].trim().to_string();
|
||||
if !thought.is_empty() {
|
||||
thinking_parts.push(thought);
|
||||
@@ -1560,7 +1654,7 @@ fn parse_groq_failed_tool_call(body: &str) -> Option<CompletionResponse> {
|
||||
let args = &call_content[brace_pos..];
|
||||
(name, args)
|
||||
} else {
|
||||
// No args — just a tool name
|
||||
// No args — just a tool name
|
||||
(call_content.trim(), "{}")
|
||||
};
|
||||
|
||||
@@ -1576,7 +1670,7 @@ fn parse_groq_failed_tool_call(body: &str) -> Option<CompletionResponse> {
|
||||
}
|
||||
|
||||
if tool_calls.is_empty() {
|
||||
// No tool calls found — the model generated plain text but Groq rejected it.
|
||||
// No tool calls found — the model generated plain text but Groq rejected it.
|
||||
// Return it as a normal text response instead of failing.
|
||||
if !failed.trim().is_empty() {
|
||||
warn!("Recovering plain text from Groq failed_generation (no tool calls)");
|
||||
@@ -1821,9 +1915,124 @@ mod tests {
|
||||
let msg: OaiResponseMessage = serde_json::from_str(json).unwrap();
|
||||
assert!(msg.content.is_none());
|
||||
assert!(msg.reasoning_content.is_none());
|
||||
assert!(msg.reasoning.is_none());
|
||||
}
|
||||
|
||||
// ── Azure OpenAI tests ──────────────────────────────────────────
|
||||
// ── Issue #1157: vLLM ≥ 0.19 reasoning field rename ─────────────────
|
||||
|
||||
/// vLLM 0.19+ (PR #33402) returns `reasoning` instead of
|
||||
/// `reasoning_content`. We must accept the new name on ingress.
|
||||
#[test]
|
||||
fn test_oai_response_message_with_vllm_reasoning_field() {
|
||||
let json =
|
||||
r#"{"content": "Answer.", "reasoning": "I weighed A vs B.", "tool_calls": null}"#;
|
||||
let msg: OaiResponseMessage = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(msg.content.as_deref(), Some("Answer."));
|
||||
assert!(msg.reasoning_content.is_none());
|
||||
assert_eq!(msg.reasoning.as_deref(), Some("I weighed A vs B."));
|
||||
// reasoning_text() must surface the new field transparently.
|
||||
assert_eq!(msg.reasoning_text(), Some("I weighed A vs B."));
|
||||
}
|
||||
|
||||
/// If a server sends both fields (during the transition), prefer the
|
||||
/// new `reasoning` name since that's what vLLM 0.19+ writes natively.
|
||||
#[test]
|
||||
fn test_reasoning_text_prefers_new_field() {
|
||||
let json = r#"{"content": null, "reasoning": "new", "reasoning_content": "old"}"#;
|
||||
let msg: OaiResponseMessage = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(msg.reasoning_text(), Some("new"));
|
||||
}
|
||||
|
||||
/// If only the legacy field is set (older vLLM, DeepSeek, Ollama),
|
||||
/// `reasoning_text()` must still return it.
|
||||
#[test]
|
||||
fn test_reasoning_text_falls_back_to_legacy_field() {
|
||||
let json = r#"{"content": null, "reasoning_content": "legacy thinking"}"#;
|
||||
let msg: OaiResponseMessage = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(msg.reasoning_text(), Some("legacy thinking"));
|
||||
}
|
||||
|
||||
/// Outbound assembler must emit BOTH `reasoning` and `reasoning_content`
|
||||
/// when a `Thinking` block carries the `reasoning_content` format hint,
|
||||
/// so the persisted thinking reaches the model regardless of whether
|
||||
/// the upstream server is pre- or post-vLLM 0.19.
|
||||
#[test]
|
||||
fn test_assemble_emits_both_reasoning_fields_for_vllm_compat() {
|
||||
let driver = OpenAIDriver::new("test".to_string(), "http://localhost:8000/v1".to_string());
|
||||
let blocks = vec![
|
||||
ContentBlock::Thinking {
|
||||
thinking: "MARKER-vllm-019".to_string(),
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({"format": "reasoning_content"})),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "final".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
];
|
||||
let msg = assemble_assistant_message(&blocks, "minimax-m2", &driver);
|
||||
assert_eq!(
|
||||
msg.reasoning_content.as_deref(),
|
||||
Some("MARKER-vllm-019"),
|
||||
"legacy reasoning_content field required for pre-0.19 vLLM and DeepSeek"
|
||||
);
|
||||
assert_eq!(
|
||||
msg.reasoning.as_deref(),
|
||||
Some("MARKER-vllm-019"),
|
||||
"new reasoning field required for vLLM ≥ 0.19 (PR #33402)"
|
||||
);
|
||||
|
||||
// Serialize and confirm the wire shape has both keys at top level.
|
||||
let json = serde_json::to_value(&msg).unwrap();
|
||||
assert_eq!(json["reasoning_content"], "MARKER-vllm-019");
|
||||
assert_eq!(json["reasoning"], "MARKER-vllm-019");
|
||||
}
|
||||
|
||||
/// Non-reasoning models (gpt-4o, claude, …) must NOT carry either
|
||||
/// reasoning field on the wire. Regression guard for the dual-emit
|
||||
/// change in #1157.
|
||||
#[test]
|
||||
fn test_assemble_no_reasoning_fields_for_plain_model() {
|
||||
let driver = OpenAIDriver::new("test".to_string(), "https://api.openai.com/v1".to_string());
|
||||
let blocks = vec![ContentBlock::Text {
|
||||
text: "hi".to_string(),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
let msg = assemble_assistant_message(&blocks, "gpt-4o", &driver);
|
||||
assert!(msg.reasoning_content.is_none());
|
||||
assert!(msg.reasoning.is_none());
|
||||
let json = serde_json::to_value(&msg).unwrap();
|
||||
assert!(json.get("reasoning").is_none());
|
||||
assert!(json.get("reasoning_content").is_none());
|
||||
}
|
||||
|
||||
/// Moonshot/Kimi legacy contract: emit empty `reasoning_content` to
|
||||
/// disable thinking on tool-call multi-turn. Issue #1157 must not
|
||||
/// regress this — `reasoning` stays absent because Moonshot doesn't
|
||||
/// understand the new name.
|
||||
#[test]
|
||||
fn test_assemble_moonshot_keeps_legacy_field_only() {
|
||||
let driver =
|
||||
OpenAIDriver::new("test".to_string(), "https://api.moonshot.cn/v1".to_string());
|
||||
let blocks = vec![ContentBlock::ToolUse {
|
||||
id: "call_1".to_string(),
|
||||
name: "search".to_string(),
|
||||
input: serde_json::json!({"q": "x"}),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
let msg = assemble_assistant_message(&blocks, "kimi-k2", &driver);
|
||||
assert_eq!(
|
||||
msg.reasoning_content.as_deref(),
|
||||
Some(""),
|
||||
"Moonshot Kimi requires empty reasoning_content on tool_calls turns"
|
||||
);
|
||||
assert!(
|
||||
msg.reasoning.is_none(),
|
||||
"Moonshot does not understand the new vLLM `reasoning` field"
|
||||
);
|
||||
}
|
||||
|
||||
// ── Azure OpenAI tests ──────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_azure_driver_creation() {
|
||||
@@ -1898,4 +2107,139 @@ mod tests {
|
||||
let url = driver.chat_url("moonshot-v1-128k");
|
||||
assert_eq!(url, "https://api.moonshot.ai/v1/chat/completions");
|
||||
}
|
||||
|
||||
// ── issue #1098: thinking-block round-trip ────────────────────────
|
||||
|
||||
/// Inline `<think>` blocks captured on ingress must be re-emitted in
|
||||
/// historical assistant turns so MiniMax-style models retain reasoning
|
||||
/// state across turns.
|
||||
#[test]
|
||||
fn test_assemble_assistant_replays_inline_think() {
|
||||
let driver = OpenAIDriver::new(
|
||||
"test".to_string(),
|
||||
"https://api.minimax.chat/v1".to_string(),
|
||||
);
|
||||
let blocks = vec![
|
||||
ContentBlock::Thinking {
|
||||
thinking: "step-by-step reasoning".to_string(),
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({"format": "inline_think"})),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "Hello, user.".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
];
|
||||
let msg = assemble_assistant_message(&blocks, "minimax-m2.5", &driver);
|
||||
let content = match msg.content {
|
||||
Some(OaiMessageContent::Text(t)) => t,
|
||||
_ => panic!("expected text content"),
|
||||
};
|
||||
assert_eq!(
|
||||
content, "<think>step-by-step reasoning</think>Hello, user.",
|
||||
"inline_think must be re-emitted as <think> wrapping prepended to text"
|
||||
);
|
||||
// No reasoning_content field should be set for non-Moonshot models.
|
||||
assert!(msg.reasoning_content.is_none());
|
||||
}
|
||||
|
||||
/// `reasoning_content`-flavoured Thinking blocks must re-emit on the
|
||||
/// `reasoning_content` field, NOT inline (DeepSeek-R1, Qwen3, MiniMax M2
|
||||
/// via LM Studio/Ollama).
|
||||
#[test]
|
||||
fn test_assemble_assistant_replays_reasoning_content_field() {
|
||||
let driver = OpenAIDriver::new(
|
||||
"test".to_string(),
|
||||
"https://api.deepseek.com/v1".to_string(),
|
||||
);
|
||||
let blocks = vec![
|
||||
ContentBlock::Thinking {
|
||||
thinking: "internal chain-of-thought".to_string(),
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({"format": "reasoning_content"})),
|
||||
},
|
||||
ContentBlock::Text {
|
||||
text: "answer".to_string(),
|
||||
provider_metadata: None,
|
||||
},
|
||||
];
|
||||
let msg = assemble_assistant_message(&blocks, "deepseek-reasoner", &driver);
|
||||
let content = match msg.content {
|
||||
Some(OaiMessageContent::Text(t)) => t,
|
||||
_ => panic!("expected text content"),
|
||||
};
|
||||
assert_eq!(
|
||||
content, "answer",
|
||||
"visible content must not include <think>"
|
||||
);
|
||||
assert_eq!(
|
||||
msg.reasoning_content.as_deref(),
|
||||
Some("internal chain-of-thought"),
|
||||
"reasoning_content field must carry the reasoning text"
|
||||
);
|
||||
}
|
||||
|
||||
/// Without thinking blocks, the outbound message should be a plain
|
||||
/// assistant message — preserve the legacy shape.
|
||||
#[test]
|
||||
fn test_assemble_assistant_no_thinking_is_plain() {
|
||||
let driver = OpenAIDriver::new("test".to_string(), "https://api.openai.com/v1".to_string());
|
||||
let blocks = vec![ContentBlock::Text {
|
||||
text: "Hi.".to_string(),
|
||||
provider_metadata: None,
|
||||
}];
|
||||
let msg = assemble_assistant_message(&blocks, "gpt-4o", &driver);
|
||||
match msg.content {
|
||||
Some(OaiMessageContent::Text(t)) => assert_eq!(t, "Hi."),
|
||||
_ => panic!("expected text content"),
|
||||
}
|
||||
assert!(msg.reasoning_content.is_none());
|
||||
}
|
||||
|
||||
/// Issue #1098 round-trip: parse a wire response with `reasoning_content`,
|
||||
/// then feed the parsed assistant turn back through the outbound path
|
||||
/// and confirm the reasoning is replayed.
|
||||
#[test]
|
||||
fn test_reasoning_content_full_round_trip() {
|
||||
// Step 1: parse server response shape.
|
||||
let json = serde_json::json!({
|
||||
"content": "Final answer.",
|
||||
"reasoning_content": "I considered options A, B, and C…",
|
||||
"tool_calls": null
|
||||
});
|
||||
let server_msg: OaiResponseMessage = serde_json::from_value(json).unwrap();
|
||||
assert_eq!(server_msg.content.as_deref(), Some("Final answer."));
|
||||
assert_eq!(
|
||||
server_msg.reasoning_content.as_deref(),
|
||||
Some("I considered options A, B, and C…")
|
||||
);
|
||||
|
||||
// Step 2: simulate the driver building blocks (mirrors the live
|
||||
// path in `complete()`).
|
||||
let mut content = Vec::new();
|
||||
if let Some(ref reasoning) = server_msg.reasoning_content {
|
||||
content.push(ContentBlock::Thinking {
|
||||
thinking: reasoning.clone(),
|
||||
signature: None,
|
||||
provider_metadata: Some(serde_json::json!({"format": "reasoning_content"})),
|
||||
});
|
||||
}
|
||||
content.push(ContentBlock::Text {
|
||||
text: server_msg.content.unwrap(),
|
||||
provider_metadata: None,
|
||||
});
|
||||
|
||||
// Step 3: replay through the outbound path.
|
||||
let driver = OpenAIDriver::new(
|
||||
"test".to_string(),
|
||||
"https://api.deepseek.com/v1".to_string(),
|
||||
);
|
||||
let outbound = assemble_assistant_message(&content, "deepseek-reasoner", &driver);
|
||||
// The reasoning_content field must round-trip verbatim.
|
||||
assert_eq!(
|
||||
outbound.reasoning_content.as_deref(),
|
||||
Some("I considered options A, B, and C…"),
|
||||
"issue #1098 regression: reasoning was stripped on resubmission"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -457,6 +457,7 @@ mod tests {
|
||||
messages: vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::text("Hello"),
|
||||
..Default::default()
|
||||
}],
|
||||
tools: vec![],
|
||||
max_tokens: 1024,
|
||||
|
||||
@@ -7,9 +7,9 @@
|
||||
//! They receive `&GuestState` (not `&mut`) and return JSON values.
|
||||
|
||||
use crate::sandbox::GuestState;
|
||||
use crate::web_fetch;
|
||||
use openfang_types::capability::{capability_matches, Capability};
|
||||
use serde_json::json;
|
||||
use std::net::ToSocketAddrs;
|
||||
use std::path::{Component, Path};
|
||||
use tracing::debug;
|
||||
|
||||
@@ -117,64 +117,9 @@ fn safe_resolve_parent(path: &str) -> Result<std::path::PathBuf, serde_json::Val
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SSRF protection
|
||||
// SSRF protection — delegates to the canonical implementation in web_fetch.rs
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// SSRF protection: check if a hostname resolves to a private/internal IP.
|
||||
/// This defeats DNS rebinding by checking the RESOLVED address, not the hostname.
|
||||
fn is_ssrf_target(url: &str) -> Result<(), serde_json::Value> {
|
||||
// Only allow http:// and https:// schemes (block file://, gopher://, ftp://)
|
||||
if !url.starts_with("http://") && !url.starts_with("https://") {
|
||||
return Err(json!({"error": "Only http:// and https:// URLs are allowed"}));
|
||||
}
|
||||
|
||||
let host = extract_host_from_url(url);
|
||||
let hostname = host.split(':').next().unwrap_or(&host);
|
||||
|
||||
// Check hostname-based blocklist first (catches metadata endpoints)
|
||||
let blocked_hostnames = [
|
||||
"localhost",
|
||||
"metadata.google.internal",
|
||||
"metadata.aws.internal",
|
||||
"instance-data",
|
||||
"169.254.169.254",
|
||||
];
|
||||
if blocked_hostnames.contains(&hostname) {
|
||||
return Err(json!({"error": format!("SSRF blocked: {hostname} is a restricted hostname")}));
|
||||
}
|
||||
|
||||
// Resolve DNS and check every returned IP
|
||||
let port = if url.starts_with("https") { 443 } else { 80 };
|
||||
let socket_addr = format!("{hostname}:{port}");
|
||||
if let Ok(addrs) = socket_addr.to_socket_addrs() {
|
||||
for addr in addrs {
|
||||
let ip = addr.ip();
|
||||
if ip.is_loopback() || ip.is_unspecified() || is_private_ip(&ip) {
|
||||
return Err(json!({"error": format!(
|
||||
"SSRF blocked: {hostname} resolves to private IP {ip}"
|
||||
)}));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn is_private_ip(ip: &std::net::IpAddr) -> bool {
|
||||
match ip {
|
||||
std::net::IpAddr::V4(v4) => {
|
||||
let octets = v4.octets();
|
||||
matches!(
|
||||
octets,
|
||||
[10, ..] | [172, 16..=31, ..] | [192, 168, ..] | [169, 254, ..]
|
||||
)
|
||||
}
|
||||
std::net::IpAddr::V6(v6) => {
|
||||
let segments = v6.segments();
|
||||
(segments[0] & 0xfe00) == 0xfc00 || (segments[0] & 0xffc0) == 0xfe80
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Always-allowed functions
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -279,13 +224,15 @@ fn host_net_fetch(state: &GuestState, params: &serde_json::Value) -> serde_json:
|
||||
.unwrap_or("GET");
|
||||
let body = params.get("body").and_then(|b| b.as_str()).unwrap_or("");
|
||||
|
||||
// SECURITY: SSRF protection — check resolved IP against private ranges
|
||||
if let Err(e) = is_ssrf_target(url) {
|
||||
return e;
|
||||
// SECURITY: SSRF protection — delegates to the canonical check in web_fetch
|
||||
// which includes the full blocklist, metadata IP detection, IPv6 support,
|
||||
// and respects the ssrf_allowed_hosts configuration.
|
||||
if let Err(msg) = web_fetch::check_ssrf(url, &state.ssrf_allowed_hosts) {
|
||||
return json!({"error": msg});
|
||||
}
|
||||
|
||||
// Extract host:port from URL for capability check
|
||||
let host = extract_host_from_url(url);
|
||||
let host = web_fetch::extract_host(url);
|
||||
if let Err(e) = check_capability(&state.capabilities, &Capability::NetConnect(host)) {
|
||||
return e;
|
||||
}
|
||||
@@ -311,22 +258,6 @@ fn host_net_fetch(state: &GuestState, params: &serde_json::Value) -> serde_json:
|
||||
})
|
||||
}
|
||||
|
||||
/// Extract host:port from a URL for capability checking.
|
||||
fn extract_host_from_url(url: &str) -> String {
|
||||
if let Some(after_scheme) = url.split("://").nth(1) {
|
||||
let host_port = after_scheme.split('/').next().unwrap_or(after_scheme);
|
||||
if host_port.contains(':') {
|
||||
host_port.to_string()
|
||||
} else if url.starts_with("https") {
|
||||
format!("{host_port}:443")
|
||||
} else {
|
||||
format!("{host_port}:80")
|
||||
}
|
||||
} else {
|
||||
url.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Shell (capability-checked)
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -501,6 +432,7 @@ mod tests {
|
||||
kernel: None,
|
||||
agent_id: "test-agent".to_string(),
|
||||
tokio_handle: tokio::runtime::Handle::current(),
|
||||
ssrf_allowed_hosts: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -618,51 +550,68 @@ mod tests {
|
||||
assert!(safe_resolve_parent("/tmp/../../etc/shadow").is_err());
|
||||
}
|
||||
|
||||
// SSRF tests now exercise the canonical implementation in web_fetch.rs,
|
||||
// which is the same code path used by host_net_fetch at runtime.
|
||||
// This verifies the integration works end-to-end for WASM host calls.
|
||||
|
||||
#[test]
|
||||
fn test_ssrf_private_ips_blocked() {
|
||||
assert!(is_ssrf_target("http://127.0.0.1:8080/secret").is_err());
|
||||
assert!(is_ssrf_target("http://localhost:3000/api").is_err());
|
||||
assert!(is_ssrf_target("http://169.254.169.254/metadata").is_err());
|
||||
assert!(is_ssrf_target("http://metadata.google.internal/v1/instance").is_err());
|
||||
let no_allow: Vec<String> = vec![];
|
||||
assert!(web_fetch::check_ssrf("http://127.0.0.1:8080/secret", &no_allow).is_err());
|
||||
assert!(web_fetch::check_ssrf("http://localhost:3000/api", &no_allow).is_err());
|
||||
assert!(web_fetch::check_ssrf("http://169.254.169.254/metadata", &no_allow).is_err());
|
||||
assert!(
|
||||
web_fetch::check_ssrf("http://metadata.google.internal/v1/instance", &no_allow)
|
||||
.is_err()
|
||||
);
|
||||
// These were previously missing from host_functions — now covered:
|
||||
assert!(web_fetch::check_ssrf("http://[::1]:8080/secret", &no_allow).is_err());
|
||||
assert!(web_fetch::check_ssrf("http://100.100.100.200/metadata", &no_allow).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ssrf_public_ips_allowed() {
|
||||
assert!(is_ssrf_target("https://api.openai.com/v1/chat").is_ok());
|
||||
assert!(is_ssrf_target("https://google.com").is_ok());
|
||||
let no_allow: Vec<String> = vec![];
|
||||
assert!(web_fetch::check_ssrf("https://api.openai.com/v1/chat", &no_allow).is_ok());
|
||||
assert!(web_fetch::check_ssrf("https://google.com", &no_allow).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ssrf_scheme_validation() {
|
||||
assert!(is_ssrf_target("file:///etc/passwd").is_err());
|
||||
assert!(is_ssrf_target("gopher://evil.com").is_err());
|
||||
assert!(is_ssrf_target("ftp://example.com").is_err());
|
||||
let no_allow: Vec<String> = vec![];
|
||||
assert!(web_fetch::check_ssrf("file:///etc/passwd", &no_allow).is_err());
|
||||
assert!(web_fetch::check_ssrf("gopher://evil.com", &no_allow).is_err());
|
||||
assert!(web_fetch::check_ssrf("ftp://example.com", &no_allow).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_private_ip() {
|
||||
use std::net::IpAddr;
|
||||
assert!(is_private_ip(&"10.0.0.1".parse::<IpAddr>().unwrap()));
|
||||
assert!(is_private_ip(&"172.16.0.1".parse::<IpAddr>().unwrap()));
|
||||
assert!(is_private_ip(&"192.168.1.1".parse::<IpAddr>().unwrap()));
|
||||
assert!(is_private_ip(&"169.254.169.254".parse::<IpAddr>().unwrap()));
|
||||
assert!(!is_private_ip(&"8.8.8.8".parse::<IpAddr>().unwrap()));
|
||||
assert!(!is_private_ip(&"1.1.1.1".parse::<IpAddr>().unwrap()));
|
||||
fn test_ssrf_allowlist_respected() {
|
||||
let allowed = vec!["192.168.1.0/24".to_string()];
|
||||
// Private IP that matches allowlist — should pass
|
||||
assert!(web_fetch::check_ssrf("http://192.168.1.100:8080/api", &allowed).is_ok());
|
||||
// Private IP outside allowlist — should still block
|
||||
let no_allow: Vec<String> = vec![];
|
||||
assert!(web_fetch::check_ssrf("http://192.168.1.100:8080/api", &no_allow).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_host_from_url() {
|
||||
fn test_extract_host_delegates_to_web_fetch() {
|
||||
assert_eq!(
|
||||
extract_host_from_url("https://api.openai.com/v1/chat"),
|
||||
web_fetch::extract_host("https://api.openai.com/v1/chat"),
|
||||
"api.openai.com:443"
|
||||
);
|
||||
assert_eq!(
|
||||
extract_host_from_url("http://localhost:8080/api"),
|
||||
web_fetch::extract_host("http://localhost:8080/api"),
|
||||
"localhost:8080"
|
||||
);
|
||||
assert_eq!(
|
||||
extract_host_from_url("http://example.com"),
|
||||
web_fetch::extract_host("http://example.com"),
|
||||
"example.com:80"
|
||||
);
|
||||
// IPv6 — previously not handled by host_functions
|
||||
assert_eq!(
|
||||
web_fetch::extract_host("http://[::1]:9090/test"),
|
||||
"[::1]:9090"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,7 +7,15 @@ use tracing::warn;
|
||||
/// Generate images via OpenAI's image generation API.
|
||||
///
|
||||
/// Requires OPENAI_API_KEY to be set.
|
||||
pub async fn generate_image(request: &ImageGenRequest) -> Result<ImageGenResult, String> {
|
||||
///
|
||||
/// `base_url_override` (sourced from `MediaConfig.image_gen_base_url`) lets
|
||||
/// callers redirect the request to a local OpenAI-compatible image service
|
||||
/// (e.g. Lemonade/Flux, LM Studio). When `None`, the hardcoded
|
||||
/// `https://api.openai.com/v1/images/generations` endpoint is used. Closes #1051.
|
||||
pub async fn generate_image(
|
||||
request: &ImageGenRequest,
|
||||
base_url_override: Option<&str>,
|
||||
) -> Result<ImageGenResult, String> {
|
||||
// Validate request
|
||||
request.validate()?;
|
||||
|
||||
@@ -30,9 +38,19 @@ pub async fn generate_image(request: &ImageGenRequest) -> Result<ImageGenResult,
|
||||
body["quality"] = serde_json::json!(request.quality);
|
||||
}
|
||||
|
||||
// `image_gen_base_url` (config.media.image_gen_base_url) overrides the
|
||||
// hardcoded provider URL when set, allowing the same OpenAI-compat JSON
|
||||
// wire format to be sent to a local image generation service
|
||||
// (Lemonade/Flux, LM Studio, etc.) instead of the cloud provider. The
|
||||
// Authorization header is still built from `OPENAI_API_KEY`; local
|
||||
// services typically accept any non-empty bearer token. Closes #1051.
|
||||
let url = base_url_override
|
||||
.map(|base| format!("{}/v1/images/generations", base.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| "https://api.openai.com/v1/images/generations".to_string());
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let response = client
|
||||
.post("https://api.openai.com/v1/images/generations")
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {}", api_key))
|
||||
.header("Content-Type", "application/json")
|
||||
.json(&body)
|
||||
@@ -201,6 +219,38 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
/// Closes #1051: when `image_gen_base_url` is set, the URL building
|
||||
/// logic must use the override (with `/v1/images/generations` appended)
|
||||
/// and strip any trailing slash from the user-supplied base. When unset,
|
||||
/// the hardcoded provider URL is used.
|
||||
#[test]
|
||||
fn test_image_gen_base_url_override_logic() {
|
||||
// Helper mirroring the URL construction in `generate_image`.
|
||||
fn build(base: Option<&str>) -> String {
|
||||
base.map(|b| format!("{}/v1/images/generations", b.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| "https://api.openai.com/v1/images/generations".to_string())
|
||||
}
|
||||
|
||||
// Default: hardcoded URL preserved (backward compatibility).
|
||||
assert_eq!(build(None), "https://api.openai.com/v1/images/generations");
|
||||
|
||||
// Override applied.
|
||||
assert_eq!(
|
||||
build(Some("http://127.0.0.1:7000")),
|
||||
"http://127.0.0.1:7000/v1/images/generations"
|
||||
);
|
||||
|
||||
// Trailing slash on the user-supplied base is stripped.
|
||||
assert_eq!(
|
||||
build(Some("http://127.0.0.1:7000/")),
|
||||
"http://127.0.0.1:7000/v1/images/generations"
|
||||
);
|
||||
assert_eq!(
|
||||
build(Some("https://images.example.com/")),
|
||||
"https://images.example.com/v1/images/generations"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_save_images_creates_dir() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
|
||||
@@ -43,6 +43,15 @@ pub trait KernelHandle: Send + Sync {
|
||||
/// Kill an agent by ID.
|
||||
fn kill_agent(&self, agent_id: &str) -> Result<(), String>;
|
||||
|
||||
/// Activate (wake up) an inactive agent by ID, flipping its state to Running.
|
||||
/// Used by orchestrator agents to dispatch work to currently inactive agents
|
||||
/// (Suspended, Crashed, or never-started). Terminated agents cannot be revived.
|
||||
/// Returns the agent's name on success.
|
||||
fn activate_agent(&self, agent_id: &str) -> Result<String, String> {
|
||||
let _ = agent_id;
|
||||
Err("Agent activation not available".to_string())
|
||||
}
|
||||
|
||||
/// Store a value in shared memory (cross-agent accessible).
|
||||
fn memory_store(&self, key: &str, value: serde_json::Value) -> Result<(), String>;
|
||||
|
||||
|
||||
@@ -100,6 +100,7 @@ impl CompletionResponse {
|
||||
self.content.iter().any(|block| match block {
|
||||
ContentBlock::Text { text, .. } => !text.is_empty(),
|
||||
ContentBlock::Thinking { thinking, .. } => !thinking.is_empty(),
|
||||
ContentBlock::RedactedThinking { data } => !data.is_empty(),
|
||||
ContentBlock::ToolUse { .. } | ContentBlock::Image { .. } => true,
|
||||
_ => false,
|
||||
})
|
||||
@@ -188,6 +189,27 @@ pub struct DriverConfig {
|
||||
/// restricts what agents can do, making this safe.
|
||||
#[serde(default = "default_skip_permissions")]
|
||||
pub skip_permissions: bool,
|
||||
|
||||
/// Per-message subprocess turn timeout in seconds.
|
||||
///
|
||||
/// Caps how long the runtime will wait for a single CLI subprocess turn
|
||||
/// (one message round-trip) before killing the process and reporting a
|
||||
/// timeout failure. When unset, the driver's own default is used
|
||||
/// (currently 300s). Long-context Opus calls with heavy tool surfaces
|
||||
/// routinely take >4 minutes, so users running large prompts may want
|
||||
/// to bump this to 480–600s.
|
||||
///
|
||||
/// Can also be overridden at runtime via the
|
||||
/// `OPENFANG_SUBPROCESS_TIMEOUT_SECS` env var, which wins over both
|
||||
/// this field and the driver default.
|
||||
///
|
||||
/// **Scope:** Currently only honored by `provider = "claude-code"`.
|
||||
/// Other providers (`default`, `qwen-code`, `openai`, `bedrock`, etc.)
|
||||
/// accept the field for forward-compatibility but silently ignore it
|
||||
/// today. As additional subprocess-based drivers are added, they will
|
||||
/// opt in to this field individually.
|
||||
#[serde(default)]
|
||||
pub subprocess_timeout_secs: Option<u64>,
|
||||
}
|
||||
|
||||
fn default_skip_permissions() -> bool {
|
||||
@@ -202,6 +224,7 @@ impl std::fmt::Debug for DriverConfig {
|
||||
.field("api_key", &self.api_key.as_ref().map(|_| "<redacted>"))
|
||||
.field("base_url", &self.base_url)
|
||||
.field("skip_permissions", &self.skip_permissions)
|
||||
.field("subprocess_timeout_secs", &self.subprocess_timeout_secs)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -247,6 +247,13 @@ impl McpConnection {
|
||||
if let Ok(path) = std::env::var("PATH") {
|
||||
cmd.env("PATH", path);
|
||||
}
|
||||
// Some stdio MCP servers launched via node/npx require a usable home
|
||||
// directory even when they do not declare any explicit secret env vars.
|
||||
for var in &["HOME", "TMP", "TEMP"] {
|
||||
if let Ok(val) = std::env::var(var) {
|
||||
cmd.env(var, val);
|
||||
}
|
||||
}
|
||||
// On Windows, npm/node need extra vars
|
||||
if cfg!(windows) {
|
||||
for var in &[
|
||||
@@ -254,9 +261,6 @@ impl McpConnection {
|
||||
"LOCALAPPDATA",
|
||||
"USERPROFILE",
|
||||
"SystemRoot",
|
||||
"TEMP",
|
||||
"TMP",
|
||||
"HOME",
|
||||
"HOMEDRIVE",
|
||||
"HOMEPATH",
|
||||
] {
|
||||
|
||||
@@ -24,6 +24,13 @@ impl MediaEngine {
|
||||
}
|
||||
}
|
||||
|
||||
/// Read-only access to the media configuration. Used by callers that
|
||||
/// need the URL overrides (e.g. image_gen_base_url for #1051) without
|
||||
/// taking ownership of the engine.
|
||||
pub fn config(&self) -> &MediaConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
/// Describe an image using a vision-capable LLM.
|
||||
/// Auto-cascade: Anthropic -> OpenAI -> Gemini (based on API key availability).
|
||||
pub async fn describe_image(
|
||||
@@ -114,16 +121,44 @@ impl MediaEngine {
|
||||
|
||||
let model = default_audio_model(provider);
|
||||
|
||||
// Build API request
|
||||
// Build API request.
|
||||
//
|
||||
// `audio_base_url` (config.media.audio_base_url) overrides the hardcoded
|
||||
// provider URL when set, allowing the same OpenAI-compatible multipart
|
||||
// wire format to be sent to a local Whisper service (speaches,
|
||||
// faster-whisper-server, LM Studio, etc.) instead of the cloud provider.
|
||||
// The Authorization header is still built from the provider's standard
|
||||
// env var (`*_API_KEY`); local services typically accept any non-empty
|
||||
// bearer token. Closes #1051.
|
||||
let (api_url, api_key) = match provider {
|
||||
"groq" => (
|
||||
"https://api.groq.com/openai/v1/audio/transcriptions",
|
||||
std::env::var("GROQ_API_KEY").map_err(|_| "GROQ_API_KEY not set")?,
|
||||
),
|
||||
"openai" => (
|
||||
"https://api.openai.com/v1/audio/transcriptions",
|
||||
std::env::var("OPENAI_API_KEY").map_err(|_| "OPENAI_API_KEY not set")?,
|
||||
),
|
||||
"groq" => {
|
||||
let url = self
|
||||
.config
|
||||
.audio_base_url
|
||||
.as_deref()
|
||||
.map(|base| format!("{}/v1/audio/transcriptions", base.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| {
|
||||
"https://api.groq.com/openai/v1/audio/transcriptions".to_string()
|
||||
});
|
||||
(
|
||||
url,
|
||||
std::env::var("GROQ_API_KEY").map_err(|_| "GROQ_API_KEY not set")?,
|
||||
)
|
||||
}
|
||||
"openai" => {
|
||||
let url = self
|
||||
.config
|
||||
.audio_base_url
|
||||
.as_deref()
|
||||
.map(|base| format!("{}/v1/audio/transcriptions", base.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| {
|
||||
"https://api.openai.com/v1/audio/transcriptions".to_string()
|
||||
});
|
||||
(
|
||||
url,
|
||||
std::env::var("OPENAI_API_KEY").map_err(|_| "OPENAI_API_KEY not set")?,
|
||||
)
|
||||
}
|
||||
other => return Err(format!("Unsupported audio provider: {}", other)),
|
||||
};
|
||||
|
||||
@@ -141,7 +176,7 @@ impl MediaEngine {
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let resp = client
|
||||
.post(api_url)
|
||||
.post(&api_url)
|
||||
.bearer_auth(&api_key)
|
||||
.multipart(form)
|
||||
.timeout(std::time::Duration::from_secs(60))
|
||||
@@ -412,6 +447,62 @@ mod tests {
|
||||
assert!(engine.semaphore.available_permits() <= 8);
|
||||
}
|
||||
|
||||
/// Closes #1051: when `audio_base_url` is set, the URL building logic
|
||||
/// must use the override (with `/v1/audio/transcriptions` appended) and
|
||||
/// strip any trailing slash from the user-supplied base. When unset, the
|
||||
/// hardcoded provider URL is used.
|
||||
#[test]
|
||||
fn test_audio_base_url_override_logic() {
|
||||
// Helper closure mirroring the URL construction in `transcribe_audio`
|
||||
// for both providers, kept in sync intentionally.
|
||||
fn build(provider: &str, base: Option<&str>) -> String {
|
||||
match provider {
|
||||
"groq" => base
|
||||
.map(|b| format!("{}/v1/audio/transcriptions", b.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| {
|
||||
"https://api.groq.com/openai/v1/audio/transcriptions".to_string()
|
||||
}),
|
||||
"openai" => base
|
||||
.map(|b| format!("{}/v1/audio/transcriptions", b.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| {
|
||||
"https://api.openai.com/v1/audio/transcriptions".to_string()
|
||||
}),
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
|
||||
// Default: hardcoded provider URLs preserved (backward compatibility).
|
||||
assert_eq!(
|
||||
build("openai", None),
|
||||
"https://api.openai.com/v1/audio/transcriptions"
|
||||
);
|
||||
assert_eq!(
|
||||
build("groq", None),
|
||||
"https://api.groq.com/openai/v1/audio/transcriptions"
|
||||
);
|
||||
|
||||
// Override applied for both providers.
|
||||
assert_eq!(
|
||||
build("openai", Some("http://127.0.0.1:8000")),
|
||||
"http://127.0.0.1:8000/v1/audio/transcriptions"
|
||||
);
|
||||
assert_eq!(
|
||||
build("groq", Some("http://localhost:9000")),
|
||||
"http://localhost:9000/v1/audio/transcriptions"
|
||||
);
|
||||
|
||||
// Trailing slash on the user-supplied base is stripped to avoid
|
||||
// double slashes in the final URL.
|
||||
assert_eq!(
|
||||
build("openai", Some("http://127.0.0.1:8000/")),
|
||||
"http://127.0.0.1:8000/v1/audio/transcriptions"
|
||||
);
|
||||
assert_eq!(
|
||||
build("openai", Some("https://whisper.example.com/")),
|
||||
"https://whisper.example.com/v1/audio/transcriptions"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_describe_image_wrong_type() {
|
||||
let engine = MediaEngine::new(MediaConfig::default());
|
||||
|
||||
@@ -10,8 +10,8 @@ use openfang_types::model_catalog::{
|
||||
HUGGINGFACE_BASE_URL, KIMI_CODING_BASE_URL, LEMONADE_BASE_URL, LMSTUDIO_BASE_URL,
|
||||
MINIMAX_BASE_URL, MISTRAL_BASE_URL, MOONSHOT_BASE_URL, NVIDIA_NIM_BASE_URL, OLLAMA_BASE_URL,
|
||||
OPENAI_BASE_URL, OPENROUTER_BASE_URL, PERPLEXITY_BASE_URL, QIANFAN_BASE_URL, QWEN_BASE_URL,
|
||||
REPLICATE_BASE_URL, SAMBANOVA_BASE_URL, TOGETHER_BASE_URL, VENICE_BASE_URL, VLLM_BASE_URL,
|
||||
VOLCENGINE_BASE_URL, VOLCENGINE_CODING_BASE_URL, XAI_BASE_URL, ZAI_BASE_URL,
|
||||
REPLICATE_BASE_URL, REQUESTY_BASE_URL, SAMBANOVA_BASE_URL, TOGETHER_BASE_URL, VENICE_BASE_URL,
|
||||
VLLM_BASE_URL, VOLCENGINE_BASE_URL, VOLCENGINE_CODING_BASE_URL, XAI_BASE_URL, ZAI_BASE_URL,
|
||||
ZAI_CODING_BASE_URL, ZHIPU_BASE_URL, ZHIPU_CODING_BASE_URL,
|
||||
};
|
||||
use std::collections::HashMap;
|
||||
@@ -339,6 +339,27 @@ impl ModelCatalog {
|
||||
}
|
||||
}
|
||||
|
||||
/// Apply environment-variable URL overrides for local providers.
|
||||
///
|
||||
/// Honours the same env vars the drivers respect (see
|
||||
/// `drivers::local_provider_url_from_env`): `OLLAMA_HOST` / `OLLAMA_BASE_URL`,
|
||||
/// `LMSTUDIO_HOST` / `LMSTUDIO_BASE_URL`, `VLLM_HOST` / `VLLM_BASE_URL`,
|
||||
/// `LEMONADE_HOST` / `LEMONADE_BASE_URL`. This keeps the dashboard's
|
||||
/// "Providers" view in sync with what the driver actually connects to,
|
||||
/// without requiring users to edit `config.toml` for remote local-LLM hosts
|
||||
/// (VPS, LXC, LAN). See issue #1154.
|
||||
pub fn apply_local_env_overrides(&mut self) {
|
||||
for provider in ["ollama", "lmstudio", "vllm", "lemonade"] {
|
||||
if let Some(url) = crate::drivers::local_provider_url_from_env(provider) {
|
||||
if let Some(p) = self.providers.iter_mut().find(|p| p.id == provider) {
|
||||
p.base_url = url;
|
||||
// A custom host indicates intentional setup, surface it as configured.
|
||||
p.auth_status = AuthStatus::Configured;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Apply a batch of provider URL overrides from config.
|
||||
///
|
||||
/// Each entry maps a provider ID to a custom base URL.
|
||||
@@ -604,6 +625,15 @@ fn builtin_providers() -> Vec<ProviderInfo> {
|
||||
auth_status: AuthStatus::Missing,
|
||||
model_count: 0,
|
||||
},
|
||||
ProviderInfo {
|
||||
id: "requesty".into(),
|
||||
display_name: "Requesty".into(),
|
||||
api_key_env: "REQUESTY_API_KEY".into(),
|
||||
base_url: REQUESTY_BASE_URL.into(),
|
||||
key_required: true,
|
||||
auth_status: AuthStatus::Missing,
|
||||
model_count: 0,
|
||||
},
|
||||
ProviderInfo {
|
||||
id: "mistral".into(),
|
||||
display_name: "Mistral AI".into(),
|
||||
@@ -1012,13 +1042,21 @@ fn builtin_aliases() -> HashMap<String, String> {
|
||||
("qwen-coder", "qwen-code/qwen3-coder"),
|
||||
("qwen-coder-plus", "qwen-code/qwen-coder-plus"),
|
||||
("qwq", "qwen-code/qwq-32b"),
|
||||
// OpenRouter free-tier aliases
|
||||
// OpenRouter free-tier aliases. Point to free models that actually support
|
||||
// tool calling on OpenRouter's free endpoints — agents send tool definitions
|
||||
// by default, so a non-tool model returns "No endpoints found that support
|
||||
// tool use" (issue #1032).
|
||||
(
|
||||
"openrouter/free",
|
||||
"openrouter/meta-llama/llama-3.1-8b-instruct:free",
|
||||
"openrouter/meta-llama/llama-3.3-70b-instruct:free",
|
||||
),
|
||||
("free", "openrouter/meta-llama/llama-3.1-8b-instruct:free"),
|
||||
("free", "openrouter/meta-llama/llama-3.3-70b-instruct:free"),
|
||||
("free-reasoning", "openrouter/deepseek/deepseek-r1:free"),
|
||||
("openrouter/free-coder", "openrouter/qwen/qwen3-coder:free"),
|
||||
(
|
||||
"openrouter/free-large",
|
||||
"openrouter/openai/gpt-oss-120b:free",
|
||||
),
|
||||
];
|
||||
pairs
|
||||
.into_iter()
|
||||
@@ -1721,7 +1759,7 @@ fn builtin_models() -> Vec<ModelCatalogEntry> {
|
||||
aliases: vec![],
|
||||
},
|
||||
// ══════════════════════════════════════════════════════════════
|
||||
// OpenRouter (10) — pass-through models using real upstream IDs
|
||||
// OpenRouter (15+) — pass-through models using real upstream IDs
|
||||
// ══════════════════════════════════════════════════════════════
|
||||
ModelCatalogEntry {
|
||||
id: "openrouter/google/gemini-2.5-flash".into(),
|
||||
@@ -1879,6 +1917,10 @@ fn builtin_models() -> Vec<ModelCatalogEntry> {
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
// NOTE: OpenRouter's free endpoint for this model rejects tool-use
|
||||
// requests ("No endpoints found that support tool use"), so we mark
|
||||
// it as no-tool to keep agents from sending tool definitions to it.
|
||||
// The paid version of llama-3.1-8b-instruct does support tools.
|
||||
id: "openrouter/meta-llama/llama-3.1-8b-instruct:free".into(),
|
||||
display_name: "Llama 3.1 8B Free (OpenRouter)".into(),
|
||||
provider: "openrouter".into(),
|
||||
@@ -1887,18 +1929,93 @@ fn builtin_models() -> Vec<ModelCatalogEntry> {
|
||||
max_output_tokens: 4_096,
|
||||
input_cost_per_m: 0.0,
|
||||
output_cost_per_m: 0.0,
|
||||
supports_tools: false,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
// Same caveat as above — OpenRouter's free 7B endpoint has no tool
|
||||
// support; use qwen3-coder:free for tool-using free workloads.
|
||||
id: "openrouter/qwen/qwen-2.5-7b-instruct:free".into(),
|
||||
display_name: "Qwen 2.5 7B Free (OpenRouter)".into(),
|
||||
provider: "openrouter".into(),
|
||||
tier: ModelTier::Fast,
|
||||
context_window: 32_768,
|
||||
max_output_tokens: 4_096,
|
||||
input_cost_per_m: 0.0,
|
||||
output_cost_per_m: 0.0,
|
||||
supports_tools: false,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
// Free models that DO support tool calling on OpenRouter's free tier.
|
||||
// Verified against `GET https://openrouter.ai/api/v1/models` —
|
||||
// `supported_parameters` includes "tools" for these IDs.
|
||||
ModelCatalogEntry {
|
||||
id: "openrouter/meta-llama/llama-3.3-70b-instruct:free".into(),
|
||||
display_name: "Llama 3.3 70B Free (OpenRouter)".into(),
|
||||
provider: "openrouter".into(),
|
||||
tier: ModelTier::Balanced,
|
||||
context_window: 65_536,
|
||||
max_output_tokens: 4_096,
|
||||
input_cost_per_m: 0.0,
|
||||
output_cost_per_m: 0.0,
|
||||
supports_tools: true,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "openrouter/qwen/qwen3-coder:free".into(),
|
||||
display_name: "Qwen3 Coder Free (OpenRouter)".into(),
|
||||
provider: "openrouter".into(),
|
||||
tier: ModelTier::Smart,
|
||||
context_window: 262_000,
|
||||
max_output_tokens: 8_192,
|
||||
input_cost_per_m: 0.0,
|
||||
output_cost_per_m: 0.0,
|
||||
supports_tools: true,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "openrouter/openai/gpt-oss-120b:free".into(),
|
||||
display_name: "GPT-OSS 120B Free (OpenRouter)".into(),
|
||||
provider: "openrouter".into(),
|
||||
tier: ModelTier::Smart,
|
||||
context_window: 131_072,
|
||||
max_output_tokens: 8_192,
|
||||
input_cost_per_m: 0.0,
|
||||
output_cost_per_m: 0.0,
|
||||
supports_tools: true,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "openrouter/openai/gpt-oss-20b:free".into(),
|
||||
display_name: "GPT-OSS 20B Free (OpenRouter)".into(),
|
||||
provider: "openrouter".into(),
|
||||
tier: ModelTier::Fast,
|
||||
context_window: 131_072,
|
||||
max_output_tokens: 4_096,
|
||||
input_cost_per_m: 0.0,
|
||||
output_cost_per_m: 0.0,
|
||||
supports_tools: true,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "openrouter/qwen/qwen-2.5-7b-instruct:free".into(),
|
||||
display_name: "Qwen 2.5 7B Free (OpenRouter)".into(),
|
||||
id: "openrouter/z-ai/glm-4.5-air:free".into(),
|
||||
display_name: "GLM 4.5 Air Free (OpenRouter)".into(),
|
||||
provider: "openrouter".into(),
|
||||
tier: ModelTier::Fast,
|
||||
context_window: 32_768,
|
||||
max_output_tokens: 4_096,
|
||||
tier: ModelTier::Smart,
|
||||
context_window: 131_072,
|
||||
max_output_tokens: 8_192,
|
||||
input_cost_per_m: 0.0,
|
||||
output_cost_per_m: 0.0,
|
||||
supports_tools: true,
|
||||
@@ -1963,6 +2080,80 @@ fn builtin_models() -> Vec<ModelCatalogEntry> {
|
||||
aliases: vec!["hunter-alpha".into()],
|
||||
},
|
||||
// ══════════════════════════════════════════════════════════════
|
||||
// Requesty (5) — router-style OpenAI-compatible gateway (issue #995)
|
||||
// Hundreds of upstream models accessible via https://router.requesty.ai/v1
|
||||
// ══════════════════════════════════════════════════════════════
|
||||
ModelCatalogEntry {
|
||||
id: "requesty/anthropic/claude-sonnet-4".into(),
|
||||
display_name: "Claude Sonnet 4 (Requesty)".into(),
|
||||
provider: "requesty".into(),
|
||||
tier: ModelTier::Smart,
|
||||
context_window: 200_000,
|
||||
max_output_tokens: 64_000,
|
||||
input_cost_per_m: 3.0,
|
||||
output_cost_per_m: 15.0,
|
||||
supports_tools: true,
|
||||
supports_vision: true,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "requesty/openai/gpt-4o".into(),
|
||||
display_name: "GPT-4o (Requesty)".into(),
|
||||
provider: "requesty".into(),
|
||||
tier: ModelTier::Smart,
|
||||
context_window: 128_000,
|
||||
max_output_tokens: 16_384,
|
||||
input_cost_per_m: 2.5,
|
||||
output_cost_per_m: 10.0,
|
||||
supports_tools: true,
|
||||
supports_vision: true,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "requesty/google/gemini-2.5-flash".into(),
|
||||
display_name: "Gemini 2.5 Flash (Requesty)".into(),
|
||||
provider: "requesty".into(),
|
||||
tier: ModelTier::Smart,
|
||||
context_window: 1_048_576,
|
||||
max_output_tokens: 65_536,
|
||||
input_cost_per_m: 0.15,
|
||||
output_cost_per_m: 0.60,
|
||||
supports_tools: true,
|
||||
supports_vision: true,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "requesty/deepseek/deepseek-chat".into(),
|
||||
display_name: "DeepSeek V3 (Requesty)".into(),
|
||||
provider: "requesty".into(),
|
||||
tier: ModelTier::Smart,
|
||||
context_window: 128_000,
|
||||
max_output_tokens: 32_768,
|
||||
input_cost_per_m: 0.14,
|
||||
output_cost_per_m: 0.28,
|
||||
supports_tools: true,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
ModelCatalogEntry {
|
||||
id: "requesty/meta-llama/llama-3.3-70b-instruct".into(),
|
||||
display_name: "Llama 3.3 70B (Requesty)".into(),
|
||||
provider: "requesty".into(),
|
||||
tier: ModelTier::Balanced,
|
||||
context_window: 128_000,
|
||||
max_output_tokens: 32_768,
|
||||
input_cost_per_m: 0.39,
|
||||
output_cost_per_m: 0.39,
|
||||
supports_tools: true,
|
||||
supports_vision: false,
|
||||
supports_streaming: true,
|
||||
aliases: vec![],
|
||||
},
|
||||
// ══════════════════════════════════════════════════════════════
|
||||
// Mistral (6)
|
||||
// ══════════════════════════════════════════════════════════════
|
||||
ModelCatalogEntry {
|
||||
@@ -3892,7 +4083,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_catalog_has_providers() {
|
||||
let catalog = ModelCatalog::new();
|
||||
assert_eq!(catalog.list_providers().len(), 41);
|
||||
assert_eq!(catalog.list_providers().len(), 42);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -4510,4 +4701,166 @@ mod tests {
|
||||
assert_eq!(found.provider, "custom_provider");
|
||||
assert_eq!(found.id, "My-Custom-LLM");
|
||||
}
|
||||
|
||||
// ── OpenRouter free-tier fixes (issue #1032) ──────────────────────────
|
||||
|
||||
/// `openrouter/free` and `free` aliases must point to a free model that
|
||||
/// actually supports tool calling on OpenRouter's free endpoints.
|
||||
/// Previously they pointed to `llama-3.1-8b-instruct:free`, which OpenRouter
|
||||
/// rejects with "No endpoints found that support tool use" when agents
|
||||
/// send tool definitions.
|
||||
#[test]
|
||||
fn test_openrouter_free_alias_supports_tools() {
|
||||
let catalog = ModelCatalog::new();
|
||||
let entry = catalog
|
||||
.find_model("openrouter/free")
|
||||
.expect("openrouter/free alias must resolve to a known model");
|
||||
assert_eq!(entry.provider, "openrouter");
|
||||
assert!(
|
||||
entry.supports_tools,
|
||||
"openrouter/free must resolve to a tool-capable model (issue #1032). \
|
||||
Resolved to {} which has supports_tools=false",
|
||||
entry.id
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_openrouter_free_short_alias_supports_tools() {
|
||||
let catalog = ModelCatalog::new();
|
||||
let entry = catalog.find_model("free").expect("free alias must resolve");
|
||||
assert_eq!(entry.provider, "openrouter");
|
||||
assert!(
|
||||
entry.supports_tools,
|
||||
"`free` alias must resolve to a tool-capable model"
|
||||
);
|
||||
}
|
||||
|
||||
/// Confirm the resolved free model's ID is one of the verified
|
||||
/// tool-supporting free endpoints on OpenRouter.
|
||||
#[test]
|
||||
fn test_openrouter_free_alias_target() {
|
||||
let catalog = ModelCatalog::new();
|
||||
let resolved = catalog
|
||||
.resolve_alias("openrouter/free")
|
||||
.expect("alias must exist");
|
||||
// Must be one of the known-good free models with tool support.
|
||||
let known_good = [
|
||||
"openrouter/meta-llama/llama-3.3-70b-instruct:free",
|
||||
"openrouter/qwen/qwen3-coder:free",
|
||||
"openrouter/openai/gpt-oss-120b:free",
|
||||
"openrouter/openai/gpt-oss-20b:free",
|
||||
"openrouter/z-ai/glm-4.5-air:free",
|
||||
];
|
||||
assert!(
|
||||
known_good.contains(&resolved),
|
||||
"openrouter/free resolves to {}, expected one of: {:?}",
|
||||
resolved,
|
||||
known_good
|
||||
);
|
||||
}
|
||||
|
||||
/// New free-tier tool-using models are present in the catalog.
|
||||
#[test]
|
||||
fn test_openrouter_free_tool_models_present() {
|
||||
let catalog = ModelCatalog::new();
|
||||
for id in [
|
||||
"openrouter/meta-llama/llama-3.3-70b-instruct:free",
|
||||
"openrouter/qwen/qwen3-coder:free",
|
||||
"openrouter/openai/gpt-oss-120b:free",
|
||||
"openrouter/openai/gpt-oss-20b:free",
|
||||
"openrouter/z-ai/glm-4.5-air:free",
|
||||
] {
|
||||
let entry = catalog
|
||||
.find_model(id)
|
||||
.unwrap_or_else(|| panic!("missing free model {}", id));
|
||||
assert_eq!(entry.provider, "openrouter");
|
||||
assert!(entry.supports_tools, "{} must support tools", id);
|
||||
assert_eq!(entry.input_cost_per_m, 0.0, "{} must be free", id);
|
||||
assert_eq!(entry.output_cost_per_m, 0.0, "{} must be free", id);
|
||||
}
|
||||
}
|
||||
|
||||
/// Free models that OpenRouter's free endpoint does NOT route to a
|
||||
/// tool-supporting backend must be marked `supports_tools=false` so
|
||||
/// agents don't send tool defs that get rejected.
|
||||
#[test]
|
||||
fn test_openrouter_free_no_tool_models_marked() {
|
||||
let catalog = ModelCatalog::new();
|
||||
let llama8b = catalog
|
||||
.find_model("openrouter/meta-llama/llama-3.1-8b-instruct:free")
|
||||
.expect("model must exist");
|
||||
assert!(
|
||||
!llama8b.supports_tools,
|
||||
"llama-3.1-8b-instruct:free has no tool-supporting free endpoint"
|
||||
);
|
||||
let qwen7b = catalog
|
||||
.find_model("openrouter/qwen/qwen-2.5-7b-instruct:free")
|
||||
.expect("model must exist");
|
||||
assert!(
|
||||
!qwen7b.supports_tools,
|
||||
"qwen-2.5-7b-instruct:free has no tool-supporting free endpoint"
|
||||
);
|
||||
}
|
||||
|
||||
// ── Requesty provider (issue #995) ────────────────────────────────────
|
||||
|
||||
/// Requesty must be registered as a provider with the correct base URL
|
||||
/// and env var, and at least one of its catalog models must resolve.
|
||||
#[test]
|
||||
fn test_requesty_provider_and_models_present() {
|
||||
let catalog = ModelCatalog::new();
|
||||
|
||||
let provider = catalog
|
||||
.list_providers()
|
||||
.iter()
|
||||
.find(|p| p.id == "requesty")
|
||||
.expect("requesty provider must be registered");
|
||||
assert_eq!(provider.display_name, "Requesty");
|
||||
assert_eq!(provider.api_key_env, "REQUESTY_API_KEY");
|
||||
assert_eq!(provider.base_url, "https://router.requesty.ai/v1");
|
||||
assert!(provider.key_required);
|
||||
assert!(
|
||||
provider.model_count >= 1,
|
||||
"requesty must have at least one model in catalog"
|
||||
);
|
||||
|
||||
let entry = catalog
|
||||
.find_model("requesty/anthropic/claude-sonnet-4")
|
||||
.expect("requesty/anthropic/claude-sonnet-4 must resolve");
|
||||
assert_eq!(entry.provider, "requesty");
|
||||
assert!(entry.supports_tools);
|
||||
}
|
||||
|
||||
// ── Issue #1154: env-var overrides for local provider URLs ──
|
||||
|
||||
/// Local guard so this catalog test doesn't clash with the driver tests
|
||||
/// that touch the same env vars. We acquire the cross-module lock from
|
||||
/// the drivers module to serialise.
|
||||
#[test]
|
||||
fn test_apply_local_env_overrides_ollama() {
|
||||
// Serialise with driver-side env tests that touch OLLAMA_*.
|
||||
let _lock = crate::drivers::env_lock_for_tests()
|
||||
.lock()
|
||||
.unwrap_or_else(|e| e.into_inner());
|
||||
let prev_base = std::env::var_os("OLLAMA_BASE_URL");
|
||||
let prev_host = std::env::var_os("OLLAMA_HOST");
|
||||
std::env::remove_var("OLLAMA_BASE_URL");
|
||||
std::env::set_var("OLLAMA_HOST", "172.16.0.10:11434");
|
||||
|
||||
let mut catalog = ModelCatalog::new();
|
||||
catalog.apply_local_env_overrides();
|
||||
let ollama = catalog.get_provider("ollama").unwrap();
|
||||
assert_eq!(ollama.base_url, "http://172.16.0.10:11434/v1");
|
||||
assert_eq!(ollama.auth_status, AuthStatus::Configured);
|
||||
|
||||
// Restore env
|
||||
if let Some(v) = prev_base {
|
||||
std::env::set_var("OLLAMA_BASE_URL", v);
|
||||
}
|
||||
if let Some(v) = prev_host {
|
||||
std::env::set_var("OLLAMA_HOST", v);
|
||||
} else {
|
||||
std::env::remove_var("OLLAMA_HOST");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -519,7 +519,7 @@ pub fn tool_category(name: &str) -> &'static str {
|
||||
|
||||
"memory_store" | "memory_recall" | "memory_delete" | "memory_list" => "Memory",
|
||||
|
||||
"agent_send" | "agent_spawn" | "agent_list" | "agent_kill" => "Agents",
|
||||
"agent_send" | "agent_spawn" | "agent_list" | "agent_kill" | "agent_activate" => "Agents",
|
||||
|
||||
"image_describe" | "image_generate" | "audio_transcribe" | "tts_speak" => "Media",
|
||||
|
||||
@@ -581,6 +581,7 @@ pub fn tool_hint(name: &str) -> &'static str {
|
||||
"agent_spawn" => "create a new agent",
|
||||
"agent_list" => "list running agents",
|
||||
"agent_kill" => "terminate an agent",
|
||||
"agent_activate" => "wake up an inactive agent so it can receive work",
|
||||
|
||||
// Media
|
||||
"image_describe" => "describe an image",
|
||||
|
||||
@@ -198,6 +198,7 @@ mod tests {
|
||||
vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::text("Hello!"),
|
||||
..Default::default()
|
||||
}],
|
||||
vec![],
|
||||
);
|
||||
@@ -216,6 +217,7 @@ mod tests {
|
||||
"Write a function that implements async file reading with struct and impl blocks:\n\
|
||||
```rust\nfn main() { }\n```"
|
||||
),
|
||||
..Default::default()
|
||||
}],
|
||||
vec![],
|
||||
);
|
||||
@@ -238,6 +240,7 @@ mod tests {
|
||||
vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::text("Use the available tools to solve this problem."),
|
||||
..Default::default()
|
||||
}],
|
||||
tools,
|
||||
);
|
||||
@@ -257,6 +260,7 @@ mod tests {
|
||||
"This is message {} with enough content to add some token weight to the conversation.",
|
||||
i
|
||||
)),
|
||||
..Default::default()
|
||||
})
|
||||
.collect();
|
||||
let request = make_request(messages, vec![]);
|
||||
@@ -353,6 +357,7 @@ mod tests {
|
||||
vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::text("Hi"),
|
||||
..Default::default()
|
||||
}],
|
||||
vec![],
|
||||
);
|
||||
@@ -363,6 +368,7 @@ mod tests {
|
||||
vec![Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::text("Hi"),
|
||||
..Default::default()
|
||||
}],
|
||||
vec![],
|
||||
);
|
||||
|
||||
@@ -42,6 +42,9 @@ pub struct SandboxConfig {
|
||||
/// Wall-clock timeout in seconds for epoch-based interruption.
|
||||
/// Defaults to 30 seconds if None.
|
||||
pub timeout_secs: Option<u64>,
|
||||
/// Hosts allowed to bypass SSRF private-IP checks.
|
||||
/// Forwarded from `[web.fetch] ssrf_allowed_hosts` in config.toml.
|
||||
pub ssrf_allowed_hosts: Vec<String>,
|
||||
}
|
||||
|
||||
impl Default for SandboxConfig {
|
||||
@@ -51,6 +54,7 @@ impl Default for SandboxConfig {
|
||||
max_memory_bytes: 16 * 1024 * 1024,
|
||||
capabilities: Vec::new(),
|
||||
timeout_secs: None,
|
||||
ssrf_allowed_hosts: Vec::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -65,6 +69,8 @@ pub struct GuestState {
|
||||
pub agent_id: String,
|
||||
/// Tokio runtime handle for async operations in sync host functions.
|
||||
pub tokio_handle: tokio::runtime::Handle,
|
||||
/// Hosts allowed to bypass SSRF private-IP checks (from config).
|
||||
pub ssrf_allowed_hosts: Vec<String>,
|
||||
}
|
||||
|
||||
/// Result of executing a WASM module.
|
||||
@@ -164,6 +170,7 @@ impl WasmSandbox {
|
||||
kernel,
|
||||
agent_id: agent_id.to_string(),
|
||||
tokio_handle,
|
||||
ssrf_allowed_hosts: config.ssrf_allowed_hosts.clone(),
|
||||
},
|
||||
);
|
||||
|
||||
|
||||
@@ -118,6 +118,7 @@ pub fn validate_and_repair_with_stats(messages: &[Message]) -> (Vec<Message>, Re
|
||||
cleaned.push(Message {
|
||||
role: msg.role,
|
||||
content: new_content,
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
|
||||
@@ -331,7 +332,7 @@ fn reorder_tool_results(messages: &mut Vec<Message>) -> usize {
|
||||
|
||||
// Insert in reverse order so indices remain valid
|
||||
let mut sorted_insertions: Vec<(usize, Vec<ContentBlock>)> = insertions.into_iter().collect();
|
||||
sorted_insertions.sort_by(|a, b| b.0.cmp(&a.0));
|
||||
sorted_insertions.sort_by_key(|b| std::cmp::Reverse(b.0));
|
||||
|
||||
for (orig_assistant_idx, blocks) in sorted_insertions {
|
||||
if let Some(¤t_idx) = current_assistant_positions.get(&orig_assistant_idx) {
|
||||
@@ -356,6 +357,7 @@ fn reorder_tool_results(messages: &mut Vec<Message>) -> usize {
|
||||
Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Blocks(blocks),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
}
|
||||
@@ -433,7 +435,7 @@ fn insert_synthetic_results(messages: &mut Vec<Message>) -> usize {
|
||||
|
||||
// Insert in reverse order so indices stay valid
|
||||
let mut sorted: Vec<(usize, Vec<ContentBlock>)> = grouped.into_iter().collect();
|
||||
sorted.sort_by(|a, b| b.0.cmp(&a.0));
|
||||
sorted.sort_by_key(|b| std::cmp::Reverse(b.0));
|
||||
|
||||
for (assistant_idx, blocks) in sorted {
|
||||
let insert_pos = assistant_idx + 1;
|
||||
@@ -456,6 +458,7 @@ fn insert_synthetic_results(messages: &mut Vec<Message>) -> usize {
|
||||
Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Blocks(blocks),
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
}
|
||||
@@ -770,6 +773,7 @@ mod tests {
|
||||
content: "some result".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Done"),
|
||||
];
|
||||
@@ -804,6 +808,7 @@ mod tests {
|
||||
Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Text(String::new()),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Hi"),
|
||||
];
|
||||
@@ -823,6 +828,7 @@ mod tests {
|
||||
input: serde_json::json!({"query": "rust"}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -832,6 +838,7 @@ mod tests {
|
||||
content: "Results found".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Here are the results"),
|
||||
];
|
||||
@@ -855,6 +862,7 @@ mod tests {
|
||||
input: serde_json::json!({"query": "rust"}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::user("While you search, I have another question"),
|
||||
Message {
|
||||
@@ -865,6 +873,7 @@ mod tests {
|
||||
content: "Search results".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Here are results"),
|
||||
];
|
||||
@@ -909,6 +918,7 @@ mod tests {
|
||||
input: serde_json::json!({"path": "/etc/hosts"}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("I tried to read the file"),
|
||||
];
|
||||
@@ -947,6 +957,7 @@ mod tests {
|
||||
input: serde_json::json!({}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -956,6 +967,7 @@ mod tests {
|
||||
content: "First result".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message {
|
||||
role: Role::User,
|
||||
@@ -965,6 +977,7 @@ mod tests {
|
||||
content: "Duplicate result".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Done"),
|
||||
];
|
||||
@@ -1014,6 +1027,7 @@ mod tests {
|
||||
input: serde_json::json!({"key": "fact1", "value": "hello"}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
// Matching ToolResult for the first call.
|
||||
Message {
|
||||
@@ -1024,6 +1038,7 @@ mod tests {
|
||||
content: "stored".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
// Second turn: assistant calls memory_store again with the SAME id
|
||||
// because Moonshot reuses the `function_name:index` format.
|
||||
@@ -1035,6 +1050,7 @@ mod tests {
|
||||
input: serde_json::json!({"key": "fact2", "value": "world"}),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
// No matching ToolResult for the second call (e.g. lost during
|
||||
// compaction or interrupted mid-execution).
|
||||
@@ -1201,11 +1217,13 @@ mod tests {
|
||||
content: "lost".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::user("World"),
|
||||
Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Text(String::new()),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Hi"),
|
||||
];
|
||||
@@ -1228,6 +1246,7 @@ mod tests {
|
||||
text: String::new(),
|
||||
provider_metadata: None,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::user("Never mind"),
|
||||
Message::assistant("OK"),
|
||||
@@ -1274,6 +1293,7 @@ mod tests {
|
||||
provider_metadata: None,
|
||||
},
|
||||
]),
|
||||
..Default::default()
|
||||
},
|
||||
// Only tu-a has a result, tu-b is missing
|
||||
Message {
|
||||
@@ -1284,6 +1304,7 @@ mod tests {
|
||||
content: "search result".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
// Orphaned result from a non-existent tool use
|
||||
Message {
|
||||
@@ -1294,11 +1315,13 @@ mod tests {
|
||||
content: "ghost result".to_string(),
|
||||
is_error: false,
|
||||
}]),
|
||||
..Default::default()
|
||||
},
|
||||
// Empty message
|
||||
Message {
|
||||
role: Role::User,
|
||||
content: MessageContent::Text(String::new()),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Done"),
|
||||
];
|
||||
@@ -1350,6 +1373,7 @@ mod tests {
|
||||
is_error: false,
|
||||
},
|
||||
]),
|
||||
..Default::default()
|
||||
},
|
||||
Message::assistant("Hi"),
|
||||
];
|
||||
|
||||
@@ -35,6 +35,12 @@ pub const SAFE_ENV_VARS_WINDOWS: &[&str] = &[
|
||||
/// - On Windows, the Windows-specific safe variables (`SAFE_ENV_VARS_WINDOWS`)
|
||||
/// - Any additional variables the caller explicitly allows via `allowed_env_vars`
|
||||
///
|
||||
/// `allowed_env_vars` accepts either explicit variable names or the special
|
||||
/// wildcard entry `"*"`, which forwards every variable present in the parent
|
||||
/// process. Use the wildcard only when the operator has explicitly opted in
|
||||
/// (e.g. `exec_policy.shell_env_passthrough = ["*"]`) — it will leak any
|
||||
/// secret the parent holds into the child.
|
||||
///
|
||||
/// Variables that are not set in the current process environment are silently
|
||||
/// skipped (rather than being set to empty strings).
|
||||
pub fn sandbox_command(cmd: &mut tokio::process::Command, allowed_env_vars: &[String]) {
|
||||
@@ -55,6 +61,14 @@ pub fn sandbox_command(cmd: &mut tokio::process::Command, allowed_env_vars: &[St
|
||||
}
|
||||
}
|
||||
|
||||
// Wildcard: forward every var from the parent process.
|
||||
if allowed_env_vars.iter().any(|v| v == "*") {
|
||||
for (key, val) in std::env::vars() {
|
||||
cmd.env(key, val);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
// Re-add caller-specified allowed vars.
|
||||
for var in allowed_env_vars {
|
||||
if let Ok(val) = std::env::var(var) {
|
||||
@@ -63,6 +77,22 @@ pub fn sandbox_command(cmd: &mut tokio::process::Command, allowed_env_vars: &[St
|
||||
}
|
||||
}
|
||||
|
||||
/// Merge two env-passthrough lists (hand-granted + exec-policy-granted),
|
||||
/// deduplicating entries. If either contains `"*"`, the result is just `["*"]`
|
||||
/// (wildcard subsumes anything else).
|
||||
pub fn merge_env_passthrough(a: &[String], b: &[String]) -> Vec<String> {
|
||||
if a.iter().any(|v| v == "*") || b.iter().any(|v| v == "*") {
|
||||
return vec!["*".to_string()];
|
||||
}
|
||||
let mut out: Vec<String> = Vec::with_capacity(a.len() + b.len());
|
||||
for v in a.iter().chain(b.iter()) {
|
||||
if !out.iter().any(|existing| existing == v) {
|
||||
out.push(v.clone());
|
||||
}
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
/// Validates that an executable path does not contain directory traversal
|
||||
/// components (`..`).
|
||||
///
|
||||
@@ -711,6 +741,40 @@ pub async fn wait_or_kill_with_idle(
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// ── Env passthrough merge (issue #1169) ────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_merge_env_passthrough_empty() {
|
||||
let merged = merge_env_passthrough(&[], &[]);
|
||||
assert!(merged.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_env_passthrough_dedup() {
|
||||
let a = vec!["TZ".to_string(), "HOME".to_string()];
|
||||
let b = vec!["TZ".to_string(), "PATH".to_string()];
|
||||
let merged = merge_env_passthrough(&a, &b);
|
||||
assert_eq!(merged, vec!["TZ", "HOME", "PATH"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_env_passthrough_wildcard_a() {
|
||||
let merged = merge_env_passthrough(&["*".to_string()], &["TZ".to_string()]);
|
||||
assert_eq!(merged, vec!["*"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_env_passthrough_wildcard_b() {
|
||||
let merged = merge_env_passthrough(&["TZ".to_string()], &["*".to_string()]);
|
||||
assert_eq!(merged, vec!["*"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_exec_policy_default_has_empty_passthrough() {
|
||||
let policy = openfang_types::config::ExecPolicy::default();
|
||||
assert!(policy.shell_env_passthrough.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_path() {
|
||||
// Clean paths should be accepted.
|
||||
|
||||
@@ -205,6 +205,7 @@ pub async fn execute_tool(
|
||||
"file_read" => tool_file_read(input, workspace_root).await,
|
||||
"file_write" => tool_file_write(input, workspace_root).await,
|
||||
"file_list" => tool_file_list(input, workspace_root).await,
|
||||
"create_directory" => tool_create_directory(input, workspace_root).await,
|
||||
"apply_patch" => tool_apply_patch(input, workspace_root).await,
|
||||
|
||||
// Web tools (upgraded: multi-provider search, SSRF-protected fetch)
|
||||
@@ -299,6 +300,7 @@ pub async fn execute_tool(
|
||||
"agent_spawn" => tool_agent_spawn(input, kernel, caller_agent_id).await,
|
||||
"agent_list" => tool_agent_list(kernel),
|
||||
"agent_kill" => tool_agent_kill(input, kernel),
|
||||
"agent_activate" => tool_agent_activate(input, kernel),
|
||||
|
||||
// Shared memory tools
|
||||
"memory_store" => tool_memory_store(input, kernel),
|
||||
@@ -330,7 +332,7 @@ pub async fn execute_tool(
|
||||
"media_transcribe" => tool_media_transcribe(input, media_engine).await,
|
||||
|
||||
// Image generation tool
|
||||
"image_generate" => tool_image_generate(input, workspace_root).await,
|
||||
"image_generate" => tool_image_generate(input, workspace_root, media_engine).await,
|
||||
|
||||
// TTS/STT tools
|
||||
"text_to_speech" => tool_text_to_speech(input, tts_engine, workspace_root).await,
|
||||
@@ -476,6 +478,11 @@ pub async fn execute_tool(
|
||||
}
|
||||
},
|
||||
|
||||
// Skill introspection tools (issue #1038)
|
||||
"skill_list" => tool_skill_list(skill_registry),
|
||||
"skill_describe" => tool_skill_describe(input, skill_registry),
|
||||
"skill_execute" => tool_skill_execute(input, skill_registry).await,
|
||||
|
||||
// Canvas / A2UI tool
|
||||
"canvas_present" => tool_canvas_present(input, workspace_root).await,
|
||||
|
||||
@@ -594,6 +601,17 @@ pub fn builtin_tool_definitions() -> Vec<ToolDefinition> {
|
||||
"required": ["path"]
|
||||
}),
|
||||
},
|
||||
ToolDefinition {
|
||||
name: "create_directory".to_string(),
|
||||
description: "Create a directory (and any missing parent directories) at the given path. Paths are relative to the agent workspace. Idempotent: succeeds if the directory already exists.".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": { "type": "string", "description": "The directory path to create" }
|
||||
},
|
||||
"required": ["path"]
|
||||
}),
|
||||
},
|
||||
ToolDefinition {
|
||||
name: "apply_patch".to_string(),
|
||||
description: "Apply a multi-hunk diff patch to add, update, move, or delete files. Use this for targeted edits instead of full file overwrites.".to_string(),
|
||||
@@ -694,6 +712,24 @@ pub fn builtin_tool_definitions() -> Vec<ToolDefinition> {
|
||||
"required": ["agent_id"]
|
||||
}),
|
||||
},
|
||||
ToolDefinition {
|
||||
name: "agent_activate".to_string(),
|
||||
description: "Activate (wake up) an inactive agent so it can receive messages \
|
||||
and process events. Use this when agent_list shows an agent in a \
|
||||
Suspended, Crashed, or Created state and you want to delegate work \
|
||||
to it via agent_send. Terminated agents cannot be revived."
|
||||
.to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"agent_id": {
|
||||
"type": "string",
|
||||
"description": "The target agent's UUID or human-readable name"
|
||||
}
|
||||
},
|
||||
"required": ["agent_id"]
|
||||
}),
|
||||
},
|
||||
// --- Shared memory tools ---
|
||||
ToolDefinition {
|
||||
name: "memory_store".to_string(),
|
||||
@@ -1277,6 +1313,42 @@ pub fn builtin_tool_definitions() -> Vec<ToolDefinition> {
|
||||
"required": ["html"]
|
||||
}),
|
||||
},
|
||||
// --- Skill introspection tools (issue #1038) ---
|
||||
// These let the agent discover and read installed skills without
|
||||
// touching the filesystem. Global skills live at ~/.openfang/skills/
|
||||
// which is outside the workspace sandbox — file_read cannot reach them.
|
||||
ToolDefinition {
|
||||
name: "skill_list".to_string(),
|
||||
description: "List all installed skills available to this agent. Returns name, version, description, runtime type, and provided tool names. Use this instead of file_list on the skills directory.".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {}
|
||||
}),
|
||||
},
|
||||
ToolDefinition {
|
||||
name: "skill_describe".to_string(),
|
||||
description: "Read the full description (SKILL.md body / prompt context) of an installed skill by name. Use this instead of file_read on a skill's SKILL.md file — global skills live outside the workspace sandbox and cannot be read with file_read.".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": { "type": "string", "description": "The skill name (as returned by skill_list)" }
|
||||
},
|
||||
"required": ["name"]
|
||||
}),
|
||||
},
|
||||
ToolDefinition {
|
||||
name: "skill_execute".to_string(),
|
||||
description: "Execute a tool provided by an installed skill. For code-runtime skills (Python/Node/Shell) this invokes the underlying script. For prompt-only skills this returns the skill's instruction body so the agent can follow it using built-in tools.".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"skill": { "type": "string", "description": "The skill name (as returned by skill_list)" },
|
||||
"tool": { "type": "string", "description": "Optional name of a tool the skill provides. Omit to invoke the skill's default behavior (returns SKILL.md body for prompt-only skills)." },
|
||||
"input": { "type": "object", "description": "Optional JSON input for the skill tool" }
|
||||
},
|
||||
"required": ["skill"]
|
||||
}),
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
@@ -1339,6 +1411,82 @@ async fn tool_file_write(
|
||||
))
|
||||
}
|
||||
|
||||
/// Resolve a directory path for creation. Unlike `resolve_file_path`, this walks
|
||||
/// up the path to find the nearest existing ancestor, canonicalizes that, and
|
||||
/// re-appends the missing segments. This lets `create_directory` accept nested
|
||||
/// paths like `a/b/c/d` even when none of `a`, `b`, `c` exist yet.
|
||||
fn resolve_directory_path_for_create(
|
||||
raw_path: &str,
|
||||
workspace_root: Option<&Path>,
|
||||
) -> Result<PathBuf, String> {
|
||||
// Reject `..` components regardless of workspace.
|
||||
let _ = validate_path(raw_path)?;
|
||||
|
||||
let Some(root) = workspace_root else {
|
||||
return Ok(PathBuf::from(raw_path));
|
||||
};
|
||||
|
||||
let path = Path::new(raw_path);
|
||||
let candidate = if path.is_absolute() {
|
||||
path.to_path_buf()
|
||||
} else {
|
||||
root.join(path)
|
||||
};
|
||||
|
||||
let canon_root = root
|
||||
.canonicalize()
|
||||
.map_err(|e| format!("Failed to resolve workspace root: {e}"))?;
|
||||
|
||||
// Walk up to find the nearest existing ancestor, canonicalize it, then
|
||||
// re-append the missing tail.
|
||||
let mut existing: PathBuf = candidate.clone();
|
||||
let mut tail: Vec<std::ffi::OsString> = Vec::new();
|
||||
while !existing.exists() {
|
||||
let parent = match existing.parent() {
|
||||
Some(p) => p.to_path_buf(),
|
||||
None => return Err("Invalid path: no existing ancestor".to_string()),
|
||||
};
|
||||
let name = match existing.file_name() {
|
||||
Some(n) => n.to_os_string(),
|
||||
None => return Err("Invalid path: no filename component".to_string()),
|
||||
};
|
||||
tail.push(name);
|
||||
existing = parent;
|
||||
}
|
||||
|
||||
let canon_existing = existing
|
||||
.canonicalize()
|
||||
.map_err(|e| format!("Failed to resolve ancestor directory: {e}"))?;
|
||||
|
||||
let mut resolved = canon_existing;
|
||||
for segment in tail.into_iter().rev() {
|
||||
resolved.push(segment);
|
||||
}
|
||||
|
||||
if !resolved.starts_with(&canon_root) {
|
||||
return Err(format!(
|
||||
"Access denied: path '{raw_path}' resolves outside workspace"
|
||||
));
|
||||
}
|
||||
|
||||
Ok(resolved)
|
||||
}
|
||||
|
||||
async fn tool_create_directory(
|
||||
input: &serde_json::Value,
|
||||
workspace_root: Option<&Path>,
|
||||
) -> Result<String, String> {
|
||||
let raw_path = input["path"].as_str().ok_or("Missing 'path' parameter")?;
|
||||
if raw_path.is_empty() {
|
||||
return Err("'path' parameter is empty".to_string());
|
||||
}
|
||||
let resolved = resolve_directory_path_for_create(raw_path, workspace_root)?;
|
||||
tokio::fs::create_dir_all(&resolved)
|
||||
.await
|
||||
.map_err(|e| format!("Failed to create directory: {e}"))?;
|
||||
Ok(format!("Created directory {}", resolved.display()))
|
||||
}
|
||||
|
||||
async fn tool_file_list(
|
||||
input: &serde_json::Value,
|
||||
workspace_root: Option<&Path>,
|
||||
@@ -1560,7 +1708,18 @@ async fn tool_shell_exec(
|
||||
|
||||
// SECURITY: Isolate environment to prevent credential leakage.
|
||||
// Hand settings may grant access to specific provider API keys.
|
||||
crate::subprocess_sandbox::sandbox_command(&mut cmd, allowed_env);
|
||||
//
|
||||
// Operators can also forward additional vars via
|
||||
// `exec_policy.shell_env_passthrough` (issue #1169). This is the path
|
||||
// Docker users hit: their container env (TZ, GOG_*, etc.) is present
|
||||
// in PID 1 but `env_clear()` strips it. Listing names (or `"*"`) here
|
||||
// re-adds them to the child.
|
||||
let policy_env_passthrough: &[String] = exec_policy
|
||||
.map(|p| p.shell_env_passthrough.as_slice())
|
||||
.unwrap_or(&[]);
|
||||
let merged_env =
|
||||
crate::subprocess_sandbox::merge_env_passthrough(allowed_env, policy_env_passthrough);
|
||||
crate::subprocess_sandbox::sandbox_command(&mut cmd, &merged_env);
|
||||
|
||||
// Ensure UTF-8 output on Windows
|
||||
#[cfg(windows)]
|
||||
@@ -1696,6 +1855,20 @@ fn tool_agent_kill(
|
||||
Ok(format!("Agent {agent_id} killed successfully."))
|
||||
}
|
||||
|
||||
fn tool_agent_activate(
|
||||
input: &serde_json::Value,
|
||||
kernel: Option<&Arc<dyn KernelHandle>>,
|
||||
) -> Result<String, String> {
|
||||
let kh = require_kernel(kernel)?;
|
||||
let agent_id = input["agent_id"]
|
||||
.as_str()
|
||||
.ok_or("Missing 'agent_id' parameter")?;
|
||||
let name = kh.activate_agent(agent_id)?;
|
||||
Ok(format!(
|
||||
"Agent '{name}' activated. It is now Running and ready to receive messages."
|
||||
))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Shared memory tools
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -2914,6 +3087,7 @@ async fn tool_media_transcribe(
|
||||
async fn tool_image_generate(
|
||||
input: &serde_json::Value,
|
||||
workspace_root: Option<&Path>,
|
||||
media_engine: Option<&crate::media_understanding::MediaEngine>,
|
||||
) -> Result<String, String> {
|
||||
let prompt = input["prompt"]
|
||||
.as_str()
|
||||
@@ -2943,7 +3117,10 @@ async fn tool_image_generate(
|
||||
count,
|
||||
};
|
||||
|
||||
let result = crate::image_gen::generate_image(&request).await?;
|
||||
// Closes #1051: route to a local OpenAI-compatible image generation
|
||||
// service when `media.image_gen_base_url` is set.
|
||||
let base_url_override = media_engine.and_then(|e| e.config().image_gen_base_url.as_deref());
|
||||
let result = crate::image_gen::generate_image(&request, base_url_override).await?;
|
||||
|
||||
// Save images to workspace if available
|
||||
let saved_paths = if let Some(workspace) = workspace_root {
|
||||
@@ -3400,6 +3577,165 @@ async fn tool_canvas_present(
|
||||
serde_json::to_string_pretty(&response).map_err(|e| format!("Serialize error: {e}"))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Skill introspection tools (issue #1038)
|
||||
//
|
||||
// Global skills live at ~/.openfang/skills/ which is outside the agent
|
||||
// workspace sandbox. Without these tools the LLM falls back to file_read /
|
||||
// shell_exec to inspect SKILL.md files — which fail with path-resolution
|
||||
// errors because file_read is workspace-scoped. These tools surface the
|
||||
// already-loaded skill registry directly to the agent.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// List all skills available to this agent, with their provided tool names.
|
||||
fn tool_skill_list(skill_registry: Option<&SkillRegistry>) -> Result<String, String> {
|
||||
let registry = match skill_registry {
|
||||
Some(r) => r,
|
||||
None => return Ok("No skill registry available.".to_string()),
|
||||
};
|
||||
let skills = registry.list();
|
||||
if skills.is_empty() {
|
||||
return Ok("No skills installed. Install skills via the dashboard or `openfang skill install <name>`.".to_string());
|
||||
}
|
||||
let entries: Vec<serde_json::Value> = skills
|
||||
.iter()
|
||||
.map(|s| {
|
||||
let tool_names: Vec<String> = s
|
||||
.manifest
|
||||
.tools
|
||||
.provided
|
||||
.iter()
|
||||
.map(|t| t.name.clone())
|
||||
.collect();
|
||||
serde_json::json!({
|
||||
"name": s.manifest.skill.name,
|
||||
"version": s.manifest.skill.version,
|
||||
"description": s.manifest.skill.description,
|
||||
"runtime": format!("{:?}", s.manifest.runtime.runtime_type),
|
||||
"enabled": s.enabled,
|
||||
"tools": tool_names,
|
||||
"has_prompt_context": s.manifest.prompt_context.as_ref().is_some_and(|c| !c.is_empty()),
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
serde_json::to_string_pretty(&serde_json::json!({
|
||||
"count": entries.len(),
|
||||
"skills": entries,
|
||||
}))
|
||||
.map_err(|e| format!("Serialize error: {e}"))
|
||||
}
|
||||
|
||||
/// Return the full description (SKILL.md body) of a named skill.
|
||||
fn tool_skill_describe(
|
||||
input: &serde_json::Value,
|
||||
skill_registry: Option<&SkillRegistry>,
|
||||
) -> Result<String, String> {
|
||||
let name = input["name"]
|
||||
.as_str()
|
||||
.ok_or("Missing 'name' parameter")?
|
||||
.trim();
|
||||
let registry = skill_registry.ok_or("No skill registry available")?;
|
||||
let skill = registry.get(name).ok_or_else(|| {
|
||||
format!("Skill '{name}' not found. Use skill_list to see installed skills.")
|
||||
})?;
|
||||
let body = skill
|
||||
.manifest
|
||||
.prompt_context
|
||||
.clone()
|
||||
.unwrap_or_else(|| "(No prompt context body — this skill provides executable tools only. Use skill_execute or call its tools directly.)".to_string());
|
||||
let tool_names: Vec<String> = skill
|
||||
.manifest
|
||||
.tools
|
||||
.provided
|
||||
.iter()
|
||||
.map(|t| t.name.clone())
|
||||
.collect();
|
||||
let response = serde_json::json!({
|
||||
"name": skill.manifest.skill.name,
|
||||
"version": skill.manifest.skill.version,
|
||||
"description": skill.manifest.skill.description,
|
||||
"runtime": format!("{:?}", skill.manifest.runtime.runtime_type),
|
||||
"tools": tool_names,
|
||||
"body": body,
|
||||
});
|
||||
serde_json::to_string_pretty(&response).map_err(|e| format!("Serialize error: {e}"))
|
||||
}
|
||||
|
||||
/// Execute a skill's tool, or for prompt-only skills return the description body.
|
||||
async fn tool_skill_execute(
|
||||
input: &serde_json::Value,
|
||||
skill_registry: Option<&SkillRegistry>,
|
||||
) -> Result<String, String> {
|
||||
let skill_name = input["skill"]
|
||||
.as_str()
|
||||
.ok_or("Missing 'skill' parameter")?
|
||||
.trim();
|
||||
let registry = skill_registry.ok_or("No skill registry available")?;
|
||||
let skill = registry.get(skill_name).ok_or_else(|| {
|
||||
format!("Skill '{skill_name}' not found. Use skill_list to see installed skills.")
|
||||
})?;
|
||||
|
||||
// If no tool name was given, default behavior depends on runtime.
|
||||
// For prompt-only skills, return the SKILL.md body (most useful response
|
||||
// for issue #1038's daily-journal style skills).
|
||||
let tool_name = input["tool"].as_str().map(|s| s.trim());
|
||||
let tool_input = input.get("input").cloned().unwrap_or(serde_json::json!({}));
|
||||
|
||||
let resolved_tool = match tool_name {
|
||||
Some(t) if !t.is_empty() => t.to_string(),
|
||||
_ => {
|
||||
// No tool specified — return SKILL.md body so the agent can act on it.
|
||||
if let Some(ref body) = skill.manifest.prompt_context {
|
||||
if !body.is_empty() {
|
||||
let response = serde_json::json!({
|
||||
"skill": skill.manifest.skill.name,
|
||||
"mode": "prompt_context",
|
||||
"body": body,
|
||||
"note": "This is a prompt-only skill. Follow the instructions in 'body' using your built-in tools.",
|
||||
});
|
||||
return serde_json::to_string_pretty(&response)
|
||||
.map_err(|e| format!("Serialize error: {e}"));
|
||||
}
|
||||
}
|
||||
// Fall through: pick the first provided tool if any
|
||||
skill
|
||||
.manifest
|
||||
.tools
|
||||
.provided
|
||||
.first()
|
||||
.map(|t| t.name.clone())
|
||||
.ok_or_else(|| {
|
||||
format!("Skill '{skill_name}' provides no tools and has no prompt body.")
|
||||
})?
|
||||
}
|
||||
};
|
||||
|
||||
match openfang_skills::loader::execute_skill_tool(
|
||||
&skill.manifest,
|
||||
&skill.path,
|
||||
&resolved_tool,
|
||||
&tool_input,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => {
|
||||
let content = serde_json::to_string_pretty(&serde_json::json!({
|
||||
"skill": skill.manifest.skill.name,
|
||||
"tool": resolved_tool,
|
||||
"output": result.output,
|
||||
"is_error": result.is_error,
|
||||
}))
|
||||
.unwrap_or_else(|_| result.output.to_string());
|
||||
if result.is_error {
|
||||
Err(content)
|
||||
} else {
|
||||
Ok(content)
|
||||
}
|
||||
}
|
||||
Err(e) => Err(format!("Skill execution failed: {e}")),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -3415,11 +3751,16 @@ mod tests {
|
||||
let names: Vec<&str> = tools.iter().map(|t| t.name.as_str()).collect();
|
||||
// Original 12
|
||||
assert!(names.contains(&"file_read"));
|
||||
assert!(names.contains(&"file_write"));
|
||||
assert!(names.contains(&"file_list"));
|
||||
assert!(names.contains(&"create_directory"));
|
||||
assert!(names.contains(&"shell_exec"));
|
||||
assert!(names.contains(&"agent_send"));
|
||||
assert!(names.contains(&"agent_spawn"));
|
||||
assert!(names.contains(&"agent_list"));
|
||||
assert!(names.contains(&"agent_kill"));
|
||||
// Issue #890 — wake up inactive agents
|
||||
assert!(names.contains(&"agent_activate"));
|
||||
assert!(names.contains(&"memory_store"));
|
||||
assert!(names.contains(&"memory_recall"));
|
||||
// 6 collaboration tools
|
||||
@@ -3468,6 +3809,67 @@ mod tests {
|
||||
assert!(names.contains(&"docker_exec"));
|
||||
// Canvas tool
|
||||
assert!(names.contains(&"canvas_present"));
|
||||
// 3 skill introspection tools (issue #1038)
|
||||
assert!(names.contains(&"skill_list"));
|
||||
assert!(names.contains(&"skill_describe"));
|
||||
assert!(names.contains(&"skill_execute"));
|
||||
}
|
||||
|
||||
/// Issue #1038: skill_list, skill_describe, skill_execute work without
|
||||
/// touching the filesystem so global skills (outside the workspace
|
||||
/// sandbox) are reachable by the agent.
|
||||
#[tokio::test]
|
||||
async fn test_skill_tools_no_filesystem_access() {
|
||||
use openfang_skills::registry::SkillRegistry;
|
||||
use tempfile::TempDir;
|
||||
|
||||
// Build a skills directory containing one prompt-only SKILL.md skill
|
||||
// (mirroring the user's daily-journal scenario from #1038).
|
||||
let global_dir = TempDir::new().unwrap();
|
||||
let skill_dir = global_dir.path().join("daily-journal");
|
||||
std::fs::create_dir_all(&skill_dir).unwrap();
|
||||
std::fs::write(
|
||||
skill_dir.join("SKILL.md"),
|
||||
"---\nname: daily-journal\ndescription: Keep a daily journal\n---\n\
|
||||
# Daily Journal\n\nWrite one paragraph per day about what you learned.",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let mut registry = SkillRegistry::new(global_dir.path().to_path_buf());
|
||||
registry.load_all().unwrap();
|
||||
assert_eq!(registry.count(), 1);
|
||||
|
||||
// skill_list returns the global skill without any filesystem call
|
||||
let list_out = tool_skill_list(Some(®istry)).unwrap();
|
||||
assert!(list_out.contains("daily-journal"));
|
||||
assert!(list_out.contains("Keep a daily journal"));
|
||||
|
||||
// skill_describe returns the SKILL.md body — no file_read needed
|
||||
let desc_out = tool_skill_describe(
|
||||
&serde_json::json!({ "name": "daily-journal" }),
|
||||
Some(®istry),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(desc_out.contains("Daily Journal"));
|
||||
assert!(desc_out.contains("Write one paragraph"));
|
||||
|
||||
// skill_execute on a prompt-only skill returns the body in 'prompt_context' mode
|
||||
let exec_out = tool_skill_execute(
|
||||
&serde_json::json!({ "skill": "daily-journal" }),
|
||||
Some(®istry),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(exec_out.contains("prompt_context"));
|
||||
assert!(exec_out.contains("Daily Journal"));
|
||||
|
||||
// skill_describe on a missing skill returns a helpful error
|
||||
let missing = tool_skill_describe(
|
||||
&serde_json::json!({ "name": "no-such-skill" }),
|
||||
Some(®istry),
|
||||
);
|
||||
assert!(missing.is_err());
|
||||
assert!(missing.unwrap_err().contains("not found"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -3584,6 +3986,98 @@ mod tests {
|
||||
assert!(result.content.contains("traversal"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_directory_path_traversal_blocked() {
|
||||
let result = execute_tool(
|
||||
"test-id",
|
||||
"create_directory",
|
||||
&serde_json::json!({"path": "../../etc/evil"}),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None, // media_engine
|
||||
None, // exec_policy
|
||||
None, // tts_engine
|
||||
None, // docker_config
|
||||
None, // process_manager
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_error);
|
||||
assert!(result.content.contains("traversal"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_directory_creates_nested() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let root = tmp.path();
|
||||
let result = tool_create_directory(&serde_json::json!({"path": "a/b/c"}), Some(root)).await;
|
||||
assert!(result.is_ok(), "Expected Ok, got: {:?}", result);
|
||||
let expected = root.join("a").join("b").join("c");
|
||||
assert!(
|
||||
expected.is_dir(),
|
||||
"Expected directory to exist: {}",
|
||||
expected.display()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_directory_idempotent() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let root = tmp.path();
|
||||
// First create
|
||||
let r1 = tool_create_directory(&serde_json::json!({"path": "data/logs"}), Some(root)).await;
|
||||
assert!(r1.is_ok());
|
||||
// Second create on existing dir should also succeed
|
||||
let r2 = tool_create_directory(&serde_json::json!({"path": "data/logs"}), Some(root)).await;
|
||||
assert!(r2.is_ok(), "Expected idempotent success, got: {:?}", r2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_directory_missing_path_param() {
|
||||
let result = tool_create_directory(&serde_json::json!({}), None).await;
|
||||
assert!(result.is_err());
|
||||
let msg = result.unwrap_err();
|
||||
assert!(msg.contains("Missing 'path'"), "got: {msg}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_directory_dispatch_via_execute_tool() {
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let root = tmp.path().to_path_buf();
|
||||
let result = execute_tool(
|
||||
"test-id",
|
||||
"create_directory",
|
||||
&serde_json::json!({"path": "nested/folder"}),
|
||||
None, // kernel
|
||||
None, // allowed_tools
|
||||
None, // caller_agent_id
|
||||
None, // skill_registry
|
||||
None, // mcp_connections
|
||||
None, // web_ctx
|
||||
None, // browser_ctx
|
||||
None, // allowed_env_vars
|
||||
Some(root.as_path()), // workspace_root
|
||||
None, // media_engine
|
||||
None, // exec_policy
|
||||
None, // tts_engine
|
||||
None, // docker_config
|
||||
None, // process_manager
|
||||
)
|
||||
.await;
|
||||
assert!(
|
||||
!result.is_error,
|
||||
"Expected success, got: {}",
|
||||
result.content
|
||||
);
|
||||
assert!(root.join("nested").join("folder").is_dir());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_file_list_path_traversal_blocked() {
|
||||
let result = execute_tool(
|
||||
|
||||
@@ -19,11 +19,38 @@ pub struct TtsResult {
|
||||
/// Text-to-speech engine.
|
||||
pub struct TtsEngine {
|
||||
config: TtsConfig,
|
||||
/// Optional override for OpenAI TTS base URL. When set, the engine POSTs
|
||||
/// to `<openai_base_url>/v1/audio/speech` instead of the hardcoded
|
||||
/// `https://api.openai.com/v1/audio/speech`. Sourced from
|
||||
/// `MediaConfig.tts_openai_base_url`. Closes #1051.
|
||||
openai_base_url: Option<String>,
|
||||
/// Optional override for ElevenLabs TTS base URL. When set, the engine
|
||||
/// POSTs to `<elevenlabs_base_url>/v1/text-to-speech/{voice_id}` instead
|
||||
/// of the hardcoded `https://api.elevenlabs.io/...`. Sourced from
|
||||
/// `MediaConfig.tts_elevenlabs_base_url`. Closes #1051.
|
||||
elevenlabs_base_url: Option<String>,
|
||||
}
|
||||
|
||||
impl TtsEngine {
|
||||
pub fn new(config: TtsConfig) -> Self {
|
||||
Self { config }
|
||||
Self {
|
||||
config,
|
||||
openai_base_url: None,
|
||||
elevenlabs_base_url: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Attach optional base-URL overrides from `MediaConfig`. Use this to
|
||||
/// route TTS calls at a local OpenAI-compatible service (e.g.
|
||||
/// Lemonade/Kokoro, LM Studio) or an ElevenLabs proxy. Closes #1051.
|
||||
pub fn with_base_urls(
|
||||
mut self,
|
||||
openai_base_url: Option<String>,
|
||||
elevenlabs_base_url: Option<String>,
|
||||
) -> Self {
|
||||
self.openai_base_url = openai_base_url;
|
||||
self.elevenlabs_base_url = elevenlabs_base_url;
|
||||
self
|
||||
}
|
||||
|
||||
/// Detect which TTS provider is available based on environment variables.
|
||||
@@ -100,9 +127,21 @@ impl TtsEngine {
|
||||
"speed": self.config.openai.speed,
|
||||
});
|
||||
|
||||
// `tts_openai_base_url` (config.media.tts_openai_base_url) overrides
|
||||
// the hardcoded provider URL when set, allowing the same OpenAI-compat
|
||||
// JSON wire format to be sent to a local TTS service (Lemonade/Kokoro,
|
||||
// LM Studio, etc.) instead of the cloud provider. The Authorization
|
||||
// header is still built from `OPENAI_API_KEY`; local services typically
|
||||
// accept any non-empty bearer token. Closes #1051.
|
||||
let url = self
|
||||
.openai_base_url
|
||||
.as_deref()
|
||||
.map(|base| format!("{}/v1/audio/speech", base.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| "https://api.openai.com/v1/audio/speech".to_string());
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let response = client
|
||||
.post("https://api.openai.com/v1/audio/speech")
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {}", api_key))
|
||||
.header("Content-Type", "application/json")
|
||||
.json(&body)
|
||||
@@ -161,7 +200,17 @@ impl TtsEngine {
|
||||
std::env::var("ELEVENLABS_API_KEY").map_err(|_| "ELEVENLABS_API_KEY not set")?;
|
||||
|
||||
let voice_id = voice_override.unwrap_or(&self.config.elevenlabs.voice_id);
|
||||
let url = format!("https://api.elevenlabs.io/v1/text-to-speech/{}", voice_id);
|
||||
// `tts_elevenlabs_base_url` (config.media.tts_elevenlabs_base_url)
|
||||
// overrides the hardcoded provider URL when set, allowing the same
|
||||
// ElevenLabs JSON wire format to be routed through a proxy or
|
||||
// self-hosted ElevenLabs-compatible gateway. The `xi-api-key` header
|
||||
// still comes from `ELEVENLABS_API_KEY`. Closes #1051.
|
||||
let base = self
|
||||
.elevenlabs_base_url
|
||||
.as_deref()
|
||||
.map(|b| b.trim_end_matches('/').to_string())
|
||||
.unwrap_or_else(|| "https://api.elevenlabs.io".to_string());
|
||||
let url = format!("{}/v1/text-to-speech/{}", base, voice_id);
|
||||
|
||||
let body = serde_json::json!({
|
||||
"text": text,
|
||||
@@ -306,4 +355,88 @@ mod tests {
|
||||
fn test_max_audio_constant() {
|
||||
assert_eq!(MAX_AUDIO_RESPONSE_BYTES, 10 * 1024 * 1024);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_with_base_urls_sets_overrides() {
|
||||
let engine = TtsEngine::new(default_config()).with_base_urls(
|
||||
Some("http://127.0.0.1:8000".to_string()),
|
||||
Some("http://127.0.0.1:9000".to_string()),
|
||||
);
|
||||
assert_eq!(
|
||||
engine.openai_base_url.as_deref(),
|
||||
Some("http://127.0.0.1:8000")
|
||||
);
|
||||
assert_eq!(
|
||||
engine.elevenlabs_base_url.as_deref(),
|
||||
Some("http://127.0.0.1:9000")
|
||||
);
|
||||
}
|
||||
|
||||
/// Closes #1051: when the OpenAI TTS base URL is overridden, the URL
|
||||
/// building logic must append `/v1/audio/speech` and strip any trailing
|
||||
/// slash. When unset, the hardcoded provider URL is used.
|
||||
#[test]
|
||||
fn test_tts_openai_base_url_override_logic() {
|
||||
// Helper mirroring the URL construction in `synthesize_openai`.
|
||||
fn build(base: Option<&str>) -> String {
|
||||
base.map(|b| format!("{}/v1/audio/speech", b.trim_end_matches('/')))
|
||||
.unwrap_or_else(|| "https://api.openai.com/v1/audio/speech".to_string())
|
||||
}
|
||||
|
||||
// Default: hardcoded URL preserved (backward compatibility).
|
||||
assert_eq!(build(None), "https://api.openai.com/v1/audio/speech");
|
||||
|
||||
// Override applied.
|
||||
assert_eq!(
|
||||
build(Some("http://127.0.0.1:8000")),
|
||||
"http://127.0.0.1:8000/v1/audio/speech"
|
||||
);
|
||||
|
||||
// Trailing slash on the user-supplied base is stripped.
|
||||
assert_eq!(
|
||||
build(Some("http://127.0.0.1:8000/")),
|
||||
"http://127.0.0.1:8000/v1/audio/speech"
|
||||
);
|
||||
assert_eq!(
|
||||
build(Some("https://tts.example.com/")),
|
||||
"https://tts.example.com/v1/audio/speech"
|
||||
);
|
||||
}
|
||||
|
||||
/// Closes #1051: when the ElevenLabs TTS base URL is overridden, the URL
|
||||
/// building logic must append `/v1/text-to-speech/{voice_id}` and strip
|
||||
/// any trailing slash. When unset, the hardcoded provider URL is used.
|
||||
#[test]
|
||||
fn test_tts_elevenlabs_base_url_override_logic() {
|
||||
fn build(base: Option<&str>, voice_id: &str) -> String {
|
||||
let b = base
|
||||
.map(|b| b.trim_end_matches('/').to_string())
|
||||
.unwrap_or_else(|| "https://api.elevenlabs.io".to_string());
|
||||
format!("{}/v1/text-to-speech/{}", b, voice_id)
|
||||
}
|
||||
|
||||
let voice = "21m00Tcm4TlvDq8ikWAM";
|
||||
|
||||
// Default: hardcoded URL preserved.
|
||||
assert_eq!(
|
||||
build(None, voice),
|
||||
format!("https://api.elevenlabs.io/v1/text-to-speech/{voice}")
|
||||
);
|
||||
|
||||
// Override applied.
|
||||
assert_eq!(
|
||||
build(Some("http://127.0.0.1:9000"), voice),
|
||||
format!("http://127.0.0.1:9000/v1/text-to-speech/{voice}")
|
||||
);
|
||||
|
||||
// Trailing slash stripped.
|
||||
assert_eq!(
|
||||
build(Some("http://127.0.0.1:9000/"), voice),
|
||||
format!("http://127.0.0.1:9000/v1/text-to-speech/{voice}")
|
||||
);
|
||||
assert_eq!(
|
||||
build(Some("https://eleven.example.com/"), voice),
|
||||
format!("https://eleven.example.com/v1/text-to-speech/{voice}")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -366,7 +366,10 @@ fn is_private_ip(ip: &IpAddr) -> bool {
|
||||
}
|
||||
|
||||
/// Extract host:port from a URL.
|
||||
fn extract_host(url: &str) -> String {
|
||||
///
|
||||
/// Handles IPv6 bracket notation (`[::1]:8080`), and infers default
|
||||
/// ports (80 for HTTP, 443 for HTTPS) when no explicit port is given.
|
||||
pub(crate) fn extract_host(url: &str) -> String {
|
||||
if let Some(after_scheme) = url.split("://").nth(1) {
|
||||
let host_port = after_scheme.split('/').next().unwrap_or(after_scheme);
|
||||
// Handle IPv6 bracket notation: [::1]:8080
|
||||
|
||||
@@ -25,3 +25,5 @@ zip = { workspace = true }
|
||||
[dev-dependencies]
|
||||
tempfile = { workspace = true }
|
||||
tokio-test = { workspace = true }
|
||||
ed25519-dalek = { workspace = true }
|
||||
rand = { workspace = true }
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user