mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(rust): integrate Python extension host
This commit is contained in:
parent
ae13524e3a
commit
8a13e2247b
27 changed files with 2561 additions and 22 deletions
|
|
@ -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
398
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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 }
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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).";
|
||||
|
||||
|
|
|
|||
16
litellm-rust/crates/python-extension-protocol/Cargo.toml
Normal file
16
litellm-rust/crates/python-extension-protocol/Cargo.toml
Normal 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 }
|
||||
18
litellm-rust/crates/python-extension-protocol/build.rs
Normal file
18
litellm-rust/crates/python-extension-protocol/build.rs
Normal 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(())
|
||||
}
|
||||
48
litellm-rust/crates/python-extension-protocol/src/lib.rs
Normal file
48
litellm-rust/crates/python-extension-protocol/src/lib.rs
Normal 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
|
||||
);
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue