From 8a13e2247b4b09cbdbee18999244a65540cdc1fc Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Tue, 1 Sep 2026 22:08:34 -0700 Subject: [PATCH] feat(rust): integrate Python extension host --- litellm-rust/AGENTS.md | 5 +- litellm-rust/Cargo.lock | 398 +++++++++++- litellm-rust/Cargo.toml | 6 + litellm-rust/crates/ai-gateway/Cargo.toml | 3 + .../src/audio_transcription/hooks.rs | 9 + .../src/integrations/custom_guardrail/mod.rs | 23 + .../integrations/custom_guardrail/types.rs | 2 + .../crates/ai-gateway/src/integrations/mod.rs | 1 + .../python_extension_host/adapters.rs | 435 +++++++++++++ .../python_extension_host/client.rs | 422 +++++++++++++ .../python_extension_host/config.rs | 106 ++++ .../integrations/python_extension_host/mod.rs | 13 + .../python_extension_host/stream.rs | 286 +++++++++ .../python_extension_host/tests.rs | 573 ++++++++++++++++++ litellm-rust/crates/ai-gateway/src/main.rs | 63 +- .../crates/ai-gateway/src/ocr/hooks.rs | 20 + .../crates/ai-gateway/src/python/config.rs | 28 +- .../ai-gateway/src/routes/messages/mod.rs | 2 + .../ai-gateway/src/routes/responses/mod.rs | 2 + litellm-rust/crates/ai-gateway/src/state.rs | 5 + .../crates/core/src/call_lifecycle/mod.rs | 72 ++- .../crates/core/src/call_lifecycle/types.rs | 2 + .../core/src/responses/instrumentation.rs | 9 + .../core/tests/workspace_crate_allowlist.rs | 16 +- .../python-extension-protocol/Cargo.toml | 16 + .../crates/python-extension-protocol/build.rs | 18 + .../python-extension-protocol/src/lib.rs | 48 ++ 27 files changed, 2561 insertions(+), 22 deletions(-) create mode 100644 litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/adapters.rs create mode 100644 litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/client.rs create mode 100644 litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/config.rs create mode 100644 litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/mod.rs create mode 100644 litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/stream.rs create mode 100644 litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/tests.rs create mode 100644 litellm-rust/crates/python-extension-protocol/Cargo.toml create mode 100644 litellm-rust/crates/python-extension-protocol/build.rs create mode 100644 litellm-rust/crates/python-extension-protocol/src/lib.rs diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 36a5ad5a8f4..96b36d807ec 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -1,6 +1,6 @@ # AGENTS.md -litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are MODULES inside the layers. +litellm-rust has exactly FOUR crates. A crate is a LAYER, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are MODULES inside the layers. ## Crates @@ -9,8 +9,9 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes ( | litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. | | litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. | | litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. | +| litellm-python-extension-protocol | Shared generated protobuf types for the external callback and guardrail host. It owns no transport or dispatch. | -Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge. +Dependency direction (acyclic): litellm-python-extension-protocol → litellm-ai-gateway, and litellm-core ← litellm-ai-gateway ← litellm-python-bridge. ## Where a route lives diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 4388e561026..274918c3eec 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -32,6 +32,12 @@ version = "1.0.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" +[[package]] +name = "anyhow" +version = "1.0.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" + [[package]] name = "arc-swap" version = "1.9.2" @@ -417,7 +423,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "edca88bc138befd0323b20752846e6587272d3b03b0343c8ea28a6f819e6e71f" dependencies = [ "async-trait", - "axum-core", + "axum-core 0.4.5", "base64", "bytes", "futures-util", @@ -427,7 +433,7 @@ dependencies = [ "hyper 1.10.1", "hyper-util", "itoa", - "matchit", + "matchit 0.7.3", "memchr", "mime", "percent-encoding", @@ -447,6 +453,31 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90" +dependencies = [ + "axum-core 0.5.6", + "bytes", + "futures-util", + "http 1.4.2", + "http-body 1.1.0", + "http-body-util", + "itoa", + "matchit 0.8.4", + "memchr", + "mime", + "percent-encoding", + "pin-project-lite", + "serde_core", + "sync_wrapper", + "tower", + "tower-layer", + "tower-service", +] + [[package]] name = "axum-core" version = "0.4.5" @@ -468,6 +499,24 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum-core" +version = "0.5.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1" +dependencies = [ + "bytes", + "futures-core", + "http 1.4.2", + "http-body 1.1.0", + "http-body-util", + "mime", + "pin-project-lite", + "sync_wrapper", + "tower-layer", + "tower-service", +] + [[package]] name = "base64" version = "0.22.1" @@ -841,6 +890,16 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + [[package]] name = "fastrand" version = "2.5.0" @@ -853,12 +912,24 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +[[package]] +name = "fixedbitset" +version = "0.5.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d674e81391d1e1ab681a28d99df07927c6d4aa5b027d7da16ba32d1d21ecd99" + [[package]] name = "fnv" version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -1021,6 +1092,15 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "foldhash", +] + [[package]] name = "hashbrown" version = "0.17.1" @@ -1202,6 +1282,19 @@ dependencies = [ "webpki-roots", ] +[[package]] +name = "hyper-timeout" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b90d566bffbce6a75bd8b09a05aa8c2cb1fabb6cb348f8840c9e4c90a0d83b0" +dependencies = [ + "hyper 1.10.1", + "hyper-util", + "pin-project-lite", + "tokio", + "tower-service", +] + [[package]] name = "hyper-util" version = "0.1.20" @@ -1335,7 +1428,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", - "hashbrown", + "hashbrown 0.17.1", ] [[package]] @@ -1386,15 +1479,22 @@ version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + [[package]] name = "litellm-ai-gateway" version = "0.1.0" dependencies = [ - "axum", + "axum 0.7.9", "base64", "futures-channel", "futures-util", "litellm-core", + "litellm-python-extension-protocol", "pyo3", "reqwest", "serde", @@ -1402,7 +1502,9 @@ dependencies = [ "sha2 0.10.9", "subtle", "tokio", + "tokio-stream", "tokio-tungstenite", + "tonic", "tower", ] @@ -1440,6 +1542,18 @@ dependencies = [ "tokio", ] +[[package]] +name = "litellm-python-extension-protocol" +version = "0.1.0" +dependencies = [ + "prost", + "prost-build", + "protoc-bin-vendored", + "tonic", + "tonic-prost", + "tonic-prost-build", +] + [[package]] name = "litemap" version = "0.8.2" @@ -1464,6 +1578,12 @@ version = "0.7.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0e7465ac9959cc2b1404e8e2367b43684a6d13790fe23056cc8c6c5a6b7bcb94" +[[package]] +name = "matchit" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3" + [[package]] name = "memchr" version = "2.8.3" @@ -1487,6 +1607,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "multimap" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d87ecb2933e8aeadb3e3a02b828fed80a7528047e68b4f424523a0981a3a084" + [[package]] name = "num-conv" version = "0.2.2" @@ -1551,6 +1677,37 @@ version = "2.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" +[[package]] +name = "petgraph" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8701b58ea97060d5e5b155d383a69952a60943f0e6dfe30b04c287beb0b27455" +dependencies = [ + "fixedbitset", + "hashbrown 0.15.5", + "indexmap", +] + +[[package]] +name = "pin-project" +version = "1.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2466b2336ed02bcdca6b294417127b90ec92038d1d5c4fbeac971a922e0e0924" +dependencies = [ + "pin-project-internal", +] + +[[package]] +name = "pin-project-internal" +version = "1.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c96395f0a926bc13b1c17622aaddda1ecb55d49c8f1bf9777e4d877800a43f8b" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "pin-project-lite" version = "0.2.17" @@ -1627,6 +1784,16 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn 2.0.119", +] + [[package]] name = "proc-macro2" version = "1.0.107" @@ -1636,6 +1803,121 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "prost" +version = "0.14.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "528ac67416ff8646872a3c02cad9cc4ee5dc9f9540c9b10771855c95cb2e5ae1" +dependencies = [ + "bytes", + "prost-derive", +] + +[[package]] +name = "prost-build" +version = "0.14.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "03da047801ff44bb6a4d407d4860c05fd70bb81714e6b2f3812603d5b145b042" +dependencies = [ + "heck", + "itertools", + "log", + "multimap", + "petgraph", + "prettyplease", + "prost", + "prost-types", + "regex", + "syn 2.0.119", + "tempfile", +] + +[[package]] +name = "prost-derive" +version = "0.14.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b570b25f7617e43d59005d0990ccb79e950a423952cea19671b7a876da390adf" +dependencies = [ + "anyhow", + "itertools", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "prost-types" +version = "0.14.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f94967dc7688f3054c7fac87473ffae4cc4c3904800e2d9f5b857246d8963b0a" +dependencies = [ + "prost", +] + +[[package]] +name = "protoc-bin-vendored" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1c381df33c98266b5f08186583660090a4ffa0889e76c7e9a5e175f645a67fa" +dependencies = [ + "protoc-bin-vendored-linux-aarch_64", + "protoc-bin-vendored-linux-ppcle_64", + "protoc-bin-vendored-linux-s390_64", + "protoc-bin-vendored-linux-x86_32", + "protoc-bin-vendored-linux-x86_64", + "protoc-bin-vendored-macos-aarch_64", + "protoc-bin-vendored-macos-x86_64", + "protoc-bin-vendored-win32", +] + +[[package]] +name = "protoc-bin-vendored-linux-aarch_64" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c350df4d49b5b9e3ca79f7e646fde2377b199e13cfa87320308397e1f37e1a4c" + +[[package]] +name = "protoc-bin-vendored-linux-ppcle_64" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a55a63e6c7244f19b5c6393f025017eb5d793fd5467823a099740a7a4222440c" + +[[package]] +name = "protoc-bin-vendored-linux-s390_64" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1dba5565db4288e935d5330a07c264a4ee8e4a5b4a4e6f4e83fad824cc32f3b0" + +[[package]] +name = "protoc-bin-vendored-linux-x86_32" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8854774b24ee28b7868cd71dccaae8e02a2365e67a4a87a6cd11ee6cdbdf9cf5" + +[[package]] +name = "protoc-bin-vendored-linux-x86_64" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b38b07546580df720fa464ce124c4b03630a6fb83e05c336fea2a241df7e5d78" + +[[package]] +name = "protoc-bin-vendored-macos-aarch_64" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "89278a9926ce312e51f1d999fee8825d324d603213344a9a706daa009f1d8092" + +[[package]] +name = "protoc-bin-vendored-macos-x86_64" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "81745feda7ccfb9471d7a4de888f0652e806d5795b61480605d4943176299756" + +[[package]] +name = "protoc-bin-vendored-win32" +version = "3.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "95067976aca6421a523e491fce939a3e65249bac4b977adee0ee9771568e8aa3" + [[package]] name = "pyo3" version = "0.29.0" @@ -1971,6 +2253,19 @@ dependencies = [ "semver", ] +[[package]] +name = "rustix" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + [[package]] name = "rustls" version = "0.21.12" @@ -2308,6 +2603,19 @@ version = "0.13.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca" +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.3", + "once_cell", + "rustix", + "windows-sys 0.61.2", +] + [[package]] name = "thiserror" version = "1.0.69" @@ -2459,6 +2767,17 @@ dependencies = [ "tokio", ] +[[package]] +name = "tokio-stream" +version = "0.1.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3d06f0b082ba57c26b79407372e57cf2a1e28124f78e9479fe80322cf53420b" +dependencies = [ + "futures-core", + "pin-project-lite", + "tokio", +] + [[package]] name = "tokio-tungstenite" version = "0.24.0" @@ -2488,6 +2807,74 @@ dependencies = [ "tokio", ] +[[package]] +name = "tonic" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac2a5518c70fa84342385732db33fb3f44bc4cc748936eb5833d2df34d6445ef" +dependencies = [ + "async-trait", + "axum 0.8.9", + "base64", + "bytes", + "h2 0.4.15", + "http 1.4.2", + "http-body 1.1.0", + "http-body-util", + "hyper 1.10.1", + "hyper-timeout", + "hyper-util", + "percent-encoding", + "pin-project", + "socket2 0.6.5", + "sync_wrapper", + "tokio", + "tokio-stream", + "tower", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tonic-build" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c68f61875ac5293cf72e6c8cf0158086428c82c37229e98c840878f1706b0322" +dependencies = [ + "prettyplease", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "tonic-prost" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50849f68853be452acf590cde0b146665b8d507b3b8af17261df47e02c209ea0" +dependencies = [ + "bytes", + "prost", + "tonic", +] + +[[package]] +name = "tonic-prost-build" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "654e5643eff75d7f8c99197ce1440ed19a3474eada74c12bbac488b2cafdae27" +dependencies = [ + "prettyplease", + "proc-macro2", + "prost-build", + "prost-types", + "quote", + "syn 2.0.119", + "tempfile", + "tonic-build", +] + [[package]] name = "tower" version = "0.5.3" @@ -2496,9 +2883,12 @@ checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" dependencies = [ "futures-core", "futures-util", + "indexmap", "pin-project-lite", + "slab", "sync_wrapper", "tokio", + "tokio-util", "tower-layer", "tower-service", "tracing", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index c17a0605fc7..a5681361565 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -3,6 +3,7 @@ members = [ "crates/core", "crates/ai-gateway", "crates/python-bridge", + "crates/python-extension-protocol", ] resolver = "2" @@ -15,7 +16,9 @@ repository = "https://github.com/BerriAI/litellm" [workspace.dependencies] litellm-core = { path = "crates/core" } litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false } +litellm-python-extension-protocol = { path = "crates/python-extension-protocol" } axum = "0.7" +prost = "0.14.4" pyo3 = "0.29.0" pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] } pythonize = "0.29.0" @@ -28,7 +31,10 @@ subtle = "2" thiserror = "2.0" tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] } tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } +tonic = { version = "0.14.6", default-features = false, features = ["codegen"] } +tonic-prost = "0.14.6" futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } +tokio-stream = "0.1" base64 = "0.22" [profile.release] diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 541beabe170..8a25363c367 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -15,6 +15,7 @@ required-features = ["server"] [dependencies] litellm-core = { workspace = true, features = ["bedrock-auth"] } +litellm-python-extension-protocol.workspace = true # reqwest (rustls + json) is used by io/ocr and ships realtime logs to the # Python proxy callbacks API. reqwest.workspace = true @@ -22,6 +23,8 @@ reqwest.workspace = true tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time", "sync"] } tokio-tungstenite.workspace = true futures-util.workspace = true +tokio-stream.workspace = true +tonic = { workspace = true, features = ["transport", "router"] } serde_json.workspace = true base64.workspace = true axum = { workspace = true, features = ["ws"], optional = true } diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs index 0c9faeda6e7..bf3537c6f7b 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -197,6 +197,7 @@ impl CallLifecycleHooks = AudioFuture<'a, PreparedAudioTranscriptionRequest>; type DuringCallFuture<'a> = AudioFuture<'a, ProviderAudioTranscriptionRequest>; + type PostCallFuture<'a> = AudioFuture<'a, Value>; type SuccessFuture<'a> = AudioLogFuture<'a>; type FailureFuture<'a> = AudioLogFuture<'a>; @@ -216,6 +217,14 @@ impl CallLifecycleHooks( + &'a self, + _context: &'a CallLifecycleContext, + response: Value, + ) -> Self::PostCallFuture<'a> { + Box::pin(async move { Ok(response) }) + } + fn async_log_success_event<'a>( &'a self, context: &'a CallLifecycleContext, diff --git a/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/mod.rs b/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/mod.rs index e5d4ce3a708..d86597d7bbf 100644 --- a/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/mod.rs @@ -39,6 +39,15 @@ pub trait CustomGuardrail: Send + Sync { ) -> GuardrailFuture<'a> { Box::pin(async move { Ok(GuardrailDecision::Allow(request)) }) } + + /// Python 1:1 name: `async_post_call_success_hook(data, user_api_key_dict, response)`. + fn async_post_call_success_hook<'a>( + &'a self, + _context: &'a GuardrailContext, + response: GuardrailRequest, + ) -> GuardrailFuture<'a> { + Box::pin(async move { Ok(GuardrailDecision::Allow(response)) }) + } } pub struct CustomGuardrailRunner { @@ -72,6 +81,15 @@ impl CustomGuardrailRunner { .await } + pub async fn run_post_call( + &self, + context: &GuardrailContext, + response: GuardrailRequest, + ) -> Result<(GuardrailRequest, GuardrailDispatchReport), GuardrailError> { + self.run_hook(GuardrailEventHook::PostCall, context, response) + .await + } + pub async fn run_before_provider( &self, event_hook: GuardrailEventHook, @@ -145,6 +163,11 @@ impl CustomGuardrailRunner { .async_moderation_hook(context, request.clone()) .await? } + GuardrailEventHook::PostCall => { + guardrail + .async_post_call_success_hook(context, request.clone()) + .await? + } }; match decision.into_request() { Ok(next_request) => request = next_request, diff --git a/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/types.rs b/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/types.rs index 825e56cc0d7..94d5500847b 100644 --- a/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/types.rs +++ b/litellm-rust/crates/ai-gateway/src/integrations/custom_guardrail/types.rs @@ -13,6 +13,7 @@ pub type GuardrailFuture<'a> = pub enum GuardrailEventHook { PreCall, DuringCall, + PostCall, } impl GuardrailEventHook { @@ -20,6 +21,7 @@ impl GuardrailEventHook { match self { Self::PreCall => "pre_call", Self::DuringCall => "during_call", + Self::PostCall => "post_call", } } } diff --git a/litellm-rust/crates/ai-gateway/src/integrations/mod.rs b/litellm-rust/crates/ai-gateway/src/integrations/mod.rs index c62f1821ef8..784fcf9f7bb 100644 --- a/litellm-rust/crates/ai-gateway/src/integrations/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/integrations/mod.rs @@ -9,4 +9,5 @@ pub mod custom_guardrail; pub mod custom_logger; pub mod litellm_python_proxy_api; +pub mod python_extension_host; pub mod types; diff --git a/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/adapters.rs b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/adapters.rs new file mode 100644 index 00000000000..87af4f7a02e --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/adapters.rs @@ -0,0 +1,435 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; + +use litellm_python_extension_protocol::{ + AuthContext, CallbackEvent, CallbackEventKind, GuardrailDecision as WireDecision, + GuardrailInvocation, HookPhase, InvocationContext, OperationResult, +}; +use serde_json::{Value, json}; + +use crate::integrations::custom_guardrail::{ + CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook, + GuardrailFuture, GuardrailRequest, +}; +use crate::integrations::custom_logger::{ + CallbackTiming, CallbackValue, CustomLogger, LogError, LogFuture, ModelCallDetails, +}; + +use super::client::PythonExtensionClient; +use super::config::{ManifestExtensionKind, PythonExtensionManifest}; + +pub struct RemoteCustomGuardrail { + name: String, + plugin_id: String, + hooks: Vec, + client: Arc, +} + +impl RemoteCustomGuardrail { + pub fn new( + name: String, + plugin_id: String, + hooks: Vec, + client: Arc, + ) -> Self { + Self { + name, + plugin_id, + hooks, + client, + } + } + + fn invoke<'a>( + &'a self, + phase: HookPhase, + context: &'a GuardrailContext, + request: GuardrailRequest, + ) -> GuardrailFuture<'a> { + Box::pin(async move { + let original = request.clone(); + let encoded = serde_json::to_vec(&request.data).map_err(|error| GuardrailError { + message: error.to_string(), + kind: "SerializationError".to_string(), + })?; + let (request_json, response_json) = if phase == HookPhase::PostCall { + (b"{}".to_vec(), Some(encoded)) + } else { + (encoded, None) + }; + let result = self + .client + .execute_guardrail(GuardrailInvocation { + context: Some(invocation_context( + self.client.manifest().revision_id.clone(), + context.call_type.as_str(), + )), + plugin_id: self.plugin_id.clone(), + hook_phase: phase.into(), + request_json, + response_json, + auth: Some(auth_context(context)), + cache: None, + }) + .await; + wire_result_to_decision( + result.operation, + result.decision, + result.request_json, + result.response_json, + result.public_error, + original, + ) + }) + } +} + +impl CustomGuardrail for RemoteCustomGuardrail { + fn guardrail_name(&self) -> &str { + &self.name + } + + fn supported_event_hooks(&self) -> &[GuardrailEventHook] { + &self.hooks + } + + fn async_pre_call_hook<'a>( + &'a self, + context: &'a GuardrailContext, + request: GuardrailRequest, + ) -> GuardrailFuture<'a> { + self.invoke(HookPhase::PreCall, context, request) + } + + fn async_moderation_hook<'a>( + &'a self, + context: &'a GuardrailContext, + request: GuardrailRequest, + ) -> GuardrailFuture<'a> { + self.invoke(HookPhase::DuringCall, context, request) + } + + fn async_post_call_success_hook<'a>( + &'a self, + context: &'a GuardrailContext, + response: GuardrailRequest, + ) -> GuardrailFuture<'a> { + self.invoke(HookPhase::PostCall, context, response) + } +} + +pub struct RemoteCustomLogger { + plugin_id: String, + success_enabled: bool, + failure_enabled: bool, + client: Arc, +} + +impl RemoteCustomLogger { + pub fn new(plugin_id: String, client: Arc) -> Self { + Self::with_events(plugin_id, true, true, client) + } + + fn with_events( + plugin_id: String, + success_enabled: bool, + failure_enabled: bool, + client: Arc, + ) -> Self { + Self { + plugin_id, + success_enabled, + failure_enabled, + client, + } + } + + fn enqueue( + &self, + kind: CallbackEventKind, + details: &ModelCallDetails, + response: Option<&CallbackValue>, + timing: CallbackTiming, + ) -> Result<(), LogError> { + let payload_json = details + .standard_logging_payload + .as_ref() + .map(serde_json::to_vec) + .transpose() + .map_err(|error| LogError { + message: error.to_string(), + kind: "SerializationError".to_string(), + })? + .unwrap_or_else(|| { + serde_json::to_vec(&json!({ + "model": details.model, + "custom_llm_provider": details.custom_llm_provider, + "call_type": details.call_type.as_str(), + })) + .unwrap_or_else(|_| b"{}".to_vec()) + }); + let response_json = response + .map(|response| serde_json::to_vec(&response.value)) + .transpose() + .map_err(|error| LogError { + message: error.to_string(), + kind: "SerializationError".to_string(), + })?; + let error_json = details.failure_error.as_ref().map(|error| { + serde_json::to_vec(&json!({"type": error.kind, "message": error.message})) + .unwrap_or_else(|_| b"{}".to_vec()) + }); + let metadata = &details.metadata; + let event = CallbackEvent { + context: Some(InvocationContext { + request_id: details.request_id.clone().unwrap_or_default(), + invocation_id: details + .litellm_call_id + .clone() + .unwrap_or_else(next_invocation_id), + active_revision: self.client.manifest().revision_id.clone(), + api_surface: details.call_type.as_str().to_string(), + call_type: details.call_type.as_str().to_string(), + trace_context: HashMap::new(), + }), + plugin_id: self.plugin_id.clone(), + kind: kind.into(), + standard_logging_payload_json: payload_json, + response_json, + error_json, + start_time_seconds: timing.start_time, + end_time_seconds: timing.end_time, + auth: Some(AuthContext { + key_hash: metadata.user_api_key_hash.clone().unwrap_or_default(), + user_id: metadata.user_api_key_user_id.clone().unwrap_or_default(), + team_id: metadata.user_api_key_team_id.clone().unwrap_or_default(), + request_metadata: HashMap::new(), + }), + cache: None, + streaming: details + .standard_logging_payload + .as_ref() + .map(|payload| payload.stream) + .unwrap_or(false), + }; + self.client.enqueue_callback(event).map_err(|reason| { + if reason.contains("full") { + LogError::channel_full() + } else { + LogError::channel_closed() + } + }) + } +} + +impl CustomLogger for RemoteCustomLogger { + fn async_log_success_event<'a>( + &'a self, + details: &'a ModelCallDetails, + response: &'a CallbackValue, + timing: CallbackTiming, + ) -> LogFuture<'a> { + Box::pin(async move { + if self.success_enabled { + self.enqueue(CallbackEventKind::Success, details, Some(response), timing) + } else { + Ok(()) + } + }) + } + + fn async_log_failure_event<'a>( + &'a self, + details: &'a ModelCallDetails, + response: Option<&'a CallbackValue>, + timing: CallbackTiming, + ) -> LogFuture<'a> { + Box::pin(async move { + if self.failure_enabled { + self.enqueue(CallbackEventKind::Failure, details, response, timing) + } else { + Ok(()) + } + }) + } +} + +pub struct RemoteExtensions { + pub guardrails: Vec>, + pub loggers: Vec>, +} + +impl RemoteExtensions { + pub fn from_manifest( + manifest: &PythonExtensionManifest, + descriptors: &[litellm_python_extension_protocol::ExtensionDescriptor], + client: Arc, + ) -> Self { + let descriptor_hooks: HashMap<&str, &[String]> = descriptors + .iter() + .map(|descriptor| (descriptor.id.as_str(), descriptor.hooks.as_slice())) + .collect(); + let mut guardrails: Vec> = Vec::new(); + let mut loggers: Vec> = Vec::new(); + for extension in &manifest.extensions { + match extension.kind { + ManifestExtensionKind::Callback => { + let events = extension + .constructor + .get("callback_events") + .and_then(Value::as_array); + let event_enabled = |name: &str| { + events.is_none_or(|events| { + events.iter().any(|event| event.as_str() == Some(name)) + }) + }; + loggers.push(Arc::new(RemoteCustomLogger::with_events( + extension.id.clone(), + event_enabled("success"), + event_enabled("failure"), + client.clone(), + ))); + } + ManifestExtensionKind::Guardrail => { + let name = extension + .constructor + .pointer("/kwargs/guardrail_name") + .and_then(Value::as_str) + .unwrap_or(&extension.id) + .to_string(); + let hooks = descriptor_hooks + .get(extension.id.as_str()) + .map(|hooks| hooks_from_descriptor(hooks)) + .filter(|hooks| !hooks.is_empty()) + .unwrap_or_else(|| { + vec![ + GuardrailEventHook::PreCall, + GuardrailEventHook::DuringCall, + GuardrailEventHook::PostCall, + ] + }); + guardrails.push(Arc::new(RemoteCustomGuardrail::new( + name, + extension.id.clone(), + hooks, + client.clone(), + ))); + } + } + } + Self { + guardrails, + loggers, + } + } +} + +fn wire_result_to_decision( + operation: Option, + decision: i32, + request_json: Option>, + response_json: Option>, + public_error: Option, + original: GuardrailRequest, +) -> Result { + if !operation.map(|operation| operation.ok).unwrap_or(false) { + return Ok(GuardrailDecision::Allow(original)); + } + match WireDecision::try_from(decision).unwrap_or(WireDecision::Error) { + WireDecision::Allow | WireDecision::Error | WireDecision::Unspecified => { + Ok(GuardrailDecision::Allow(original)) + } + WireDecision::ReplaceRequest | WireDecision::ReplaceResponse => { + let replacement = request_json + .or(response_json) + .ok_or_else(|| GuardrailError { + message: "extension replacement omitted JSON body".to_string(), + kind: "ExtensionProtocolError".to_string(), + })?; + let data = serde_json::from_slice(&replacement).map_err(|error| GuardrailError { + message: error.to_string(), + kind: "SerializationError".to_string(), + })?; + Ok(GuardrailDecision::Mask(GuardrailRequest::new(data))) + } + WireDecision::Block => Ok(GuardrailDecision::Block(GuardrailError::blocked( + public_error + .map(|error| error.message) + .unwrap_or_else(|| "blocked by Python extension".to_string()), + ))), + } +} + +fn invocation_context(revision_id: String, call_type: &str) -> InvocationContext { + let invocation_id = next_invocation_id(); + InvocationContext { + request_id: invocation_id.clone(), + invocation_id, + active_revision: revision_id, + api_surface: call_type.to_string(), + call_type: call_type.to_string(), + trace_context: HashMap::new(), + } +} + +fn auth_context(context: &GuardrailContext) -> AuthContext { + let request_metadata = context + .metadata + .iter() + .filter(|(name, _)| !is_sensitive_name(name)) + .filter_map(|(name, value)| scalar_string(value).map(|value| (name.clone(), value))) + .collect(); + AuthContext { + key_hash: context.user_api_key_hash.clone().unwrap_or_default(), + user_id: context.user_api_key_user_id.clone().unwrap_or_default(), + team_id: context.user_api_key_team_id.clone().unwrap_or_default(), + request_metadata, + } +} + +fn is_sensitive_name(name: &str) -> bool { + let name = name.to_ascii_lowercase(); + [ + "authorization", + "api_key", + "token", + "cookie", + "secret", + "password", + ] + .iter() + .any(|part| name.contains(part)) +} + +fn scalar_string(value: &Value) -> Option { + match value { + Value::String(value) => Some(value.clone()), + Value::Number(value) => Some(value.to_string()), + Value::Bool(value) => Some(value.to_string()), + _ => None, + } +} + +fn hooks_from_descriptor(hooks: &[String]) -> Vec { + hooks + .iter() + .filter_map(|hook| match hook.as_str() { + "async_pre_call_hook" => Some(GuardrailEventHook::PreCall), + "async_moderation_hook" => Some(GuardrailEventHook::DuringCall), + "async_post_call_success_hook" => Some(GuardrailEventHook::PostCall), + _ => None, + }) + .collect() +} + +pub(crate) fn next_invocation_id() -> String { + static COUNTER: AtomicU64 = AtomicU64::new(1); + let sequence = COUNTER.fetch_add(1, Ordering::Relaxed); + let timestamp = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_nanos()) + .unwrap_or(0); + format!("extension-{timestamp}-{sequence}") +} diff --git a/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/client.rs b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/client.rs new file mode 100644 index 00000000000..399dad96d9c --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/client.rs @@ -0,0 +1,422 @@ +use std::collections::HashMap; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex, RwLock}; +use std::time::Duration; + +use litellm_python_extension_protocol::python_extension_host_client::PythonExtensionHostClient; +use litellm_python_extension_protocol::{ + CallbackEvent, CommitRevisionRequest, ErrorCode, ExtensionDescriptor, GetCapabilitiesRequest, + GuardrailDecision, GuardrailInvocation, GuardrailResult, PrepareRevisionRequest, + PublishCallbackEventsRequest, RetireRevisionRequest, StreamFrame, +}; +use tokio::sync::mpsc; +use tonic::metadata::{Ascii, MetadataValue}; +use tonic::service::Interceptor; +use tonic::service::interceptor::InterceptedService; +use tonic::transport::{Channel, Endpoint}; +use tonic::{Code, Request, Status, Streaming}; + +use super::config::ManifestExtensionKind; +use super::config::{PythonExtensionManifest, PythonExtensionSettings}; + +const PROTOCOL_MAJOR: u32 = 1; +const PROTOCOL_MINOR: u32 = 0; +const TOKEN_METADATA_KEY: &str = "x-litellm-extension-token"; + +type HostStub = PythonExtensionHostClient>; + +#[derive(Clone)] +struct TokenInterceptor { + token: MetadataValue, +} + +impl Interceptor for TokenInterceptor { + fn call(&mut self, mut request: Request<()>) -> Result, Status> { + request + .metadata_mut() + .insert(TOKEN_METADATA_KEY, self.token.clone()); + Ok(request) + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ExtensionHostHealth { + pub healthy: bool, + pub reason: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum ActivationState { + Active(Vec), + Degraded(String), +} + +#[derive(Debug)] +pub enum InitializationError { + InvalidConfiguration(String), + Rejected(String), +} + +impl std::fmt::Display for InitializationError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::InvalidConfiguration(message) | Self::Rejected(message) => { + formatter.write_str(message) + } + } + } +} + +impl std::error::Error for InitializationError {} + +#[derive(Clone)] +pub struct PythonExtensionClient { + stub: HostStub, + settings: PythonExtensionSettings, + manifest: PythonExtensionManifest, + callback_tx: mpsc::Sender, + health: Arc>, + bypass_counts: Arc>>, + recovering: Arc, +} + +impl PythonExtensionClient { + pub async fn connect( + settings: PythonExtensionSettings, + manifest: PythonExtensionManifest, + ) -> Result<(Arc, ActivationState), InitializationError> { + let token = MetadataValue::try_from(settings.token.as_str()).map_err(|error| { + InitializationError::InvalidConfiguration(format!("invalid extension token: {error}")) + })?; + let endpoint = Endpoint::from_shared(settings.endpoint.clone()).map_err(|error| { + InitializationError::InvalidConfiguration(format!( + "invalid extension endpoint: {error}" + )) + })?; + let channel = endpoint + .connect_timeout(settings.connect_timeout) + .connect_lazy(); + let stub = PythonExtensionHostClient::with_interceptor(channel, TokenInterceptor { token }); + let (callback_tx, callback_rx) = mpsc::channel(settings.callback_queue_size); + let client = Arc::new(Self { + stub, + settings, + manifest, + callback_tx, + health: Arc::new(RwLock::new(ExtensionHostHealth { + healthy: false, + reason: Some("not connected".to_string()), + })), + bypass_counts: Arc::new(Mutex::new(HashMap::new())), + recovering: Arc::new(AtomicBool::new(false)), + }); + tokio::spawn(client.clone().callback_worker(callback_rx)); + let activation = match client.activate().await { + Ok(descriptors) => ActivationState::Active(descriptors), + Err(status) if is_transient(&status) => { + let reason = status.code().to_string(); + client.mark_unhealthy(reason.clone()); + client.schedule_recovery(); + ActivationState::Degraded(reason) + } + Err(status) => { + return Err(InitializationError::Rejected(format!( + "extension host rejected startup: {}", + status.message() + ))); + } + }; + Ok((client, activation)) + } + + pub fn manifest(&self) -> &PythonExtensionManifest { + &self.manifest + } + + pub fn health(&self) -> ExtensionHostHealth { + self.health + .read() + .map(|health| health.clone()) + .unwrap_or(ExtensionHostHealth { + healthy: false, + reason: Some("health lock poisoned".to_string()), + }) + } + + pub fn bypass_counts(&self) -> HashMap<(String, String, String), u64> { + self.bypass_counts + .lock() + .map(|counts| counts.clone()) + .unwrap_or_default() + } + + pub async fn execute_guardrail(&self, invocation: GuardrailInvocation) -> GuardrailResult { + let plugin_id = invocation.plugin_id.clone(); + let hook = invocation.hook_phase.to_string(); + let mut request = Request::new(invocation); + request.set_timeout(self.settings.hook_timeout); + let mut stub = self.stub.clone(); + match stub.execute_guardrail(request).await { + Ok(response) => { + let result = response.into_inner(); + if let Some(operation) = result.operation.as_ref().filter(|operation| !operation.ok) + { + self.record_bypass(&plugin_id, &hook, &operation_reason(operation)); + } else { + self.mark_healthy(); + } + result + } + Err(status) => { + self.record_bypass(&plugin_id, &hook, status.code().description()); + self.schedule_recovery(); + GuardrailResult { + operation: Some(litellm_python_extension_protocol::OperationResult { + ok: true, + ..Default::default() + }), + decision: GuardrailDecision::Allow.into(), + ..Default::default() + } + } + } + } + + pub fn enqueue_callback(&self, event: CallbackEvent) -> Result<(), &'static str> { + let plugin_id = event.plugin_id.clone(); + match self.callback_tx.try_send(event) { + Ok(()) => Ok(()), + Err(mpsc::error::TrySendError::Full(_)) => { + self.record_bypass(&plugin_id, "callback", "queue_full"); + Err("callback queue is full") + } + Err(mpsc::error::TrySendError::Closed(_)) => { + self.record_bypass(&plugin_id, "callback", "queue_closed"); + Err("callback queue is closed") + } + } + } + + pub(crate) fn record_stream_bypass(&self, reason: &str) { + self.record_bypass("stream", "transform", reason); + self.schedule_recovery(); + } + + pub async fn transform_stream(&self, frames: S) -> Result, Status> + where + S: futures_util::Stream + Send + 'static, + { + let mut request = Request::new(frames); + request.set_timeout(self.settings.hook_timeout); + let mut stub = self.stub.clone(); + stub.transform_stream(request) + .await + .map(tonic::Response::into_inner) + } + + pub async fn retire(&self, revision_id: String) { + let mut request = Request::new(RetireRevisionRequest { revision_id }); + request.set_timeout(self.settings.hook_timeout); + let mut stub = self.stub.clone(); + let _ = stub.retire_revision(request).await; + } + + async fn activate(&self) -> Result, Status> { + let mut capabilities_request = Request::new(GetCapabilitiesRequest { + protocol_major: PROTOCOL_MAJOR, + protocol_minor: PROTOCOL_MINOR, + }); + capabilities_request.set_timeout(self.settings.connect_timeout); + let mut stub = self.stub.clone(); + let capabilities = stub + .get_capabilities(capabilities_request) + .await? + .into_inner(); + if capabilities.protocol_major != PROTOCOL_MAJOR { + return Err(Status::failed_precondition(format!( + "protocol major {} does not match {PROTOCOL_MAJOR}", + capabilities.protocol_major + ))); + } + let has_callbacks = self + .manifest + .extensions + .iter() + .any(|extension| extension.kind == ManifestExtensionKind::Callback); + if has_callbacks && !capabilities.supports_callback_batching { + return Err(Status::failed_precondition( + "extension host does not support callback batching", + )); + } + if capabilities.max_callback_batch_size > 0 + && self.settings.callback_batch_size > capabilities.max_callback_batch_size as usize + { + return Err(Status::failed_precondition( + "callback batch size exceeds extension host capability", + )); + } + let extensions = self + .manifest + .specs() + .map_err(|error| Status::invalid_argument(error.to_string()))?; + let mut prepare_request = Request::new(PrepareRevisionRequest { + revision_id: self.manifest.revision_id.clone(), + extensions, + }); + prepare_request.set_timeout(self.settings.hook_timeout); + let prepared = stub.prepare_revision(prepare_request).await?.into_inner(); + let operation = prepared + .operation + .ok_or_else(|| Status::internal("PrepareRevision omitted operation"))?; + if !operation.ok && operation.error_code != ErrorCode::AlreadyExists as i32 { + return Err(Status::failed_precondition(operation.error_message)); + } + if !capabilities.supports_duplex_streaming + && prepared.extensions.iter().any(|descriptor| { + descriptor.hooks.iter().any(|hook| { + matches!( + hook.as_str(), + "async_post_call_streaming_hook" + | "async_post_call_streaming_iterator_hook" + ) + }) + }) + { + return Err(Status::failed_precondition( + "extension host does not support required duplex streaming hooks", + )); + } + let mut commit_request = Request::new(CommitRevisionRequest { + revision_id: self.manifest.revision_id.clone(), + }); + commit_request.set_timeout(self.settings.hook_timeout); + let committed = stub.commit_revision(commit_request).await?.into_inner(); + if !committed.ok { + return Err(Status::failed_precondition(committed.error_message)); + } + self.mark_healthy(); + Ok(prepared.extensions) + } + + async fn callback_worker(self: Arc, mut receiver: mpsc::Receiver) { + while let Some(first) = receiver.recv().await { + let mut events = vec![first]; + while events.len() < self.settings.callback_batch_size { + match receiver.try_recv() { + Ok(event) => events.push(event), + Err(_) => break, + } + } + let mut request = Request::new(PublishCallbackEventsRequest { + events: events.clone(), + }); + request.set_timeout(self.settings.hook_timeout); + let mut stub = self.stub.clone(); + match stub.publish_callback_events(request).await { + Ok(response) => { + let operations = response.into_inner().operations; + let mut all_ok = operations.len() == events.len(); + for (index, event) in events.iter().enumerate() { + match operations.get(index) { + Some(operation) if operation.ok => {} + Some(operation) => { + all_ok = false; + self.record_bypass( + &event.plugin_id, + "callback", + &operation_reason(operation), + ); + } + None => { + all_ok = false; + self.record_bypass( + &event.plugin_id, + "callback", + "missing_operation", + ); + } + } + } + if all_ok { + self.mark_healthy(); + } + } + Err(status) => { + for event in events { + self.record_bypass( + &event.plugin_id, + "callback", + status.code().description(), + ); + } + self.schedule_recovery(); + } + } + } + } + + fn schedule_recovery(&self) { + if self.recovering.swap(true, Ordering::AcqRel) { + return; + } + let client = self.clone(); + tokio::spawn(async move { + let mut delay = Duration::from_millis(250); + loop { + match client.activate().await { + Ok(_) => break, + Err(status) => { + client.mark_unhealthy(status.code().to_string()); + tokio::time::sleep(delay).await; + delay = (delay * 2).min(Duration::from_secs(5)); + } + } + } + client.recovering.store(false, Ordering::Release); + }); + } + + fn record_bypass(&self, plugin_id: &str, hook: &str, reason: &str) { + if let Ok(mut counts) = self.bypass_counts.lock() { + *counts + .entry((plugin_id.to_string(), hook.to_string(), reason.to_string())) + .or_insert(0) += 1; + } + self.mark_unhealthy(reason.to_string()); + eprintln!("python_extension_host_bypass plugin={plugin_id} hook={hook} reason={reason}"); + } + + fn mark_healthy(&self) { + if let Ok(mut health) = self.health.write() { + *health = ExtensionHostHealth { + healthy: true, + reason: None, + }; + } + } + + fn mark_unhealthy(&self, reason: String) { + if let Ok(mut health) = self.health.write() { + *health = ExtensionHostHealth { + healthy: false, + reason: Some(reason), + }; + } + } +} + +fn is_transient(status: &Status) -> bool { + matches!( + status.code(), + Code::Unavailable | Code::DeadlineExceeded | Code::Cancelled | Code::Unknown + ) +} + +fn operation_reason(operation: &litellm_python_extension_protocol::OperationResult) -> String { + let code = ErrorCode::try_from(operation.error_code).unwrap_or(ErrorCode::Unspecified); + if operation.error_message.is_empty() { + format!("{code:?}") + } else { + format!("{code:?}:{}", operation.error_message) + } +} diff --git a/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/config.rs b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/config.rs new file mode 100644 index 00000000000..6a385aa599d --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/config.rs @@ -0,0 +1,106 @@ +use litellm_python_extension_protocol::{ExtensionKind, ExtensionSpec}; +use serde::Deserialize; + +#[derive(Clone, Debug, Deserialize)] +pub struct PythonExtensionManifest { + pub revision_id: String, + pub extensions: Vec, +} + +#[derive(Clone, Debug, Deserialize)] +pub struct ManifestExtension { + pub id: String, + pub kind: ManifestExtensionKind, + pub entrypoint: String, + #[serde(default)] + pub constructor: serde_json::Value, +} + +#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub enum ManifestExtensionKind { + Callback, + Guardrail, +} + +impl PythonExtensionManifest { + pub fn specs(&self) -> Result, serde_json::Error> { + self.extensions + .iter() + .map(|extension| { + Ok(ExtensionSpec { + id: extension.id.clone(), + kind: match extension.kind { + ManifestExtensionKind::Callback => ExtensionKind::Callback.into(), + ManifestExtensionKind::Guardrail => ExtensionKind::Guardrail.into(), + }, + entrypoint: extension.entrypoint.clone(), + constructor_json: serde_json::to_vec(&extension.constructor)?, + }) + }) + .collect() + } +} + +#[derive(Clone, Debug)] +pub struct PythonExtensionSettings { + pub endpoint: String, + pub token: String, + pub connect_timeout: std::time::Duration, + pub hook_timeout: std::time::Duration, + pub callback_queue_size: usize, + pub callback_batch_size: usize, +} + +impl PythonExtensionSettings { + pub fn from_env() -> Result, String> { + let Some(endpoint) = non_empty_env("LITELLM_PYTHON_EXTENSION_HOST_ENDPOINT") else { + return Ok(None); + }; + let token = non_empty_env("LITELLM_PYTHON_EXTENSION_HOST_TOKEN").ok_or_else(|| { + "LITELLM_PYTHON_EXTENSION_HOST_TOKEN is required when the endpoint is configured" + .to_string() + })?; + Ok(Some(Self { + endpoint, + token, + connect_timeout: seconds_env("LITELLM_PYTHON_EXTENSION_CONNECT_TIMEOUT_SECONDS", 5.0)?, + hook_timeout: seconds_env("LITELLM_PYTHON_EXTENSION_HOOK_TIMEOUT_SECONDS", 30.0)?, + callback_queue_size: usize_env("LITELLM_PYTHON_EXTENSION_CALLBACK_QUEUE_SIZE", 1_000)?, + callback_batch_size: usize_env("LITELLM_PYTHON_EXTENSION_CALLBACK_BATCH_SIZE", 50)?, + })) + } +} + +fn non_empty_env(name: &str) -> Option { + std::env::var(name) + .ok() + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + +fn seconds_env(name: &str, default: f64) -> Result { + let value = match non_empty_env(name) { + Some(value) => value + .parse::() + .map_err(|error| format!("{name} must be a number: {error}"))?, + None => default, + }; + if !value.is_finite() || value <= 0.0 { + return Err(format!("{name} must be greater than zero")); + } + Ok(std::time::Duration::from_secs_f64(value)) +} + +fn usize_env(name: &str, default: usize) -> Result { + let value = match non_empty_env(name) { + Some(value) => value + .parse::() + .map_err(|error| format!("{name} must be an integer: {error}"))?, + None => default, + }; + if value == 0 { + return Err(format!("{name} must be greater than zero")); + } + Ok(value) +} diff --git a/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/mod.rs b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/mod.rs new file mode 100644 index 00000000000..18365658832 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/mod.rs @@ -0,0 +1,13 @@ +mod adapters; +mod client; +pub mod config; +mod stream; + +pub use adapters::{RemoteCustomGuardrail, RemoteCustomLogger, RemoteExtensions}; +pub use client::{ + ActivationState, ExtensionHostHealth, InitializationError, PythonExtensionClient, +}; +pub use stream::RemoteStreamTransformer; + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/stream.rs b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/stream.rs new file mode 100644 index 00000000000..5a9edefe2fb --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/stream.rs @@ -0,0 +1,286 @@ +use std::collections::VecDeque; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex}; + +use futures_util::{Stream, StreamExt}; +use litellm_core::CoreError; +use litellm_python_extension_protocol::{ + AuthContext, InvocationContext, PublicError, StreamFrame, StreamFrameKind, StreamOpen, +}; +use serde_json::Value; +use tokio::sync::mpsc; +use tokio_stream::wrappers::ReceiverStream; + +use super::adapters::next_invocation_id; +use super::client::PythonExtensionClient; + +#[derive(Clone)] +pub struct RemoteStreamTransformer { + plugin_id: String, + client: Arc, + iterator_hook: bool, +} + +impl RemoteStreamTransformer { + pub fn new(plugin_id: String, client: Arc, iterator_hook: bool) -> Self { + Self { + plugin_id, + client, + iterator_hook, + } + } + + pub fn transform( + &self, + request: Value, + auth: AuthContext, + input: S, + ) -> ReceiverStream> + where + S: Stream> + Send + Unpin + 'static, + { + let stream_id = next_invocation_id(); + let (frame_tx, frame_rx) = mpsc::channel(8); + let (output_tx, output_rx) = mpsc::channel(8); + let pending = Arc::new(Mutex::new(VecDeque::new())); + let failed = Arc::new(AtomicBool::new(false)); + let terminal_error = Arc::new(Mutex::new(None)); + let producer = tokio::spawn(produce_frames( + frame_tx, + output_tx.clone(), + input, + pending.clone(), + failed.clone(), + terminal_error.clone(), + StreamOpen { + context: Some(InvocationContext { + request_id: stream_id.clone(), + invocation_id: stream_id.clone(), + active_revision: self.client.manifest().revision_id.clone(), + api_surface: "stream".to_string(), + call_type: "stream".to_string(), + trace_context: Default::default(), + }), + plugin_id: self.plugin_id.clone(), + request_json: serde_json::to_vec(&request).unwrap_or_else(|_| b"{}".to_vec()), + auth: Some(auth), + cache: None, + iterator_hook: self.iterator_hook, + }, + stream_id.clone(), + )); + let client = self.client.clone(); + let consumer_output = output_tx.clone(); + tokio::spawn(async move { + let result = client.transform_stream(ReceiverStream::new(frame_rx)).await; + let outcome = match result { + Ok(mut output) => { + consume_output(&mut output, &consumer_output, &terminal_error).await + } + Err(_) => ConsumeOutcome::Failed, + }; + match outcome { + ConsumeOutcome::Complete | ConsumeOutcome::Cancelled => producer.abort(), + ConsumeOutcome::Failed => { + failed.store(true, Ordering::Release); + client.record_stream_bypass("remote_stream_failed"); + let originals = pending + .lock() + .map(|mut values| values.drain(..).collect::>()) + .unwrap_or_default(); + for original in originals { + if consumer_output.send(Ok(original)).await.is_err() { + break; + } + } + let upstream_error = terminal_error + .lock() + .ok() + .and_then(|mut error| error.take()); + if let Some(error) = upstream_error { + let _ = consumer_output.send(Err(error)).await; + } + } + } + }); + drop(output_tx); + ReceiverStream::new(output_rx) + } +} + +async fn produce_frames( + sender: mpsc::Sender, + output: mpsc::Sender>, + mut input: S, + pending: Arc>>, + failed: Arc, + terminal_error: Arc>>, + open: StreamOpen, + stream_id: String, +) where + S: Stream> + Send + Unpin + 'static, +{ + if sender + .send(StreamFrame { + kind: StreamFrameKind::Open.into(), + stream_id: stream_id.clone(), + open: Some(open), + ..Default::default() + }) + .await + .is_err() + { + failed.store(true, Ordering::Release); + forward_remaining(&mut input, &output).await; + return; + } + while let Some(chunk) = input.next().await { + match chunk { + Ok(chunk) => { + if failed.load(Ordering::Acquire) { + forward_pending(&pending, &output).await; + if output.send(Ok(chunk)).await.is_err() { + return; + } + continue; + } + if let Ok(mut values) = pending.lock() { + values.push_back(chunk.clone()); + } + if sender + .send(StreamFrame { + kind: StreamFrameKind::InputChunk.into(), + stream_id: stream_id.clone(), + chunk_json: serde_json::to_vec(&chunk).ok(), + ..Default::default() + }) + .await + .is_err() + { + failed.store(true, Ordering::Release); + let originals = pending + .lock() + .map(|mut values| values.drain(..).collect::>()) + .unwrap_or_default(); + for original in originals { + if output.send(Ok(original)).await.is_err() { + return; + } + } + } else if failed.load(Ordering::Acquire) { + forward_pending(&pending, &output).await; + } + } + Err(error) => { + let message = error.to_string(); + if let Ok(mut terminal) = terminal_error.lock() { + *terminal = Some(error); + } + if failed.load(Ordering::Acquire) { + let upstream_error = terminal_error + .lock() + .ok() + .and_then(|mut value| value.take()); + if let Some(error) = upstream_error { + let _ = output.send(Err(error)).await; + } + return; + } + let _ = sender + .send(StreamFrame { + kind: StreamFrameKind::Error.into(), + stream_id, + error: Some(PublicError { + r#type: "upstream_error".to_string(), + message, + ..Default::default() + }), + ..Default::default() + }) + .await; + return; + } + } + } + let _ = sender + .send(StreamFrame { + kind: StreamFrameKind::End.into(), + stream_id, + ..Default::default() + }) + .await; +} + +async fn forward_remaining(input: &mut S, output: &mpsc::Sender>) +where + S: Stream> + Send + Unpin + 'static, +{ + while let Some(chunk) = input.next().await { + if output.send(chunk).await.is_err() { + return; + } + } +} + +async fn forward_pending( + pending: &Arc>>, + output: &mpsc::Sender>, +) { + let originals = pending + .lock() + .map(|mut values| values.drain(..).collect::>()) + .unwrap_or_default(); + for original in originals { + if output.send(Ok(original)).await.is_err() { + return; + } + } +} + +enum ConsumeOutcome { + Complete, + Failed, + Cancelled, +} + +async fn consume_output( + output: &mut tonic::Streaming, + sender: &mpsc::Sender>, + terminal_error: &Arc>>, +) -> ConsumeOutcome { + while let Some(frame) = output.next().await { + let Ok(frame) = frame else { + return ConsumeOutcome::Failed; + }; + match StreamFrameKind::try_from(frame.kind).unwrap_or(StreamFrameKind::Error) { + StreamFrameKind::OutputChunk => { + let Some(chunk_json) = frame.chunk_json else { + return ConsumeOutcome::Failed; + }; + let Ok(chunk) = serde_json::from_slice(&chunk_json) else { + return ConsumeOutcome::Failed; + }; + if sender.send(Ok(chunk)).await.is_err() { + return ConsumeOutcome::Cancelled; + } + } + StreamFrameKind::End => return ConsumeOutcome::Complete, + StreamFrameKind::Error => { + let upstream_error = terminal_error + .lock() + .ok() + .and_then(|mut error| error.take()); + if let Some(error) = upstream_error { + return if sender.send(Err(error)).await.is_ok() { + ConsumeOutcome::Complete + } else { + ConsumeOutcome::Cancelled + }; + } + return ConsumeOutcome::Failed; + } + _ => return ConsumeOutcome::Failed, + } + } + ConsumeOutcome::Failed +} diff --git a/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/tests.rs b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/tests.rs new file mode 100644 index 00000000000..d6c1c4b7380 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/integrations/python_extension_host/tests.rs @@ -0,0 +1,573 @@ +use std::pin::Pin; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::time::Duration; + +use futures_util::{Stream, StreamExt, stream}; +use litellm_python_extension_protocol::python_extension_host_server::{ + PythonExtensionHost, PythonExtensionHostServer, +}; +use litellm_python_extension_protocol::*; +use serde_json::json; +use tokio::sync::Notify; +use tonic::{Request, Response, Status}; + +use crate::integrations::custom_guardrail::{ + CustomGuardrail, CustomGuardrailRunner, GuardrailContext, GuardrailEventHook, GuardrailRequest, +}; +use crate::integrations::custom_logger::{ + CallType, CallbackTiming, CallbackValue, CustomLogger, ModelCallDetails, +}; + +use super::adapters::{RemoteCustomGuardrail, RemoteCustomLogger}; +use super::client::{ActivationState, PythonExtensionClient}; +use super::config::{ + ManifestExtension, ManifestExtensionKind, PythonExtensionManifest, PythonExtensionSettings, +}; +use super::stream::RemoteStreamTransformer; + +const TOKEN: &str = "rust-test-token"; + +#[derive(Default)] +struct MockHost { + callback_count: AtomicUsize, + callback_notify: Notify, + fail_operations: AtomicBool, +} + +#[tonic::async_trait] +impl PythonExtensionHost for MockHost { + async fn get_capabilities( + &self, + request: Request, + ) -> Result, Status> { + assert_token(&request); + Ok(Response::new(HostCapabilities { + protocol_major: 1, + protocol_minor: 0, + supported_hooks: vec![ + "async_pre_call_hook".to_string(), + "async_moderation_hook".to_string(), + "async_post_call_success_hook".to_string(), + ], + supports_duplex_streaming: true, + supports_callback_batching: true, + max_callback_batch_size: 50, + ..Default::default() + })) + } + + async fn prepare_revision( + &self, + request: Request, + ) -> Result, Status> { + assert_token(&request); + let descriptors = request + .get_ref() + .extensions + .iter() + .map(|extension| ExtensionDescriptor { + id: extension.id.clone(), + kind: extension.kind, + hooks: vec![ + "async_pre_call_hook".to_string(), + "async_moderation_hook".to_string(), + "async_post_call_success_hook".to_string(), + ], + ..Default::default() + }) + .collect(); + Ok(Response::new(PrepareRevisionResponse { + operation: Some(ok()), + extensions: descriptors, + })) + } + + async fn commit_revision( + &self, + request: Request, + ) -> Result, Status> { + assert_token(&request); + Ok(Response::new(ok())) + } + + async fn retire_revision( + &self, + request: Request, + ) -> Result, Status> { + assert_token(&request); + Ok(Response::new(ok())) + } + + async fn execute_guardrail( + &self, + request: Request, + ) -> Result, Status> { + assert_token(&request); + let invocation = request.into_inner(); + if self.fail_operations.load(Ordering::SeqCst) { + return Ok(Response::new(GuardrailResult { + operation: Some(operation_error()), + decision: GuardrailDecision::Error.into(), + ..Default::default() + })); + } + let body = if invocation.hook_phase == HookPhase::PostCall as i32 { + invocation.response_json.unwrap_or_default() + } else { + invocation.request_json + }; + let mut value: serde_json::Value = serde_json::from_slice(&body).unwrap(); + if value.get("block") == Some(&json!(true)) { + return Ok(Response::new(GuardrailResult { + operation: Some(ok()), + decision: GuardrailDecision::Block.into(), + public_error: Some(PublicError { + r#type: "GuardrailRaisedException".to_string(), + message: "blocked by mock".to_string(), + status_code: Some(400), + ..Default::default() + }), + ..Default::default() + })); + } + value["remote"] = json!(true); + let replacement = serde_json::to_vec(&value).unwrap(); + Ok(Response::new(GuardrailResult { + operation: Some(ok()), + decision: if invocation.hook_phase == HookPhase::PostCall as i32 { + GuardrailDecision::ReplaceResponse.into() + } else { + GuardrailDecision::ReplaceRequest.into() + }, + request_json: (invocation.hook_phase != HookPhase::PostCall as i32) + .then_some(replacement.clone()), + response_json: (invocation.hook_phase == HookPhase::PostCall as i32) + .then_some(replacement), + ..Default::default() + })) + } + + async fn publish_callback_events( + &self, + request: Request, + ) -> Result, Status> { + assert_token(&request); + let count = request.get_ref().events.len(); + self.callback_count.fetch_add(count, Ordering::SeqCst); + self.callback_notify.notify_one(); + Ok(Response::new(PublishCallbackEventsResponse { + operations: (0..count) + .map(|_| { + if self.fail_operations.load(Ordering::SeqCst) { + operation_error() + } else { + ok() + } + }) + .collect(), + })) + } + + type TransformStreamStream = Pin> + Send>>; + + async fn transform_stream( + &self, + request: Request>, + ) -> Result, Status> { + assert_token(&request); + let mut input = request.into_inner(); + let (sender, receiver) = tokio::sync::mpsc::channel(8); + tokio::spawn(async move { + while let Some(Ok(frame)) = input.next().await { + match StreamFrameKind::try_from(frame.kind).unwrap_or(StreamFrameKind::Error) { + StreamFrameKind::Open => { + let fail_stream = frame + .open + .as_ref() + .and_then(|open| { + serde_json::from_slice::(&open.request_json).ok() + }) + .and_then(|request| request.get("fail_stream").cloned()) + == Some(json!(true)); + if fail_stream { + let _ = sender + .send(Ok(StreamFrame { + kind: StreamFrameKind::Error.into(), + stream_id: frame.stream_id, + error: Some(PublicError { + r#type: "plugin_error".to_string(), + message: "plugin failed".to_string(), + ..Default::default() + }), + ..Default::default() + })) + .await; + break; + } + } + StreamFrameKind::InputChunk => { + let mut value: serde_json::Value = + serde_json::from_slice(&frame.chunk_json.unwrap()).unwrap(); + value["transformed"] = json!(true); + let _ = sender + .send(Ok(StreamFrame { + kind: StreamFrameKind::OutputChunk.into(), + stream_id: frame.stream_id, + chunk_json: Some(serde_json::to_vec(&value).unwrap()), + ..Default::default() + })) + .await; + } + StreamFrameKind::End => { + let _ = sender + .send(Ok(StreamFrame { + kind: StreamFrameKind::End.into(), + stream_id: frame.stream_id, + ..Default::default() + })) + .await; + break; + } + StreamFrameKind::Error => { + let _ = sender.send(Ok(frame)).await; + break; + } + _ => break, + } + } + }); + Ok(Response::new(Box::pin( + tokio_stream::wrappers::ReceiverStream::new(receiver), + ))) + } +} + +#[tokio::test] +async fn remote_guardrail_blocks_before_provider_and_mutates_all_phases() { + let (host, client, _) = start_client().await; + let guardrail = Arc::new(RemoteCustomGuardrail::new( + "remote".to_string(), + "guardrail-1".to_string(), + vec![ + GuardrailEventHook::PreCall, + GuardrailEventHook::DuringCall, + GuardrailEventHook::PostCall, + ], + client, + )); + let runner = CustomGuardrailRunner::new(vec![guardrail.clone()]); + let context = GuardrailContext::new(CallType::Ocr); + let provider_called = Arc::new(AtomicBool::new(false)); + let called = provider_called.clone(); + let result = runner + .run_before_provider( + GuardrailEventHook::PreCall, + &context, + GuardrailRequest::new(json!({"block": true})), + move |_| async move { + called.store(true, Ordering::SeqCst); + Ok(()) + }, + ) + .await; + assert!(result.is_err()); + assert!(!provider_called.load(Ordering::SeqCst)); + + let (request, _) = runner + .run_pre_call(&context, GuardrailRequest::new(json!({"model": "ocr"}))) + .await + .unwrap(); + assert_eq!(request.data["remote"], json!(true)); + let (response, _) = runner + .run_post_call(&context, GuardrailRequest::new(json!({"id": "response"}))) + .await + .unwrap(); + assert_eq!(response.data["remote"], json!(true)); + drop(host); +} + +#[tokio::test] +async fn remote_logger_batches_terminal_event_to_same_host() { + let (host, client, _) = start_client().await; + let logger = RemoteCustomLogger::new("callback-1".to_string(), client); + let details = ModelCallDetails::new("model", "provider", CallType::Ocr); + logger + .async_log_success_event( + &details, + &CallbackValue::new("ocr", json!({"id": "response"})), + CallbackTiming::new(1.0, 2.0), + ) + .await + .unwrap(); + tokio::time::timeout(Duration::from_secs(2), host.callback_notify.notified()) + .await + .unwrap(); + assert_eq!(host.callback_count.load(Ordering::SeqCst), 1); +} + +#[tokio::test] +async fn remote_stream_transformer_uses_one_duplex_rpc() { + let (_host, client, _) = start_client().await; + let transformer = RemoteStreamTransformer::new("callback-1".to_string(), client, true); + let output = transformer + .transform( + json!({"model": "test"}), + AuthContext::default(), + stream::iter(vec![Ok(json!({"value": "hello"}))]), + ) + .collect::>() + .await; + assert_eq!(output.len(), 1); + assert_eq!(output[0].as_ref().unwrap()["transformed"], json!(true)); +} + +#[tokio::test] +async fn unavailable_stream_fails_open_without_losing_buffered_chunks() { + let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + drop(listener); + let (client, activation) = + PythonExtensionClient::connect(settings(format!("http://{address}")), manifest()) + .await + .unwrap(); + assert!(matches!(activation, ActivationState::Degraded(_))); + let transformer = RemoteStreamTransformer::new("callback-1".to_string(), client, true); + let originals = (0..16) + .map(|value| json!({"value": value})) + .collect::>(); + let output = tokio::time::timeout( + Duration::from_secs(2), + transformer + .transform( + json!({"model": "test"}), + AuthContext::default(), + stream::iter(originals.clone().into_iter().map(Ok)), + ) + .collect::>(), + ) + .await + .unwrap(); + assert_eq!( + output.into_iter().collect::, _>>().unwrap(), + originals + ); +} + +#[tokio::test] +async fn upstream_stream_failure_is_preserved() { + let (_host, client, _) = start_client().await; + let transformer = RemoteStreamTransformer::new("callback-1".to_string(), client, true); + let output = transformer + .transform( + json!({"model": "test"}), + AuthContext::default(), + stream::iter(vec![ + Ok(json!({"value": "hello"})), + Err(litellm_core::CoreError::Network( + "upstream closed".to_string(), + )), + ]), + ) + .collect::>() + .await; + assert_eq!(output.len(), 2); + assert_eq!(output[0].as_ref().unwrap()["transformed"], json!(true)); + assert!(matches!( + &output[1], + Err(litellm_core::CoreError::Network(message)) if message == "upstream closed" + )); +} + +#[tokio::test] +async fn plugin_stream_failure_passes_through_original_chunks() { + let (_host, client, _) = start_client().await; + let transformer = RemoteStreamTransformer::new("callback-1".to_string(), client, true); + let originals = vec![json!({"value": 1}), json!({"value": 2})]; + let output = transformer + .transform( + json!({"fail_stream": true}), + AuthContext::default(), + stream::iter(originals.clone().into_iter().map(Ok)), + ) + .collect::>() + .await; + assert_eq!( + output.into_iter().collect::, _>>().unwrap(), + originals + ); +} + +#[tokio::test] +async fn dropping_transformed_stream_cancels_upstream_production() { + let (_host, client, _) = start_client().await; + let transformer = RemoteStreamTransformer::new("callback-1".to_string(), client, true); + let consumed = Arc::new(AtomicUsize::new(0)); + let observed = consumed.clone(); + let input = stream::iter(0..10_000).map(move |value| { + observed.fetch_add(1, Ordering::SeqCst); + Ok(json!({"value": value})) + }); + let mut output = transformer.transform(json!({}), AuthContext::default(), input); + assert!(output.next().await.is_some()); + drop(output); + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(consumed.load(Ordering::SeqCst) < 10_000); +} + +#[tokio::test] +async fn unavailable_host_fails_open_and_records_bypass() { + let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + drop(listener); + let settings = settings(format!("http://{address}")); + let (client, activation) = PythonExtensionClient::connect(settings, manifest()) + .await + .unwrap(); + assert!(matches!(activation, ActivationState::Degraded(_))); + let guardrail = RemoteCustomGuardrail::new( + "remote".to_string(), + "guardrail-1".to_string(), + vec![GuardrailEventHook::PreCall], + client.clone(), + ); + let decision = guardrail + .async_pre_call_hook( + &GuardrailContext::new(CallType::Ocr), + GuardrailRequest::new(json!({"model": "ocr"})), + ) + .await + .unwrap(); + assert!(matches!( + decision, + crate::integrations::custom_guardrail::GuardrailDecision::Allow(_) + )); + assert!(!client.bypass_counts().is_empty()); +} + +#[tokio::test] +async fn plugin_operation_errors_fail_open_and_record_bypasses() { + let (host, client, _) = start_client().await; + host.fail_operations.store(true, Ordering::SeqCst); + let guardrail = RemoteCustomGuardrail::new( + "remote".to_string(), + "guardrail-1".to_string(), + vec![GuardrailEventHook::PreCall], + client.clone(), + ); + let decision = guardrail + .async_pre_call_hook( + &GuardrailContext::new(CallType::Ocr), + GuardrailRequest::new(json!({"model": "ocr"})), + ) + .await + .unwrap(); + assert!(matches!( + decision, + crate::integrations::custom_guardrail::GuardrailDecision::Allow(_) + )); + + let logger = RemoteCustomLogger::new("callback-1".to_string(), client.clone()); + logger + .async_log_success_event( + &ModelCallDetails::new("model", "provider", CallType::Ocr), + &CallbackValue::new("ocr", json!({"id": "response"})), + CallbackTiming::new(1.0, 2.0), + ) + .await + .unwrap(); + tokio::time::timeout(Duration::from_secs(2), host.callback_notify.notified()) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(10)).await; + let counts = client.bypass_counts(); + assert!(counts.keys().any(|(plugin, hook, _)| { + plugin == "guardrail-1" && hook == &(HookPhase::PreCall as i32).to_string() + })); + assert!( + counts + .keys() + .any(|(plugin, hook, _)| plugin == "callback-1" && hook == "callback") + ); +} + +async fn start_client() -> ( + Arc, + Arc, + tokio::task::JoinHandle<()>, +) { + let host = Arc::new(MockHost::default()); + let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + drop(listener); + let service = PythonExtensionHostServer::from_arc(host.clone()); + let server = tokio::spawn(async move { + tonic::transport::Server::builder() + .add_service(service) + .serve(address) + .await + .unwrap(); + }); + let (client, activation) = + PythonExtensionClient::connect(settings(format!("http://{address}")), manifest()) + .await + .unwrap(); + assert!(matches!(activation, ActivationState::Active(_))); + (host, client, server) +} + +fn manifest() -> PythonExtensionManifest { + PythonExtensionManifest { + revision_id: "rust-test-revision".to_string(), + extensions: vec![ + ManifestExtension { + id: "guardrail-1".to_string(), + kind: ManifestExtensionKind::Guardrail, + entrypoint: "fixture.Guardrail".to_string(), + constructor: json!({"kwargs": {"guardrail_name": "remote"}}), + }, + ManifestExtension { + id: "callback-1".to_string(), + kind: ManifestExtensionKind::Callback, + entrypoint: "fixture.callback".to_string(), + constructor: json!({"callback_events": ["success", "failure"]}), + }, + ], + } +} + +fn settings(endpoint: String) -> PythonExtensionSettings { + PythonExtensionSettings { + endpoint, + token: TOKEN.to_string(), + connect_timeout: Duration::from_millis(200), + hook_timeout: Duration::from_millis(200), + callback_queue_size: 8, + callback_batch_size: 4, + } +} + +fn ok() -> OperationResult { + OperationResult { + ok: true, + ..Default::default() + } +} + +fn operation_error() -> OperationResult { + OperationResult { + ok: false, + error_code: ErrorCode::ExtensionFailed.into(), + error_message: "plugin failed".to_string(), + } +} + +fn assert_token(request: &Request) { + assert_eq!( + request + .metadata() + .get("x-litellm-extension-token") + .and_then(|value| value.to_str().ok()), + Some(TOKEN) + ); +} diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs index da3a486d4ee..6a6088b25cc 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -16,8 +16,15 @@ use litellm_ai_gateway::routes; use litellm_ai_gateway::state::AppState; use litellm_core::router::{Deployment, LiteLLMParams, Router}; +use litellm_ai_gateway::integrations::custom_guardrail::CustomGuardrail; use litellm_ai_gateway::integrations::custom_logger::CustomLogger; use litellm_ai_gateway::integrations::litellm_python_proxy_api::LiteLLMPythonProxyAPILogger; +use litellm_ai_gateway::integrations::python_extension_host::config::{ + PythonExtensionManifest, PythonExtensionSettings, +}; +use litellm_ai_gateway::integrations::python_extension_host::{ + ActivationState, PythonExtensionClient, RemoteExtensions, +}; #[cfg(feature = "python-config")] use litellm_ai_gateway::python; @@ -45,9 +52,14 @@ async fn main() { // Python proxy's /v1/callbacks/logs). Built here so the spawn lands on the // tokio runtime. `from_env` reads LITELLM_PROXY_BASE_URL + LITELLM_MASTER_KEY. let proxy_logger = LiteLLMPythonProxyAPILogger::from_env(); - let loggers: Vec> = vec![proxy_logger]; + let mut loggers: Vec> = vec![proxy_logger]; - let router = Arc::new(build_router()); + let (router, extension_manifest) = build_gateway_config(); + let router = Arc::new(router); + let (python_extension_host, remote_extensions) = + initialize_python_extensions(extension_manifest).await; + loggers.extend(remote_extensions.loggers); + let guardrails: Vec> = remote_extensions.guardrails; // Build the pre-warmed realtime pool and register each deployment's upstream // so the background replenisher starts warming it. `REALTIME_POOL_SIZE=0` @@ -71,6 +83,8 @@ async fn main() { router, master_key, loggers: Arc::new(loggers), + guardrails: Arc::new(guardrails), + python_extension_host, realtime_pool, }; @@ -121,20 +135,55 @@ fn resolve_port() -> u16 { /// Build the router. With the `python-config` feature and `LITELLM_CONFIG_PATH` /// set, load the resolved `model_list` from the proxy config via the embedded /// Python reader (load time only). Otherwise fall back to the env stand-in. -fn build_router() -> Router { +fn build_gateway_config() -> (Router, Option) { #[cfg(feature = "python-config")] if let Ok(config_path) = std::env::var("LITELLM_CONFIG_PATH") { - match python::config::load_router_from_config(&config_path) { - Ok(router) => { + match python::config::load_gateway_config_from_config(&config_path) { + Ok(config) => { eprintln!("loaded model_list from {config_path} via python config reader"); - return router; + return (config.router, Some(config.extension_manifest)); } Err(err) => { + if std::env::var("LITELLM_PYTHON_EXTENSION_HOST_ENDPOINT") + .is_ok_and(|endpoint| !endpoint.trim().is_empty()) + { + panic!("config load failed while Python extensions are enabled: {err}"); + } eprintln!("config load failed ({err}); falling back to env deployment"); } } } - build_router_from_env() + (build_router_from_env(), None) +} + +async fn initialize_python_extensions( + manifest: Option, +) -> (Option>, RemoteExtensions) { + let empty = RemoteExtensions { + guardrails: Vec::new(), + loggers: Vec::new(), + }; + let settings = PythonExtensionSettings::from_env() + .unwrap_or_else(|error| panic!("invalid Python extension settings: {error}")); + let Some(settings) = settings else { + return (None, empty); + }; + let manifest = manifest.unwrap_or(PythonExtensionManifest { + revision_id: "rust-empty-v1".to_string(), + extensions: Vec::new(), + }); + let (client, activation) = PythonExtensionClient::connect(settings, manifest.clone()) + .await + .unwrap_or_else(|error| panic!("Python extension host initialization failed: {error}")); + let descriptors = match activation { + ActivationState::Active(descriptors) => descriptors, + ActivationState::Degraded(reason) => { + eprintln!("Python extension host unavailable at startup; fail-open active: {reason}"); + Vec::new() + } + }; + let extensions = RemoteExtensions::from_manifest(&manifest, &descriptors, client.clone()); + (Some(client), extensions) } /// Build a minimal single-deployment `model_list` from the environment. diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index 95df566dc53..f7507ed6248 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -172,6 +172,7 @@ impl OcrLifecycleHooks { impl CallLifecycleHooks for OcrLifecycleHooks { type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; type DuringCallFuture<'a> = OcrFuture<'a, ProviderOcrRequest>; + type PostCallFuture<'a> = OcrFuture<'a, Value>; type SuccessFuture<'a> = OcrLogFuture<'a>; type FailureFuture<'a> = OcrLogFuture<'a>; @@ -191,6 +192,25 @@ impl CallLifecycleHooks for OcrLi Box::pin(async move { self.prepare_provider_request(request).await }) } + fn async_post_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + response: Value, + ) -> Self::PostCallFuture<'a> { + Box::pin(async move { + if self.guardrail_runner.is_empty() { + return Ok(response); + } + let context = guardrail_context(&self.request_metadata); + let (response, _) = self + .guardrail_runner + .run_post_call(&context, GuardrailRequest::new(response)) + .await + .map_err(guardrail_error_to_core_error)?; + Ok(response.data) + }) + } + fn async_log_success_event<'a>( &'a self, context: &'a CallLifecycleContext, diff --git a/litellm-rust/crates/ai-gateway/src/python/config.rs b/litellm-rust/crates/ai-gateway/src/python/config.rs index c028d3d6b51..8ae61d4535d 100644 --- a/litellm-rust/crates/ai-gateway/src/python/config.rs +++ b/litellm-rust/crates/ai-gateway/src/python/config.rs @@ -13,9 +13,19 @@ use litellm_core::router::{Deployment, Router}; use pyo3::prelude::*; use crate::gil; +use crate::integrations::python_extension_host::config::PythonExtensionManifest; + +pub struct LoadedGatewayConfig { + pub router: Router, + pub extension_manifest: PythonExtensionManifest, +} /// Load the router's `model_list` from `config_path` via the Python reader. pub fn load_router_from_config(config_path: &str) -> CoreResult { + load_gateway_config_from_config(config_path).map(|config| config.router) +} + +pub fn load_gateway_config_from_config(config_path: &str) -> CoreResult { gil::record_acquisition(); Python::attach(|py| { let model_list = py @@ -34,6 +44,22 @@ pub fn load_router_from_config(config_path: &str) -> CoreResult { let deployments: Vec = serde_json::from_str(&model_list_json) .map_err(|err| CoreError::Routing(format!("parsing model_list failed: {err}")))?; - Ok(Router::new(deployments)) + let manifest_json: String = py + .import("litellm.extensions.manifest") + .and_then(|module| module.getattr("manifest_json_from_config_path")) + .and_then(|reader| reader.call1((config_path,))) + .and_then(|encoded| encoded.extract()) + .map_err(|err| { + CoreError::Routing(format!("reading Python extension manifest failed: {err}")) + })?; + let extension_manifest: PythonExtensionManifest = serde_json::from_str(&manifest_json) + .map_err(|err| { + CoreError::Routing(format!("parsing Python extension manifest failed: {err}")) + })?; + + Ok(LoadedGatewayConfig { + router: Router::new(deployments), + extension_manifest, + }) }) } diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index 7e38d10c6ff..f1d7435bf40 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -167,6 +167,8 @@ mod tests { }])), master_key: master_key.map(Arc::from), loggers: Arc::new(Vec::new()), + guardrails: Arc::new(Vec::new()), + python_extension_host: None, realtime_pool: RealtimePool::disabled(), } } diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs index a94853e106d..427be3d79e6 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs @@ -317,6 +317,8 @@ mod tests { router: Arc::new(ModelRouter::default()), master_key: Some(Arc::from("master-key")), loggers: Arc::new(Vec::new()), + guardrails: Arc::new(Vec::new()), + python_extension_host: None, realtime_pool: RealtimePool::disabled(), } } diff --git a/litellm-rust/crates/ai-gateway/src/state.rs b/litellm-rust/crates/ai-gateway/src/state.rs index 3b61d8309ea..7d25355f661 100644 --- a/litellm-rust/crates/ai-gateway/src/state.rs +++ b/litellm-rust/crates/ai-gateway/src/state.rs @@ -3,7 +3,9 @@ use std::sync::Arc; use crate::io::realtime_pool::RealtimePool; use litellm_core::router::Router; +use crate::integrations::custom_guardrail::CustomGuardrail; use crate::integrations::custom_logger::CustomLogger; +use crate::integrations::python_extension_host::PythonExtensionClient; /// Shared application state handed to every route handler. #[derive(Clone)] @@ -14,6 +16,9 @@ pub struct AppState { pub master_key: Option>, /// Logging callbacks fanned out at the end of each realtime session. pub loggers: Arc>>, + pub guardrails: Arc>>, + /// Shared long-lived HTTP/2 client. `None` keeps the extension feature fully disabled. + pub python_extension_host: Option>, /// Pre-warmed upstream realtime connection pool. Disabled /// (`RealtimePool::disabled()`) when `REALTIME_POOL_SIZE=0`, in which case /// every realtime connect fresh-dials exactly as before. diff --git a/litellm-rust/crates/core/src/call_lifecycle/mod.rs b/litellm-rust/crates/core/src/call_lifecycle/mod.rs index d9b68a1b726..bc8f4f28aec 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/mod.rs @@ -25,6 +25,13 @@ pub trait CallLifecycleHooks: Send + Sync { ProviderReq: 'a, Resp: 'a; + type PostCallFuture<'a>: Future> + Send + 'a + where + Self: 'a, + InitialReq: 'a, + ProviderReq: 'a, + Resp: 'a; + type SuccessFuture<'a>: Future + Send + 'a where Self: 'a, @@ -46,6 +53,12 @@ pub trait CallLifecycleHooks: Send + Sync { request: InitialReq, ) -> Self::DuringCallFuture<'a>; + fn async_post_call_hook<'a>( + &'a self, + context: &'a CallLifecycleContext, + response: Resp, + ) -> Self::PostCallFuture<'a>; + fn async_log_success_event<'a>( &'a self, context: &'a CallLifecycleContext, @@ -144,6 +157,25 @@ impl<'a> CallLifecycle<'a> { let result = provider_call(provider_request).await; phases.push(self.finish_phase(&context, provider_phase)); + let result = match result { + Ok(response) => { + let post_call = self.start_phase(&context, CallLifecyclePhase::PostCall); + match hooks.async_post_call_hook(&context, response).await { + Ok(response) => { + phases.push(self.finish_phase(&context, post_call)); + Ok(response) + } + Err(error) => { + phases.push(self.finish_phase(&context, post_call)); + self.log_failure(&context, hooks, &error, call_start, &mut phases) + .await; + return Err(error); + } + } + } + Err(error) => Err(error), + }; + match &result { Ok(response) => { let success_phase = self.start_phase(&context, CallLifecyclePhase::SuccessCallback); @@ -253,6 +285,7 @@ mod tests { impl CallLifecycleHooks for RecordingHooks { type PreCallFuture<'a> = BoxFuture<'a, CoreResult>; type DuringCallFuture<'a> = BoxFuture<'a, CoreResult>; + type PostCallFuture<'a> = BoxFuture<'a, CoreResult>; type SuccessFuture<'a> = BoxFuture<'a, ()>; type FailureFuture<'a> = BoxFuture<'a, ()>; @@ -278,6 +311,17 @@ mod tests { }) } + fn async_post_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + response: String, + ) -> Self::PostCallFuture<'a> { + Box::pin(async move { + self.events.lock().unwrap().push("post_call"); + Ok(format!("{response}:post")) + }) + } + fn async_log_success_event<'a>( &'a self, _context: &'a CallLifecycleContext, @@ -286,7 +330,7 @@ mod tests { ) -> Self::SuccessFuture<'a> { Box::pin(async move { assert!(timing.end_time >= timing.start_time); - assert_eq!(timing.phases.len(), 3); + assert_eq!(timing.phases.len(), 4); self.events.lock().unwrap().push("success"); }) } @@ -306,6 +350,7 @@ mod tests { impl CallLifecycleHooks for RecordingHooks { type PreCallFuture<'a> = BoxFuture<'a, CoreResult>; type DuringCallFuture<'a> = BoxFuture<'a, CoreResult>; + type PostCallFuture<'a> = BoxFuture<'a, CoreResult>; type SuccessFuture<'a> = BoxFuture<'a, ()>; type FailureFuture<'a> = BoxFuture<'a, ()>; @@ -331,6 +376,17 @@ mod tests { }) } + fn async_post_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + response: String, + ) -> Self::PostCallFuture<'a> { + Box::pin(async move { + self.events.lock().unwrap().push("post_call"); + Ok(format!("{response}:post")) + }) + } + fn async_log_success_event<'a>( &'a self, _context: &'a CallLifecycleContext, @@ -370,8 +426,11 @@ mod tests { .await .expect("call succeeds"); - assert_eq!(response, "response"); - assert_eq!(hooks.events(), vec!["pre_call", "during_call", "success"]); + assert_eq!(response, "response:post"); + assert_eq!( + hooks.events(), + vec!["pre_call", "during_call", "post_call", "success"] + ); } #[tokio::test] @@ -408,7 +467,10 @@ mod tests { .await .expect("call succeeds"); - assert_eq!(response, "response"); - assert_eq!(hooks.events(), vec!["pre_call", "during_call", "success"]); + assert_eq!(response, "response:post"); + assert_eq!( + hooks.events(), + vec!["pre_call", "during_call", "post_call", "success"] + ); } } diff --git a/litellm-rust/crates/core/src/call_lifecycle/types.rs b/litellm-rust/crates/core/src/call_lifecycle/types.rs index 8819c8830d2..8e52172779e 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/types.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/types.rs @@ -33,6 +33,7 @@ pub enum CallLifecyclePhase { PreCall, DuringCall, ProviderCall, + PostCall, SuccessCallback, FailureCallback, } @@ -43,6 +44,7 @@ impl CallLifecyclePhase { Self::PreCall => "pre_call", Self::DuringCall => "during_call", Self::ProviderCall => "provider_call", + Self::PostCall => "post_call", Self::SuccessCallback => "success_callback", Self::FailureCallback => "failure_callback", } diff --git a/litellm-rust/crates/core/src/responses/instrumentation.rs b/litellm-rust/crates/core/src/responses/instrumentation.rs index ec04571da14..7d22d8e9bec 100644 --- a/litellm-rust/crates/core/src/responses/instrumentation.rs +++ b/litellm-rust/crates/core/src/responses/instrumentation.rs @@ -210,6 +210,7 @@ type LifecycleFuture<'a, T> = Pin> + Send impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation { type PreCallFuture<'a> = LifecycleFuture<'a, ()>; type DuringCallFuture<'a> = LifecycleFuture<'a, ()>; + type PostCallFuture<'a> = LifecycleFuture<'a, ()>; type SuccessFuture<'a> = Pin + Send + 'a>>; type FailureFuture<'a> = Pin + Send + 'a>>; @@ -229,6 +230,14 @@ impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation { Box::pin(async move { Ok(request) }) } + fn async_post_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + response: (), + ) -> Self::PostCallFuture<'a> { + Box::pin(async move { Ok(response) }) + } + fn async_log_success_event<'a>( &'a self, _context: &'a CallLifecycleContext, diff --git a/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs b/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs index 656ba033b62..e7610efafed 100644 --- a/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs +++ b/litellm-rust/crates/core/tests/workspace_crate_allowlist.rs @@ -1,4 +1,4 @@ -//! Enforcement: the litellm-rust workspace has exactly three crates. +//! Enforcement: the litellm-rust workspace has exactly four crates. //! //! `core` (pure translation), `ai-gateway` (routes + all network I/O), and //! `python-bridge` (the PyO3 cdylib). Adding or removing a crate must be a @@ -16,10 +16,20 @@ use std::path::{Path, PathBuf}; /// The one true crate set. Update BOTH this and `litellm-rust/AGENTS.md` when the /// workspace legitimately gains or loses a crate. -const EXPECTED_MEMBERS: &[&str] = &["crates/core", "crates/ai-gateway", "crates/python-bridge"]; +const EXPECTED_MEMBERS: &[&str] = &[ + "crates/core", + "crates/ai-gateway", + "crates/python-bridge", + "crates/python-extension-protocol", +]; /// The crate subdirectory names that must exist under `crates/`. -const EXPECTED_CRATE_DIRS: &[&str] = &["core", "ai-gateway", "python-bridge"]; +const EXPECTED_CRATE_DIRS: &[&str] = &[ + "core", + "ai-gateway", + "python-bridge", + "python-extension-protocol", +]; const MISMATCH: &str = "litellm-rust crate set changed — update this allowlist AND litellm-rust/AGENTS.md, and justify the crate per the rule (crate = layer needing independent compilation / its own deps / a separate artifact)."; diff --git a/litellm-rust/crates/python-extension-protocol/Cargo.toml b/litellm-rust/crates/python-extension-protocol/Cargo.toml new file mode 100644 index 00000000000..f5dcdfa5b2e --- /dev/null +++ b/litellm-rust/crates/python-extension-protocol/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "litellm-python-extension-protocol" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +prost.workspace = true +tonic.workspace = true +tonic-prost.workspace = true + +[build-dependencies] +prost-build = "0.14.4" +protoc-bin-vendored = "3.2.0" +tonic-prost-build = { version = "0.14.6", default-features = false } diff --git a/litellm-rust/crates/python-extension-protocol/build.rs b/litellm-rust/crates/python-extension-protocol/build.rs new file mode 100644 index 00000000000..caea5a557c6 --- /dev/null +++ b/litellm-rust/crates/python-extension-protocol/build.rs @@ -0,0 +1,18 @@ +use std::error::Error; +use std::path::PathBuf; + +fn main() -> Result<(), Box> { + let protocol_root = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("../../../proto"); + let protocol = protocol_root.join("litellm/python_extension/v1/extension_host.proto"); + let mut prost_config = prost_build::Config::new(); + prost_config.protoc_executable(protoc_bin_vendored::protoc_bin_path()?); + tonic_prost_build::configure() + .build_transport(false) + .compile_with_config( + prost_config, + &[protocol.as_path()], + &[protocol_root.as_path()], + )?; + println!("cargo:rerun-if-changed={}", protocol.display()); + Ok(()) +} diff --git a/litellm-rust/crates/python-extension-protocol/src/lib.rs b/litellm-rust/crates/python-extension-protocol/src/lib.rs new file mode 100644 index 00000000000..1926c369ed8 --- /dev/null +++ b/litellm-rust/crates/python-extension-protocol/src/lib.rs @@ -0,0 +1,48 @@ +pub mod generated { + tonic::include_proto!("litellm.python_extension.v1"); +} + +pub use generated::*; + +#[cfg(test)] +mod tests { + use prost::Message; + + use super::{CacheRef, GuardrailDecision, StreamFrame, StreamFrameKind}; + + #[test] + fn cache_reference_round_trips_without_gateway_objects() -> Result<(), prost::DecodeError> { + let reference = CacheRef { + invocation_id: "invocation-1".to_string(), + opaque_handle: "opaque".to_string(), + }; + assert_eq!( + CacheRef::decode(reference.encode_to_vec().as_slice())?, + reference + ); + Ok(()) + } + + #[test] + fn duplex_frame_round_trips() -> Result<(), prost::DecodeError> { + let frame = StreamFrame { + kind: StreamFrameKind::InputChunk.into(), + stream_id: "stream-1".to_string(), + chunk_json: Some(br#"{"value":"hello"}"#.to_vec()), + ..Default::default() + }; + assert_eq!( + StreamFrame::decode(frame.encode_to_vec().as_slice())?, + frame + ); + Ok(()) + } + + #[test] + fn block_and_transport_error_are_distinct_outcomes() { + assert_ne!( + GuardrailDecision::Block as i32, + GuardrailDecision::Error as i32 + ); + } +}