diff --git a/Cargo.lock b/Cargo.lock index 84f2be6..52d719c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -119,9 +119,10 @@ dependencies = [ "opentelemetry_sdk", "percent-encoding", "prometheus", + "ra2a", "regex", "regorus", - "reqwest", + "reqwest 0.12.28", "rusqlite", "rustls", "rustls-pemfile", @@ -479,6 +480,12 @@ dependencies = [ "shlex", ] +[[package]] +name = "cesu8" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6d43a04d8753f35258c91f8ec639f792891f748a1edbd759cf1dcea3382ad83c" + [[package]] name = "cfg-if" version = "1.0.4" @@ -511,6 +518,7 @@ dependencies = [ "iana-time-zone", "js-sys", "num-traits", + "serde", "wasm-bindgen", "windows-link", ] @@ -580,6 +588,16 @@ version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" +[[package]] +name = "combine" +version = "4.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba5a308b75df32fe02788e748662718f03fde005016435c444eea572398219fd" +dependencies = [ + "bytes", + "memchr", +] + [[package]] name = "const-oid" version = "0.9.6" @@ -612,6 +630,26 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c74b8349d32d297c9134b8c88677813a227df8f779daa29bfc29c183fe3dca6" +[[package]] +name = "core-foundation" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91e195e091a93c46f7102ec7818a2aa394e1e1771c3ab4825963fa03e45afb8f" +dependencies = [ + "core-foundation-sys", + "libc", +] + +[[package]] +name = "core-foundation" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2a6cd9ae233e7f62ba4e9353e81a88df7fc8a5987b8d445b4d90c879bd156f6" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "core-foundation-sys" version = "0.8.7" @@ -828,6 +866,15 @@ dependencies = [ "serde", ] +[[package]] +name = "encoding_rs" +version = "0.8.35" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75030f3c4f45dafd7586dd6780965a8c7e8e285a5ecb86713e63a79c5b2766f3" +dependencies = [ + "cfg-if", +] + [[package]] name = "env_home" version = "0.1.0" @@ -976,6 +1023,21 @@ version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" +[[package]] +name = "futures" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +dependencies = [ + "futures-channel", + "futures-core", + "futures-executor", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + [[package]] name = "futures-channel" version = "0.3.32" @@ -983,6 +1045,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" dependencies = [ "futures-core", + "futures-sink", ] [[package]] @@ -1037,6 +1100,7 @@ version = "0.3.32" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" dependencies = [ + "futures-channel", "futures-core", "futures-io", "futures-macro", @@ -1325,9 +1389,11 @@ dependencies = [ "percent-encoding", "pin-project-lite", "socket2 0.6.3", + "system-configuration", "tokio", "tower-service", "tracing", + "windows-registry", ] [[package]] @@ -1548,6 +1614,50 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jni" +version = "0.21.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1a87aa2bb7d2af34197c04845522473242e1aa17c12f4935d5856491a7fb8c97" +dependencies = [ + "cesu8", + "cfg-if", + "combine", + "jni-sys 0.3.1", + "log", + "thiserror 1.0.69", + "walkdir", + "windows-sys 0.45.0", +] + +[[package]] +name = "jni-sys" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41a652e1f9b6e0275df1f15b32661cf0d4b78d4d87ddec5e0c3c20f097433258" +dependencies = [ + "jni-sys 0.4.1", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn", +] + [[package]] name = "jobserver" version = "0.1.34" @@ -1783,6 +1893,16 @@ version = "0.3.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" +[[package]] +name = "mime_guess" +version = "2.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f7c44f8e672c00fe5308fa235f821cb4198414e1c77935c1ab6948d3fd78550e" +dependencies = [ + "mime", + "unicase", +] + [[package]] name = "minimal-lexical" version = "0.2.1" @@ -1951,6 +2071,12 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" +[[package]] +name = "openssl-probe" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" + [[package]] name = "opentelemetry" version = "0.26.0" @@ -2328,6 +2454,7 @@ version = "0.11.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "434b42fec591c96ef50e21e886936e66d3cc3f737104fdb9b737c40ffb94c098" dependencies = [ + "aws-lc-rs", "bytes", "getrandom 0.3.4", "lru-slab", @@ -2378,6 +2505,25 @@ version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "ra2a" +version = "0.9.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "17b68f1bed65f88d09cbdde48311f2f1aaa9d6ce761acfd4dde176df49f30202" +dependencies = [ + "axum 0.8.8", + "base64", + "chrono", + "futures", + "reqwest 0.13.2", + "serde", + "serde_json", + "thiserror 2.0.18", + "tokio", + "tracing", + "uuid", +] + [[package]] name = "rand" version = "0.8.5" @@ -2593,11 +2739,55 @@ dependencies = [ "url", "wasm-bindgen", "wasm-bindgen-futures", - "wasm-streams", + "wasm-streams 0.4.2", "web-sys", "webpki-roots", ] +[[package]] +name = "reqwest" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab3f43e3283ab1488b624b44b0e988d0acea0b3214e694730a055cb6b2efa801" +dependencies = [ + "base64", + "bytes", + "encoding_rs", + "futures-core", + "futures-util", + "h2", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-rustls", + "hyper-util", + "js-sys", + "log", + "mime", + "mime_guess", + "percent-encoding", + "pin-project-lite", + "quinn", + "rustls", + "rustls-pki-types", + "rustls-platform-verifier", + "serde", + "serde_json", + "sync_wrapper", + "tokio", + "tokio-rustls", + "tokio-util", + "tower 0.5.3", + "tower-http", + "tower-service", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "wasm-streams 0.5.0", + "web-sys", +] + [[package]] name = "rfc6979" version = "0.4.0" @@ -2709,6 +2899,18 @@ dependencies = [ "zeroize", ] +[[package]] +name = "rustls-native-certs" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "612460d5f7bea540c490b2b6395d8e34a953e52b491accd6c86c8164c5932a63" +dependencies = [ + "openssl-probe", + "rustls-pki-types", + "schannel", + "security-framework", +] + [[package]] name = "rustls-pemfile" version = "2.2.0" @@ -2728,6 +2930,33 @@ dependencies = [ "zeroize", ] +[[package]] +name = "rustls-platform-verifier" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d99feebc72bae7ab76ba994bb5e121b8d83d910ca40b36e0921f53becc41784" +dependencies = [ + "core-foundation 0.10.1", + "core-foundation-sys", + "jni", + "log", + "once_cell", + "rustls", + "rustls-native-certs", + "rustls-platform-verifier-android", + "rustls-webpki", + "security-framework", + "security-framework-sys", + "webpki-root-certs", + "windows-sys 0.61.2", +] + +[[package]] +name = "rustls-platform-verifier-android" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f" + [[package]] name = "rustls-webpki" version = "0.103.10" @@ -2752,6 +2981,24 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + +[[package]] +name = "schannel" +version = "0.1.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91c1b7e4904c873ef0710c1f407dde2e6287de2bebc1bbbf7d430bb7cbffd939" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "scientific" version = "0.5.3" @@ -2792,6 +3039,29 @@ dependencies = [ "zeroize", ] +[[package]] +name = "security-framework" +version = "3.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" +dependencies = [ + "bitflags", + "core-foundation 0.10.1", + "core-foundation-sys", + "libc", + "security-framework-sys", +] + +[[package]] +name = "security-framework-sys" +version = "2.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2691df843ecc5d231c0b14ece2acc3efb62c0a398c7e1d875f3983ce020e3" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "semver" version = "1.0.27" @@ -3049,6 +3319,27 @@ dependencies = [ "syn", ] +[[package]] +name = "system-configuration" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a13f3d0daba03132c0aa9767f98351b3488edc2c100cda2d2ec2b04f3d8d3c8b" +dependencies = [ + "bitflags", + "core-foundation 0.9.4", + "system-configuration-sys", +] + +[[package]] +name = "system-configuration-sys" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e1d1b10ced5ca923a1fcb8d03e96b8d3268065d724548c0211415ff6ac6bac4" +dependencies = [ + "core-foundation-sys", + "libc", +] + [[package]] name = "tempfile" version = "3.27.0" @@ -3430,6 +3721,12 @@ version = "1.19.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb" +[[package]] +name = "unicase" +version = "2.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142" + [[package]] name = "unicode-ident" version = "1.0.24" @@ -3496,6 +3793,7 @@ dependencies = [ "getrandom 0.4.2", "js-sys", "rand 0.10.0", + "serde_core", "wasm-bindgen", ] @@ -3534,6 +3832,16 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5c3082ca00d5a5ef149bb8b555a72ae84c9c59f7250f013ac822ac2e49b19c64" +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + [[package]] name = "want" version = "0.3.1" @@ -3657,6 +3965,19 @@ dependencies = [ "web-sys", ] +[[package]] +name = "wasm-streams" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1ec4f6517c9e11ae630e200b2b65d193279042e28edd4a2cda233e46670bbb" +dependencies = [ + "futures-util", + "js-sys", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + [[package]] name = "wasmparser" version = "0.244.0" @@ -3703,6 +4024,15 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "webpki-root-certs" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "804f18a4ac2676ffb4e8b5b5fa9ae38af06df08162314f96a68d2a363e21a8ca" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "webpki-roots" version = "1.0.3" @@ -3724,6 +4054,15 @@ dependencies = [ "winsafe", ] +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "windows-core" version = "0.62.2" @@ -3765,6 +4104,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-registry" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720" +dependencies = [ + "windows-link", + "windows-result", + "windows-strings", +] + [[package]] name = "windows-result" version = "0.4.1" @@ -3783,13 +4133,22 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-sys" +version = "0.45.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75283be5efb2831d37ea142365f009c02ec203cd29a3ebecbc093d52315b66d0" +dependencies = [ + "windows-targets 0.42.2", +] + [[package]] name = "windows-sys" version = "0.52.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" dependencies = [ - "windows-targets", + "windows-targets 0.52.6", ] [[package]] @@ -3801,34 +4160,67 @@ dependencies = [ "windows-link", ] +[[package]] +name = "windows-targets" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e5180c00cd44c9b1c88adb3693291f1cd93605ded80c250a75d472756b4d071" +dependencies = [ + "windows_aarch64_gnullvm 0.42.2", + "windows_aarch64_msvc 0.42.2", + "windows_i686_gnu 0.42.2", + "windows_i686_msvc 0.42.2", + "windows_x86_64_gnu 0.42.2", + "windows_x86_64_gnullvm 0.42.2", + "windows_x86_64_msvc 0.42.2", +] + [[package]] name = "windows-targets" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" dependencies = [ - "windows_aarch64_gnullvm", - "windows_aarch64_msvc", - "windows_i686_gnu", + "windows_aarch64_gnullvm 0.52.6", + "windows_aarch64_msvc 0.52.6", + "windows_i686_gnu 0.52.6", "windows_i686_gnullvm", - "windows_i686_msvc", - "windows_x86_64_gnu", - "windows_x86_64_gnullvm", - "windows_x86_64_msvc", + "windows_i686_msvc 0.52.6", + "windows_x86_64_gnu 0.52.6", + "windows_x86_64_gnullvm 0.52.6", + "windows_x86_64_msvc 0.52.6", ] +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "597a5118570b68bc08d8d59125332c54f1ba9d9adeedeef5b99b02ba2b0698f8" + [[package]] name = "windows_aarch64_gnullvm" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" +[[package]] +name = "windows_aarch64_msvc" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e08e8864a60f06ef0d0ff4ba04124db8b0fb3be5776a5cd47641e942e58c4d43" + [[package]] name = "windows_aarch64_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" +[[package]] +name = "windows_i686_gnu" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c61d927d8da41da96a81f029489353e68739737d3beca43145c8afec9a31a84f" + [[package]] name = "windows_i686_gnu" version = "0.52.6" @@ -3841,24 +4233,48 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" +[[package]] +name = "windows_i686_msvc" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "44d840b6ec649f480a41c8d80f9c65108b92d89345dd94027bfe06ac444d1060" + [[package]] name = "windows_i686_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" +[[package]] +name = "windows_x86_64_gnu" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8de912b8b8feb55c064867cf047dda097f92d51efad5b491dfb98f6bbb70cb36" + [[package]] name = "windows_x86_64_gnu" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26d41b46a36d453748aedef1486d5c7a85db22e56aff34643984ea85514e94a3" + [[package]] name = "windows_x86_64_gnullvm" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" +[[package]] +name = "windows_x86_64_msvc" +version = "0.42.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9aec5da331524158c6d1a4ac0ab1541149c0b9505fde06423b02f5ef0106b9f0" + [[package]] name = "windows_x86_64_msvc" version = "0.52.6" diff --git a/Cargo.toml b/Cargo.toml index 1702054..5765d1a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -46,6 +46,7 @@ opentelemetry-otlp = { version = "0.26", features = ["trace", "metrics", "logs", opentelemetry-appender-tracing = "0.26" jsonwebtoken = { version = "10", features = ["rust_crypto"] } subtle = "2" +ra2a = { version = "0.9.3", features = ["server"] } chrono = { version = "0.4", default-features = false, features = ["clock"] } base64 = "0.22" percent-encoding = "2" diff --git a/src/a2a/executor.rs b/src/a2a/executor.rs new file mode 100644 index 0000000..0305d51 --- /dev/null +++ b/src/a2a/executor.rs @@ -0,0 +1,143 @@ +//! A2A proxy executor — forwards A2A messages to an upstream agent. + +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +use ra2a::{ + error::{A2AError, Result}, + server::{AgentExecutor, EventQueue, RequestContext}, + types::{SendMessageRequest, SendMessageResponse, StreamResponse, Task, TaskState, TaskStatus}, +}; +use reqwest::Client; +use serde_json::json; + +/// Proxies incoming A2A `message/send` requests to an upstream A2A agent. +/// +/// The `A2aPolicyInterceptor` runs before this executor and enforces rate limits, +/// API key authentication, and payload filtering. This executor only handles the +/// mechanical proxy step: forward the message to the upstream and write the +/// response to the event queue. +pub struct A2aProxyExecutor { + upstream_url: String, + client: Arc, +} + +impl A2aProxyExecutor { + /// Creates a new proxy executor pointing at the given upstream A2A JSON-RPC endpoint. + /// + /// `upstream_url` should be the base URL without trailing slash + /// (e.g., `"http://localhost:4001"`). The executor will POST to this URL. + pub fn new(upstream_url: impl Into) -> anyhow::Result { + let client = Client::builder() + .timeout(std::time::Duration::from_secs(60)) + .build() + .map_err(|e| anyhow::anyhow!("failed to build A2A proxy client: {e}"))?; + + Ok(Self { + upstream_url: upstream_url.into(), + client: Arc::new(client), + }) + } +} + +impl AgentExecutor for A2aProxyExecutor { + fn execute<'a>( + &'a self, + ctx: &'a RequestContext, + queue: &'a EventQueue, + ) -> Pin> + Send + 'a>> { + Box::pin(async move { + let message = ctx + .message + .clone() + .ok_or_else(|| A2AError::InvalidParams("no message in request context".into()))?; + + let send_req = SendMessageRequest::new(message); + + // Build the JSON-RPC 2.0 request envelope. + let rpc_body = json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "message/send", + "params": send_req, + }); + + let resp = self + .client + .post(&self.upstream_url) + .json(&rpc_body) + .send() + .await + .map_err(|e| A2AError::Other(format!("upstream A2A request failed: {e}")))?; + + if !resp.status().is_success() { + let status = resp.status(); + return Err(A2AError::Other(format!( + "upstream A2A agent returned HTTP {status}" + ))); + } + + let body: serde_json::Value = resp + .json() + .await + .map_err(|e| A2AError::Other(format!("failed to parse upstream response: {e}")))?; + + // Propagate JSON-RPC errors from the upstream agent. + if let Some(error) = body.get("error") { + let code = error["code"].as_i64().unwrap_or(-32603) as i32; + let message_str = error["message"] + .as_str() + .unwrap_or("upstream agent error") + .to_string(); + return Err(A2AError::JsonRpc(ra2a::error::JsonRpcError { + code, + message: message_str, + data: None, + })); + } + + let result = body + .get("result") + .ok_or_else(|| A2AError::Other("upstream response missing 'result'".into()))?; + + let upstream_response: SendMessageResponse = serde_json::from_value(result.clone()) + .map_err(|e| { + A2AError::Other(format!("failed to parse upstream SendMessageResponse: {e}")) + })?; + + let event = match upstream_response { + SendMessageResponse::Message(m) => StreamResponse::Message(m), + SendMessageResponse::Task(t) => StreamResponse::Task(t), + }; + + queue.send(event)?; + + Ok(()) + }) + } + + fn cancel<'a>( + &'a self, + ctx: &'a RequestContext, + queue: &'a EventQueue, + ) -> Pin> + Send + 'a>> { + Box::pin(async move { + let mut task = Task::new(&ctx.task_id, &ctx.context_id); + task.status = TaskStatus::new(TaskState::Canceled); + queue.send(StreamResponse::Task(task))?; + Ok(()) + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn executor_constructs_without_error() { + let result = A2aProxyExecutor::new("http://localhost:4001"); + assert!(result.is_ok()); + } +} diff --git a/src/a2a/interceptor.rs b/src/a2a/interceptor.rs new file mode 100644 index 0000000..ebb23d4 --- /dev/null +++ b/src/a2a/interceptor.rs @@ -0,0 +1,386 @@ +//! A2A policy interceptor — enforces Arbitus agent policies on A2A requests. +//! +//! This interceptor runs before every A2A handler method and applies: +//! 1. Agent identity extraction from the `x-arbitus-agent` header +//! 2. API key authentication if the agent has `api_key` configured +//! 3. Per-agent sliding-window rate limiting +//! 4. Payload filtering for blocked patterns on message text parts + +use std::collections::HashMap; +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use ra2a::{ + error::{A2AError, Result}, + server::{AuthenticatedUser, CallContext, CallInterceptor, Request, Response}, + types::SendMessageRequest, +}; +use tokio::sync::{Mutex, watch}; + +use crate::live_config::LiveConfig; + +/// Header used by callers to declare their agent identity. +pub const AGENT_ID_HEADER: &str = "x-arbitus-agent"; +/// Header used for API key authentication. +pub const API_KEY_HEADER: &str = "x-api-key"; + +/// Per-agent request timestamp store for sliding-window rate limiting. +type RateCounts = Arc>>>; + +/// Enforces Arbitus agent policies on incoming A2A requests. +pub struct A2aPolicyInterceptor { + config: watch::Receiver>, + counts: RateCounts, +} + +impl A2aPolicyInterceptor { + /// Creates a new interceptor backed by the given live config receiver. + pub fn new(config: watch::Receiver>) -> Self { + let counts: RateCounts = Arc::new(Mutex::new(HashMap::new())); + + // Background task: prune stale entries every 5 minutes to prevent unbounded growth. + { + let counts = Arc::clone(&counts); + tokio::spawn(async move { + let mut interval = tokio::time::interval(Duration::from_secs(300)); + interval.tick().await; + loop { + interval.tick().await; + let window = Duration::from_secs(60); + let now = Instant::now(); + let mut m = counts.lock().await; + m.retain(|_, ts: &mut Vec| { + ts.retain(|t| now.duration_since(*t) < window); + !ts.is_empty() + }); + } + }); + } + + Self { config, counts } + } + + /// Returns `true` if the agent has consumed fewer than `limit` requests in the last 60s. + /// Appends the current timestamp to the window when returning `true`. + async fn check_rate_limit(&self, agent_id: &str, limit: usize) -> bool { + let now = Instant::now(); + let window = Duration::from_secs(60); + let mut m = self.counts.lock().await; + let ts = m.entry(agent_id.to_string()).or_default(); + ts.retain(|t| now.duration_since(*t) < window); + if ts.len() >= limit { + return false; + } + ts.push(now); + true + } +} + +impl CallInterceptor for A2aPolicyInterceptor { + fn before<'a>( + &'a self, + ctx: &'a mut CallContext, + req: &'a mut Request, + ) -> Pin> + Send + 'a>> { + Box::pin(async move { + // ── 1. Agent identity ───────────────────────────────────────────── + let agent_id = ctx + .request_meta() + .get(AGENT_ID_HEADER) + .and_then(|v| v.first()) + .cloned() + .ok_or_else(|| { + A2AError::Other(format!( + "missing '{AGENT_ID_HEADER}' header — agent identity required" + )) + })?; + + // Reject suspiciously long agent IDs to prevent log injection. + if agent_id.len() > 128 { + return Err(A2AError::Other("agent ID exceeds maximum length".into())); + } + + // Clone the Arc immediately so we don't hold the + // watch::Ref (RwLockReadGuard) across any .await points. + let cfg: Arc = Arc::clone(&*self.config.borrow()); + + let (rate_limit, expected_api_key, patterns) = { + let policy = match cfg.agents.get(&agent_id) { + Some(p) => p, + None => match cfg.default_policy.as_ref() { + Some(p) => p, + None => { + return Err(A2AError::Other(format!( + "agent '{agent_id}' is not configured" + ))); + } + }, + }; + ( + policy.rate_limit, + policy.api_key.clone(), + Arc::clone(&cfg.block_patterns), + ) + }; + drop(cfg); // release Arc before any await + + // ── 2. API key authentication ───────────────────────────────────── + if let Some(expected_key) = &expected_api_key { + let provided = ctx + .request_meta() + .get(API_KEY_HEADER) + .and_then(|v| v.first()) + .map(String::as_str) + .unwrap_or(""); + + // Constant-time comparison to prevent timing attacks. + use subtle::ConstantTimeEq; + let ok: bool = expected_key.as_bytes().ct_eq(provided.as_bytes()).into(); + if !ok { + return Err(A2AError::Other("invalid or missing API key".into())); + } + } + + // ── 3. Rate limiting ────────────────────────────────────────────── + if !self.check_rate_limit(&agent_id, rate_limit).await { + return Err(A2AError::Other(format!( + "rate limit exceeded for agent '{agent_id}'" + ))); + } + + // ── 4. Payload filtering ────────────────────────────────────────── + // Only `message/send` and `message/stream` carry a message payload. + if let Some(send_req) = req.downcast_ref::() + && !patterns.is_empty() + { + // Collect all text content from the message parts. + let text_content: String = send_req + .message + .parts + .iter() + .filter_map(|p| p.as_text()) + .collect::>() + .join(" "); + + for pattern in patterns.iter() { + if pattern.is_match(&text_content) { + return Err(A2AError::Other( + "request blocked: sensitive data detected".into(), + )); + } + } + } + + // ── 5. Mark request as authenticated ───────────────────────────── + ctx.user = Arc::new(AuthenticatedUser::new(agent_id)); + + Ok(()) + }) + } + + fn after<'a>( + &'a self, + _ctx: &'a CallContext, + _resp: &'a mut Response, + ) -> Pin> + Send + 'a>> { + Box::pin(async { Ok(()) }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::{AgentPolicy, FilterMode}; + use std::collections::HashMap; + use tokio::sync::watch; + + fn policy(rate_limit: usize) -> AgentPolicy { + AgentPolicy { + allowed_tools: None, + denied_tools: vec![], + rate_limit, + tool_rate_limits: HashMap::new(), + upstream: None, + api_key: None, + timeout_secs: None, + approval_required: vec![], + hitl_timeout_secs: 60, + shadow_tools: vec![], + federate: false, + allowed_resources: None, + denied_resources: vec![], + allowed_prompts: None, + denied_prompts: vec![], + mtls_identity: None, + } + } + + fn policy_with_key(rate_limit: usize, key: &str) -> AgentPolicy { + AgentPolicy { + api_key: Some(key.to_string()), + ..policy(rate_limit) + } + } + + fn make_interceptor(agents: HashMap) -> A2aPolicyInterceptor { + let live = Arc::new(LiveConfig::new( + agents, + vec![], + vec![], + None, + FilterMode::Block, + None, + )); + let (_, rx) = watch::channel(live); + A2aPolicyInterceptor::new(rx) + } + + fn make_call_ctx(agent: Option<&str>, api_key: Option<&str>) -> CallContext { + use ra2a::server::RequestMeta; + use std::collections::HashMap; + let mut meta: HashMap> = HashMap::new(); + if let Some(a) = agent { + meta.insert(AGENT_ID_HEADER.to_string(), vec![a.to_string()]); + } + if let Some(k) = api_key { + meta.insert(API_KEY_HEADER.to_string(), vec![k.to_string()]); + } + CallContext::new("message/send", RequestMeta::new(meta)) + } + + fn send_req(text: &str) -> ra2a::types::SendMessageRequest { + use ra2a::types::{Message, MessageId, Part, Role}; + let msg = Message { + message_id: MessageId::from("msg-1"), + role: Role::User, // serialized as "ROLE_USER" by ra2a + parts: vec![Part::text(text)], + task_id: None, + context_id: None, + reference_task_ids: vec![], + metadata: None, + extensions: vec![], + }; + ra2a::types::SendMessageRequest::new(msg) + } + + #[tokio::test] + async fn missing_agent_header_is_rejected() { + let interceptor = make_interceptor(HashMap::new()); + let mut ctx = make_call_ctx(None, None); + let mut req = Request::new(send_req("hello")); + let result = interceptor.before(&mut ctx, &mut req).await; + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("missing")); + } + + #[tokio::test] + async fn unknown_agent_is_rejected() { + let interceptor = make_interceptor(HashMap::new()); + let mut ctx = make_call_ctx(Some("ghost"), None); + let mut req = Request::new(send_req("hello")); + let result = interceptor.before(&mut ctx, &mut req).await; + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("not configured")); + } + + #[tokio::test] + async fn known_agent_without_key_passes() { + let mut agents = HashMap::new(); + agents.insert("cursor".to_string(), policy(60)); + let interceptor = make_interceptor(agents); + let mut ctx = make_call_ctx(Some("cursor"), None); + let mut req = Request::new(send_req("hello")); + assert!(interceptor.before(&mut ctx, &mut req).await.is_ok()); + } + + #[tokio::test] + async fn wrong_api_key_is_rejected() { + let mut agents = HashMap::new(); + agents.insert("secured".to_string(), policy_with_key(60, "secret-key")); + let interceptor = make_interceptor(agents); + let mut ctx = make_call_ctx(Some("secured"), Some("wrong-key")); + let mut req = Request::new(send_req("hello")); + let result = interceptor.before(&mut ctx, &mut req).await; + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("API key")); + } + + #[tokio::test] + async fn correct_api_key_passes() { + let mut agents = HashMap::new(); + agents.insert("secured".to_string(), policy_with_key(60, "secret-key")); + let interceptor = make_interceptor(agents); + let mut ctx = make_call_ctx(Some("secured"), Some("secret-key")); + let mut req = Request::new(send_req("hello")); + assert!(interceptor.before(&mut ctx, &mut req).await.is_ok()); + } + + #[tokio::test] + async fn rate_limit_blocks_when_exceeded() { + let mut agents = HashMap::new(); + agents.insert("limited".to_string(), policy(2)); + let interceptor = make_interceptor(agents); + + for _ in 0..2 { + let mut ctx = make_call_ctx(Some("limited"), None); + let mut req = Request::new(send_req("hello")); + assert!(interceptor.before(&mut ctx, &mut req).await.is_ok()); + } + + let mut ctx = make_call_ctx(Some("limited"), None); + let mut req = Request::new(send_req("hello")); + let result = interceptor.before(&mut ctx, &mut req).await; + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("rate limit")); + } + + #[tokio::test] + async fn blocked_pattern_in_message_is_rejected() { + use crate::live_config::LiveConfig; + use regex::Regex; + + let mut agents = HashMap::new(); + agents.insert("cursor".to_string(), policy(60)); + let live = Arc::new(LiveConfig::new( + agents, + vec![Regex::new("private_key").unwrap()], + vec![], + None, + FilterMode::Block, + None, + )); + let (_, rx) = watch::channel(live); + let interceptor = A2aPolicyInterceptor::new(rx); + + let mut ctx = make_call_ctx(Some("cursor"), None); + let mut req = Request::new(send_req("my private_key=AAABBB")); + let result = interceptor.before(&mut ctx, &mut req).await; + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("sensitive data")); + } + + #[tokio::test] + async fn clean_message_passes_filter() { + use crate::live_config::LiveConfig; + use regex::Regex; + + let mut agents = HashMap::new(); + agents.insert("cursor".to_string(), policy(60)); + let live = Arc::new(LiveConfig::new( + agents, + vec![Regex::new("private_key").unwrap()], + vec![], + None, + FilterMode::Block, + None, + )); + let (_, rx) = watch::channel(live); + let interceptor = A2aPolicyInterceptor::new(rx); + + let mut ctx = make_call_ctx(Some("cursor"), None); + let mut req = Request::new(send_req("hello world, no secrets here")); + assert!(interceptor.before(&mut ctx, &mut req).await.is_ok()); + } +} diff --git a/src/a2a/mod.rs b/src/a2a/mod.rs new file mode 100644 index 0000000..bab6181 --- /dev/null +++ b/src/a2a/mod.rs @@ -0,0 +1,33 @@ +//! A2A (Agent-to-Agent) protocol support. +//! +//! Implements a security proxy that enforces Arbitus policies on A2A protocol +//! requests before forwarding them to an upstream A2A agent. +//! +//! ## Protocol +//! The A2A protocol is a JSON-RPC 2.0 interface for agent-to-agent communication. +//! Arbitus sits between the caller and the upstream agent, enforcing: +//! - Per-agent rate limits (using the `rate_limit` field from agent policy) +//! - API key authentication (using the `api_key` field from agent policy) +//! - Payload filtering for blocked patterns (using `rules.block_patterns`) +//! +//! ## Agent Identity +//! Callers identify themselves via the `x-arbitus-agent` HTTP header. +//! The value must match an agent name in the Arbitus config. +//! +//! ## Configuration +//! ```yaml +//! a2a: +//! upstream: "http://localhost:4001" +//! mount: "/a2a" +//! agent_card: +//! name: "My Agent" +//! description: "Proxied via Arbitus" +//! url: "http://localhost:4000/a2a" +//! version: "1.0.0" +//! ``` + +pub mod executor; +pub mod interceptor; + +pub use executor::A2aProxyExecutor; +pub use interceptor::A2aPolicyInterceptor; diff --git a/src/bin/arbitus.rs b/src/bin/arbitus.rs index 4f99596..089ea38 100644 --- a/src/bin/arbitus.rs +++ b/src/bin/arbitus.rs @@ -1,3 +1,4 @@ +use arbitus::a2a::{A2aPolicyInterceptor, A2aProxyExecutor}; use arbitus::live_config::OpaPolicy; use arbitus::{ audit::{ @@ -365,6 +366,33 @@ async fn cmd_start(config_path: String) -> anyhow::Result<()> { Arc::new(MultiJwtValidator::new(configs)) }); + // ── A2A server state ────────────────────────────────────────────────────── + let a2a_server_state = if let Some(a2a_cfg) = &config.a2a { + let executor = A2aProxyExecutor::new(&a2a_cfg.upstream)?; + let interceptor = std::sync::Arc::new(A2aPolicyInterceptor::new(config_rx.clone())); + let card_cfg = &a2a_cfg.agent_card; + let description = card_cfg.description.clone().unwrap_or_default(); + let endpoint_url = card_cfg + .url + .clone() + .unwrap_or_else(|| a2a_cfg.upstream.clone()); + let interface = ra2a::types::AgentInterface::new( + &endpoint_url, + ra2a::types::TransportProtocol::new(ra2a::types::TransportProtocol::JSONRPC), + ); + let mut card = ra2a::types::AgentCard::new(&card_cfg.name, description, vec![interface]); + card.version = card_cfg.version.clone(); + card.documentation_url = Some(endpoint_url); + let handler = ra2a::server::HandlerBuilder::new(executor, card.clone()) + .with_call_interceptor(interceptor) + .build(); + let state = ra2a::server::ServerState::new(std::sync::Arc::new(handler), card); + tracing::info!(upstream = %a2a_cfg.upstream, mount = %a2a_cfg.mount, "A2A protocol enabled"); + Some((a2a_cfg.mount.clone(), state)) + } else { + None + }; + match config.transport { TransportConfig::Http { addr, @@ -388,7 +416,7 @@ async fn cmd_start(config_path: String) -> anyhow::Result<()> { config_rx.clone(), schema_cache, )); - HttpTransport::new( + let mut transport = HttpTransport::new( addr, session_ttl_secs, tls, @@ -399,9 +427,11 @@ async fn cmd_start(config_path: String) -> anyhow::Result<()> { config.admin_token, hitl_store, oauth_manager, - ) - .serve(gateway) - .await?; + ); + if let Some((mount, state)) = a2a_server_state { + transport = transport.with_a2a(mount, state); + } + transport.serve(gateway).await?; } TransportConfig::StreamableHttp { addr, @@ -425,7 +455,7 @@ async fn cmd_start(config_path: String) -> anyhow::Result<()> { config_rx.clone(), schema_cache, )); - StreamableHttpTransport::new( + let mut transport = StreamableHttpTransport::new( addr, session_ttl_secs, tls, @@ -436,9 +466,11 @@ async fn cmd_start(config_path: String) -> anyhow::Result<()> { config.admin_token, hitl_store, oauth_manager, - ) - .serve(gateway) - .await?; + ); + if let Some((mount, state)) = a2a_server_state { + transport = transport.with_a2a(mount, state); + } + transport.serve(gateway).await?; } TransportConfig::Stdio { server, verify } => { tracing::info!(server = %server.join(" "), "stdio mode"); diff --git a/src/config.rs b/src/config.rs index 5f1dabb..184ac65 100644 --- a/src/config.rs +++ b/src/config.rs @@ -34,6 +34,61 @@ pub struct Config { /// them into the config before the gateway starts. #[serde(default)] pub secrets: Option, + /// A2A (Agent-to-Agent) protocol endpoint. When present, Arbitus mounts an A2A + /// JSON-RPC proxy at `a2a.mount` and enforces per-agent policies on incoming requests. + #[serde(default)] + pub a2a: Option, +} + +// ── A2A ────────────────────────────────────────────────────────────────────── + +/// Configuration for the A2A (Agent-to-Agent) protocol endpoint. +#[derive(Debug, Deserialize, Clone)] +pub struct A2aConfig { + /// URL of the upstream A2A agent (e.g. `"http://localhost:4001"`). + pub upstream: String, + /// Path at which to mount the A2A router. Default: `/a2a`. + #[serde(default = "default_a2a_mount")] + pub mount: String, + /// Agent card metadata served at `/.well-known/agent.json`. + #[serde(default)] + pub agent_card: A2aAgentCardConfig, +} + +/// Metadata for the A2A agent card. +#[derive(Debug, Deserialize, Clone)] +pub struct A2aAgentCardConfig { + #[serde(default = "default_a2a_name")] + pub name: String, + #[serde(default)] + pub description: Option, + #[serde(default)] + pub url: Option, + #[serde(default = "default_a2a_version")] + pub version: String, +} + +impl Default for A2aAgentCardConfig { + fn default() -> Self { + Self { + name: default_a2a_name(), + description: None, + url: None, + version: default_a2a_version(), + } + } +} + +fn default_a2a_mount() -> String { + "/a2a".to_string() +} + +fn default_a2a_name() -> String { + "Arbitus A2A Proxy".to_string() +} + +fn default_a2a_version() -> String { + "1.0.0".to_string() } // ── Transport ──────────────────────────────────────────────────────────────── @@ -819,6 +874,7 @@ mod tests { admin_token: None, telemetry: None, secrets: None, + a2a: None, } } diff --git a/src/lib.rs b/src/lib.rs index fca952c..5be26c1 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,4 @@ +pub mod a2a; pub mod audit; pub mod config; pub mod cost; diff --git a/src/transport/http.rs b/src/transport/http.rs index 5a29ec2..93bb951 100644 --- a/src/transport/http.rs +++ b/src/transport/http.rs @@ -107,6 +107,8 @@ pub struct HttpTransport { /// Operator kill switch — tool names in this set are immediately blocked /// regardless of agent policy. Managed via the dashboard UI. kill_switch: Arc>>, + /// Optional A2A server state: (mount_path, server_state). + a2a: Option<(String, ra2a::server::ServerState)>, } impl HttpTransport { @@ -135,8 +137,20 @@ impl HttpTransport { hitl_store, oauth_manager, kill_switch: Arc::new(std::sync::Mutex::new(std::collections::HashSet::new())), + a2a: None, } } + + /// Attach an A2A server state to this transport. + /// When set, the A2A router is mounted at `mount_path` alongside the MCP routes. + pub fn with_a2a( + mut self, + mount_path: impl Into, + state: ra2a::server::ServerState, + ) -> Self { + self.a2a = Some((mount_path.into(), state)); + self + } } struct HttpState { @@ -172,7 +186,7 @@ impl Transport for HttpTransport { kill_switch: Arc::clone(&self.kill_switch), }); - let app = Router::new() + let mut app = Router::new() .route("/mcp", post(handle_mcp)) .route("/mcp", get(handle_sse)) .route("/mcp", delete(handle_delete_session)) @@ -190,6 +204,11 @@ impl Transport for HttpTransport { .route("/oauth/callback", get(handle_oauth_callback)) .with_state(state); + if let Some((mount, a2a_state)) = &self.a2a { + app = app.nest(mount, ra2a::server::a2a_router(a2a_state.clone())); + tracing::info!(mount, "A2A endpoint enabled"); + } + if let Some(tls) = &self.tls { let mode = if tls.client_ca.is_some() { "HTTPS+mTLS" @@ -239,6 +258,8 @@ pub struct StreamableHttpTransport { hitl_store: Arc, oauth_manager: Arc, kill_switch: Arc>>, + /// Optional A2A server state: (mount_path, server_state). + a2a: Option<(String, ra2a::server::ServerState)>, } impl StreamableHttpTransport { @@ -267,8 +288,19 @@ impl StreamableHttpTransport { hitl_store, oauth_manager, kill_switch: Arc::new(std::sync::Mutex::new(std::collections::HashSet::new())), + a2a: None, } } + + /// Attach an A2A server state to this transport. + pub fn with_a2a( + mut self, + mount_path: impl Into, + state: ra2a::server::ServerState, + ) -> Self { + self.a2a = Some((mount_path.into(), state)); + self + } } #[async_trait] @@ -287,7 +319,7 @@ impl Transport for StreamableHttpTransport { kill_switch: Arc::clone(&self.kill_switch), }); - let app = Router::new() + let mut app = Router::new() .route("/mcp", post(handle_streamable_post)) .route("/mcp", get(handle_sse)) .route("/mcp", delete(handle_delete_session)) @@ -305,6 +337,11 @@ impl Transport for StreamableHttpTransport { .route("/oauth/callback", get(handle_oauth_callback)) .with_state(state); + if let Some((mount, a2a_state)) = &self.a2a { + app = app.nest(mount, ra2a::server::a2a_router(a2a_state.clone())); + tracing::info!(mount, "A2A endpoint enabled"); + } + if let Some(tls) = &self.tls { let mode = if tls.client_ca.is_some() { "HTTPS+mTLS" diff --git a/tests/a2a.rs b/tests/a2a.rs new file mode 100644 index 0000000..de434d0 --- /dev/null +++ b/tests/a2a.rs @@ -0,0 +1,199 @@ +mod common; + +use common::*; +use serde_json::json; + +// ── Agent card discovery ────────────────────────────────────────────────────── + +#[tokio::test] +async fn agent_card_is_served_at_well_known_path() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + let card = h.agent_card().await; + assert_eq!(card["name"].as_str().unwrap(), "Test Agent"); + assert!(!card["supportedInterfaces"].as_array().unwrap().is_empty()); +} + +// ── message/send proxy ──────────────────────────────────────────────────────── + +#[tokio::test] +async fn message_send_proxied_to_upstream() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + let resp = h.send_message("cursor", "hello world", None).await; + assert!(resp.status().is_success()); + let body: serde_json::Value = resp.json().await.unwrap(); + // Result should contain the upstream's echo reply. + assert!( + body["result"]["message"]["parts"][0]["text"] + .as_str() + .unwrap_or("") + .contains("echo: hello world"), + "unexpected response: {body}" + ); +} + +// ── Agent identity enforcement ──────────────────────────────────────────────── + +#[tokio::test] +async fn missing_agent_header_is_rejected() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + // Send without x-arbitus-agent header. + let body = json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "message/send", + "params": { + "message": { + "messageId": "msg-1", + "role": "ROLE_USER", + "parts": [{ "text": "hello" }], + "extensions": [] + } + } + }); + let resp = h + .client + .post(h.url("/a2a")) + .json(&body) + .send() + .await + .unwrap(); + // ra2a returns 200 with a JSON-RPC error for interceptor rejections. + let resp_body: serde_json::Value = resp.json().await.unwrap(); + assert!( + resp_body["error"].is_object(), + "expected error, got: {resp_body}" + ); +} + +#[tokio::test] +async fn unknown_agent_is_rejected() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + let resp = h.send_message("ghost-agent", "hello", None).await; + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["error"].is_object(), + "expected error for unconfigured agent, got: {body}" + ); +} + +// ── API key authentication ──────────────────────────────────────────────────── + +#[tokio::test] +async fn correct_api_key_allows_request() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + let resp = h + .send_message("secured-agent", "hello", Some("test-key-123")) + .await; + assert!(resp.status().is_success()); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!(body["result"].is_object(), "expected result, got: {body}"); +} + +#[tokio::test] +async fn wrong_api_key_is_rejected() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + let resp = h + .send_message("secured-agent", "hello", Some("wrong-key")) + .await; + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["error"].is_object(), + "expected error for wrong API key, got: {body}" + ); +} + +#[tokio::test] +async fn missing_api_key_is_rejected() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + // secured-agent requires api_key but we send none. + let resp = h.send_message("secured-agent", "hello", None).await; + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["error"].is_object(), + "expected error for missing API key, got: {body}" + ); +} + +// ── Payload filtering ───────────────────────────────────────────────────────── + +#[tokio::test] +async fn message_with_blocked_pattern_is_rejected() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + // DEFAULT_CONFIG has block_patterns: ["password=", "private_key"] + let resp = h + .send_message("cursor", "my private_key=AAABBBCCC", None) + .await; + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["error"].is_object(), + "expected block for sensitive pattern, got: {body}" + ); +} + +#[tokio::test] +async fn clean_message_passes_payload_filter() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + let resp = h + .send_message("cursor", "this is a harmless message", None) + .await; + assert!(resp.status().is_success()); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["result"].is_object(), + "expected result for clean message, got: {body}" + ); +} + +// ── Rate limiting ───────────────────────────────────────────────────────────── + +#[tokio::test] +async fn rate_limit_blocks_after_limit_exceeded() { + // rate-test agent has rate_limit: 3 + let h = harness_with_a2a(DEFAULT_CONFIG).await; + + for _ in 0..3 { + let resp = h.send_message("rate-test", "hello", None).await; + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["result"].is_object(), + "expected success within limit, got: {body}" + ); + } + + // 4th request should be blocked. + let resp = h.send_message("rate-test", "hello", None).await; + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["error"].is_object(), + "expected rate limit error on 4th request, got: {body}" + ); +} + +// ── MCP endpoint unaffected ─────────────────────────────────────────────────── + +#[tokio::test] +async fn mcp_endpoint_still_works_when_a2a_is_configured() { + let h = harness_with_a2a(DEFAULT_CONFIG).await; + // MCP initialize should still work normally. + let mcp_body = json!({ + "jsonrpc": "2.0", "id": 1, "method": "initialize", + "params": { + "protocolVersion": "2025-03-26", + "capabilities": {}, + "clientInfo": { "name": "cursor", "version": "1.0.0" } + } + }); + let resp = h + .client + .post(h.url("/mcp")) + .json(&mcp_body) + .send() + .await + .unwrap(); + assert!(resp.status().is_success()); + let body: serde_json::Value = resp.json().await.unwrap(); + assert!( + body["result"]["serverInfo"].is_object(), + "expected MCP initialize result, got: {body}" + ); +} diff --git a/tests/common/mod.rs b/tests/common/mod.rs index 7ec812e..c9c3a84 100644 --- a/tests/common/mod.rs +++ b/tests/common/mod.rs @@ -167,6 +167,176 @@ pub async fn start_dummy() -> (u16, tokio::task::AbortHandle) { (port, handle.abort_handle()) } +// ── In-process dummy A2A server ─────────────────────────────────────────────── + +async fn dummy_a2a(Json(msg): Json) -> impl IntoResponse { + let method = msg["method"].as_str().unwrap_or(""); + let id = &msg["id"]; + + match method { + "message/send" => { + let input_text = msg["params"]["message"]["parts"] + .as_array() + .and_then(|parts| parts.iter().find_map(|p| p["text"].as_str())) + .unwrap_or("(no text)"); + + Json(json!({ + "jsonrpc": "2.0", + "id": id, + "result": { + "message": { + "messageId": "resp-1", + "role": "ROLE_AGENT", + "parts": [{ "text": format!("echo: {input_text}") }], + "extensions": [] + } + } + })) + .into_response() + } + _ => Json(json!({ + "jsonrpc": "2.0", + "id": id, + "error": { "code": -32601, "message": format!("unknown method '{method}'") } + })) + .into_response(), + } +} + +/// Start an in-process dummy A2A upstream server. Returns (port, abort_handle). +pub async fn start_dummy_a2a() -> (u16, tokio::task::AbortHandle) { + let listener = TcpListener::bind("0.0.0.0:0").await.unwrap(); + let port = listener.local_addr().unwrap().port(); + let app = Router::new().route("/", post(dummy_a2a)); + let handle = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + (port, handle.abort_handle()) +} + +// ── A2A Harness ─────────────────────────────────────────────────────────────── + +/// Harness for A2A protocol tests. +/// The gateway exposes the A2A endpoint at `/a2a`. +pub struct A2aHarness { + pub gw_port: u16, + pub client: Client, + pub config_path: String, + _dummy_a2a: tokio::task::AbortHandle, + _dummy_mcp: tokio::task::AbortHandle, + _gw: tokio::process::Child, +} + +impl Drop for A2aHarness { + fn drop(&mut self) { + self._dummy_a2a.abort(); + self._dummy_mcp.abort(); + let _ = self._gw.start_kill(); + let _ = std::fs::remove_file(&self.config_path); + } +} + +impl A2aHarness { + pub fn url(&self, path: &str) -> String { + format!("http://127.0.0.1:{}{}", self.gw_port, path) + } + + /// POST to the A2A endpoint with a `message/send` request. + /// `agent` is sent as the `x-arbitus-agent` header. + /// `api_key` is sent as the `x-api-key` header (optional). + pub async fn send_message( + &self, + agent: &str, + text: &str, + api_key: Option<&str>, + ) -> reqwest::Response { + let body = json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "message/send", + "params": { + "message": { + "messageId": "msg-1", + "role": "ROLE_USER", + "parts": [{ "text": text }], + "extensions": [] + } + } + }); + let mut req = self + .client + .post(self.url("/a2a")) + .header("x-arbitus-agent", agent) + .json(&body); + if let Some(key) = api_key { + req = req.header("x-api-key", key); + } + req.send().await.unwrap() + } + + /// GET the agent card from the well-known endpoint. + pub async fn agent_card(&self) -> Value { + self.client + .get(self.url("/a2a/.well-known/agent-card.json")) + .send() + .await + .unwrap() + .json() + .await + .unwrap() + } +} + +/// Spin up a gateway binary with an A2A endpoint pointing at a dummy A2A upstream. +/// +/// The gateway listens on an auto-assigned port with both MCP (`/mcp`) and A2A (`/a2a`). +/// `agents_config` provides the `agents:` and `rules:` sections. +pub async fn harness_with_a2a(agents_config: &str) -> A2aHarness { + let (a2a_port, dummy_a2a_abort) = start_dummy_a2a().await; + let (mcp_port, dummy_mcp_abort) = start_dummy().await; + let gw_port = free_port().await; + + let config = format!( + r#"transport: + type: http + addr: "0.0.0.0:{gw_port}" + upstream: "http://127.0.0.1:{mcp_port}/mcp" + session_ttl_secs: 3600 +audit: + type: stdout +a2a: + upstream: "http://127.0.0.1:{a2a_port}" + mount: "/a2a" + agent_card: + name: "Test Agent" + description: "Arbitus A2A proxy for tests" + url: "http://127.0.0.1:{gw_port}/a2a" + version: "1.0.0" +{agents_config}"# + ); + + let config_path = format!("/tmp/arbitus-a2a-test-{gw_port}.yml"); + std::fs::write(&config_path, &config).unwrap(); + + let gw = tokio::process::Command::new(GATEWAY_BIN) + .arg(&config_path) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .spawn() + .unwrap(); + + wait_for_port(gw_port).await; + + A2aHarness { + gw_port, + client: Client::new(), + config_path, + _dummy_a2a: dummy_a2a_abort, + _dummy_mcp: dummy_mcp_abort, + _gw: gw, + } +} + // ── Harness ─────────────────────────────────────────────────────────────────── pub struct Harness {