feat(rust): integrate Python extension host

This commit is contained in:
Yujong Lee 2026-09-01 22:08:34 -07:00
parent ae13524e3a
commit 8a13e2247b
27 changed files with 2561 additions and 22 deletions

View file

@ -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

398
litellm-rust/Cargo.lock generated
View file

@ -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",

View file

@ -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]

View file

@ -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 }

View file

@ -197,6 +197,7 @@ impl CallLifecycleHooks<PreparedAudioTranscriptionRequest, ProviderAudioTranscri
{
type PreCallFuture<'a> = 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<PreparedAudioTranscriptionRequest, ProviderAudioTranscri
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 { Ok(response) })
}
fn async_log_success_event<'a>(
&'a self,
context: &'a CallLifecycleContext,

View file

@ -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<F, Fut, T>(
&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,

View file

@ -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",
}
}
}

View file

@ -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;

View file

@ -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<GuardrailEventHook>,
client: Arc<PythonExtensionClient>,
}
impl RemoteCustomGuardrail {
pub fn new(
name: String,
plugin_id: String,
hooks: Vec<GuardrailEventHook>,
client: Arc<PythonExtensionClient>,
) -> 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<PythonExtensionClient>,
}
impl RemoteCustomLogger {
pub fn new(plugin_id: String, client: Arc<PythonExtensionClient>) -> Self {
Self::with_events(plugin_id, true, true, client)
}
fn with_events(
plugin_id: String,
success_enabled: bool,
failure_enabled: bool,
client: Arc<PythonExtensionClient>,
) -> 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<Arc<dyn CustomGuardrail>>,
pub loggers: Vec<Arc<dyn CustomLogger>>,
}
impl RemoteExtensions {
pub fn from_manifest(
manifest: &PythonExtensionManifest,
descriptors: &[litellm_python_extension_protocol::ExtensionDescriptor],
client: Arc<PythonExtensionClient>,
) -> Self {
let descriptor_hooks: HashMap<&str, &[String]> = descriptors
.iter()
.map(|descriptor| (descriptor.id.as_str(), descriptor.hooks.as_slice()))
.collect();
let mut guardrails: Vec<Arc<dyn CustomGuardrail>> = Vec::new();
let mut loggers: Vec<Arc<dyn CustomLogger>> = 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<OperationResult>,
decision: i32,
request_json: Option<Vec<u8>>,
response_json: Option<Vec<u8>>,
public_error: Option<litellm_python_extension_protocol::PublicError>,
original: GuardrailRequest,
) -> Result<GuardrailDecision, GuardrailError> {
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<String> {
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<GuardrailEventHook> {
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}")
}

View file

@ -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<InterceptedService<Channel, TokenInterceptor>>;
#[derive(Clone)]
struct TokenInterceptor {
token: MetadataValue<Ascii>,
}
impl Interceptor for TokenInterceptor {
fn call(&mut self, mut request: Request<()>) -> Result<Request<()>, 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<String>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ActivationState {
Active(Vec<ExtensionDescriptor>),
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<CallbackEvent>,
health: Arc<RwLock<ExtensionHostHealth>>,
bypass_counts: Arc<Mutex<HashMap<(String, String, String), u64>>>,
recovering: Arc<AtomicBool>,
}
impl PythonExtensionClient {
pub async fn connect(
settings: PythonExtensionSettings,
manifest: PythonExtensionManifest,
) -> Result<(Arc<Self>, 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<S>(&self, frames: S) -> Result<Streaming<StreamFrame>, Status>
where
S: futures_util::Stream<Item = StreamFrame> + 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<Vec<ExtensionDescriptor>, 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<Self>, mut receiver: mpsc::Receiver<CallbackEvent>) {
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)
}
}

View file

@ -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<ManifestExtension>,
}
#[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<Vec<ExtensionSpec>, 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<Option<Self>, 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<String> {
std::env::var(name)
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
fn seconds_env(name: &str, default: f64) -> Result<std::time::Duration, String> {
let value = match non_empty_env(name) {
Some(value) => value
.parse::<f64>()
.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<usize, String> {
let value = match non_empty_env(name) {
Some(value) => value
.parse::<usize>()
.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)
}

View file

@ -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;

View file

@ -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<PythonExtensionClient>,
iterator_hook: bool,
}
impl RemoteStreamTransformer {
pub fn new(plugin_id: String, client: Arc<PythonExtensionClient>, iterator_hook: bool) -> Self {
Self {
plugin_id,
client,
iterator_hook,
}
}
pub fn transform<S>(
&self,
request: Value,
auth: AuthContext,
input: S,
) -> ReceiverStream<Result<Value, CoreError>>
where
S: Stream<Item = Result<Value, CoreError>> + 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::<Vec<_>>())
.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<S>(
sender: mpsc::Sender<StreamFrame>,
output: mpsc::Sender<Result<Value, CoreError>>,
mut input: S,
pending: Arc<Mutex<VecDeque<Value>>>,
failed: Arc<AtomicBool>,
terminal_error: Arc<Mutex<Option<CoreError>>>,
open: StreamOpen,
stream_id: String,
) where
S: Stream<Item = Result<Value, CoreError>> + 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::<Vec<_>>())
.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<S>(input: &mut S, output: &mpsc::Sender<Result<Value, CoreError>>)
where
S: Stream<Item = Result<Value, CoreError>> + Send + Unpin + 'static,
{
while let Some(chunk) = input.next().await {
if output.send(chunk).await.is_err() {
return;
}
}
}
async fn forward_pending(
pending: &Arc<Mutex<VecDeque<Value>>>,
output: &mpsc::Sender<Result<Value, CoreError>>,
) {
let originals = pending
.lock()
.map(|mut values| values.drain(..).collect::<Vec<_>>())
.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<StreamFrame>,
sender: &mpsc::Sender<Result<Value, CoreError>>,
terminal_error: &Arc<Mutex<Option<CoreError>>>,
) -> 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
}

View file

@ -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<GetCapabilitiesRequest>,
) -> Result<Response<HostCapabilities>, 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<PrepareRevisionRequest>,
) -> Result<Response<PrepareRevisionResponse>, 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<CommitRevisionRequest>,
) -> Result<Response<OperationResult>, Status> {
assert_token(&request);
Ok(Response::new(ok()))
}
async fn retire_revision(
&self,
request: Request<RetireRevisionRequest>,
) -> Result<Response<OperationResult>, Status> {
assert_token(&request);
Ok(Response::new(ok()))
}
async fn execute_guardrail(
&self,
request: Request<GuardrailInvocation>,
) -> Result<Response<GuardrailResult>, 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<PublishCallbackEventsRequest>,
) -> Result<Response<PublishCallbackEventsResponse>, 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<Box<dyn Stream<Item = Result<StreamFrame, Status>> + Send>>;
async fn transform_stream(
&self,
request: Request<tonic::Streaming<StreamFrame>>,
) -> Result<Response<Self::TransformStreamStream>, 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::<serde_json::Value>(&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::<Vec<_>>()
.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::<Vec<_>>();
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::<Vec<_>>(),
)
.await
.unwrap();
assert_eq!(
output.into_iter().collect::<Result<Vec<_>, _>>().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::<Vec<_>>()
.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::<Vec<_>>()
.await;
assert_eq!(
output.into_iter().collect::<Result<Vec<_>, _>>().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<MockHost>,
Arc<PythonExtensionClient>,
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<T>(request: &Request<T>) {
assert_eq!(
request
.metadata()
.get("x-litellm-extension-token")
.and_then(|value| value.to_str().ok()),
Some(TOKEN)
);
}

View file

@ -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<Arc<dyn CustomLogger>> = vec![proxy_logger];
let mut loggers: Vec<Arc<dyn CustomLogger>> = 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<Arc<dyn CustomGuardrail>> = 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<PythonExtensionManifest>) {
#[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<PythonExtensionManifest>,
) -> (Option<Arc<PythonExtensionClient>>, 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.

View file

@ -172,6 +172,7 @@ impl OcrLifecycleHooks {
impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> 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<PreparedOcrRequest, ProviderOcrRequest, Value> 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,

View file

@ -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<Router> {
load_gateway_config_from_config(config_path).map(|config| config.router)
}
pub fn load_gateway_config_from_config(config_path: &str) -> CoreResult<LoadedGatewayConfig> {
gil::record_acquisition();
Python::attach(|py| {
let model_list = py
@ -34,6 +44,22 @@ pub fn load_router_from_config(config_path: &str) -> CoreResult<Router> {
let deployments: Vec<Deployment> = 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,
})
})
}

View file

@ -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(),
}
}

View file

@ -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(),
}
}

View file

@ -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<Arc<str>>,
/// Logging callbacks fanned out at the end of each realtime session.
pub loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
pub guardrails: Arc<Vec<Arc<dyn CustomGuardrail>>>,
/// Shared long-lived HTTP/2 client. `None` keeps the extension feature fully disabled.
pub python_extension_host: Option<Arc<PythonExtensionClient>>,
/// 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.

View file

@ -25,6 +25,13 @@ pub trait CallLifecycleHooks<InitialReq, ProviderReq, Resp>: Send + Sync {
ProviderReq: 'a,
Resp: 'a;
type PostCallFuture<'a>: Future<Output = CoreResult<Resp>> + Send + 'a
where
Self: 'a,
InitialReq: 'a,
ProviderReq: 'a,
Resp: 'a;
type SuccessFuture<'a>: Future<Output = ()> + Send + 'a
where
Self: 'a,
@ -46,6 +53,12 @@ pub trait CallLifecycleHooks<InitialReq, ProviderReq, Resp>: 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<String, String, String> for RecordingHooks {
type PreCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
type DuringCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
type PostCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
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<RecordingRequest, String, String> for RecordingHooks {
type PreCallFuture<'a> = BoxFuture<'a, CoreResult<RecordingRequest>>;
type DuringCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
type PostCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
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"]
);
}
}

View file

@ -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",
}

View file

@ -210,6 +210,7 @@ type LifecycleFuture<'a, T> = Pin<Box<dyn Future<Output = CoreResult<T>> + Send
impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation {
type PreCallFuture<'a> = LifecycleFuture<'a, ()>;
type DuringCallFuture<'a> = LifecycleFuture<'a, ()>;
type PostCallFuture<'a> = LifecycleFuture<'a, ()>;
type SuccessFuture<'a> = Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
type FailureFuture<'a> = Pin<Box<dyn Future<Output = ()> + 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,

View file

@ -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).";

View file

@ -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 }

View file

@ -0,0 +1,18 @@
use std::error::Error;
use std::path::PathBuf;
fn main() -> Result<(), Box<dyn Error>> {
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(())
}

View file

@ -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
);
}
}