refactor(rust): remove gateway, config, router, realtime, and trace-parity infrastructure

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yujong Lee 2026-09-16 16:00:07 +00:00
parent 9cd787386e
commit 96baeb8b04
162 changed files with 334 additions and 12315 deletions

View file

@ -1,73 +0,0 @@
name: ai-gateway image
on:
push:
paths:
- "litellm-rust/**"
- "litellm/**"
- "enterprise/**"
- "litellm-proxy-extras/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/workflows/ai-gateway-image.yml"
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "litellm-rust/**"
- "litellm/**"
- "enterprise/**"
- "litellm-proxy-extras/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/workflows/ai-gateway-image.yml"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
ai-gateway-image:
name: ai-gateway release image
runs-on: ubuntu-latest
timeout-minutes: 60
permissions:
contents: read
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Build the release image
run: docker build -f litellm-rust/crates/ai-gateway/Dockerfile -t litellm-ai-gateway:${{ github.sha }} .
- name: Start the gateway and wait for readiness
env:
IMAGE: litellm-ai-gateway:${{ github.sha }}
run: |
docker run -d --name ai-gateway -p 4001:4001 \
-e LITELLM_MASTER_KEY=sk-ci-not-a-real-key \
-e OPENAI_API_KEY=sk-ci-not-a-real-key \
"$IMAGE"
for _ in $(seq 1 60); do
if curl -fsS http://127.0.0.1:4001/health/readiness; then
echo "gateway is serving readiness"
exit 0
fi
sleep 2
done
echo "gateway never became ready" >&2
docker logs ai-gateway >&2
exit 1
- name: Assert the gateway loaded the baked config
run: |
docker logs ai-gateway 2>&1 | tee gateway.log
grep 'via python config reader' gateway.log
- name: Stop the gateway
if: always()
run: docker rm -f ai-gateway || true

View file

@ -97,8 +97,6 @@ jobs:
- run: cargo clippy -p litellm-core --all-targets --features bedrock-auth --locked -- -D warnings
- run: cargo clippy -p litellm-ai-gateway --all-targets --all-features --locked -- -D warnings
rust-test:
runs-on: ubuntu-latest
timeout-minutes: 30
@ -134,9 +132,6 @@ jobs:
- run: cargo test -p litellm-core --features bedrock-auth --locked
working-directory: litellm-rust
- run: cargo test -p litellm-ai-gateway --features server --locked
working-directory: litellm-rust
- run: uv build --wheel --out-dir dist
- run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl

138
litellm-rust/Cargo.lock generated
View file

@ -462,64 +462,6 @@ dependencies = [
"tracing",
]
[[package]]
name = "axum"
version = "0.7.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "edca88bc138befd0323b20752846e6587272d3b03b0343c8ea28a6f819e6e71f"
dependencies = [
"async-trait",
"axum-core",
"base64 0.22.1",
"bytes",
"futures-util",
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
"hyper 1.10.1",
"hyper-util",
"itoa",
"matchit",
"memchr",
"mime",
"percent-encoding",
"pin-project-lite",
"rustversion",
"serde",
"serde_json",
"serde_path_to_error",
"serde_urlencoded",
"sha1",
"sync_wrapper",
"tokio",
"tokio-tungstenite",
"tower",
"tower-layer",
"tower-service",
"tracing",
]
[[package]]
name = "axum-core"
version = "0.4.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "09f2bd6146b97ae3359fa0cc6d6b376d9539582c7b4220f041a33ec24c226199"
dependencies = [
"async-trait",
"bytes",
"futures-util",
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
"mime",
"pin-project-lite",
"rustversion",
"sync_wrapper",
"tower-layer",
"tower-service",
"tracing",
]
[[package]]
name = "azure_core"
version = "1.1.0"
@ -1582,7 +1524,6 @@ dependencies = [
"http 1.4.2",
"http-body 1.1.0",
"httparse",
"httpdate",
"itoa",
"pin-project-lite",
"smallvec",
@ -1890,51 +1831,12 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "lazy_static"
version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe"
[[package]]
name = "libc"
version = "0.2.186"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66"
[[package]]
name = "litellm-ai-gateway"
version = "0.1.0"
dependencies = [
"axum",
"base64 0.22.1",
"futures-channel",
"futures-util",
"litellm-config",
"litellm-core",
"reqwest 0.12.28",
"rustls 0.23.42",
"rustls-native-certs",
"serde",
"serde_json",
"sha2 0.10.9",
"subtle",
"tokio",
"tokio-tungstenite",
"tower",
"tracing",
]
[[package]]
name = "litellm-config"
version = "0.1.0"
dependencies = [
"litellm-core",
"pyo3",
"serde_json",
"thiserror 2.0.19",
]
[[package]]
name = "litellm-core"
version = "0.1.0"
@ -1968,8 +1870,6 @@ dependencies = [
"thiserror 2.0.19",
"tokio",
"tokio-tungstenite",
"tracing",
"tracing-subscriber",
"url",
"veil",
]
@ -1990,7 +1890,6 @@ dependencies = [
"serde_json",
"tokio",
"tokio-tungstenite",
"tracing",
]
[[package]]
@ -2065,12 +1964,6 @@ version = "0.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fc04a4c58212d57930a24bf47d3fa87485264a3a054e9c10e042eb373573ad3c"
[[package]]
name = "matchit"
version = "0.7.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0e7465ac9959cc2b1404e8e2367b43684a6d13790fe23056cc8c6c5a6b7bcb94"
[[package]]
name = "memchr"
version = "2.8.3"
@ -3150,15 +3043,6 @@ dependencies = [
"digest 0.11.3",
]
[[package]]
name = "sharded-slab"
version = "0.1.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6"
dependencies = [
"lazy_static",
]
[[package]]
name = "shlex"
version = "2.0.1"
@ -3380,15 +3264,6 @@ dependencies = [
"syn 3.0.0",
]
[[package]]
name = "thread_local"
version = "1.1.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070"
dependencies = [
"cfg-if",
]
[[package]]
name = "time"
version = "0.3.53"
@ -3606,7 +3481,6 @@ dependencies = [
"tokio",
"tower-layer",
"tower-service",
"tracing",
]
[[package]]
@ -3650,7 +3524,6 @@ version = "0.1.44"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100"
dependencies = [
"log",
"pin-project-lite",
"tracing-attributes",
"tracing-core",
@ -3686,17 +3559,6 @@ dependencies = [
"tracing",
]
[[package]]
name = "tracing-subscriber"
version = "0.3.23"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319"
dependencies = [
"sharded-slab",
"thread_local",
"tracing-core",
]
[[package]]
name = "try-lock"
version = "0.2.5"

View file

@ -2,8 +2,6 @@
members = [
"crates/core",
"crates/token-counter",
"crates/config",
"crates/ai-gateway",
"crates/python-interop",
"crates/python-bridge",
]
@ -17,14 +15,9 @@ repository = "https://github.com/BerriAI/litellm"
[workspace.dependencies]
bytes = "1"
tracing = "0.1"
tracing-subscriber = { version = "0.3", default-features = false, features = ["registry", "std"] }
litellm-core = { path = "crates/core" }
litellm-token-counter = { path = "crates/token-counter" }
litellm-config = { path = "crates/config" }
litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false }
litellm-python-interop = { path = "crates/python-interop" }
axum = "0.7"
pyo3 = "0.29.2"
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
pythonize = "0.29.0"

View file

@ -1,56 +0,0 @@
[package]
name = "litellm-ai-gateway"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[lib]
name = "litellm_ai_gateway"
[[bin]]
name = "litellm-ai-gateway"
path = "src/main.rs"
required-features = ["server"]
[[bin]]
name = "trace-parity-gateway"
path = "src/bin/trace_parity_gateway.rs"
required-features = ["trace-parity"]
[dependencies]
tracing.workspace = true
litellm-core = { workspace = true, features = ["bedrock-auth"] }
litellm-config.workspace = true
# reqwest (rustls + json) is used by io/ocr and ships realtime logs to the
# Python proxy callbacks API.
reqwest.workspace = true
# rustls and its root store are direct dependencies so `io::tls` can build the
# one TLS config the outbound dials use; see that module for why it has to.
rustls.workspace = true
rustls-native-certs.workspace = true
# `sync` powers the bounded mpsc channel the realtime logger drains.
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time", "sync"] }
tokio-tungstenite.workspace = true
futures-util.workspace = true
serde_json.workspace = true
base64.workspace = true
axum = { workspace = true, features = ["ws"], optional = true }
serde.workspace = true
subtle = { workspace = true, optional = true }
# sha2 hashes the master key into user_api_key_hash (matches the proxy's
# SHA-256 hash_token) so the plaintext credential never enters a log payload.
sha2 = { workspace = true, optional = true }
tower = { version = "0.5.3", features = ["util"], optional = true }
[features]
default = []
server = ["dep:axum", "dep:subtle", "dep:sha2"]
# Build the gateway's config from the proxy YAML via an embedded Python
# interpreter (links libpython; requires `litellm` importable at runtime).
python-config = ["litellm-config/python"]
trace-parity = ["server", "dep:tower", "litellm-core/observability"]
[dev-dependencies]
futures-channel = "0.3"
tower = { version = "0.5.3", features = ["util"] }

View file

@ -1,109 +0,0 @@
# Multi-stage build for the LiteLLM Rust AI Gateway (realtime WebSocket proxy).
#
# Build context is the **repo root** so we can install `litellm` from this repo's
# source (the gateway loads its model_list via litellm.proxy.read_model_list,
# which is not in any PyPI release yet) AND build the rust workspace under
# litellm-rust/.
#
# docker build -f litellm-rust/crates/ai-gateway/Dockerfile -t litellm-ai-gateway .
#
# No secrets live in this file. Runtime config (LITELLM_MASTER_KEY,
# OPENAI_API_KEY referenced by config.yaml, etc.) is injected as environment
# variables at deploy time.
# ---- Chef -------------------------------------------------------------------
# cargo-chef caches the dependency build so only the gateway crate recompiles on
# a source-only change. python3-dev is present in every rust stage because the
# `python-config` feature links libpython via pyo3 (even in the cook step), and
# python3-pip builds the litellm wheel in the builder stage.
FROM rust:1.98-slim-bookworm AS chef
ENV PYO3_PYTHON=python3.11
# rustup reads rust-toolchain.toml from any parent of the working directory, so
# copying it in is what keeps every cargo call below on the repo's pinned
# channel rather than on whatever the base image happens to ship.
COPY rust-toolchain.toml /build/rust-toolchain.toml
WORKDIR /build/litellm-rust
RUN apt-get update \
&& apt-get install -y --no-install-recommends \
python3 python3-dev python3-pip pkg-config libssl-dev clang \
&& rm -rf /var/lib/apt/lists/* \
&& cargo install cargo-chef --locked --version 0.1.77
# ---- Planner ----------------------------------------------------------------
# Produce the dependency recipe from the rust workspace manifests + Cargo.lock.
FROM chef AS planner
COPY litellm-rust/ .
RUN cargo chef prepare --recipe-path recipe.json
# ---- Builder ----------------------------------------------------------------
FROM chef AS builder
# Cook (compile) just the dependencies first — this layer is cached and reused
# whenever only gateway source changes.
COPY --from=planner /build/litellm-rust/recipe.json recipe.json
RUN cargo chef cook --locked --release \
-p litellm-ai-gateway --features server,python-config \
--recipe-path recipe.json
# Now copy the real sources and build the gateway binary. Deps are already cooked
# above, so this step only recompiles the gateway crate.
COPY litellm-rust/ .
RUN cargo build --locked --release -p litellm-ai-gateway --bin litellm-ai-gateway --features server,python-config
# The root pyproject builds with maturin against litellm-rust/crates/python-bridge,
# so the wheel is built here, next to the crate sources and the cargo toolchain,
# and the runtime stage installs the artifact instead of compiling anything.
# litellm[proxy] pins litellm-enterprise and litellm-proxy-extras to the versions
# in this repo, and those hit PyPI hours after every version bump merges, so both
# wheels are built from the repo too instead of being resolved from PyPI.
COPY pyproject.toml README.md LICENSE /build/
COPY litellm/ /build/litellm/
COPY enterprise/ /build/enterprise/
COPY litellm-proxy-extras/ /build/litellm-proxy-extras/
RUN pip3 wheel --no-cache-dir --no-deps --wheel-dir /build/dist \
/build /build/enterprise /build/litellm-proxy-extras
# ---- Runtime ----------------------------------------------------------------
# python:3.11-slim-bookworm ships libpython3.11, matching the builder's PyO3
# 3.11 ABI so the embedded interpreter links and imports cleanly.
FROM python:3.11-slim-bookworm AS runtime
# CA certificates for outbound TLS to the OpenAI realtime endpoint.
RUN apt-get update \
&& apt-get install -y --no-install-recommends ca-certificates \
&& rm -rf /var/lib/apt/lists/*
WORKDIR /app
# Install litellm (with proxy extras) FROM THIS REPO'S SOURCE so
# `import litellm.proxy.read_model_list` works — it is not on PyPI yet. The two
# sibling wheels come from the builder as well, so the pins in litellm[proxy]
# resolve against them and never wait on a PyPI publish.
COPY --from=builder /build/dist/*.whl /tmp/wheels/
RUN wheel="$(ls /tmp/wheels/litellm-*.whl)" \
&& pip install --no-cache-dir \
/tmp/wheels/litellm_enterprise-*.whl \
/tmp/wheels/litellm_proxy_extras-*.whl \
"${wheel}[proxy]" \
&& rm -rf /tmp/wheels
# The compiled gateway binary (pure-Rust realtime hot path; Python is load-time
# only).
COPY --from=builder /build/litellm-rust/target/release/litellm-ai-gateway /usr/local/bin/litellm-ai-gateway
# Default config.yaml. A real deploy can override this (e.g. mount a Render
# secret file at the same path) — never bake secrets into the image.
COPY litellm-rust/crates/ai-gateway/config.yaml /app/config.yaml
# Bind to all interfaces (Render routes to 0.0.0.0:$PORT) and load the model_list
# from config.yaml via the embedded python config reader.
ENV HOST=0.0.0.0 \
LITELLM_CONFIG_PATH=/app/config.yaml
# Drop to a non-root user. The realtime hot path needs no root privileges, so
# running unprivileged limits blast radius if the process is ever compromised.
# The binary in /usr/local/bin is world-executable (COPY default mode 755); we
# only need /app (and the config.yaml it reads) owned by the unprivileged user.
RUN useradd --system --no-create-home --uid 10001 appuser \
&& chown -R appuser:appuser /app
USER appuser
ENTRYPOINT ["/usr/local/bin/litellm-ai-gateway"]

View file

@ -1,54 +0,0 @@
# Dockerfile-specific ignore-file for the Rust AI Gateway build.
#
# The build context is the repo root (so the image can pip install litellm from
# source AND build the rust workspace). BuildKit honors `<Dockerfile>.dockerignore`
# next to the Dockerfile and it takes precedence over the repo-root `.dockerignore`,
# so this file shrinks the (large) repo-root context for THIS build only without
# touching the root `.dockerignore` used by the main litellm images.
#
# Strategy: ignore everything, then re-include only what the build needs:
# - litellm/ (pip install . needs the full package + proxy reader)
# - litellm-rust/ (the rust workspace; Cargo.lock + crate sources)
# - enterprise/ (litellm/proxy/enterprise symlinks into it; maturin walks it)
# - litellm-proxy-extras/ (built into a wheel alongside enterprise/ for litellm[proxy])
# - pyproject.toml / README.md / LICENSE (packaging metadata for the wheel build)
# - rust-toolchain.toml (the pinned channel every cargo call in the build uses)
*
# --- re-include the build inputs ---
!litellm/
!litellm-rust/
!enterprise/
!litellm-proxy-extras/
!pyproject.toml
!rust-toolchain.toml
!README.md
!LICENSE
# --- prune heavy / irrelevant subpaths back out of the re-included trees ---
# Rust build artifacts (huge; regenerated in the builder).
**/target/
# Committed python distribution artifacts; the wheel build does not read them.
enterprise/dist/
litellm-proxy-extras/dist/
# Python caches and compiled bytecode.
**/__pycache__/
**/*.pyc
**/*.pyo
**/.pytest_cache/
**/.ruff_cache/
**/.mypy_cache/
# Node / UI build output bundled under the python package (not needed to import
# litellm.proxy.read_model_list).
**/node_modules/
litellm/proxy/_experimental/out/
# Tests, logs, and local scratch.
**/tests/
**/test/
*.log
log.txt
*.tgz
# VCS / editor / CI metadata that may live under re-included trees.
**/.git/
.git/
**/.DS_Store

View file

@ -1,206 +0,0 @@
# LiteLLM Rust AI Gateway
A minimal Axum service that fronts OpenAI's realtime API. Clients open a
WebSocket to `GET /v1/realtime`; the gateway authenticates, selects a deployment,
dials OpenAI upstream, and splices the two sockets frame-by-frame.
## Crates
`litellm-rust` has six crates. A crate is a layer or shared foundation, not a route:
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust — per-route entrypoints (`messages::messages()`) that resolve the provider, transform, and make the call; plus types, provider transforms, and the router. |
| litellm-token-counter | Standalone input token counting shared by host integrations without pulling in the full SDK. |
| litellm-config | Config-loading boundary. Returns resolved deployments and optionally delegates loading to Python. |
| litellm-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. |
Dependency direction is acyclic: config depends on core, the gateway depends on config and core, and the Python bridge depends on the domain layers, token counter, and Python interop.
- **Client endpoint:** `wss://<host>/v1/realtime?model=<model>` (WebSocket)
- **Auth:** `Authorization: Bearer $LITELLM_MASTER_KEY` (fails closed if unset)
- **Health:** `GET /health/readiness`, `GET /health/liveness`
- **Request logs:** POSTed to a LiteLLM proxy at `/v1/rust_control_plane/logs` (see [Request logging](#request-logging))
> **Realtime serving is pure Rust.** Python is used at **load time only** — to
> read the config once at boot. The realtime hot path never touches Python.
The former `/health/gil` route and its acquisition counter were removed. They
only observed the single startup config load and did not prove that every GIL
acquisition was instrumented
## Configuration (config.yaml)
The gateway loads its `model_list` from a **config.yaml**, the same as the
LiteLLM proxy. Point `LITELLM_CONFIG_PATH` at the file:
```yaml
# config.yaml
model_list:
- model_name: gpt-realtime
litellm_params:
model: openai/gpt-realtime
api_key: os.environ/OPENAI_API_KEY
```
```bash
LITELLM_CONFIG_PATH=./config.yaml ./litellm-ai-gateway
```
At boot `litellm-config` calls into `litellm.proxy.read_model_list` and returns
resolved deployments to the gateway, which constructs the router. The Python
backend still reuses the **real proxy config reader** (`ProxyConfig.get_config`),
so everything the proxy supports in config.yaml works here too:
- `include:` to merge in other config files,
- `os.environ/VAR` secret references (resolved via the secret manager, never
inlined),
- DB-stored models (when a database is configured).
Secrets stay out of the config — reference them with `os.environ/...` and set
the env var at deploy time. The shipped Docker image is built with the
`python-config` feature and **bundles litellm**, so config loading works out of
the box; the default baked config lives at `/app/config.yaml` and can be
overridden at deploy time (e.g. a Render secret file mounted at the same path).
### Environment variables
| Var | Required | Default | Purpose |
|---|---|---|---|
| `LITELLM_CONFIG_PATH` | yes (config mode) | — | Path to the config.yaml the gateway loads its `model_list` from. The Docker image defaults this to `/app/config.yaml`. |
| `LITELLM_MASTER_KEY` | yes | — | Bearer token clients must send. Unset ⇒ all `/v1/realtime` requests are rejected (fail closed). |
| `OPENAI_API_KEY` | yes | — | Upstream OpenAI key. Referenced by config.yaml as `os.environ/OPENAI_API_KEY` for the gateway→OpenAI dial. |
| `HOST` | no | `127.0.0.1` | **Set to `0.0.0.0` in any container/deploy** or external traffic is refused. |
| `PORT` | no | `4001` | Listen port. Render and most PaaS inject this automatically. |
| `LITELLM_PROXY_BASE_URL` | no | `http://localhost:4000` | LiteLLM proxy that request logs are POSTed to. See [Request logging](#request-logging). |
> Secrets (`LITELLM_MASTER_KEY`, `OPENAI_API_KEY`) are never baked into the image
> or `render.yaml` — inject them at deploy time only.
### Lean env stand-in (fallback)
If the binary is built **without** `python-config` (default features), or
`LITELLM_CONFIG_PATH` is unset, the gateway falls back to a single-deployment
stand-in built from the environment:
| Var | Default | Purpose |
|---|---|---|
| `OPENAI_REALTIME_MODEL` | `gpt-realtime` | The single deployment's model name (also the `?model=` clients pass). |
The default workspace build links no libpython and needs no config file. This
fallback mode only supports one hard-coded OpenAI deployment. **config.yaml is the recommended path** — use the
stand-in only for the leanest possible build.
## Request logging
The gateway runs no spend logic. When a session ends it builds one
`StandardLoggingPayload` and POSTs it to `{LITELLM_PROXY_BASE_URL}/v1/rust_control_plane/logs`
(admin-only, bearer = `LITELLM_MASTER_KEY`), and the proxy replays it through its
normal callbacks (spend logs, Langfuse, etc.). The POST is non-blocking: a bounded
channel drained by a background worker, dropping with a counter if the proxy is
down. It sends one payload per session. Both env vars are in the table above.
Worker tuning, rarely needed: `LITELLM_LOG_CHANNEL_CAPACITY` (4096),
`LITELLM_LOG_BATCH_SIZE` (256), `LITELLM_LOG_FLUSH_INTERVAL_MS` (500).
## Build & run with Docker
The image is built `--features server,python-config` and installs litellm **from this
repo's source** (the config reader is newer than any PyPI release), so the build
**context is the repo root**:
```bash
# from the repo root
docker build -f litellm-rust/crates/ai-gateway/Dockerfile -t litellm-ai-gateway .
docker run --rm -p 4001:4001 \
-e HOST=0.0.0.0 -e PORT=4001 \
-e LITELLM_MASTER_KEY=sk-local \
-e OPENAI_API_KEY=$OPENAI_API_KEY \
litellm-ai-gateway # LITELLM_CONFIG_PATH defaults to /app/config.yaml
# smoke test
curl -s -o /dev/null -w '%{http_code}\n' localhost:4001/health/readiness # -> 200
curl -s -o /dev/null -w '%{http_code}\n' localhost:4001/v1/realtime # -> 401 (auth fails closed)
```
On boot you should see `loaded model_list from /app/config.yaml via python
config reader` — that confirms the config path (not the env stand-in fallback).
To use your own config, mount it over the default:
```bash
docker run --rm -p 4001:4001 \
-e HOST=0.0.0.0 -e LITELLM_MASTER_KEY=sk-local -e OPENAI_API_KEY=$OPENAI_API_KEY \
-v $(pwd)/my-config.yaml:/app/config.yaml:ro \
litellm-ai-gateway
```
### Cargo-only (no Docker)
```bash
# config.yaml mode — needs litellm importable in the active python env
LITELLM_CONFIG_PATH=./crates/ai-gateway/config.yaml \
cargo run --release -p litellm-ai-gateway --features server,python-config
# env stand-in mode — no python, no config
cargo run --release -p litellm-ai-gateway --features server
```
## Deploy on Render
The service is a Docker **web service**; Render terminates TLS and supports
WebSockets, so the public endpoint is `wss://<service>.onrender.com/v1/realtime`.
### Option A — Blueprint (`render.yaml`)
`crates/ai-gateway/render.yaml` describes the service (Docker runtime,
`healthCheckPath: /health/readiness`, repo-root `dockerContext: .`,
`dockerfilePath: ./litellm-rust/crates/ai-gateway/Dockerfile`,
`LITELLM_CONFIG_PATH: /app/config.yaml`). `LITELLM_MASTER_KEY` and
`OPENAI_API_KEY` are `sync: false` — set them in the dashboard after the first
deploy. To use a non-default model_list, mount a **Render Secret File** at
`/app/config.yaml`. Point a Render Blueprint at this repo/branch and apply.
### Option B — Render API
```bash
# create a Docker web service from this repo+branch, then set env vars:
curl -X POST https://api.render.com/v1/services \
-H "Authorization: Bearer $RENDER_API_KEY" -H "Content-Type: application/json" \
-d '{
"type": "web_service", "name": "litellm-rust-ai-gateway",
"ownerId": "<owner-id>", "repo": "https://github.com/BerriAI/litellm",
"branch": "<branch-with-this-dockerfile>",
"serviceDetails": {
"env": "docker",
"envSpecificDetails": {
"dockerfilePath": "./litellm-rust/crates/ai-gateway/Dockerfile",
"dockerContext": "."
},
"healthCheckPath": "/health/readiness"
}
}'
# then set env vars LITELLM_MASTER_KEY, OPENAI_API_KEY, HOST=0.0.0.0,
# LITELLM_CONFIG_PATH=/app/config.yaml
```
Health check path **must** be `/health/readiness`. `autoDeploy` is off by default
in the blueprint — trigger deploys manually (or flip it on) to pick up new commits.
## Scaling
Concurrency is what matters, not total connections: each in-flight session holds
one client socket + one upstream socket. To scale, raise the instance count /
enable autoscaling on the Render service (e.g. baseline 10, max 100). Each
instance needs file descriptors for `2 × peak_concurrent_sessions` — raise
`ulimit -n` if you push very high concurrency.
## Latency note
The gateway adds the cost of one extra hop: client→gateway, then a fresh
gateway→OpenAI realtime handshake (TLS + WS upgrade + `session.created`). In
benchmarks this is ~100150 ms of added session-establishment time; first-audio
and steady-state streaming add no measurable overhead. To minimize it, deploy the
gateway in the Render region with the lowest RTT to OpenAI's realtime endpoint.

View file

@ -1,13 +0,0 @@
# Sample realtime config for the LiteLLM Rust AI Gateway.
#
# litellm-config resolves this model_list at boot through the Python config
# reader (litellm.proxy.read_model_list), then the gateway builds its router.
# Includes, environment secrets, and database-stored models still work.
#
# Secrets are referenced (never inlined) via os.environ/. A real deploy can
# override this file (e.g. mount a Render secret file at LITELLM_CONFIG_PATH).
model_list:
- model_name: gpt-realtime
litellm_params:
model: openai/gpt-realtime
api_key: os.environ/OPENAI_API_KEY

View file

@ -1,35 +0,0 @@
# Render blueprint for the LiteLLM Rust AI Gateway (realtime WebSocket proxy).
#
# Single instance for now (no autoscaling). The public endpoint is a
# WebSocket served over TLS: wss://<service>.onrender.com/v1/realtime
#
# Paths are relative to the **repo root** (Render's convention). The build
# context is the repo root so the image can install litellm from source — the
# gateway loads its model_list via litellm.proxy.read_model_list at boot.
#
# Secrets (LITELLM_MASTER_KEY, OPENAI_API_KEY) are marked sync: false — set
# them in the Render dashboard or via the API, never inline here.
services:
- type: web
name: litellm-rust-ai-gateway
runtime: docker
plan: standard
dockerfilePath: ./litellm-rust/crates/ai-gateway/Dockerfile
dockerContext: .
healthCheckPath: /health/readiness
numInstances: 1
envVars:
# The gateway loads its model_list from this config.yaml via the embedded
# python config reader. The image bakes a default config at /app/config.yaml;
# a real deploy can override it by mounting a Render secret file at this
# same path (Dashboard → Environment → Secret Files) — never inline secrets.
- key: LITELLM_CONFIG_PATH
value: /app/config.yaml
- key: HOST
value: 0.0.0.0
# Bearer token clients must send on /v1/realtime (fail closed if unset).
- key: LITELLM_MASTER_KEY
sync: false
# Referenced by config.yaml as os.environ/OPENAI_API_KEY for the upstream dial.
- key: OPENAI_API_KEY
sync: false

View file

@ -1,288 +0,0 @@
use litellm_core::audio_transcription::{
AudioTranscriptionRequest as CoreAudioTranscriptionRequest, ProviderAudioTranscriptionRequest,
prepare_audio_transcription_provider_call,
};
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::Error;
use serde_json::{Map, Value, json};
use std::future::Future;
use std::pin::Pin;
use super::types::PreparedAudioTranscriptionRequest;
use crate::integrations::custom_guardrail::{
CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest,
};
use crate::integrations::custom_logger::{
CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload,
};
pub(crate) struct AudioTranscriptionLifecycleHooks {
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
}
type AudioFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
type AudioLogFuture<'a> = Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
impl AudioTranscriptionLifecycleHooks {
pub(crate) fn new(
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
) -> Self {
Self {
logger_runner,
guardrail_runner,
request_metadata,
}
}
async fn run_pre_call_guardrails(
&self,
request: PreparedAudioTranscriptionRequest,
) -> Result<PreparedAudioTranscriptionRequest, Error> {
if self.guardrail_runner.is_empty() {
return Ok(request);
}
let (guardrail_request, _) = self
.guardrail_runner
.run_pre_call(
&guardrail_context(&self.request_metadata),
GuardrailRequest::new(json!({
"model": request.model,
"custom_llm_provider": request.custom_llm_provider,
"audio": request.audio,
"optional_params": request.optional_params,
})),
)
.await
.map_err(guardrail_error_to_core_error)?;
let Value::Object(mut data) = guardrail_request.data else {
return Err(Error::InvalidRequest(
"audio transcription pre_call guardrail must return an object".to_string(),
));
};
let audio = data.remove("audio").ok_or_else(|| {
Error::InvalidRequest("audio transcription guardrail removed audio".to_string())
})?;
let optional_params = match data.remove("optional_params") {
Some(Value::Object(value)) => value,
Some(_) => {
return Err(Error::InvalidRequest(
"audio transcription optional_params must be an object".to_string(),
));
}
None => Map::new(),
};
Ok(PreparedAudioTranscriptionRequest {
audio,
optional_params,
..request
})
}
async fn prepare_provider_request(
&self,
request: PreparedAudioTranscriptionRequest,
) -> Result<ProviderAudioTranscriptionRequest, Error> {
let PreparedAudioTranscriptionRequest {
model,
custom_llm_provider,
audio,
api_key,
api_base,
extra_headers,
optional_params,
timeout,
..
} = request;
let provider_request =
prepare_audio_transcription_provider_call(CoreAudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: Some(&custom_llm_provider),
extra_headers,
optional_params,
timeout,
})?;
self.run_during_call_guardrails(provider_request).await
}
async fn run_during_call_guardrails(
&self,
request: ProviderAudioTranscriptionRequest,
) -> Result<ProviderAudioTranscriptionRequest, Error> {
if self.guardrail_runner.is_empty() {
return Ok(request);
}
let (guardrail_request, _) = self
.guardrail_runner
.run_during_call(
&guardrail_context(&self.request_metadata),
GuardrailRequest::new(json!({
"model": request.model(),
"custom_llm_provider": request.custom_llm_provider(),
"url": request.url(),
"body": request.body(),
})),
)
.await
.map_err(guardrail_error_to_core_error)?;
let Value::Object(mut data) = guardrail_request.data else {
return Err(Error::InvalidRequest(
"audio transcription during_call guardrail must return an object".to_string(),
));
};
let body = data.remove("body").ok_or_else(|| {
Error::InvalidRequest("audio transcription guardrail removed body".to_string())
})?;
Ok(request.with_body(body))
}
fn logging_payload(
&self,
context: &CallLifecycleContext,
timing: &CallLifecycleTiming,
) -> StandardLoggingPayload {
StandardLoggingPayload {
id: context.litellm_call_id.clone(),
litellm_call_id: context.litellm_call_id.clone(),
call_type: context.call_type.clone(),
model: context.model.clone(),
custom_llm_provider: context.custom_llm_provider.clone(),
response_cost: 0.0,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
start_time: timing.start_time,
end_time: timing.end_time,
stream: false,
metadata: StandardLoggingMetadata {
user_api_key_hash: self.request_metadata.user_api_key_hash.clone(),
user_api_key_user_id: self.request_metadata.user_api_key_user_id.clone(),
user_api_key_team_id: self.request_metadata.user_api_key_team_id.clone(),
..Default::default()
},
messages: None,
}
}
}
impl CallLifecycleHooks<PreparedAudioTranscriptionRequest, ProviderAudioTranscriptionRequest, Value>
for AudioTranscriptionLifecycleHooks
{
type PreCallFuture<'a> = AudioFuture<'a, PreparedAudioTranscriptionRequest>;
type DuringCallFuture<'a> = AudioFuture<'a, ProviderAudioTranscriptionRequest>;
type SuccessFuture<'a> = AudioLogFuture<'a>;
type FailureFuture<'a> = AudioLogFuture<'a>;
fn async_pre_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: PreparedAudioTranscriptionRequest,
) -> Self::PreCallFuture<'a> {
Box::pin(async move { self.run_pre_call_guardrails(request).await })
}
fn async_during_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: PreparedAudioTranscriptionRequest,
) -> Self::DuringCallFuture<'a> {
Box::pin(async move { self.prepare_provider_request(request).await })
}
fn async_log_success_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
response: &'a Value,
timing: &'a CallLifecycleTiming,
) -> Self::SuccessFuture<'a> {
Box::pin(async move {
if self.logger_runner.is_empty() {
return;
}
self.logger_runner
.async_log_success_event(
&ModelCallDetails::from_standard_logging_payload(
self.logging_payload(context, timing),
),
&CallbackValue::new("audio_transcription", response.clone()),
CallbackTiming::new(timing.start_time, timing.end_time),
)
.await;
})
}
fn async_log_failure_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
error: &'a Error,
timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
Box::pin(async move {
if self.logger_runner.is_empty() {
return;
}
let logging_error = LoggingError {
message: error.to_string(),
kind: core_error_kind(error).to_string(),
};
self.logger_runner
.async_log_failure_event(
&ModelCallDetails::from_standard_logging_payload(
self.logging_payload(context, timing),
)
.with_failure_error(logging_error.clone()),
Some(&CallbackValue::new(
"error",
json!({"message": logging_error.message, "kind": logging_error.kind}),
)),
CallbackTiming::new(timing.start_time, timing.end_time),
)
.await;
})
}
}
fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
GuardrailContext {
call_type: CallType::Other("audio_transcription".to_string()),
selected_guardrails: Vec::new(),
metadata: std::collections::HashMap::new(),
user_api_key_hash: metadata.user_api_key_hash.clone(),
user_api_key_user_id: metadata.user_api_key_user_id.clone(),
user_api_key_team_id: metadata.user_api_key_team_id.clone(),
trace_parent: None,
}
}
fn guardrail_error_to_core_error(error: GuardrailError) -> Error {
Error::InvalidRequest(format!("{}: {}", error.kind, error.message))
}
fn core_error_kind(error: &Error) -> &'static str {
match error {
Error::Auth(_)
| Error::MissingApiKey { .. }
| Error::MissingAzureAiCredentials
| Error::MissingAzureDocumentIntelligenceCredentials
| Error::MissingReductoApiKey => "AuthError",
Error::InvalidProvider(_) => "InvalidProvider",
Error::InvalidRequest(_) => "InvalidRequest",
Error::InvalidType { .. } => "InvalidType",
Error::MissingField(_) | Error::MissingDocumentUrl => "MissingField",
Error::Http { .. } => "HttpError",
Error::InvalidResponse(_) => "InvalidResponse",
Error::Network(_) => "NetworkError",
Error::Connect(_) => "ConnectError",
Error::Routing(_) => "RoutingError",
Error::Unsupported(_) => "UnsupportedRequest",
}
}

View file

@ -1,23 +0,0 @@
use litellm_core::Error;
use litellm_core::audio_transcription::execute_audio_transcription_provider_call;
use litellm_core::call_lifecycle::CallLifecycle;
use serde_json::Value;
mod hooks;
mod prepare;
mod types;
pub use types::AudioTranscriptionRequest;
use prepare::{PreparedAudioTranscriptionCall, prepare_audio_transcription_call};
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
let PreparedAudioTranscriptionCall { request, hooks } =
prepare_audio_transcription_call(request);
CallLifecycle::default()
.run_request(request, &hooks, execute_audio_transcription_provider_call)
.await
}
#[cfg(test)]
mod tests;

View file

@ -1,55 +0,0 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::hooks::AudioTranscriptionLifecycleHooks;
use super::types::{AudioTranscriptionRequest, PreparedAudioTranscriptionRequest};
use crate::integrations::custom_guardrail::CustomGuardrailRunner;
use crate::integrations::custom_logger::CustomLoggerRunner;
pub(crate) struct PreparedAudioTranscriptionCall {
pub(crate) request: PreparedAudioTranscriptionRequest,
pub(crate) hooks: AudioTranscriptionLifecycleHooks,
}
pub(crate) fn prepare_audio_transcription_call(
request: AudioTranscriptionRequest<'_>,
) -> PreparedAudioTranscriptionCall {
let call_id = request
.litellm_call_id
.map(str::to_string)
.unwrap_or_else(new_audio_transcription_call_id);
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "bedrock",
});
PreparedAudioTranscriptionCall {
request: PreparedAudioTranscriptionRequest {
model: provider_info.model.to_string(),
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
litellm_call_id: call_id,
audio: request.audio,
api_key: request.api_key.map(str::to_string),
api_base: request.api_base.map(str::to_string),
extra_headers: request.extra_headers,
optional_params: request.optional_params,
timeout: request.timeout,
},
hooks: AudioTranscriptionLifecycleHooks::new(
CustomLoggerRunner::new(request.callbacks),
CustomGuardrailRunner::new(request.guardrails),
request.request_metadata,
),
}
}
fn new_audio_transcription_call_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_or(0, |duration| duration.as_nanos());
format!("audio-transcription-{timestamp}-{sequence}")
}

View file

@ -1,53 +0,0 @@
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
use serde_json::{Map, json};
use super::{AudioTranscriptionRequest, audio_transcription};
#[tokio::test]
async fn bedrock_request_is_signed_and_contains_audio() {
let listener = TcpListener::bind("127.0.0.1:0").expect("listener");
let address = listener.local_addr().expect("address");
let server = thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("connection");
let mut request = Vec::new();
let mut buffer = [0_u8; 16_384];
let count = stream.read(&mut buffer).expect("request");
request.extend_from_slice(&buffer[..count]);
let request = String::from_utf8_lossy(&request);
assert!(request.contains("POST /model/mistral.voxtral-mini-3b-2507/converse"));
assert!(request.contains("authorization: AWS4-HMAC-SHA256"));
assert!(request.contains("x-amz-date:"));
assert!(request.contains("\"bytes\":\"AQI=\""));
assert!(request.contains("Transcribe the audio. Respond with only the transcript."));
let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}";
stream.write_all(response).expect("response");
});
let optional_params = Map::from_iter([
("aws_access_key_id".to_string(), json!("access-key")),
("aws_secret_access_key".to_string(), json!("secret-key")),
("aws_region_name".to_string(), json!("us-east-1")),
]);
let api_base = format!("http://{address}");
let response = audio_transcription(AudioTranscriptionRequest {
model: "mistral.voxtral-mini-3b-2507",
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
api_key: None,
api_base: Some(&api_base),
custom_llm_provider: Some("bedrock"),
extra_headers: None,
optional_params,
timeout: None,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
})
.await
.expect("transcription");
assert_eq!(response, json!({"text": "hello"}));
server.join().expect("server");
}

View file

@ -1,47 +0,0 @@
use std::sync::Arc;
use std::time::Duration;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
use serde_json::{Map, Value};
use crate::integrations::custom_guardrail::CustomGuardrail;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
pub struct AudioTranscriptionRequest<'a> {
pub model: &'a str,
pub audio: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub callbacks: Vec<Arc<dyn CustomLogger>>,
pub guardrails: Vec<Arc<dyn CustomGuardrail>>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
}
pub(crate) struct PreparedAudioTranscriptionRequest {
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) litellm_call_id: String,
pub(crate) audio: Value,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) optional_params: Map<String, Value>,
pub(crate) timeout: Option<Duration>,
}
impl CallLifecycleRequest for PreparedAudioTranscriptionRequest {
fn lifecycle_context(&self) -> CallLifecycleContext {
CallLifecycleContext::new(
"audio_transcription",
self.model.clone(),
self.custom_llm_provider.clone(),
self.litellm_call_id.clone(),
)
}
}

View file

@ -1,93 +0,0 @@
//! Gateway authentication, as an axum **extractor** (the idiomatic pattern —
//! keeps handlers clean and auth testable).
//!
//! For now this is a single **master key**: any caller presenting it as
//! `Authorization: Bearer <key>` may invoke the gateway. Per-key auth, budgets,
//! and rate limits are delegated to the Python proxy in a later phase.
//!
//! A handler opts in by adding [`RequireMasterKey`] to its arguments; auth then
//! runs during extraction, before the handler body. Routes never re-implement it.
use axum::extract::FromRequestParts;
use axum::http::StatusCode;
use axum::http::header::AUTHORIZATION;
use axum::http::request::Parts;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
use crate::state::AppState;
/// SHA-256 hex digest of a token — the exact transform the Python proxy applies
/// (`litellm.proxy.utils.hash_token`).
///
/// STRICT REQUIREMENT: a raw key (`LITELLM_MASTER_KEY`, a virtual key, …) must
/// **never** leave this gateway in a log payload. Spend logs and every callback
/// integration receive `user_api_key_hash`, so that field must be this hash, not
/// the credential. Hashing here also means the value matches the key's hash in
/// `LiteLLM_SpendLogs.api_key`, so realtime spend joins with the rest of LiteLLM.
pub fn hash_token(token: &str) -> String {
let digest = Sha256::digest(token.as_bytes());
let mut hex = String::with_capacity(digest.len() * 2);
for byte in digest {
use std::fmt::Write;
let _ = write!(hex, "{byte:02x}");
}
hex
}
/// Extractor that requires the configured master key as a bearer token.
///
/// Rejections: `500` when no master key is configured (permanent
/// misconfiguration, not a transient outage); `401` on a missing/incorrect
/// token. The comparison is constant-time.
pub struct RequireMasterKey;
#[axum::async_trait]
impl FromRequestParts<AppState> for RequireMasterKey {
type Rejection = (StatusCode, String);
async fn from_request_parts(
parts: &mut Parts,
state: &AppState,
) -> Result<Self, Self::Rejection> {
let Some(expected) = state.master_key.as_deref() else {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
"gateway auth not configured (set LITELLM_MASTER_KEY)".to_string(),
));
};
let provided = parts
.headers
.get(AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
.map(str::trim);
match provided {
Some(token) if bool::from(token.as_bytes().ct_eq(expected.as_bytes())) => Ok(Self),
_ => Err((
StatusCode::UNAUTHORIZED,
"missing or invalid bearer token".to_string(),
)),
}
}
}
#[cfg(test)]
mod tests {
use super::hash_token;
#[test]
fn hash_token_matches_python_sha256_hexdigest() {
// Must equal hashlib.sha256("sk-1234".encode()).hexdigest() — the value
// the proxy stores in LiteLLM_SpendLogs.api_key.
assert_eq!(
hash_token("sk-1234"),
"88dc28d0f030c55ed4ab77ed8faf098196cb1c05df778539800c9f1243fe6b4b"
);
// 64 lowercase hex chars, and never the raw input.
let h = hash_token("sk-secret");
assert_eq!(h.len(), 64);
assert!(h.chars().all(|c| c.is_ascii_hexdigit()));
assert_ne!(h, "sk-secret");
}
}

View file

@ -1,42 +0,0 @@
use std::io::Read;
use serde::Deserialize;
use serde_json::Value;
#[derive(Deserialize)]
struct Input {
path: String,
model_alias: String,
provider_model: String,
api_base: String,
body: Value,
}
#[tokio::main]
async fn main() {
let mut input = String::new();
if let Err(error) = std::io::stdin().read_to_string(&mut input) {
fail(error);
}
let input: Input = match serde_json::from_str(&input) {
Ok(input) => input,
Err(error) => fail(error),
};
let result = litellm_ai_gateway::trace_parity::traced_request(
input.path,
input.model_alias,
input.provider_model,
input.api_base,
input.body,
)
.await;
match serde_json::to_string(&result) {
Ok(result) => println!("{result}"),
Err(error) => fail(error),
}
}
fn fail(error: impl std::fmt::Display) -> ! {
eprintln!("{error}");
std::process::exit(1)
}

View file

@ -1,14 +0,0 @@
use std::sync::OnceLock;
use std::time::Duration;
const HTTP_CLIENT_TIMEOUT_SECS: u64 = 600;
pub(crate) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(HTTP_CLIENT_TIMEOUT_SECS))
.build()
.expect("failed to build reqwest client")
})
}

View file

@ -1,42 +0,0 @@
//! Crate-level constants for the ai-gateway.
//!
//! Per `litellm-rust/CLAUDE.md`, magic numbers and fixed strings live here
//! (the Rust mirror of Python's `litellm/constants.py`), not inline in feature
//! modules. Env-overridable tunables keep their `DEFAULT_*` value here; the env
//! read + fallback happens at the host/config layer.
/// Default LiteLLM control-plane base URL for request-log egress when
/// `LITELLM_PROXY_BASE_URL` is unset.
pub(crate) const DEFAULT_PROXY_BASE_URL: &str = "http://localhost:4000";
/// The logs ingest path appended to the proxy base. Not a tunable; it is the
/// proxy's API contract (the rust-control-plane router on the Python proxy).
pub(crate) const RUST_CONTROL_PLANE_LOGS_PATH: &str = "/v1/rust_control_plane/logs";
/// Default bounded channel depth for the log-egress worker.
/// Override: `LITELLM_LOG_CHANNEL_CAPACITY`.
pub(crate) const DEFAULT_CHANNEL_CAPACITY: usize = 4096;
/// Default max records POSTed per request to the control plane.
/// Override: `LITELLM_LOG_BATCH_SIZE`.
pub(crate) const DEFAULT_MAX_BATCH_SIZE: usize = 256;
/// Default partial-batch flush cadence, in ms.
/// Override: `LITELLM_LOG_FLUSH_INTERVAL_MS`.
pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
/// Provider attributed to realtime sessions in the logging payload.
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
/// HTTP path for the non-streaming Anthropic Messages route.
#[cfg(feature = "server")]
pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages";
/// Request headers owned by the gateway and never forwarded upstream.
#[cfg(feature = "server")]
pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] =
&["authorization", "connection", "content-length", "host"];

View file

@ -1,127 +0,0 @@
# LiteLLM Rust integrations
This directory contains Rust-native equivalents of LiteLLM integration hooks.
The first supported surfaces are terminal custom loggers and pre/during-call
custom guardrails.
## File layout
Every integration is a folder:
- `mod.rs` contains the implementation, trait, runner, or adapter
- `types.rs` contains the integration-local request, response, error, and future
types
Do not add new flat integration files such as `custom_logger.rs`. Shared wire
contracts that are used by multiple integrations can stay in
`integrations/types.rs`.
Call ordering and lifecycle timing live in `litellm-core/src/call_lifecycle`.
Call-type modules, such as OCR, adapt their request and response shapes into
that generic lifecycle runner.
## CustomLogger
Implement `CustomLogger` when Rust code needs to observe terminal success or
failure events. Method names intentionally match Python `CustomLogger` names.
```rust
use litellm_ai_gateway::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails,
};
struct RecordingLogger;
impl CustomLogger for RecordingLogger {
fn async_log_success_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
response_obj: &'a CallbackValue,
timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
let model = &model_call_details.model;
let provider = &model_call_details.custom_llm_provider;
let call_type = model_call_details.call_type.to_string();
let request_id = model_call_details.request_id.as_deref();
let response_object = &response_obj.object;
let duration = timing.end_time - timing.start_time;
let standard_payload = model_call_details.standard_logging_payload.as_ref();
Ok(())
})
}
fn async_log_failure_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
response_obj: Option<&'a CallbackValue>,
timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
let error = model_call_details.failure_error.as_ref();
let response_object = response_obj.map(|value| value.object.as_str());
let duration = timing.end_time - timing.start_time;
Ok(())
})
}
}
```
Use `CustomLoggerRunner` to fan out terminal events to configured loggers. The
runner is a no-op when no loggers are configured, which is the expected fast
path for requests without callbacks.
## CustomGuardrail
Implement `CustomGuardrail` when Rust code needs to run pre-call or native
during-call checks. Method names intentionally match Python `CustomGuardrail`
entrypoints inherited from Python `CustomLogger`.
```rust
use litellm_ai_gateway::integrations::custom_guardrail::{
CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailEventHook,
GuardrailFuture, GuardrailRequest,
};
struct BlocklistedPromptGuardrail;
impl CustomGuardrail for BlocklistedPromptGuardrail {
fn guardrail_name(&self) -> &str {
"blocklisted-prompt"
}
fn supported_event_hooks(&self) -> &[GuardrailEventHook] {
&[GuardrailEventHook::PreCall]
}
fn async_pre_call_hook<'a>(
&'a self,
_context: &'a GuardrailContext,
request: GuardrailRequest,
) -> GuardrailFuture<'a> {
Box::pin(async move {
if request.data.to_string().contains("blocked phrase") {
return Ok(GuardrailDecision::Block(
litellm_ai_gateway::integrations::custom_guardrail::GuardrailError::blocked(
"blocked phrase detected",
),
));
}
Ok(GuardrailDecision::Allow(request))
})
}
}
```
Use `CustomGuardrailRunner::run_pre_call` for `pre_call` guardrails and
`CustomGuardrailRunner::run_during_call` for `during_call` guardrails. A
`GuardrailDecision::Mask` continues with modified request data.
`GuardrailDecision::Block` short-circuits the provider call.
## Current boundary
These are Rust-only primitives. Python callback and guardrail adapters are a
separate layer that should implement these Rust traits instead of changing the
runner interfaces.

View file

@ -1,468 +0,0 @@
//! Rust mirror of Python `CustomGuardrail` entrypoints used by the proxy.
//!
//! This module is intentionally Rust-only: Python/PyO3 adapters are a later
//! layer that should implement this trait rather than changing the runner.
use std::future::Future;
use std::sync::Arc;
use crate::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
pub mod types;
pub use types::{
GuardrailContext, GuardrailDecision, GuardrailDispatchReport, GuardrailError,
GuardrailEventHook, GuardrailFuture, GuardrailRequest,
};
pub trait CustomGuardrail: Send + Sync {
fn guardrail_name(&self) -> &str;
fn supported_event_hooks(&self) -> &[GuardrailEventHook];
/// Python 1:1 name: `async_pre_call_hook(user_api_key_dict, cache, data, call_type)`.
fn async_pre_call_hook<'a>(
&'a self,
_context: &'a GuardrailContext,
request: GuardrailRequest,
) -> GuardrailFuture<'a> {
Box::pin(async move { Ok(GuardrailDecision::Allow(request)) })
}
/// Python 1:1 name: `async_moderation_hook(data, user_api_key_dict, call_type)`.
fn async_moderation_hook<'a>(
&'a self,
_context: &'a GuardrailContext,
request: GuardrailRequest,
) -> GuardrailFuture<'a> {
Box::pin(async move { Ok(GuardrailDecision::Allow(request)) })
}
}
pub struct CustomGuardrailRunner {
guardrails: Vec<Arc<dyn CustomGuardrail>>,
}
impl CustomGuardrailRunner {
pub fn new(guardrails: Vec<Arc<dyn CustomGuardrail>>) -> Self {
Self { guardrails }
}
pub fn is_empty(&self) -> bool {
self.guardrails.is_empty()
}
pub async fn run_pre_call(
&self,
context: &GuardrailContext,
request: GuardrailRequest,
) -> Result<(GuardrailRequest, GuardrailDispatchReport), GuardrailError> {
self.run_hook(GuardrailEventHook::PreCall, context, request)
.await
}
pub async fn run_during_call(
&self,
context: &GuardrailContext,
request: GuardrailRequest,
) -> Result<(GuardrailRequest, GuardrailDispatchReport), GuardrailError> {
self.run_hook(GuardrailEventHook::DuringCall, context, request)
.await
}
pub async fn run_before_provider<F, Fut, T>(
&self,
event_hook: GuardrailEventHook,
context: &GuardrailContext,
request: GuardrailRequest,
provider: F,
) -> Result<T, GuardrailError>
where
F: FnOnce(GuardrailRequest) -> Fut,
Fut: Future<Output = Result<T, GuardrailError>>,
{
let (request, _) = self.run_hook(event_hook, context, request).await?;
provider(request).await
}
pub async fn run_pre_call_with_failure_logging(
&self,
context: &GuardrailContext,
request: GuardrailRequest,
logger_runner: &CustomLoggerRunner,
model_call_details: &ModelCallDetails,
timing: CallbackTiming,
) -> Result<(GuardrailRequest, GuardrailDispatchReport), GuardrailError> {
match self.run_pre_call(context, request).await {
Ok(result) => Ok(result),
Err(error) => {
let failure_details = model_call_details.clone().with_failure_error(LoggingError {
message: error.message.clone(),
kind: error.kind.clone(),
});
let response_obj = CallbackValue::new(
"guardrail_error",
serde_json::json!({
"message": error.message,
"kind": error.kind,
}),
);
logger_runner
.async_log_failure_event(&failure_details, Some(&response_obj), timing)
.await;
Err(error)
}
}
}
async fn run_hook(
&self,
event_hook: GuardrailEventHook,
context: &GuardrailContext,
mut request: GuardrailRequest,
) -> Result<(GuardrailRequest, GuardrailDispatchReport), GuardrailError> {
if self.guardrails.is_empty() {
return Ok((request, GuardrailDispatchReport::default()));
}
let mut report = GuardrailDispatchReport::default();
for guardrail in &self.guardrails {
if !self.should_run(guardrail.as_ref(), event_hook, context) {
continue;
}
report.invoked += 1;
let decision = match event_hook {
GuardrailEventHook::PreCall => {
guardrail
.async_pre_call_hook(context, request.clone())
.await?
}
GuardrailEventHook::DuringCall => {
guardrail
.async_moderation_hook(context, request.clone())
.await?
}
};
match decision.into_request() {
Ok(next_request) => request = next_request,
Err(error) => return Err(error),
}
}
Ok((request, report))
}
fn should_run(
&self,
guardrail: &dyn CustomGuardrail,
event_hook: GuardrailEventHook,
context: &GuardrailContext,
) -> bool {
let supports_hook = guardrail.supported_event_hooks().contains(&event_hook);
let selected = context.selected_guardrails.is_empty()
|| context
.selected_guardrails
.iter()
.any(|name| name == guardrail.guardrail_name());
supports_hook && selected
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::integrations::custom_logger::{CallType, CallbackValue, CustomLogger, LogFuture};
use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload};
use serde_json::json;
use std::sync::Mutex;
#[derive(Clone)]
enum TestDecision {
Allow,
Mask,
Block,
}
struct RecordingCustomGuardrail {
name: String,
hooks: Vec<GuardrailEventHook>,
decision: TestDecision,
calls: Mutex<Vec<&'static str>>,
}
impl RecordingCustomGuardrail {
fn new(name: &str, hooks: Vec<GuardrailEventHook>, decision: TestDecision) -> Self {
Self {
name: name.to_string(),
hooks,
decision,
calls: Mutex::new(Vec::new()),
}
}
fn calls(&self) -> Vec<&'static str> {
self.calls.lock().unwrap().clone()
}
fn decision(&self, mut request: GuardrailRequest) -> GuardrailDecision {
match self.decision {
TestDecision::Allow => GuardrailDecision::Allow(request),
TestDecision::Mask => {
request.data["masked"] = json!(true);
GuardrailDecision::Mask(request)
}
TestDecision::Block => {
GuardrailDecision::Block(GuardrailError::blocked("blocked by guardrail"))
}
}
}
}
impl CustomGuardrail for RecordingCustomGuardrail {
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> {
Box::pin(async move {
self.calls.lock().unwrap().push("async_pre_call_hook");
Ok(self.decision(request))
})
}
fn async_moderation_hook<'a>(
&'a self,
_context: &'a GuardrailContext,
request: GuardrailRequest,
) -> GuardrailFuture<'a> {
Box::pin(async move {
self.calls.lock().unwrap().push("async_moderation_hook");
Ok(self.decision(request))
})
}
}
#[tokio::test]
async fn pre_call_dispatches_to_async_pre_call_hook() {
let guardrail = Arc::new(RecordingCustomGuardrail::new(
"pre",
vec![GuardrailEventHook::PreCall],
TestDecision::Allow,
));
let runner = CustomGuardrailRunner::new(vec![guardrail.clone()]);
let context =
GuardrailContext::new(CallType::Ocr).with_selected_guardrails(vec!["pre".to_string()]);
let request = GuardrailRequest::new(json!({"messages": ["hello"]}));
let (result, report) = runner
.run_pre_call(&context, request)
.await
.expect("guardrail allows request");
assert_eq!(report.invoked, 1);
assert_eq!(result.data["messages"], json!(["hello"]));
assert_eq!(guardrail.calls(), vec!["async_pre_call_hook"]);
}
#[tokio::test]
async fn during_call_dispatches_to_async_moderation_hook() {
let guardrail = Arc::new(RecordingCustomGuardrail::new(
"during",
vec![GuardrailEventHook::DuringCall],
TestDecision::Allow,
));
let runner = CustomGuardrailRunner::new(vec![guardrail.clone()]);
let context = GuardrailContext::new(CallType::Completion)
.with_selected_guardrails(vec!["during".to_string()]);
let request = GuardrailRequest::new(json!({"prompt": "hello"}));
let (_result, report) = runner
.run_during_call(&context, request)
.await
.expect("guardrail allows request");
assert_eq!(report.invoked, 1);
assert_eq!(guardrail.calls(), vec!["async_moderation_hook"]);
}
#[tokio::test]
async fn mask_decision_continues_with_updated_request() {
let guardrail = Arc::new(RecordingCustomGuardrail::new(
"masker",
vec![GuardrailEventHook::PreCall],
TestDecision::Mask,
));
let runner = CustomGuardrailRunner::new(vec![guardrail]);
let context = GuardrailContext::new(CallType::Ocr);
let request = GuardrailRequest::new(json!({"document": "secret"}));
let (result, report) = runner
.run_pre_call(&context, request)
.await
.expect("mask continues");
assert_eq!(report.invoked, 1);
assert_eq!(result.data["masked"], json!(true));
}
#[tokio::test]
async fn block_decision_short_circuits_and_logs_failure() {
struct RecordingFailureLogger {
errors: Mutex<Vec<String>>,
}
impl CustomLogger for RecordingFailureLogger {
fn async_log_failure_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
_response_obj: Option<&'a CallbackValue>,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
self.errors.lock().unwrap().push(
model_call_details
.failure_error
.as_ref()
.map(|error| error.kind.clone())
.unwrap_or_default(),
);
Ok(())
})
}
}
let guardrail = Arc::new(RecordingCustomGuardrail::new(
"blocker",
vec![GuardrailEventHook::PreCall],
TestDecision::Block,
));
let guardrail_runner = CustomGuardrailRunner::new(vec![guardrail]);
let logger = Arc::new(RecordingFailureLogger {
errors: Mutex::new(Vec::new()),
});
let logger_runner = CustomLoggerRunner::new(vec![logger.clone()]);
let context = GuardrailContext::new(CallType::Ocr);
let details = ModelCallDetails::from_standard_logging_payload(StandardLoggingPayload {
id: "req_ocr".to_string(),
litellm_call_id: "req_ocr".to_string(),
call_type: "ocr".to_string(),
model: "mistral-ocr-latest".to_string(),
custom_llm_provider: "mistral".to_string(),
response_cost: 0.0,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
start_time: 1.0,
end_time: 1.0,
stream: false,
metadata: StandardLoggingMetadata::default(),
messages: None,
});
let err = guardrail_runner
.run_pre_call_with_failure_logging(
&context,
GuardrailRequest::new(json!({"document": "bad"})),
&logger_runner,
&details,
CallbackTiming::new(1.0, 2.0),
)
.await
.expect_err("guardrail blocks request");
assert_eq!(err.kind, "GuardrailBlocked");
assert_eq!(
logger.errors.lock().unwrap().as_slice(),
["GuardrailBlocked"]
);
}
#[tokio::test]
async fn block_decision_short_circuits_later_guardrails_and_provider_work() {
let blocking_guardrail = Arc::new(RecordingCustomGuardrail::new(
"blocker",
vec![GuardrailEventHook::PreCall],
TestDecision::Block,
));
let later_guardrail = Arc::new(RecordingCustomGuardrail::new(
"later",
vec![GuardrailEventHook::PreCall],
TestDecision::Allow,
));
let runner =
CustomGuardrailRunner::new(vec![blocking_guardrail.clone(), later_guardrail.clone()]);
let provider_called = Arc::new(Mutex::new(false));
let provider_called_for_closure = provider_called.clone();
let result = runner
.run_before_provider(
GuardrailEventHook::PreCall,
&GuardrailContext::new(CallType::Completion),
GuardrailRequest::new(json!({"prompt": "blocked"})),
move |_request| async move {
*provider_called_for_closure.lock().unwrap() = true;
Ok("provider response")
},
)
.await;
assert!(result.is_err());
assert_eq!(blocking_guardrail.calls(), vec!["async_pre_call_hook"]);
assert_eq!(later_guardrail.calls(), Vec::<&'static str>::new());
assert!(!*provider_called.lock().unwrap());
}
#[tokio::test]
async fn run_before_provider_returns_provider_guardrail_error_directly() {
let guardrail = Arc::new(RecordingCustomGuardrail::new(
"allow",
vec![GuardrailEventHook::PreCall],
TestDecision::Allow,
));
let runner = CustomGuardrailRunner::new(vec![guardrail]);
let result = runner
.run_before_provider(
GuardrailEventHook::PreCall,
&GuardrailContext::new(CallType::Completion),
GuardrailRequest::new(json!({"prompt": "allowed"})),
|_request| async move {
Err::<&'static str, GuardrailError>(GuardrailError::blocked(
"provider-side guardrail error",
))
},
)
.await;
let err = result.expect_err("provider error is returned directly");
assert_eq!(err.kind, "GuardrailBlocked");
assert_eq!(err.message, "provider-side guardrail error");
}
#[tokio::test]
async fn no_guardrails_fast_path_dispatches_nothing() {
let runner = CustomGuardrailRunner::new(Vec::new());
let context = GuardrailContext::new(CallType::Ocr);
let request = GuardrailRequest::new(json!({"document": "ok"}));
let (result, report) = runner
.run_pre_call(&context, request)
.await
.expect("no guardrails allow request");
assert!(runner.is_empty());
assert_eq!(report, GuardrailDispatchReport::default());
assert_eq!(result.data["document"], json!("ok"));
}
}

View file

@ -1,110 +0,0 @@
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use serde_json::Value;
use crate::integrations::custom_logger::CallType;
pub type GuardrailFuture<'a> =
Pin<Box<dyn Future<Output = Result<GuardrailDecision, GuardrailError>> + Send + 'a>>;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum GuardrailEventHook {
PreCall,
DuringCall,
}
impl GuardrailEventHook {
pub fn as_str(&self) -> &'static str {
match self {
Self::PreCall => "pre_call",
Self::DuringCall => "during_call",
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct GuardrailError {
pub message: String,
pub kind: String,
}
impl GuardrailError {
pub fn blocked(message: impl Into<String>) -> Self {
Self {
message: message.into(),
kind: "GuardrailBlocked".to_string(),
}
}
}
impl std::fmt::Display for GuardrailError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}: {}", self.kind, self.message)
}
}
impl std::error::Error for GuardrailError {}
#[derive(Clone, Debug)]
pub struct GuardrailContext {
pub call_type: CallType,
pub selected_guardrails: Vec<String>,
pub metadata: HashMap<String, Value>,
pub user_api_key_hash: Option<String>,
pub user_api_key_user_id: Option<String>,
pub user_api_key_team_id: Option<String>,
pub trace_parent: Option<String>,
}
impl GuardrailContext {
pub fn new(call_type: CallType) -> Self {
Self {
call_type,
selected_guardrails: Vec::new(),
metadata: HashMap::new(),
user_api_key_hash: None,
user_api_key_user_id: None,
user_api_key_team_id: None,
trace_parent: None,
}
}
pub fn with_selected_guardrails(mut self, selected_guardrails: Vec<String>) -> Self {
self.selected_guardrails = selected_guardrails;
self
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct GuardrailRequest {
pub data: Value,
}
impl GuardrailRequest {
pub fn new(data: Value) -> Self {
Self { data }
}
}
#[derive(Clone, Debug, PartialEq)]
pub enum GuardrailDecision {
Allow(GuardrailRequest),
Mask(GuardrailRequest),
Block(GuardrailError),
}
impl GuardrailDecision {
pub(super) fn into_request(self) -> Result<GuardrailRequest, GuardrailError> {
match self {
Self::Allow(request) | Self::Mask(request) => Ok(request),
Self::Block(error) => Err(error),
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct GuardrailDispatchReport {
pub invoked: usize,
}

View file

@ -1,317 +0,0 @@
//! The `CustomLogger` trait — the Rust mirror of Python
//! `litellm/integrations/custom_logger.py::CustomLogger`.
//!
//! The Python-named async terminal methods are the public Rust callback shape.
use std::sync::Arc;
pub mod types;
pub use types::{
CallType, CallbackDispatchReport, CallbackTiming, CallbackValue, LogError, LogFuture,
LoggingError, ModelCallDetails,
};
pub trait CustomLogger: Send + Sync {
/// Python 1:1 name: `async_log_success_event(model_call_details, response_obj, start_time, end_time)`.
fn async_log_success_event<'a>(
&'a self,
_model_call_details: &'a ModelCallDetails,
_response_obj: &'a CallbackValue,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async { Ok(()) })
}
/// Python 1:1 name: `async_log_failure_event(model_call_details, response_obj, start_time, end_time)`.
fn async_log_failure_event<'a>(
&'a self,
_model_call_details: &'a ModelCallDetails,
_response_obj: Option<&'a CallbackValue>,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async { Ok(()) })
}
}
pub struct CustomLoggerRunner {
loggers: Vec<Arc<dyn CustomLogger>>,
}
impl CustomLoggerRunner {
pub fn new(loggers: Vec<Arc<dyn CustomLogger>>) -> Self {
Self { loggers }
}
pub fn is_empty(&self) -> bool {
self.loggers.is_empty()
}
pub async fn async_log_success_event(
&self,
model_call_details: &ModelCallDetails,
response_obj: &CallbackValue,
timing: CallbackTiming,
) -> CallbackDispatchReport {
if self.loggers.is_empty() {
return CallbackDispatchReport::default();
}
let mut report = CallbackDispatchReport::default();
for logger in &self.loggers {
report.invoked += 1;
if let Err(err) = logger
.async_log_success_event(model_call_details, response_obj, timing)
.await
{
report.dropped += 1;
eprintln!("litellm-ai-gateway: async_log_success_event dropped: {err}");
}
}
report
}
pub async fn async_log_failure_event(
&self,
model_call_details: &ModelCallDetails,
response_obj: Option<&CallbackValue>,
timing: CallbackTiming,
) -> CallbackDispatchReport {
if self.loggers.is_empty() {
return CallbackDispatchReport::default();
}
let mut report = CallbackDispatchReport::default();
for logger in &self.loggers {
report.invoked += 1;
if let Err(err) = logger
.async_log_failure_event(model_call_details, response_obj, timing)
.await
{
report.dropped += 1;
eprintln!("litellm-ai-gateway: async_log_failure_event dropped: {err}");
}
}
report
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload};
use serde_json::json;
use std::sync::Mutex;
#[derive(Clone, Debug, PartialEq)]
struct RecordedEvent {
hook: &'static str,
model: String,
provider: String,
call_type: String,
request_id: Option<String>,
litellm_call_id: Option<String>,
user_id: Option<String>,
response_object: Option<String>,
error_kind: Option<String>,
start_time: f64,
end_time: f64,
standard_logging_model: Option<String>,
}
#[derive(Default)]
struct RecordingCustomLogger {
events: Mutex<Vec<RecordedEvent>>,
}
impl RecordingCustomLogger {
fn events(&self) -> Vec<RecordedEvent> {
self.events.lock().unwrap().clone()
}
}
impl CustomLogger for RecordingCustomLogger {
fn async_log_success_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
response_obj: &'a CallbackValue,
timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push(RecordedEvent {
hook: "async_log_success_event",
model: model_call_details.model.clone(),
provider: model_call_details.custom_llm_provider.clone(),
call_type: model_call_details.call_type.to_string(),
request_id: model_call_details.request_id.clone(),
litellm_call_id: model_call_details.litellm_call_id.clone(),
user_id: model_call_details.metadata.user_api_key_user_id.clone(),
response_object: Some(response_obj.object.clone()),
error_kind: None,
start_time: timing.start_time,
end_time: timing.end_time,
standard_logging_model: model_call_details
.standard_logging_payload
.as_ref()
.map(|payload| payload.model.clone()),
});
Ok(())
})
}
fn async_log_failure_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
response_obj: Option<&'a CallbackValue>,
timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push(RecordedEvent {
hook: "async_log_failure_event",
model: model_call_details.model.clone(),
provider: model_call_details.custom_llm_provider.clone(),
call_type: model_call_details.call_type.to_string(),
request_id: model_call_details.request_id.clone(),
litellm_call_id: model_call_details.litellm_call_id.clone(),
user_id: model_call_details.metadata.user_api_key_user_id.clone(),
response_object: response_obj.map(|value| value.object.clone()),
error_kind: model_call_details
.failure_error
.as_ref()
.map(|error| error.kind.clone()),
start_time: timing.start_time,
end_time: timing.end_time,
standard_logging_model: model_call_details
.standard_logging_payload
.as_ref()
.map(|payload| payload.model.clone()),
});
Ok(())
})
}
}
fn payload(call_type: &str, model: &str, provider: &str) -> StandardLoggingPayload {
StandardLoggingPayload {
id: format!("req_{call_type}"),
litellm_call_id: format!("call_{call_type}"),
call_type: call_type.to_string(),
model: model.to_string(),
custom_llm_provider: provider.to_string(),
response_cost: 0.25,
prompt_tokens: 3,
completion_tokens: 4,
total_tokens: 7,
start_time: 10.0,
end_time: 11.5,
stream: false,
metadata: StandardLoggingMetadata {
user_api_key_hash: Some("hash".to_string()),
user_api_key_user_id: Some("user".to_string()),
user_api_key_team_id: Some("team".to_string()),
..Default::default()
},
messages: Some(json!([{"role": "user", "content": "read this"}])),
}
}
#[tokio::test]
async fn rust_custom_logger_reads_success_payload_for_ocr() {
let logger = Arc::new(RecordingCustomLogger::default());
let runner = CustomLoggerRunner::new(vec![logger.clone()]);
let details = ModelCallDetails::from_standard_logging_payload(payload(
"ocr",
"mistral-ocr-latest",
"mistral",
));
let response = CallbackValue::new("ocr", json!({"pages": [{"markdown": "ok"}]}));
let report = runner
.async_log_success_event(&details, &response, CallbackTiming::new(10.0, 11.5))
.await;
assert_eq!(report.invoked, 1);
assert_eq!(report.dropped, 0);
assert_eq!(
logger.events(),
vec![RecordedEvent {
hook: "async_log_success_event",
model: "mistral-ocr-latest".to_string(),
provider: "mistral".to_string(),
call_type: "ocr".to_string(),
request_id: Some("req_ocr".to_string()),
litellm_call_id: Some("call_ocr".to_string()),
user_id: Some("user".to_string()),
response_object: Some("ocr".to_string()),
error_kind: None,
start_time: 10.0,
end_time: 11.5,
standard_logging_model: Some("mistral-ocr-latest".to_string()),
}]
);
}
#[tokio::test]
async fn rust_custom_logger_reads_failure_payload_for_non_ocr_call_type() {
let logger = Arc::new(RecordingCustomLogger::default());
let runner = CustomLoggerRunner::new(vec![logger.clone()]);
let details = ModelCallDetails::from_standard_logging_payload(payload(
"acompletion",
"gpt-4.1-mini",
"openai",
))
.with_failure_error(LoggingError {
message: "provider failed".to_string(),
kind: "ProviderError".to_string(),
});
let response = CallbackValue::new("error", json!({"message": "provider failed"}));
let report = runner
.async_log_failure_event(&details, Some(&response), CallbackTiming::new(2.0, 3.0))
.await;
assert_eq!(report.invoked, 1);
assert_eq!(report.dropped, 0);
assert_eq!(
logger.events(),
vec![RecordedEvent {
hook: "async_log_failure_event",
model: "gpt-4.1-mini".to_string(),
provider: "openai".to_string(),
call_type: "acompletion".to_string(),
request_id: Some("req_acompletion".to_string()),
litellm_call_id: Some("call_acompletion".to_string()),
user_id: Some("user".to_string()),
response_object: Some("error".to_string()),
error_kind: Some("ProviderError".to_string()),
start_time: 2.0,
end_time: 3.0,
standard_logging_model: Some("gpt-4.1-mini".to_string()),
}]
);
}
#[tokio::test]
async fn no_callback_fast_path_dispatches_nothing() {
let runner = CustomLoggerRunner::new(Vec::new());
let details = ModelCallDetails::new("mistral-ocr-latest", "mistral", CallType::Ocr);
let response = CallbackValue::new("ocr", json!({}));
let report = runner
.async_log_success_event(&details, &response, CallbackTiming::new(1.0, 1.5))
.await;
assert!(runner.is_empty());
assert_eq!(report, CallbackDispatchReport::default());
}
#[test]
fn with_standard_logging_payload_keeps_top_level_fields_in_sync() {
let details = ModelCallDetails::new("old-model", "old-provider", CallType::Completion)
.with_standard_logging_payload(payload("ocr", "mistral-ocr-latest", "mistral"));
assert_eq!(details.model, "mistral-ocr-latest");
assert_eq!(details.custom_llm_provider, "mistral");
assert_eq!(details.call_type, CallType::Ocr);
assert_eq!(details.request_id, Some("req_ocr".to_string()));
assert_eq!(details.litellm_call_id, Some("call_ocr".to_string()));
}
}

View file

@ -1,194 +0,0 @@
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use serde_json::Value;
use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload};
pub type LogFuture<'a> = Pin<Box<dyn Future<Output = Result<(), LogError>> + Send + 'a>>;
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct CallbackDispatchReport {
pub invoked: usize,
pub dropped: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum CallType {
Ocr,
Realtime,
Completion,
Acompletion,
ChatCompletion,
Other(String),
}
impl CallType {
pub fn as_str(&self) -> &str {
match self {
Self::Ocr => "ocr",
Self::Realtime => "realtime",
Self::Completion => "completion",
Self::Acompletion => "acompletion",
Self::ChatCompletion => "chat_completion",
Self::Other(value) => value.as_str(),
}
}
}
impl From<&str> for CallType {
fn from(value: &str) -> Self {
match value {
"ocr" => Self::Ocr,
"realtime" => Self::Realtime,
"completion" => Self::Completion,
"acompletion" => Self::Acompletion,
"chat_completion" => Self::ChatCompletion,
other => Self::Other(other.to_string()),
}
}
}
impl std::fmt::Display for CallType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct CallbackTiming {
pub start_time: f64,
pub end_time: f64,
}
impl CallbackTiming {
pub fn new(start_time: f64, end_time: f64) -> Self {
Self {
start_time,
end_time,
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct CallbackValue {
pub object: String,
pub value: Value,
}
impl CallbackValue {
pub fn new(object: impl Into<String>, value: Value) -> Self {
Self {
object: object.into(),
value,
}
}
}
#[derive(Clone, Debug)]
pub struct ModelCallDetails {
pub model: String,
pub custom_llm_provider: String,
pub call_type: CallType,
pub metadata: StandardLoggingMetadata,
pub extra_metadata: HashMap<String, Value>,
pub request_id: Option<String>,
pub litellm_call_id: Option<String>,
pub response_cost: Option<f64>,
pub standard_logging_payload: Option<StandardLoggingPayload>,
pub failure_error: Option<LoggingError>,
}
impl ModelCallDetails {
pub fn new(
model: impl Into<String>,
custom_llm_provider: impl Into<String>,
call_type: CallType,
) -> Self {
Self {
model: model.into(),
custom_llm_provider: custom_llm_provider.into(),
call_type,
metadata: StandardLoggingMetadata::default(),
extra_metadata: HashMap::new(),
request_id: None,
litellm_call_id: None,
response_cost: None,
standard_logging_payload: None,
failure_error: None,
}
}
pub fn from_standard_logging_payload(payload: StandardLoggingPayload) -> Self {
let request_id = Some(payload.id.clone());
let litellm_call_id = Some(payload.litellm_call_id.clone());
let response_cost = Some(payload.response_cost);
let metadata = payload.metadata.clone();
Self {
model: payload.model.clone(),
custom_llm_provider: payload.custom_llm_provider.clone(),
call_type: CallType::from(payload.call_type.as_str()),
metadata,
extra_metadata: HashMap::new(),
request_id,
litellm_call_id,
response_cost,
standard_logging_payload: Some(payload),
failure_error: None,
}
}
pub fn with_standard_logging_payload(mut self, payload: StandardLoggingPayload) -> Self {
self.model = payload.model.clone();
self.custom_llm_provider = payload.custom_llm_provider.clone();
self.call_type = CallType::from(payload.call_type.as_str());
self.request_id = Some(payload.id.clone());
self.litellm_call_id = Some(payload.litellm_call_id.clone());
self.response_cost = Some(payload.response_cost);
self.metadata = payload.metadata.clone();
self.standard_logging_payload = Some(payload);
self
}
pub fn with_failure_error(mut self, error: LoggingError) -> Self {
self.failure_error = Some(error);
self
}
}
#[derive(Clone, Debug)]
pub struct LoggingError {
pub message: String,
pub kind: String,
}
#[derive(Clone, Debug)]
pub struct LogError {
pub message: String,
pub kind: String,
}
impl LogError {
pub fn channel_full() -> Self {
Self {
message: "logging channel is full; dropping record".to_string(),
kind: "ChannelFull".to_string(),
}
}
pub fn channel_closed() -> Self {
Self {
message: "logging channel is closed; worker has shut down".to_string(),
kind: "ChannelClosed".to_string(),
}
}
}
impl std::fmt::Display for LogError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}: {}", self.kind, self.message)
}
}
impl std::error::Error for LogError {}

View file

@ -1,197 +0,0 @@
//! A `CustomLogger` that ships finished events to the LiteLLM Python proxy's
//! `/v1/rust_control_plane/logs` endpoint.
//!
//! The callback path is non-blocking: `async_log_success_event` /
//! `async_log_failure_event`
//! build a `LogRecord` and `try_send` it onto a bounded channel, returning a
//! `LogError` (never panicking, never awaiting) if the channel is full or the
//! worker has gone away. A spawned background worker drains the channel, batches
//! records into `{"records":[...]}`, and POSTs them to the proxy with a pooled
//! `reqwest::Client`.
use std::sync::Arc;
use std::time::Duration;
use reqwest::Client;
use tokio::sync::mpsc::{self, Receiver, Sender};
use tokio::time::interval;
use crate::constants::{DEFAULT_PROXY_BASE_URL, RUST_CONTROL_PLANE_LOGS_PATH};
use crate::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, LogError, LogFuture, LoggingError,
ModelCallDetails,
};
use types::{CallbackLogsRequest, EgressTunables, LogRecord};
pub mod types;
/// Ships realtime logging events to the LiteLLM Python proxy.
pub struct LiteLLMPythonProxyAPILogger {
sink: Sender<LogRecord>,
}
impl LiteLLMPythonProxyAPILogger {
/// Spawn the background worker and return a logger handle. `base` is the
/// proxy base URL (no trailing path); `master_key` is sent as a bearer token.
pub fn start(base: String, master_key: String) -> Arc<Self> {
let tunables = EgressTunables::from_env();
let (sink, receiver) = mpsc::channel::<LogRecord>(tunables.channel_capacity);
let url = format!(
"{}{}",
base.trim_end_matches('/'),
RUST_CONTROL_PLANE_LOGS_PATH
);
let client = Client::new();
tokio::spawn(worker_loop(
receiver,
client,
url,
master_key,
tunables.max_batch_size,
tunables.flush_interval,
));
Arc::new(Self { sink })
}
/// Build a logger from the environment: `LITELLM_PROXY_BASE_URL` (default
/// `http://localhost:4000`) and `LITELLM_MASTER_KEY`.
///
/// `LITELLM_PROXY_BASE_URL` is treated as the full base and the route is
/// appended verbatim, so if the proxy runs under a `SERVER_ROOT_PATH`
/// (e.g. served at `https://host/litellm`), include it in the base
/// (`LITELLM_PROXY_BASE_URL=https://host/litellm`) and the POST lands at
/// `https://host/litellm/v1/rust_control_plane/logs`.
pub fn from_env() -> Arc<Self> {
let base = std::env::var("LITELLM_PROXY_BASE_URL")
.ok()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| DEFAULT_PROXY_BASE_URL.to_string());
let key = std::env::var("LITELLM_MASTER_KEY").unwrap_or_default();
Self::start(base, key)
}
fn enqueue(&self, record: LogRecord) -> Result<(), LogError> {
self.sink.try_send(record).map_err(|err| match err {
mpsc::error::TrySendError::Full(_) => LogError::channel_full(),
mpsc::error::TrySendError::Closed(_) => LogError::channel_closed(),
})
}
}
impl CustomLogger for LiteLLMPythonProxyAPILogger {
fn async_log_success_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
_response_obj: &'a CallbackValue,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
if let Some(payload) = &model_call_details.standard_logging_payload {
self.enqueue(LogRecord {
status: "success".to_string(),
payload: payload.clone(),
error: None,
})?;
}
Ok(())
})
}
fn async_log_failure_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
_response_obj: Option<&'a CallbackValue>,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
if let Some(payload) = &model_call_details.standard_logging_payload {
let fallback_error;
let error = match &model_call_details.failure_error {
Some(error) => error,
None => {
fallback_error = LoggingError {
message: "callback failure event".to_string(),
kind: "CallbackFailure".to_string(),
};
&fallback_error
}
};
self.enqueue(LogRecord {
status: "failure".to_string(),
payload: payload.clone(),
error: Some(format!("{}: {}", error.kind, error.message)),
})?;
}
Ok(())
})
}
}
/// Drain the channel, batching records and POSTing them to the proxy. Exits when
/// the channel is closed (all senders dropped) and drained.
async fn worker_loop(
mut receiver: Receiver<LogRecord>,
client: Client,
url: String,
master_key: String,
max_batch_size: usize,
flush_interval: Duration,
) {
let mut ticker = interval(flush_interval);
let mut batch: Vec<LogRecord> = Vec::with_capacity(max_batch_size);
loop {
tokio::select! {
maybe_record = receiver.recv() => {
match maybe_record {
Some(record) => {
batch.push(record);
if batch.len() >= max_batch_size {
flush(&client, &url, &master_key, &mut batch).await;
}
}
None => {
// Channel closed: flush remaining and exit.
flush(&client, &url, &master_key, &mut batch).await;
break;
}
}
}
_ = ticker.tick() => {
flush(&client, &url, &master_key, &mut batch).await;
}
}
}
}
/// POST the current batch (if any), clearing it. Errors are logged, not fatal.
async fn flush(client: &Client, url: &str, master_key: &str, batch: &mut Vec<LogRecord>) {
if batch.is_empty() {
return;
}
let records = std::mem::take(batch)
.into_iter()
.map(LogRecord::into_callback_record)
.collect();
let body = CallbackLogsRequest { records };
let response = client
.post(url)
.bearer_auth(master_key)
.json(&body)
.send()
.await;
match response {
Ok(resp) if resp.status().is_success() => {}
Ok(resp) => {
eprintln!(
"litellm-ai-gateway: callback logs POST returned {} to {url}",
resp.status()
);
}
Err(err) => {
eprintln!("litellm-ai-gateway: callback logs POST failed to {url}: {err}");
}
}
}

View file

@ -1,72 +0,0 @@
use std::time::Duration;
use serde::Serialize;
use crate::constants::{
DEFAULT_CHANNEL_CAPACITY, DEFAULT_FLUSH_INTERVAL_MS, DEFAULT_MAX_BATCH_SIZE,
};
use crate::integrations::types::StandardLoggingPayload;
#[derive(Serialize)]
pub struct CallbackLogsRequest {
pub records: Vec<CallbackLogRecord>,
}
#[derive(Serialize)]
pub struct CallbackLogRecord {
pub status: String,
pub standard_logging_payload: StandardLoggingPayload,
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
}
#[derive(Clone, Debug)]
pub struct LogRecord {
pub status: String,
pub payload: StandardLoggingPayload,
pub error: Option<String>,
}
impl LogRecord {
pub fn into_callback_record(self) -> CallbackLogRecord {
CallbackLogRecord {
status: self.status,
standard_logging_payload: self.payload,
error: self.error,
}
}
}
pub(super) struct EgressTunables {
pub channel_capacity: usize,
pub max_batch_size: usize,
pub flush_interval: Duration,
}
impl EgressTunables {
pub fn from_env() -> Self {
Self {
channel_capacity: env_positive(
"LITELLM_LOG_CHANNEL_CAPACITY",
DEFAULT_CHANNEL_CAPACITY,
),
max_batch_size: env_positive("LITELLM_LOG_BATCH_SIZE", DEFAULT_MAX_BATCH_SIZE),
flush_interval: Duration::from_millis(env_positive(
"LITELLM_LOG_FLUSH_INTERVAL_MS",
DEFAULT_FLUSH_INTERVAL_MS,
)),
}
}
}
fn env_positive<T>(name: &str, default: T) -> T
where
T: std::str::FromStr + PartialOrd + From<u8>,
{
let zero = T::from(0u8);
std::env::var(name)
.ok()
.and_then(|value| value.trim().parse::<T>().ok())
.filter(|n| *n > zero)
.unwrap_or(default)
}

View file

@ -1,12 +0,0 @@
//! Pure-Rust logging integrations. Names map 1:1 to Python
//! `litellm/integrations/`:
//! - [`custom_guardrail::CustomGuardrail`] — the guardrail callback trait
//! - [`custom_logger::CustomLogger`] — the callback trait
//! - [`litellm_python_proxy_api::LiteLLMPythonProxyAPILogger`] — ships events
//! to the Python proxy's `/v1/rust_control_plane/logs` endpoint
//! - [`types`] — the typed `StandardLoggingPayload` wire contract
pub mod custom_guardrail;
pub mod custom_logger;
pub mod litellm_python_proxy_api;
pub mod types;

View file

@ -1,83 +0,0 @@
//! Typed payloads for the LiteLLM `/v1/callbacks/logs` realtime-logging contract.
//!
//! Field names below are the EXACT JSON keys the Python replay path + spend-logs
//! builder read. Note the deliberate mix:
//! - `startTime` / `endTime` are camelCase (epoch f64 seconds)
//! - `response_cost` / `prompt_tokens` / etc. are snake_case
//!
//! Mirrors Python `litellm/integrations/` + the proxy `CallbackLogsRequest`
//! contract 1:1.
use serde::Serialize;
use serde_json::Value;
use std::collections::HashMap;
/// Cumulative token usage for a realtime session.
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Usage {
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub total_tokens: u64,
}
/// Cost-attribution metadata threaded from the authenticated request.
#[derive(Clone, Debug, Default)]
pub struct RequestMetadata {
pub user_api_key_hash: Option<String>,
pub user_api_key_user_id: Option<String>,
pub user_api_key_team_id: Option<String>,
}
/// The self-describing payload. Field names are the EXACT JSON keys the Python
/// replay path + spend-logs builder read.
#[derive(Clone, Debug, Serialize)]
pub struct StandardLoggingPayload {
pub id: String,
pub litellm_call_id: String,
/// e.g. "realtime", "acompletion". Falls back to "acompletion" if absent.
pub call_type: String,
pub model: String,
pub custom_llm_provider: String,
/// Spend ($) written to LiteLLM_SpendLogs.spend.
pub response_cost: f64,
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub total_tokens: u64,
/// EPOCH SECONDS as float — camelCase keys, NOT snake_case.
#[serde(rename = "startTime")]
pub start_time: f64,
#[serde(rename = "endTime")]
pub end_time: f64,
pub stream: bool,
pub metadata: StandardLoggingMetadata,
/// Optional; stored as request input on the spend log row.
#[serde(skip_serializing_if = "Option::is_none")]
pub messages: Option<Value>,
}
/// Cost-attribution keys. The replayer maps these into litellm_params.metadata,
/// which the spend-logs builder reads to set user / team_id / organization_id.
#[derive(Clone, Debug, Serialize, Default)]
pub struct StandardLoggingMetadata {
pub user_api_key_hash: Option<String>, // -> SpendLogs.api_key
pub user_api_key_user_id: Option<String>, // -> SpendLogs.user
pub user_api_key_team_id: Option<String>, // -> SpendLogs.team_id
// Optional but read by the builder; include when known:
#[serde(skip_serializing_if = "Option::is_none")]
pub user_api_key_alias: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_api_key_org_id: Option<String>, // -> SpendLogs.organization_id
#[serde(skip_serializing_if = "Option::is_none")]
pub user_api_key_end_user_id: Option<String>, // -> SpendLogs.end_user
#[serde(skip_serializing_if = "Option::is_none")]
pub spend_logs_metadata: Option<HashMap<String, Value>>,
}

View file

@ -1 +0,0 @@
pub use crate::audio_transcription::{AudioTranscriptionRequest, audio_transcription};

View file

@ -1,6 +0,0 @@
pub mod audio_transcription;
pub mod ocr;
pub mod realtime;
pub mod realtime_pool;
pub mod responses_ws;
pub(crate) mod tls;

View file

@ -1 +0,0 @@
pub use crate::ocr::{OcrRequest, ocr};

View file

@ -1,418 +0,0 @@
//! End-to-end OpenAI realtime invocation.
//!
//! The host-facing entry point opens the WebSocket to OpenAI, then splices a
//! client realtime stream to the upstream, driving typed events through the pure
//! `OPENAI_REALTIME_CONFIG` transforms.
//! Network, auth header, key resolution, and wire (de)serialization live here so
//! the `transformation` module stays pure and typed.
//!
//! The dial and splice steps are factored out ([`dial_upstream`], [`splice`]) so
//! the connection pool ([`crate::io::realtime_pool`]) can pre-establish an upstream,
//! buffer its `session.created`, and later hand the live socket to the same
//! splice loop a fresh dial uses.
use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use litellm_core::AuthError;
use litellm_core::auth::error::MissingCredential;
use litellm_core::error::Error;
use litellm_core::realtime::transformation::RealtimeProviderConfig;
use litellm_core::realtime::types::RealtimeEvent;
use tokio::net::TcpStream;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
use litellm_core::providers::openai::realtime::transformation::OPENAI_REALTIME_CONFIG;
use crate::io::tls::connect_upstream;
/// Environment variable holding the OpenAI API key (last-resort fallback).
const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
/// Default **idle** timeout: if neither side sends a frame for this long, the
/// session is reaped. It resets on any activity, so it does not cap a healthy
/// (continuously streaming) session — it only frees a stalled one (e.g. a
/// half-open upstream that keeps the socket open but stops sending).
const DEFAULT_IDLE_TIMEOUT_SECS: u64 = 300;
/// The concrete upstream WebSocket type (TLS or plain). Shared by the dial path
/// and the pool so warm sockets and fresh sockets are the exact same type.
pub type UpstreamWs = WebSocketStream<MaybeTlsStream<TcpStream>>;
pub(crate) type UpstreamTx = SplitSink<UpstreamWs, Message>;
pub(crate) type UpstreamRx = SplitStream<UpstreamWs>;
/// Resolve the OpenAI API key from the explicit param or the environment.
///
/// Blank/whitespace values are treated as absent (guard at resolution time).
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result<String, Error> {
api_key
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
std::env::var(OPENAI_API_KEY_ENV)
.ok()
.filter(|key| !key.trim().is_empty())
})
.ok_or_else(|| Error::from(AuthError::from(MissingCredential::OpenAiRealtimeApiKey)))
}
/// Open the upstream WebSocket to OpenAI for `(model, api_key, api_base)`.
///
/// This is the dial half of [`realtime`], factored out so the pool can
/// pre-establish sockets ahead of any client. `api_key` here is already resolved
/// (non-blank) — the pool resolves it once when it is created.
pub(crate) async fn dial_upstream(
model: &str,
api_key: &str,
api_base: Option<&str>,
) -> Result<UpstreamWs, Error> {
let url = OPENAI_REALTIME_CONFIG.complete_url(api_base, model);
let mut request = url
.as_str()
.into_client_request()
.map_err(|err| Error::Network(err.to_string()))?;
// GA realtime: only Authorization. The legacy OpenAI-Beta header triggers
// beta_api_shape_disabled, so we do not send it.
request.headers_mut().insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {api_key}"))
.map_err(|err| Error::Auth(err.to_string()))?,
);
let (upstream, _response) = connect_upstream(request)
.await
.map_err(|err| Error::Network(err.to_string()))?;
Ok(upstream)
}
/// Read the next text frame from the upstream and decode it as a typed event.
///
/// Used by the pool to pre-read OpenAI's unprompted `session.created`. Returns an
/// error on a non-text frame, a closed socket, or undecodable JSON so the pool can
/// discard a misbehaving socket rather than warm it.
pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> Result<RealtimeEvent, Error> {
loop {
let message = upstream_rx
.next()
.await
.ok_or_else(|| Error::Network("upstream closed before first event".to_string()))?
.map_err(|err| Error::Network(err.to_string()))?;
match message {
Message::Text(text) => {
return serde_json::from_str(&text)
.map_err(|err| Error::InvalidResponse(err.to_string()));
}
// Ignore protocol frames (ping/pong) while waiting for the first event.
Message::Ping(_) | Message::Pong(_) => continue,
Message::Close(_) => {
return Err(Error::Network(
"upstream closed before first event".to_string(),
));
}
_ => continue,
}
}
}
/// Splice an already-connected upstream to the client streams.
///
/// `prelude` is relayed to the client first (the pool passes the buffered
/// `session.created` here; the fresh-dial path passes `None` and lets the upstream
/// deliver it). Then a single select loop forwards both directions through the
/// transforms until either side closes or the idle timeout fires.
/// `observe` is invoked on **upstream→client** events only (the trusted side that
/// carries `session.created` and `response.done` usage) — never on client events,
/// so a client cannot fabricate usage into its own logs.
#[allow(clippy::too_many_arguments)]
pub(crate) async fn splice<In, Out>(
model: &str,
mut upstream_tx: UpstreamTx,
mut upstream_rx: UpstreamRx,
prelude: Option<RealtimeEvent>,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&RealtimeEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::Error: std::fmt::Display,
{
let config = &OPENAI_REALTIME_CONFIG;
// Relay a buffered backend event (warm handoff's session.created) first, so a
// warm session looks identical to a fresh one from the client's view.
if let Some(event) = prelude {
for outbound in config.transform_realtime_response(&event, model)?.events {
client_out
.send(outbound)
.await
.map_err(|err| Error::Network(err.to_string()))?;
}
}
let idle = idle_timeout.unwrap_or(Duration::from_secs(DEFAULT_IDLE_TIMEOUT_SECS));
// One loop forwarding both directions. The `sleep(idle)` arm is rebuilt every
// iteration, so any frame (either way) resets it — it fires only when the
// session has been fully idle for `idle`, reaping a stalled connection
// (task + upstream TCP socket) instead of leaking it.
loop {
tokio::select! {
// client -> upstream
client_event = client_in.next() => {
let Some(event) = client_event else { break }; // client disconnected
// NOTE: do NOT observe client events. session.created / response.done
// (carrying usage) are server→client events; observing the client arm
// would let an authenticated client POST a fabricated response.done and
// inflate its own spend log. Logging observes upstream events only.
for outbound in config.transform_realtime_request(&event, model)?.events {
let payload = serde_json::to_string(&outbound)
.map_err(|err| Error::InvalidResponse(err.to_string()))?;
upstream_tx
.send(Message::Text(payload))
.await
.map_err(|err| Error::Network(err.to_string()))?;
}
}
// upstream -> client
upstream_message = upstream_rx.next() => {
let Some(message) = upstream_message else { break }; // upstream closed
match message.map_err(|err| Error::Network(err.to_string()))? {
Message::Text(text) => {
let event: RealtimeEvent = serde_json::from_str(&text)
.map_err(|err| Error::InvalidResponse(err.to_string()))?;
observe(&event);
for outbound in config.transform_realtime_response(&event, model)?.events {
client_out
.send(outbound)
.await
.map_err(|err| Error::Network(err.to_string()))?;
}
}
Message::Close(_) => break,
_ => {}
}
}
// idle timeout: no activity from either side within `idle`
_ = tokio::time::sleep(idle) => break,
}
}
Ok(())
}
/// Splice a client realtime stream to OpenAI: forward client events upstream
/// (via `transform_realtime_request`) and backend events downstream (via
/// `transform_realtime_response`). Returns when either side closes.
///
/// Generic over the client transport (typed events) so this crate stays
/// framework-agnostic; the gateway adapts its axum socket to these. This is the
/// fresh-dial path: dial, then splice. The pool's warm-handoff path skips the dial
/// and calls [`splice`] directly with a buffered `session.created`.
#[allow(clippy::too_many_arguments)]
pub async fn realtime<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::Error: std::fmt::Display,
{
let api_key = resolve_api_key(api_key)?;
let upstream = dial_upstream(model, &api_key, api_base).await?;
let (upstream_tx, upstream_rx) = upstream.split();
splice(
model,
upstream_tx,
upstream_rx,
None,
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
/// Splice a pre-warmed upstream (taken from [`crate::io::realtime_pool`]) to the
/// client. Relays the buffered `session.created` first, then splices exactly like
/// the fresh-dial path — so a warm session is indistinguishable from a fresh one.
#[allow(clippy::too_many_arguments)]
pub async fn realtime_warm<In, Out>(
model: &str,
handoff: crate::io::realtime_pool::WarmHandoff,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::Error: std::fmt::Display,
{
splice(
model,
handoff.tx,
handoff.rx,
Some(handoff.session_created),
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
fn event(raw: &str) -> RealtimeEvent {
serde_json::from_str(raw).expect("valid event json")
}
/// The realtime dial has to reach a `wss://` upstream without a process-wide
/// crypto provider installed, which is what dialing through `io::tls` buys.
#[tokio::test]
async fn dial_upstream_over_wss_reports_an_error_instead_of_panicking() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind a loopback port");
let port = listener
.local_addr()
.expect("read the bound address")
.port();
tokio::spawn(async move {
while let Ok((stream, _peer)) = listener.accept().await {
drop(stream);
}
});
let result = dial_upstream(
"gpt-realtime",
"sk-test",
Some(&format!("wss://127.0.0.1:{port}")),
)
.await;
assert!(matches!(result, Err(Error::Network(_))));
}
#[test]
fn resolve_api_key_prefers_param_then_blank_falls_through() {
assert_eq!(resolve_api_key(Some("sk-test")).unwrap(), "sk-test");
// A blank param with no env set should error.
if std::env::var(OPENAI_API_KEY_ENV).is_err() {
assert!(resolve_api_key(Some(" ")).is_err());
}
}
/// Live end-to-end check against OpenAI. Ignored by default (CI never runs
/// it); run explicitly with `OPENAI_API_KEY` set:
/// `cargo test -p litellm-ai-gateway --features server realtime_invokes_openai -- --ignored --nocapture`
#[tokio::test]
#[ignore = "hits the live OpenAI realtime API; needs OPENAI_API_KEY"]
async fn realtime_invokes_openai_and_responds() {
use futures_channel::mpsc;
let key =
std::env::var(OPENAI_API_KEY_ENV).expect("set OPENAI_API_KEY to run this ignored test");
// client -> provider (we hold `client_tx` to push events upstream)
let (mut client_tx, client_in) = mpsc::unbounded::<RealtimeEvent>();
// provider -> client (we hold `backend_rx` to read backend events)
let (client_out, mut backend_rx) = mpsc::unbounded::<RealtimeEvent>();
// Clone the key so the spawned task owns its `String` (no borrow across await).
let key_owned = key.clone();
let call = tokio::spawn(async move {
realtime(
"gpt-realtime",
Some(&key_owned),
None,
None,
|_| {},
client_in,
client_out,
)
.await
});
// 1. First backend event should be session.created.
let first = tokio::time::timeout(Duration::from_secs(30), backend_rx.next())
.await
.expect("timed out waiting for session.created")
.expect("backend stream closed before session.created");
assert_eq!(
first.event_type, "session.created",
"expected session.created, got: {}",
first.event_type
);
// 2. Ask for a short audio response.
client_tx
.send(event(
r#"{"type":"conversation.item.create","item":{"type":"message","role":"user","content":[{"type":"input_text","text":"Say hi."}]}}"#,
))
.await
.expect("send conversation.item.create");
client_tx
.send(event(r#"{"type":"response.create"}"#))
.await
.expect("send response.create");
// 3. Read backend events; require a non-empty audio delta, then response.done.
let mut saw_audio_delta = false;
let mut saw_done = false;
for _ in 0..500 {
let next = tokio::time::timeout(Duration::from_secs(30), backend_rx.next()).await;
let event = match next {
Ok(Some(event)) => event,
Ok(None) => break,
Err(_) => panic!("timed out waiting for backend events"),
};
match event.event_type.as_str() {
"response.output_audio.delta" => {
let delta = event
.data
.get("delta")
.and_then(|value| value.as_str())
.unwrap_or("");
if !delta.is_empty() {
saw_audio_delta = true;
}
}
"response.done" => {
saw_done = true;
break;
}
_ => {}
}
}
assert!(
saw_audio_delta,
"expected a response.output_audio.delta with non-empty delta"
);
assert!(saw_done, "expected a response.done event");
// Drop the client sender so the provider's to_upstream side finishes.
drop(client_tx);
let _ = call.await;
}
}

View file

@ -1,712 +0,0 @@
//! Pre-warmed upstream realtime connection pool.
//!
//! The gateway's realtime overhead lives entirely in session establishment: on
//! every client connect it dials a fresh upstream WS to OpenAI and waits for
//! `session.created` before it can serve. This pool keeps a small set of upstream
//! sockets **already connected and already past `session.created`** so a connect
//! can be served from a warm socket and the handshake is off the critical path.
//!
//! Layering: this lives in the gateway's `io` module next to the dial/splice it
//! reuses. The gateway holds an `Arc<RealtimePool>` in its state and asks for a
//! warm socket per connect; on a miss it fresh-dials exactly as before. The pool
//! is a latency optimization, never a correctness dependency — see the gateway's
//! `src/routes/realtime/README.md`.
//!
//! ## Caveats (enforced here)
//! - One warm socket serves exactly one session (realtime isn't multiplexed), so
//! the pool is sized to the connect *rate*, not concurrent connections.
//! - `session.created` is pre-read once and buffered; nothing else is read from a
//! warm socket before handoff, so a warm session starts at OpenAI defaults just
//! like a fresh one (`session.update` semantics unchanged).
//! - Warm sockets are short-lived (`max_idle`) and liveness-checked at handoff to
//! bound idle billing / dodge OpenAI's idle timeout.
//! - On miss or dead socket the caller fresh-dials; the pool never blocks or fails
//! a connect because it is empty.
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use futures_util::StreamExt;
use litellm_core::Error;
use litellm_core::realtime::types::RealtimeEvent;
use crate::io::realtime::{
UpstreamRx, UpstreamTx, UpstreamWs, dial_upstream, read_event, resolve_api_key,
};
/// Default target warm sockets per key when pooling is enabled.
pub const DEFAULT_POOL_SIZE: usize = 4;
/// Default max time a warm socket may sit before it is closed and replaced.
pub const DEFAULT_MAX_IDLE: Duration = Duration::from_secs(30);
/// Env var: target warm sockets per key. `0` disables pooling (fresh-dial only).
pub const POOL_SIZE_ENV: &str = "REALTIME_POOL_SIZE";
/// Env var: max warm-socket idle lifetime, in seconds.
pub const MAX_IDLE_ENV: &str = "REALTIME_POOL_MAX_IDLE_SECS";
/// How often the background replenisher wakes to top up and reap stale sockets.
const REPLENISH_TICK: Duration = Duration::from_millis(250);
/// Backoff floor after a key's warm-up dials all fail. The first failed pass
/// waits this long before retrying that key.
const BACKOFF_BASE: Duration = Duration::from_millis(500);
/// Backoff ceiling. A key that keeps failing (invalid credentials, an
/// unreachable upstream) is retried at most once per this interval — instead of
/// firing `needed` concurrent TLS dials every 250 ms tick, which would hammer
/// the upstream and risk rate-limit exhaustion that degrades valid cold-path
/// traffic. Backoff resets the moment a dial for the key succeeds.
const BACKOFF_MAX: Duration = Duration::from_secs(30);
/// Identifies an upstream connection: the tuple that fully determines the dial.
/// `api_key` is included so a warm socket is only ever reused for the same key
/// (no cross-tenant reuse).
#[derive(Clone, PartialEq, Eq, Hash)]
pub struct UpstreamKey {
pub model: String,
pub api_key: String,
pub api_base: Option<String>,
}
impl std::fmt::Debug for UpstreamKey {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UpstreamKey")
.field("model", &self.model)
.field("api_key", &"[REDACTED]")
.field("api_base", &self.api_base)
.finish()
}
}
/// A warm upstream: split halves + the buffered `session.created` + when it was
/// warmed (for `max_idle` expiry).
struct WarmConnection {
tx: UpstreamTx,
rx: UpstreamRx,
session_created: RealtimeEvent,
warmed_at: Instant,
}
/// A live upstream taken from the pool, ready to splice. The caller relays
/// `session_created` to the client first, then splices `(tx, rx)` as usual.
pub struct WarmHandoff {
pub tx: UpstreamTx,
pub rx: UpstreamRx,
pub session_created: RealtimeEvent,
}
/// Pool configuration, resolved once at startup from the environment.
#[derive(Clone, Copy, Debug)]
pub struct PoolConfig {
/// Target warm sockets per key. `0` disables pooling.
pub target_size: usize,
/// Max time a warm socket may sit before it is closed and replaced.
pub max_idle: Duration,
}
impl Default for PoolConfig {
fn default() -> Self {
Self {
target_size: DEFAULT_POOL_SIZE,
max_idle: DEFAULT_MAX_IDLE,
}
}
}
impl PoolConfig {
/// Read config from the environment, falling back to defaults. An invalid
/// value warns and uses the default rather than failing startup.
pub fn from_env() -> Self {
let target_size = match std::env::var(POOL_SIZE_ENV) {
Ok(raw) => raw.trim().parse().unwrap_or_else(|_| {
eprintln!("warning: {POOL_SIZE_ENV}={raw:?} is not a valid size; using {DEFAULT_POOL_SIZE}");
DEFAULT_POOL_SIZE
}),
Err(_) => DEFAULT_POOL_SIZE,
};
let max_idle = match std::env::var(MAX_IDLE_ENV) {
Ok(raw) => raw
.trim()
.parse()
.map(Duration::from_secs)
.unwrap_or_else(|_| {
eprintln!(
"warning: {MAX_IDLE_ENV}={raw:?} is not a valid number of seconds; using {}s",
DEFAULT_MAX_IDLE.as_secs()
);
DEFAULT_MAX_IDLE
}),
Err(_) => DEFAULT_MAX_IDLE,
};
Self {
target_size,
max_idle,
}
}
/// Whether pooling is on (`target_size > 0`).
pub fn enabled(&self) -> bool {
self.target_size > 0
}
}
/// Per-key warm sockets, behind a single `Mutex`. Realtime warm sockets are few
/// (the pool is small), so a plain mutex over a `VecDeque`-ish `Vec` is simpler
/// and faster than sharding; contention is negligible at this scale.
type Warm = HashMap<UpstreamKey, Vec<WarmConnection>>;
/// Per-key replenish backoff. Absent (or `consecutive_failures == 0`) means the
/// key is healthy and replenished every tick. After a pass whose dials all fail,
/// `retry_after` is pushed out with exponential backoff so a broken key (invalid
/// credentials, unreachable upstream) is not re-dialed on every 250 ms tick.
#[derive(Default)]
struct Backoff {
/// Don't attempt warm-up dials for this key until this instant. `None` =
/// eligible now.
retry_after: Option<Instant>,
consecutive_failures: u32,
}
type Backoffs = HashMap<UpstreamKey, Backoff>;
/// Pre-warmed upstream realtime connection pool.
///
/// Cheap to clone-via-`Arc`. The background replenisher is spawned by
/// [`RealtimePool::spawn`]; a pool built with [`RealtimePool::disabled`] never
/// warms anything and every `take` misses (callers fresh-dial).
pub struct RealtimePool {
config: PoolConfig,
warm: Mutex<Warm>,
/// Per-key replenish backoff so a broken key doesn't trigger unbounded
/// concurrent dials every tick. Separate lock from `warm` so the request
/// hot path (`take`) never contends on it.
backoff: Mutex<Backoffs>,
}
impl RealtimePool {
/// A disabled pool: no background task, every `take` returns `None`.
pub fn disabled() -> Arc<Self> {
Arc::new(Self {
config: PoolConfig {
target_size: 0,
..PoolConfig::default()
},
warm: Mutex::new(HashMap::new()),
backoff: Mutex::new(HashMap::new()),
})
}
/// Build a pool from config **without** the background replenisher. The pool
/// only warms when [`RealtimePool::warm_now`] is called. Used by deterministic
/// unit tests; production uses [`RealtimePool::spawn`].
#[cfg(test)]
fn new_unspawned(config: PoolConfig) -> Arc<Self> {
Arc::new(Self {
config,
warm: Mutex::new(HashMap::new()),
backoff: Mutex::new(HashMap::new()),
})
}
/// Build a pool from config and, if enabled, spawn the background replenisher.
/// Returns the shared handle the gateway stores in its state.
pub fn spawn(config: PoolConfig) -> Arc<Self> {
let pool = Arc::new(Self {
config,
warm: Mutex::new(HashMap::new()),
backoff: Mutex::new(HashMap::new()),
});
if config.enabled() {
let weak = Arc::downgrade(&pool);
tokio::spawn(async move {
let mut tick = tokio::time::interval(REPLENISH_TICK);
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tick.tick().await;
// Stop once the gateway has dropped its handle.
let Some(pool) = weak.upgrade() else { break };
pool.replenish_all().await;
}
});
}
pool
}
/// Resolved config (test/inspection).
pub fn config(&self) -> PoolConfig {
self.config
}
/// Register a key so the replenisher starts warming it. Idempotent. The
/// gateway calls this once per known deployment at startup; the pool only
/// warms keys it has seen, so it never dials a model nobody asked for.
pub fn register(&self, key: UpstreamKey) {
if !self.config.enabled() {
return;
}
self.warm.lock().unwrap().entry(key).or_default();
}
/// Take a warm, live socket for `key`, or `None` on miss / dead socket.
///
/// Pops the freshest non-expired socket and liveness-checks it; a socket that
/// is too old or already dead is dropped (closing it) and the next candidate
/// tried. Never blocks: if nothing warm is live, returns `None` so the caller
/// fresh-dials.
pub fn take(&self, key: &UpstreamKey) -> Option<WarmHandoff> {
if !self.config.enabled() {
return None;
}
loop {
let mut candidate = {
let mut warm = self.warm.lock().unwrap();
let bucket = warm.get_mut(key)?;
bucket.pop()?
};
// Discard sockets past their warm lifetime (idle-billing guard).
if candidate.warmed_at.elapsed() > self.config.max_idle {
continue; // drops `candidate`, closing the socket
}
// Liveness: a non-blocking check that the socket hasn't already
// delivered a Close/Err. A warm socket should be silent after
// session.created, so anything pending means it is unhealthy.
if is_dead(&mut candidate.rx) {
continue;
}
return Some(WarmHandoff {
tx: candidate.tx,
rx: candidate.rx,
session_created: candidate.session_created,
});
}
}
/// One replenish pass over every registered key: reap stale sockets, then
/// dial up to `target_size`. Dials run concurrently; failures are swallowed
/// (a key that can't be warmed just keeps fresh-dialing on the request path)
/// and put the key into exponential backoff so a broken key isn't re-dialed
/// on every tick.
async fn replenish_all(&self) {
let keys: Vec<UpstreamKey> = { self.warm.lock().unwrap().keys().cloned().collect() };
for key in keys {
self.reap_stale(&key);
// Skip keys still in backoff from a prior all-failed pass — this is
// what bounds dials against an invalid/unreachable key to once per
// `BACKOFF_MAX` instead of `needed` dials every 250 ms tick.
if self.in_backoff(&key) {
continue;
}
let needed = {
let warm = self.warm.lock().unwrap();
let have = warm.get(&key).map(Vec::len).unwrap_or(0);
self.config.target_size.saturating_sub(have)
};
if needed == 0 {
continue;
}
// Dial the missing sockets CONCURRENTLY. A sequential loop here makes
// a full refill cost `needed × handshake` (~needed × 350 ms), which
// can't keep up with a high connect rate — the pool drains faster
// than it refills and most connects miss. Firing the dials together
// refills in ~one handshake window, keeping warm supply ≈ peak
// concurrent connects so the sub-ms warm handoff becomes the median,
// not the lucky-hit tail.
let dials = (0..needed).map(|_| warm_one(&key));
let results = futures_util::future::join_all(dials).await;
let mut any_ok = false;
// `.flatten()` keeps only the successful dials; a key that can't be
// warmed just keeps fresh-dialing on the request path.
for conn in results.into_iter().flatten() {
any_ok = true;
self.warm
.lock()
.unwrap()
.entry(key.clone())
.or_default()
.push(conn);
}
// Reset backoff on any success; otherwise grow it. We only ever enter
// backoff when a pass that *attempted* dials produced none — a `needed
// == 0` pass is handled by the `continue` above and never touches it.
self.record_replenish_outcome(&key, any_ok);
}
}
/// Whether `key` is currently in a backoff window (a prior pass failed and
/// the retry time hasn't arrived). Eligible keys are pruned from the backoff
/// map so it doesn't grow unbounded for healthy keys.
fn in_backoff(&self, key: &UpstreamKey) -> bool {
let mut backoff = self.backoff.lock().unwrap();
match backoff.get(key).and_then(|b| b.retry_after) {
Some(retry_after) if Instant::now() < retry_after => true,
Some(_) => {
// Window elapsed — allow the attempt. Keep the failure count so a
// still-broken key backs off further, but clear the gate so this
// tick proceeds.
if let Some(b) = backoff.get_mut(key) {
b.retry_after = None;
}
false
}
None => false,
}
}
/// Update a key's backoff after a replenish attempt. Success clears it;
/// failure grows the retry delay exponentially up to `BACKOFF_MAX`.
fn record_replenish_outcome(&self, key: &UpstreamKey, any_ok: bool) {
let mut backoff = self.backoff.lock().unwrap();
if any_ok {
backoff.remove(key);
return;
}
let entry = backoff.entry(key.clone()).or_default();
entry.consecutive_failures = entry.consecutive_failures.saturating_add(1);
// Exponential: BASE * 2^(failures-1), saturating at MAX. `min` of the
// shift exponent keeps the doubling from overflowing.
let shift = (entry.consecutive_failures - 1).min(16);
let delay = BACKOFF_BASE.saturating_mul(1u32 << shift).min(BACKOFF_MAX);
entry.retry_after = Some(Instant::now() + delay);
}
/// Drop sockets past `max_idle` or already dead for a key.
fn reap_stale(&self, key: &UpstreamKey) {
let mut warm = self.warm.lock().unwrap();
if let Some(bucket) = warm.get_mut(key) {
bucket.retain_mut(|conn| {
conn.warmed_at.elapsed() <= self.config.max_idle && !is_dead(&mut conn.rx)
});
}
}
/// Test/inspection: number of warm sockets currently held for `key`.
#[cfg(test)]
pub fn warm_len(&self, key: &UpstreamKey) -> usize {
self.warm
.lock()
.unwrap()
.get(key)
.map(Vec::len)
.unwrap_or(0)
}
/// Test/inspection: consecutive replenish failures recorded for `key` (0 if
/// the key is healthy / has no backoff entry).
#[cfg(test)]
pub fn backoff_failures(&self, key: &UpstreamKey) -> u32 {
self.backoff
.lock()
.unwrap()
.get(key)
.map(|b| b.consecutive_failures)
.unwrap_or(0)
}
/// Test helper: synchronously warm `target_size` sockets for `key` (no
/// background task). Lets tests assert handoff behavior deterministically.
#[cfg(test)]
pub async fn warm_now(&self, key: &UpstreamKey) {
let needed = {
let warm = self.warm.lock().unwrap();
let have = warm.get(key).map(Vec::len).unwrap_or(0);
self.config.target_size.saturating_sub(have)
};
for _ in 0..needed {
if let Ok(conn) = warm_one(key).await {
self.warm
.lock()
.unwrap()
.entry(key.clone())
.or_default()
.push(conn);
}
}
}
/// Test helper: insert an already-built warm connection (used to inject a
/// dead socket and assert it is discarded at handoff).
#[cfg(test)]
fn insert_warm(&self, key: UpstreamKey, conn: WarmConnection) {
self.warm.lock().unwrap().entry(key).or_default().push(conn);
}
}
/// Dial one upstream and pre-read its `session.created` into a [`WarmConnection`].
///
/// `key.api_key` is already resolved (non-blank). The first frame OpenAI sends
/// unprompted is `session.created`; we buffer exactly that and read nothing more.
async fn warm_one(key: &UpstreamKey) -> Result<WarmConnection, Error> {
let upstream: UpstreamWs =
dial_upstream(&key.model, &key.api_key, key.api_base.as_deref()).await?;
let (tx, mut rx) = upstream.split();
let session_created = read_event(&mut rx).await?;
Ok(WarmConnection {
tx,
rx,
session_created,
warmed_at: Instant::now(),
})
}
/// Resolve a deployment's API key into the pool key, returning `None` when no key
/// can be resolved (those deployments simply aren't pooled — the request path
/// still fresh-dials and surfaces the auth error there).
pub fn upstream_key(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
) -> Option<UpstreamKey> {
let api_key = resolve_api_key(api_key).ok()?;
Some(UpstreamKey {
model: model.to_string(),
api_key,
api_base: api_base.map(str::to_string),
})
}
/// Non-blocking liveness check: poll the upstream once. A warm socket is silent
/// after `session.created`, so a pending `Close`/`Err`/`None` means it is dead.
/// A pending data frame (shouldn't happen pre-handoff) is also treated as
/// unhealthy — we'd rather discard and fresh-dial than hand over a socket in an
/// unexpected state. `Pending` (the healthy case) returns `false`.
fn is_dead(rx: &mut UpstreamRx) -> bool {
use futures_util::Stream;
use futures_util::task::noop_waker_ref;
use std::pin::Pin;
use std::task::{Context, Poll};
let mut cx = Context::from_waker(noop_waker_ref());
match Pin::new(rx).poll_next(&mut cx) {
Poll::Pending => false,
Poll::Ready(None) => true,
Poll::Ready(Some(Err(_))) => true,
// Any frame arriving before handoff is unexpected for a silent warm
// socket; treat it as unhealthy.
Poll::Ready(Some(Ok(_))) => true,
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures_util::SinkExt;
use std::net::SocketAddr;
use tokio::net::TcpListener;
use tokio_tungstenite::tungstenite::Message;
/// An in-process fake OpenAI realtime WS server. On connect it sends
/// `session.created`; on `response.create` it sends `response.created` +
/// `response.output_audio.delta` + `response.done`. Returns its `ws://` base.
async fn spawn_fake_openai() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr: SocketAddr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
tokio::spawn(handle_fake_conn(stream));
}
});
format!("ws://{addr}")
}
async fn handle_fake_conn(stream: tokio::net::TcpStream) {
let mut ws = match tokio_tungstenite::accept_async(stream).await {
Ok(ws) => ws,
Err(_) => return,
};
// Unprompted session.created, exactly like OpenAI.
let _ = ws
.send(Message::Text(
r#"{"type":"session.created","session":{"id":"sess_fake"}}"#.to_string(),
))
.await;
while let Some(Ok(msg)) = ws.next().await {
if let Message::Text(text) = msg
&& text.contains("response.create")
{
for frame in [
r#"{"type":"response.created"}"#,
r#"{"type":"response.output_audio.delta","delta":"AAAA"}"#,
r#"{"type":"response.done"}"#,
] {
let _ = ws.send(Message::Text(frame.to_string())).await;
}
}
}
}
fn test_config() -> PoolConfig {
PoolConfig {
target_size: 2,
max_idle: Duration::from_secs(30),
}
}
fn key_for(base: &str) -> UpstreamKey {
UpstreamKey {
model: "gpt-realtime".to_string(),
api_key: "sk-test".to_string(),
api_base: Some(base.to_string()),
}
}
#[tokio::test]
async fn warm_handoff_relays_buffered_session_created() {
let base = spawn_fake_openai().await;
let pool = RealtimePool::new_unspawned(test_config());
let key = key_for(&base);
pool.register(key.clone());
pool.warm_now(&key).await;
assert_eq!(pool.warm_len(&key), 2);
let handoff = pool.take(&key).expect("a warm socket should be available");
assert_eq!(handoff.session_created.event_type, "session.created");
assert_eq!(
handoff
.session_created
.data
.get("session")
.and_then(|s| s.get("id"))
.and_then(|v| v.as_str()),
Some("sess_fake")
);
// Taking one leaves one.
assert_eq!(pool.warm_len(&key), 1);
}
#[tokio::test]
async fn pool_miss_returns_none_for_fresh_dial_fallback() {
let base = spawn_fake_openai().await;
let pool = RealtimePool::new_unspawned(test_config());
let key = key_for(&base);
// Registered but never warmed → empty bucket → miss.
pool.register(key.clone());
assert!(pool.take(&key).is_none());
// Unknown key → miss.
let other = key_for("ws://127.0.0.1:1");
assert!(pool.take(&other).is_none());
}
#[tokio::test]
async fn disabled_pool_never_hands_off() {
let pool = RealtimePool::disabled();
let key = key_for("ws://127.0.0.1:1");
pool.register(key.clone());
assert_eq!(pool.warm_len(&key), 0);
assert!(pool.take(&key).is_none());
}
#[tokio::test]
async fn dead_warm_socket_is_discarded() {
let base = spawn_fake_openai().await;
let pool = RealtimePool::new_unspawned(test_config());
let key = key_for(&base);
pool.register(key.clone());
// Build one real warm connection, then kill the upstream by dropping the
// server side: easiest is to dial, read session.created, then close our
// own rx's peer. Instead we forge "dead" via an already-closed socket:
// dial a connection and immediately send a Close from the client side so
// the server closes back, then warm it. Simpler: warm normally, then
// mark it stale by backdating warmed_at past max_idle and confirm it's
// dropped — that exercises the same discard path.
let mut conn = warm_one(&key).await.expect("warm one");
conn.warmed_at = Instant::now() - Duration::from_secs(3600); // past max_idle
pool.insert_warm(key.clone(), conn);
assert_eq!(pool.warm_len(&key), 1);
// take() must discard the stale socket and report a miss.
assert!(pool.take(&key).is_none());
assert_eq!(pool.warm_len(&key), 0);
}
#[tokio::test]
async fn background_replenisher_tops_up_registered_key() {
let base = spawn_fake_openai().await;
let pool = RealtimePool::spawn(test_config());
let key = key_for(&base);
pool.register(key.clone());
// Wait (bounded) for the background task to reach the target size.
let mut warmed = 0;
for _ in 0..40 {
tokio::time::sleep(Duration::from_millis(50)).await;
warmed = pool.warm_len(&key);
if warmed >= test_config().target_size {
break;
}
}
assert_eq!(
warmed,
test_config().target_size,
"background replenisher should warm up to target_size"
);
let handoff = pool.take(&key).expect("a warm socket should be available");
assert_eq!(handoff.session_created.event_type, "session.created");
}
#[tokio::test]
async fn closed_upstream_socket_is_detected_dead() {
// A genuinely dead socket: dial the fake, read session.created, then drop
// the server by closing from our side and waiting for the close to land.
let base = spawn_fake_openai().await;
let pool = RealtimePool::new_unspawned(test_config());
let key = key_for(&base);
pool.register(key.clone());
let mut conn = warm_one(&key).await.expect("warm one");
// Close the upstream from the client side; the server echoes a close.
let _ = conn.tx.send(Message::Close(None)).await;
// Give the close a moment to arrive on rx.
tokio::time::sleep(Duration::from_millis(50)).await;
pool.insert_warm(key.clone(), conn);
// Liveness check at take() should detect the close and discard it.
assert!(pool.take(&key).is_none());
assert_eq!(pool.warm_len(&key), 0);
}
#[tokio::test]
async fn broken_key_backs_off_instead_of_dialing_every_tick() {
// A key whose upstream is unreachable: every warm-up dial fails.
let pool = RealtimePool::new_unspawned(test_config());
let key = key_for("ws://127.0.0.1:1"); // nothing listens here
pool.register(key.clone());
// First pass attempts dials, they all fail → key enters backoff, no warm
// sockets, one recorded failure.
pool.replenish_all().await;
assert_eq!(pool.warm_len(&key), 0);
assert_eq!(pool.backoff_failures(&key), 1);
assert!(
pool.in_backoff(&key),
"a key whose dials all failed must be in backoff"
);
// An immediate next pass must be SKIPPED (still in the backoff window), so
// it does NOT fire another round of dials — the failure count is unchanged.
pool.replenish_all().await;
assert_eq!(
pool.backoff_failures(&key),
1,
"replenish during the backoff window must not re-dial the broken key"
);
}
#[tokio::test]
async fn healthy_key_never_enters_backoff_and_clears_after_recovery() {
let base = spawn_fake_openai().await;
let pool = RealtimePool::new_unspawned(test_config());
let key = key_for(&base);
pool.register(key.clone());
// A reachable upstream: the pass succeeds, so the key is never backed off.
pool.replenish_all().await;
assert_eq!(pool.warm_len(&key), test_config().target_size);
assert_eq!(pool.backoff_failures(&key), 0);
assert!(!pool.in_backoff(&key));
}
}

View file

@ -1,485 +0,0 @@
use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use litellm_core::AuthError;
use litellm_core::Error;
use litellm_core::auth::error::MissingCredential;
use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG;
use litellm_core::responses::types::ResponsesWsEvent;
use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
use litellm_core::responses::websocket::{ResponsesUpstreamWs, connect_upstream};
use crate::constants::{
DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS,
};
const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
type UpstreamTx = SplitSink<ResponsesUpstreamWs, Message>;
type UpstreamRx = SplitStream<ResponsesUpstreamWs>;
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result<String, Error> {
api_key
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.or_else(|| {
std::env::var(OPENAI_API_KEY_ENV)
.ok()
.filter(|value| !value.trim().is_empty())
})
.ok_or_else(|| Error::from(AuthError::from(MissingCredential::OpenAiResponsesApiKey)))
}
async fn dial_upstream(
model: &str,
api_key: &str,
api_base: Option<&str>,
) -> Result<ResponsesUpstreamWs, Error> {
let url = OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, model);
let mut request = url
.as_str()
.into_client_request()
.map_err(|error| Error::Network(error.to_string()))?;
request.headers_mut().insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {api_key}"))
.map_err(|error| Error::Auth(error.to_string()))?,
);
let result = tokio::time::timeout(
Duration::from_secs(DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS),
connect_upstream(request),
)
.await
.map_err(|_| Error::Network("Responses WebSocket connection timed out".to_string()))?;
result
.map(|(socket, _)| socket)
.map_err(|error| match *error {
tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http {
status: response.status().as_u16(),
body: String::new(),
},
other => Error::Network(other.to_string()),
})
}
pub struct ResponsesWebSocketStreaming;
impl ResponsesWebSocketStreaming {
pub async fn bidirectional_forward<In, Out>(
model: &str,
upstream_tx: UpstreamTx,
upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
splice(
model,
upstream_tx,
upstream_rx,
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
}
pub(crate) async fn splice<In, Out>(
model: &str,
mut upstream_tx: UpstreamTx,
mut upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let idle =
idle_timeout.unwrap_or_else(|| Duration::from_secs(DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS));
loop {
tokio::select! {
event = client_in.next() => {
let Some(event) = event else { break };
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_request(&event, model)?
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx.send(Message::Text(payload))
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
message = upstream_rx.next() => {
let Some(message) = message else { break };
match message.map_err(|error| Error::Network(error.to_string()))? {
Message::Text(text) => {
let event = serde_json::from_str::<ResponsesWsEvent>(&text)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
observe(&event);
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_response(&event, model)?
.events
{
client_out.send(outbound)
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
Message::Close(_) => break,
_ => {}
}
}
_ = tokio::time::sleep(idle) => break,
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub async fn async_responses_websocket<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let key = resolve_api_key(api_key)?;
let upstream = dial_upstream(model, &key, api_base).await?;
let (mut upstream_tx, upstream_rx) = upstream.split();
if let Some(first_frame) = first_frame {
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_request(&first_frame, model)?
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx
.send(Message::Text(payload))
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
ResponsesWebSocketStreaming::bidirectional_forward(
model,
upstream_tx,
upstream_rx,
idle_timeout,
&mut observe,
client_in,
client_out,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn responses_ws<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
async_responses_websocket(
model,
api_key,
api_base,
first_frame,
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
use futures_channel::mpsc;
use futures_util::{SinkExt, StreamExt};
use litellm_core::responses::types::ResponsesWsEventType;
use serde_json::json;
use tokio::io::AsyncWriteExt;
use tokio::net::TcpListener;
use tokio_tungstenite::accept_async;
/// The Responses dial has to reach a `wss://` upstream without a process-wide
/// crypto provider installed, which is what dialing through `io::tls` buys.
#[tokio::test]
async fn dial_upstream_over_wss_reports_an_error_instead_of_panicking() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind a loopback port");
let port = listener
.local_addr()
.expect("read the bound address")
.port();
tokio::spawn(async move {
while let Ok((stream, _peer)) = listener.accept().await {
drop(stream);
}
});
let result =
dial_upstream("gpt-5", "sk-test", Some(&format!("wss://127.0.0.1:{port}"))).await;
assert!(matches!(result, Err(Error::Network(_))));
}
async fn websocket_base() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("local address");
let task = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let mut socket = accept_async(stream).await.expect("websocket handshake");
while let Some(Ok(Message::Text(text))) = socket.next().await {
let request: serde_json::Value = serde_json::from_str(&text).expect("request json");
let model = request
.get("model")
.and_then(serde_json::Value::as_str)
.or_else(|| {
request
.get("response")
.and_then(serde_json::Value::as_object)
.and_then(|response| {
response.get("model").and_then(serde_json::Value::as_str)
})
})
.expect("enforced model");
socket
.send(Message::Text(
json!({
"type": "response.created",
"response": {
"id": format!("resp-{model}"),
"model": model,
"extra": "preserved"
}
})
.to_string(),
))
.await
.expect("created event");
socket
.send(Message::Text(
json!({
"type": "response.completed",
"response": {
"id": format!("resp-{model}"),
"model": model,
"usage": {
"input_tokens": 1,
"output_tokens": 2,
"total_tokens": 3
}
}
})
.to_string(),
))
.await
.expect("completed event");
}
});
(format!("http://{address}"), task)
}
fn event(value: serde_json::Value) -> ResponsesWsEvent {
serde_json::from_value(value).expect("event")
}
#[test]
fn explicit_nonblank_key_wins() {
assert_eq!(
resolve_api_key(Some(" explicit ")).expect("key"),
"explicit"
);
}
#[test]
fn blank_key_is_not_accepted_without_environment_key() {
if std::env::var(OPENAI_API_KEY_ENV).is_err() {
assert!(resolve_api_key(Some(" ")).is_err());
}
}
#[tokio::test]
async fn forwards_events_sequentially_and_enforces_model() {
let (api_base, server) = websocket_base().await;
let (client_tx, client_rx) = mpsc::unbounded();
let (output_tx, mut output_rx) = mpsc::unbounded();
let (observed_tx, observed_rx) = mpsc::unbounded();
client_tx
.unbounded_send(event(json!({
"type": "response.create",
"model": "wrong"
})))
.expect("first request");
client_tx
.unbounded_send(event(json!({
"type": "response.create",
"response": {"model": "also-wrong"}
})))
.expect("second request");
let task = tokio::spawn(async move {
responses_ws(
"authorized-model",
Some("test-key"),
Some(&api_base),
None,
Some(Duration::from_secs(1)),
move |event| {
observed_tx
.unbounded_send(event.clone())
.expect("observe event");
},
client_rx,
output_tx,
)
.await
});
let first = output_rx.next().await.expect("first output");
let second = output_rx.next().await.expect("second output");
let third = output_rx.next().await.expect("third output");
let fourth = output_rx.next().await.expect("fourth output");
drop(client_tx);
task.await.expect("splice task").expect("successful splice");
server.await.expect("server task");
assert_eq!(first.event_type, ResponsesWsEventType::ResponseCreated);
assert_eq!(first.model(), Some("authorized-model"));
assert_eq!(first.data["response"]["extra"], "preserved");
assert_eq!(second.event_type, ResponsesWsEventType::ResponseCompleted);
assert_eq!(third.event_type, ResponsesWsEventType::ResponseCreated);
assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted);
let observed: Vec<_> = observed_rx.collect().await;
assert_eq!(observed.len(), 4);
assert!(
observed
.iter()
.all(|event| event.event_type != ResponsesWsEventType::ResponseCreate)
);
}
#[tokio::test]
async fn idle_timeout_ends_without_upstream_events() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let _socket = accept_async(stream).await.expect("handshake");
tokio::time::sleep(Duration::from_secs(1)).await;
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, mut output_rx) = mpsc::unbounded();
let result = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await;
assert!(result.is_ok());
assert!(output_rx.next().await.is_none());
server.abort();
}
#[tokio::test]
async fn dial_http_status_is_preserved() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
stream
.write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n")
.await
.expect("response");
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, _output_rx) = mpsc::unbounded();
let error = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await
.expect_err("status error");
assert!(matches!(error, Error::Http { status: 401, .. }));
server.await.expect("server task");
}
#[tokio::test]
async fn dial_http_500_status_is_preserved() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
stream
.write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n")
.await
.expect("response");
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, _output_rx) = mpsc::unbounded();
let error = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await
.expect_err("status error");
assert!(matches!(error, Error::Http { status: 500, .. }));
server.await.expect("server task");
}
}

View file

@ -1,80 +0,0 @@
//! Outbound WebSocket dials over a TLS config this crate builds once and owns.
//!
//! `reqwest/rustls-tls` enables `rustls/ring` and `litellm-core`'s `bedrock-auth`
//! enables `rustls/aws-lc-rs`, so the bare `ClientConfig::builder()` that
//! `tokio-tungstenite` uses when handed no connector panics rather than guess
//! between them. Naming ring on a connector of our own settles that for these
//! dials without touching the process-wide default, and building the config
//! once keeps the platform trust store, which `tokio-tungstenite` would
//! otherwise re-read on every dial, off the dial path.
use std::io;
use std::sync::{Arc, OnceLock};
use rustls::{ClientConfig, RootCertStore};
use tokio::net::TcpStream;
use tokio_tungstenite::tungstenite::Error;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::error::TlsError;
use tokio_tungstenite::tungstenite::handshake::client::Response;
use tokio_tungstenite::{
Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config,
};
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
fn build_config() -> Result<ClientConfig, Box<Error>> {
let native = rustls_native_certs::load_native_certs();
let roots = {
let mut store = RootCertStore::empty();
let (added, _ignored) = store.add_parsable_certificates(native.certs);
if added == 0 {
return Err(Box::new(Error::Io(io::Error::other(format!(
"no usable native root certificates: {:?}",
native.errors
)))));
}
store
};
ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider()))
.with_safe_default_protocol_versions()
.map(|builder| builder.with_root_certificates(roots).with_no_client_auth())
.map_err(|error| Box::new(Error::Tls(TlsError::Rustls(error))))
}
fn tls_config() -> Result<Arc<ClientConfig>, Box<Error>> {
if let Some(config) = TLS_CONFIG.get() {
return Ok(Arc::clone(config));
}
let built = Arc::new(build_config()?);
Ok(Arc::clone(TLS_CONFIG.get_or_init(|| built)))
}
pub(crate) async fn connect_upstream<R>(
request: R,
) -> Result<(WebSocketStream<MaybeTlsStream<TcpStream>>, Response), Box<Error>>
where
R: IntoClientRequest + Unpin,
{
let request = request.into_client_request().map_err(Box::new)?;
let connector = match request.uri().scheme_str() {
Some("wss") => Some(Connector::Rustls(tls_config()?)),
_ => None,
};
connect_async_tls_with_config(request, None, false, connector)
.await
.map_err(Box::new)
}
#[cfg(test)]
mod tests {
use super::build_config;
#[test]
fn builds_a_usable_config_with_both_provider_features_enabled() {
let config = build_config().expect("a client config");
assert!(!config.crypto_provider().cipher_suites.is_empty());
}
}

View file

@ -1,32 +0,0 @@
//! LiteLLM AI Gateway library.
//!
//! Two layers, split by feature so the Python `cdylib` can depend on the I/O
//! without pulling in the HTTP server:
//!
//! - Call-type modules such as [`ocr`]: provider transforms, lifecycle hooks,
//! and provider I/O. Always available — no feature required. These predate the
//! rule that a route's entrypoint and handler live in `litellm-core` (see
//! `litellm_core::messages`) and move there as they are touched.
//! - [`io`]: compatibility exports and realtime WebSocket splice helpers.
//! - The server modules ([`auth`], [`routes`], [`state`]) and anything pulling
//! `axum` are gated behind the `server` feature, which the `litellm-ai-gateway`
//! binary turns on.
pub mod audio_transcription;
mod client;
pub mod io;
pub mod ocr;
#[cfg(feature = "server")]
pub mod auth;
#[cfg(feature = "server")]
pub mod routes;
#[cfg(feature = "server")]
pub mod state;
#[cfg(feature = "trace-parity")]
pub mod trace_parity;
mod constants;
pub mod integrations;
#[cfg(feature = "server")]
mod realtime;

View file

@ -1,162 +0,0 @@
//! LiteLLM AI Gateway — a minimal Axum server fronting the Rust router.
//!
//! Flow: client → `POST /v1/realtime` → `router.realtime()` selects a deployment
//! (simple-shuffle) → `io::realtime::realtime()` invokes OpenAI. The
//! server owns transport + config; routing lives in the `router` crate.
//!
//! The binary requires the `server` feature (declared in `Cargo.toml` via
//! `required-features`), so cargo skips it unless that feature is on. Everything
//! the binary needs lives in the library (`litellm_ai_gateway`); `main` just
//! wires startup.
use std::sync::Arc;
use litellm_ai_gateway::io::realtime_pool::{PoolConfig, RealtimePool, upstream_key};
use litellm_ai_gateway::routes;
use litellm_ai_gateway::state::AppState;
#[cfg(feature = "python-config")]
use litellm_config::load_model_list;
use litellm_core::router::{Deployment, LiteLLMParams, Router};
use litellm_ai_gateway::integrations::custom_logger::CustomLogger;
use litellm_ai_gateway::integrations::litellm_python_proxy_api::LiteLLMPythonProxyAPILogger;
/// Bind to localhost by default so the gateway is not a public, unauthenticated
/// provider proxy out of the box. Override with `HOST` (e.g. `0.0.0.0`).
const DEFAULT_HOST: &str = "127.0.0.1";
const DEFAULT_PORT: u16 = 4001;
#[tokio::main]
async fn main() {
// Trim before storing so it matches the trimmed bearer token in `auth`
// (avoids a silent auth failure when the env var has surrounding whitespace).
let master_key: Option<Arc<str>> = std::env::var("LITELLM_MASTER_KEY")
.ok()
.map(|key| key.trim().to_string())
.filter(|key| !key.is_empty())
.map(Arc::from);
if master_key.is_none() {
eprintln!(
"warning: LITELLM_MASTER_KEY is not set; /v1/realtime will reject all requests (fail closed)"
);
}
// Spawn the realtime-logging worker (drains a channel → POSTs batches to the
// 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 router = Arc::new(build_router());
// Build the pre-warmed realtime pool and register each deployment's upstream
// so the background replenisher starts warming it. `REALTIME_POOL_SIZE=0`
// yields a disabled pool → every connect fresh-dials (original behavior).
let pool_config = PoolConfig::from_env();
let realtime_pool = RealtimePool::spawn(pool_config);
if pool_config.enabled() {
register_deployments(&router, &realtime_pool);
eprintln!(
"realtime connection pool enabled: target {} warm sockets/key, max idle {}s",
pool_config.target_size,
pool_config.max_idle.as_secs()
);
} else {
eprintln!(
"realtime connection pool disabled (REALTIME_POOL_SIZE=0); fresh-dialing each connect"
);
}
let state = AppState {
router,
master_key,
loggers: Arc::new(loggers),
realtime_pool,
};
let host = std::env::var("HOST").unwrap_or_else(|_| DEFAULT_HOST.to_string());
let port = resolve_port();
let listener = tokio::net::TcpListener::bind((host.as_str(), port))
.await
.expect("failed to bind listener");
eprintln!("litellm-ai-gateway listening on {host}:{port}");
axum::serve(listener, routes::app(state))
.await
.expect("server error");
}
/// Register every deployment's upstream key with the pool so the replenisher
/// pre-warms it. Mirrors `service::run`'s key derivation (strip `openai/`, resolve
/// api_key); deployments whose key can't be resolved are skipped (they fresh-dial
/// and surface the auth error on the request path, as before).
fn register_deployments(router: &Router, pool: &RealtimePool) {
for deployment in router.deployments() {
let params = &deployment.litellm_params;
let provider_model = params
.model
.strip_prefix("openai/")
.unwrap_or(&params.model);
if let Some(key) = upstream_key(
provider_model,
params.api_key.as_deref(),
params.api_base.as_deref(),
) {
pool.register(key);
}
}
}
/// Resolve `PORT`, warning (rather than silently defaulting) on an invalid value.
fn resolve_port() -> u16 {
match std::env::var("PORT") {
Ok(raw) => raw.parse().unwrap_or_else(|_| {
eprintln!("warning: PORT={raw:?} is not a valid port; using {DEFAULT_PORT}");
DEFAULT_PORT
}),
Err(_) => DEFAULT_PORT,
}
}
/// 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 {
#[cfg(feature = "python-config")]
if let Ok(config_path) = std::env::var("LITELLM_CONFIG_PATH") {
match load_model_list(std::path::Path::new(&config_path)) {
Ok(deployments) => {
eprintln!("loaded model_list from {config_path} via python config reader");
return Router::new(deployments);
}
Err(err) => {
eprintln!("config load failed ({err}); falling back to env deployment");
}
}
}
build_router_from_env()
}
/// Build a minimal single-deployment `model_list` from the environment.
///
/// A real deployment loads `model_list` from config; this is the minimal stand-in
/// so the gateway has one OpenAI deployment to route to.
fn build_router_from_env() -> Router {
let model =
std::env::var("OPENAI_REALTIME_MODEL").unwrap_or_else(|_| "gpt-realtime".to_string());
let api_key = std::env::var("OPENAI_API_KEY").ok();
if api_key.is_none() {
eprintln!(
"warning: OPENAI_API_KEY is not set; realtime requests will fail with auth errors"
);
}
let deployment = Deployment {
model_name: model.clone(),
litellm_params: LiteLLMParams {
model,
api_key,
api_base: None,
},
};
Router::new(vec![deployment])
}

View file

@ -1,127 +0,0 @@
use litellm_core::Error;
use litellm_core::ocr::{
OcrClient,
wire::{OcrWireRequest, decode_request},
};
use serde_json::Value;
mod types;
pub use types::OcrRequest;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn ocr(request: OcrRequest<'_>) -> Result<Value, Error> {
core_ocr(request).await
}
async fn core_ocr(request: OcrRequest<'_>) -> Result<Value, Error> {
validate_host_hooks(&request)?;
let client = OcrClient::new(crate::client::http_client().clone())?;
let core_request = decode_request(OcrWireRequest {
model: request.model.to_string(),
document: request.document,
api_key: request.api_key.map(str::to_string),
api_base: request.api_base.map(str::to_string),
custom_llm_provider: request.custom_llm_provider.map(str::to_string),
extra_headers: request.extra_headers,
optional_params: request.optional_params,
input_sources: Default::default(),
timeout_seconds: request.timeout.map(|timeout| timeout.as_secs_f64()),
})?;
client
.perform(core_request)
.await
.map(|response| response.into_json())
}
fn validate_host_hooks(request: &OcrRequest<'_>) -> Result<(), Error> {
if !request.guardrails.is_empty() {
return Err(Error::Unsupported(
"OCR host guardrails are not wired to the core path",
));
}
if !request.callbacks.is_empty() {
return Err(Error::Unsupported(
"OCR host callbacks are not wired to the core path",
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use litellm_core::ocr::wire::is_supported_request;
use serde_json::{Map, json};
use super::{OcrRequest, validate_host_hooks};
use crate::integrations::custom_guardrail::{CustomGuardrail, GuardrailEventHook};
use crate::integrations::custom_logger::CustomLogger;
struct TestGuardrail;
impl CustomGuardrail for TestGuardrail {
fn guardrail_name(&self) -> &str {
"test"
}
fn supported_event_hooks(&self) -> &[GuardrailEventHook] {
&[]
}
}
struct TestLogger;
impl CustomLogger for TestLogger {}
fn request() -> OcrRequest<'static> {
OcrRequest {
model: "model",
document: json!({"type":"image_url","image_url":"data:image/png;base64,YQ=="}),
api_key: None,
api_base: None,
custom_llm_provider: Some("mistral"),
extra_headers: None,
optional_params: Map::new(),
timeout: None,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
}
}
#[test]
fn core_activation_includes_migrated_providers() {
assert!(is_supported_request("model", Some("mistral")));
assert!(is_supported_request("pixtral-12b", Some("azure_ai")));
assert!(is_supported_request(
"doc-intelligence/prebuilt-layout",
Some("azure_ai")
));
assert!(is_supported_request("parse-v3", Some("reducto")));
assert!(is_supported_request("mistral-ocr", Some("vertex_ai")));
assert!(is_supported_request("deepseek-ocr", Some("vertex_ai")));
}
#[test]
fn core_path_rejects_unwired_guardrails() {
let request = OcrRequest {
guardrails: vec![Arc::new(TestGuardrail)],
..request()
};
let error = validate_host_hooks(&request).unwrap_err();
assert!(error.to_string().contains("guardrails are not wired"));
}
#[test]
fn core_path_rejects_unwired_callbacks() {
let request = OcrRequest {
callbacks: vec![Arc::new(TestLogger)],
..request()
};
let error = validate_host_hooks(&request).unwrap_err();
assert!(error.to_string().contains("callbacks are not wired"));
}
}

View file

@ -1,23 +0,0 @@
use std::sync::Arc;
use std::time::Duration;
use serde_json::{Map, Value};
use crate::integrations::custom_guardrail::CustomGuardrail;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
pub struct OcrRequest<'a> {
pub model: &'a str,
pub document: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub callbacks: Vec<Arc<dyn CustomLogger>>,
pub guardrails: Vec<Arc<dyn CustomGuardrail>>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
}

View file

@ -1,4 +0,0 @@
//! Realtime logging collector. Observes the realtime event stream and emits a
//! `StandardLoggingPayload` to the registered callbacks on session close.
pub mod streaming;

View file

@ -1,414 +0,0 @@
//! `RealTimeStreaming` — the realtime logging collector.
//!
//! Mirrors Python `litellm.realtime_api.main.RealTimeStreaming`: it observes the
//! event stream in O(1) (never buffering frames), accumulating just the fields
//! the spend log needs (model, id, cumulative usage), then on session close
//! builds a `StandardLoggingPayload` and fans it out to every registered
//! `CustomLogger`.
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::realtime::types::RealtimeEvent;
use serde_json::Value;
use crate::constants::DEFAULT_PROVIDER;
use crate::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, Usage,
};
/// Current wall-clock time as epoch seconds (float), matching the Python
/// `startTime`/`endTime` contract.
fn epoch_seconds() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs_f64())
.unwrap_or(0.0)
}
/// Status of a finished realtime session, mapped to the callback record status.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SessionStatus {
Success,
Failure,
}
/// Accumulates realtime session state and emits a logging payload on close.
pub struct RealTimeStreaming {
callbacks: Vec<Arc<dyn CustomLogger>>,
/// REQUEST-ID RULE: the SpendLogs `request_id` == the OpenAI realtime session
/// id (`sess_…`), captured from `session.created`. Both `id` and
/// `litellm_call_id` are set to that value so the Python writer logs the same
/// id regardless of which field it reads. The gateway-generated `rt-…` id
/// (the constructor seed) is only a fallback for sessions that fail before
/// `session.created` arrives.
litellm_call_id: String,
/// See the request-id rule above — mirrors `litellm_call_id`.
id: String,
model: String,
custom_llm_provider: String,
usage: Usage,
response_cost: f64,
start_time: f64,
end_time: f64,
metadata: RequestMetadata,
/// Count of logging callbacks that failed to enqueue (non-fatal).
dropped: u64,
}
impl RealTimeStreaming {
/// Create a collector for one session. `litellm_call_id` is the gateway's
/// per-connection id; `model` is the requested model (a sane default until
/// `session.created` reports the upstream model).
pub fn new(
callbacks: Vec<Arc<dyn CustomLogger>>,
litellm_call_id: String,
model: String,
metadata: RequestMetadata,
) -> Self {
let now = epoch_seconds();
Self {
callbacks,
id: litellm_call_id.clone(),
litellm_call_id,
model,
custom_llm_provider: DEFAULT_PROVIDER.to_string(),
usage: Usage::default(),
response_cost: 0.0,
start_time: now,
end_time: now,
metadata,
dropped: 0,
}
}
/// Number of logging callbacks that failed to enqueue so far (test/observ.).
#[allow(dead_code)]
pub fn dropped(&self) -> u64 {
self.dropped
}
/// Observe one realtime event. O(1): updates accumulated state only; never
/// buffers frames. Safe to call on every event in either direction.
pub fn observe(&mut self, event: &RealtimeEvent) {
match event.event_type.as_str() {
"session.created" | "session.updated" => self.on_session(event),
"response.done" => self.on_response_done(event),
_ => {}
}
}
/// `session.created` / `session.updated` → capture upstream id + model.
/// Per the request-id rule, the OpenAI session id becomes BOTH `id` and
/// `litellm_call_id`, replacing the gateway-generated fallback.
fn on_session(&mut self, event: &RealtimeEvent) {
let session = event.data.get("session").and_then(Value::as_object);
if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str)
&& !id.is_empty()
{
self.id = id.to_string();
self.litellm_call_id = id.to_string();
}
if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str)
&& !model.is_empty()
{
self.model = model.to_string();
}
}
/// `response.done` → add this response's usage to the cumulative totals.
fn on_response_done(&mut self, event: &RealtimeEvent) {
let usage = event
.data
.get("response")
.and_then(Value::as_object)
.and_then(|r| r.get("usage"))
.and_then(Value::as_object);
let Some(usage) = usage else { return };
let input = usage.get("input_tokens").and_then(Value::as_u64);
let output = usage.get("output_tokens").and_then(Value::as_u64);
let total = usage.get("total_tokens").and_then(Value::as_u64);
if let Some(input) = input {
self.usage.prompt_tokens += input;
}
if let Some(output) = output {
self.usage.completion_tokens += output;
}
// Prefer the upstream-reported total; otherwise derive it.
match total {
Some(total) => self.usage.total_tokens += total,
None => {
self.usage.total_tokens += input.unwrap_or(0) + output.unwrap_or(0);
}
}
}
/// Set the per-session response cost ($). Cost computation is Python-side in
/// the proxy; the gateway forwards 0.0 by default and lets the proxy price.
/// Public API (exercised in tests) for the future path where the gateway
/// prices realtime sessions itself.
#[allow(dead_code)]
pub fn set_response_cost(&mut self, cost: f64) {
self.response_cost = cost;
}
/// Build the `StandardLoggingPayload` from accumulated state.
pub fn build_payload(&self) -> StandardLoggingPayload {
StandardLoggingPayload {
id: self.id.clone(),
litellm_call_id: self.litellm_call_id.clone(),
call_type: "realtime".to_string(),
model: self.model.clone(),
custom_llm_provider: self.custom_llm_provider.clone(),
response_cost: self.response_cost,
prompt_tokens: self.usage.prompt_tokens,
completion_tokens: self.usage.completion_tokens,
total_tokens: self.usage.total_tokens,
start_time: self.start_time,
end_time: self.end_time,
stream: true,
metadata: StandardLoggingMetadata {
user_api_key_hash: self.metadata.user_api_key_hash.clone(),
user_api_key_user_id: self.metadata.user_api_key_user_id.clone(),
user_api_key_team_id: self.metadata.user_api_key_team_id.clone(),
..Default::default()
},
messages: None,
}
}
/// Finish the session: stamp the end time and fan the payload out to every
/// callback. On a logger enqueue error we bump a non-fatal counter (the
/// realtime session has already ended; a dropped log must never propagate).
pub async fn log_messages(&mut self, status: SessionStatus) {
self.end_time = epoch_seconds();
let payload = self.build_payload();
let timing = CallbackTiming::new(payload.start_time, payload.end_time);
let runner = CustomLoggerRunner::new(self.callbacks.clone());
match status {
SessionStatus::Success => {
let response = CallbackValue::new("realtime", serde_json::Value::Null);
let report = runner
.async_log_success_event(
&ModelCallDetails::from_standard_logging_payload(payload),
&response,
timing,
)
.await;
self.dropped += report.dropped as u64;
}
SessionStatus::Failure => {
let error = LoggingError {
message: "realtime session ended in failure".to_string(),
kind: "RealtimeSessionError".to_string(),
};
let response = CallbackValue::new(
"error",
serde_json::json!({
"message": error.message,
"kind": error.kind,
}),
);
let report = runner
.async_log_failure_event(
&ModelCallDetails::from_standard_logging_payload(payload)
.with_failure_error(error),
Some(&response),
timing,
)
.await;
self.dropped += report.dropped as u64;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::integrations::custom_logger::LogError;
use crate::integrations::custom_logger::LogFuture;
use std::sync::atomic::{AtomicU64, Ordering};
fn event(raw: &str) -> RealtimeEvent {
serde_json::from_str(raw).expect("valid event json")
}
/// A test logger that records the last payload it saw.
#[derive(Default)]
struct CapturingLogger {
calls: AtomicU64,
last_model: std::sync::Mutex<Option<String>>,
last_total_tokens: AtomicU64,
}
impl CustomLogger for CapturingLogger {
fn async_log_success_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
_response_obj: &'a CallbackValue,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
let payload = model_call_details
.standard_logging_payload
.as_ref()
.expect("standard logging payload");
self.calls.fetch_add(1, Ordering::SeqCst);
*self.last_model.lock().unwrap() = Some(payload.model.clone());
self.last_total_tokens
.store(payload.total_tokens, Ordering::SeqCst);
Ok(())
})
}
}
#[tokio::test]
async fn observe_accumulates_model_and_tokens_then_logs() {
let logger = Arc::new(CapturingLogger::default());
let callbacks: Vec<Arc<dyn CustomLogger>> = vec![logger.clone()];
let mut streaming = RealTimeStreaming::new(
callbacks,
"call_abc".to_string(),
"gpt-realtime".to_string(),
RequestMetadata {
user_api_key_hash: Some("hash123".to_string()),
user_api_key_user_id: Some("user-1".to_string()),
user_api_key_team_id: Some("team-1".to_string()),
},
);
streaming.observe(&event(
r#"{"type":"session.created","session":{"id":"sess_001","model":"gpt-realtime-2025"}}"#,
));
streaming.observe(&event(
r#"{"type":"response.done","response":{"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15}}}"#,
));
// A second response.done accumulates.
streaming.observe(&event(
r#"{"type":"response.done","response":{"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}"#,
));
let payload = streaming.build_payload();
assert_eq!(payload.model, "gpt-realtime-2025");
// Request-id rule: session.created's id becomes BOTH id and
// litellm_call_id (replacing the "call_abc" gateway fallback), so the
// SpendLogs request_id is always the OpenAI session id.
assert_eq!(payload.id, "sess_001");
assert_eq!(payload.litellm_call_id, "sess_001");
assert_eq!(payload.prompt_tokens, 13);
assert_eq!(payload.completion_tokens, 7);
assert_eq!(payload.total_tokens, 20);
assert_eq!(payload.response_cost, 0.0);
assert_eq!(payload.call_type, "realtime");
assert_eq!(payload.custom_llm_provider, "openai");
assert_eq!(
payload.metadata.user_api_key_hash.as_deref(),
Some("hash123")
);
streaming.log_messages(SessionStatus::Success).await;
assert_eq!(logger.calls.load(Ordering::SeqCst), 1);
assert_eq!(
logger.last_model.lock().unwrap().as_deref(),
Some("gpt-realtime-2025")
);
assert_eq!(logger.last_total_tokens.load(Ordering::SeqCst), 20);
assert_eq!(streaming.dropped(), 0);
}
#[test]
fn blank_session_id_and_model_keep_the_gateway_fallbacks() {
let mut streaming = RealTimeStreaming::new(
Vec::new(),
"call_fallback".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
);
streaming.observe(&event(
r#"{"type":"session.created","session":{"id":"","model":""}}"#,
));
let payload = streaming.build_payload();
assert_eq!(payload.id, "call_fallback");
assert_eq!(payload.litellm_call_id, "call_fallback");
assert_eq!(payload.model, "gpt-realtime");
streaming.observe(&event(
r#"{"type":"session.updated","session":{"id":"sess_002","model":""}}"#,
));
let payload = streaming.build_payload();
assert_eq!(payload.id, "sess_002");
assert_eq!(payload.litellm_call_id, "sess_002");
assert_eq!(payload.model, "gpt-realtime");
}
#[test]
fn payload_serializes_with_camelcase_times_and_realtime_call_type() {
let mut streaming = RealTimeStreaming::new(
Vec::new(),
"call_xyz".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
);
streaming.observe(&event(
r#"{"type":"response.done","response":{"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}"#,
));
streaming.set_response_cost(0.0042);
let payload = streaming.build_payload();
let json = serde_json::to_string(&payload).expect("serialize payload");
assert!(json.contains("\"startTime\""), "missing startTime: {json}");
assert!(json.contains("\"endTime\""), "missing endTime: {json}");
assert!(
json.contains("\"call_type\":\"realtime\""),
"missing call_type realtime: {json}"
);
assert!(
json.contains("\"response_cost\""),
"missing response_cost: {json}"
);
assert_eq!(payload.response_cost, 0.0042);
}
/// A logger whose enqueue always fails should bump the dropped counter, not
/// panic or propagate.
#[tokio::test]
async fn failing_logger_bumps_dropped_counter() {
struct FailingLogger;
impl CustomLogger for FailingLogger {
fn async_log_success_event<'a>(
&'a self,
_model_call_details: &'a ModelCallDetails,
_response_obj: &'a CallbackValue,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async { Err(LogError::channel_full()) })
}
fn async_log_failure_event<'a>(
&'a self,
_model_call_details: &'a ModelCallDetails,
_response_obj: Option<&'a CallbackValue>,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async { Err(LogError::channel_closed()) })
}
}
let callbacks: Vec<Arc<dyn CustomLogger>> = vec![Arc::new(FailingLogger)];
let mut streaming = RealTimeStreaming::new(
callbacks,
"call_1".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
);
streaming.log_messages(SessionStatus::Success).await;
assert_eq!(streaming.dropped(), 1);
}
}

View file

@ -1,43 +0,0 @@
# routes/ — the route template
Every route follows the **same shape** so the layout is predictable. The rule:
> **Each route module exposes `pub fn router() -> Router<AppState>`.**
> `routes/mod.rs::app` merges them all and applies state once. Adding a route is:
> create the module, then add one `.merge(<name>::router())` line.
## Default: one file
A route is a single file containing `router()` + its handler(s) (handlers stay
private). This is the norm — don't split until it hurts.
```
pub fn router() -> Router<AppState> { Router::new().route(PATH, get(handle)) }
async fn handle(...) -> impl IntoResponse { ... }
```
`health.rs` is the example.
## Split out `service` when there's real logic
When a route has business logic worth testing without axum, put it in a sibling
`service` (a file, or a folder if the route grows). The route file stays the
**axum surface** (router + handler + any socket/SSE adapter); `service` is plain
Rust with **no axum types**, and its job is to pick the deployment and call the
`core` route entrypoint (see `messages/service.rs` calling
`litellm_core::messages::messages`). Never build a provider request, resolve a
key, or perform the provider call here. `realtime/` is the older example:
```
realtime/
mod.rs # axum surface: router() + handler + the WS<->events adapter
service.rs # pure logic: select deployment + call provider (no axum) — testable
```
Split `service` further (or add `transport`, `repo`, …) only once a single file
genuinely gets hard to read.
## Invariants
- **Auth is an extractor, not a manual call.** A handler requires auth by adding
`crate::auth::RequireMasterKey` to its arguments; it runs during extraction.
Never re-implement the check per route.
- **Handlers contain no business logic; `service` contains no axum types.**
- **No provider handlers in this crate.** Transforms, auth headers, and the
provider HTTP call live in `core/src/<route>/`.
- A route owns its paths in its own `router()`; `mod.rs` only merges.
- Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`,
not duplicated in handlers.

View file

@ -1,24 +0,0 @@
//! Health probes. Simple-route template: a `router()` plus its handlers, in one file.
use axum::Router;
use axum::http::StatusCode;
use axum::routing::get;
use crate::state::AppState;
/// This route's contribution to the app router.
pub fn router() -> Router<AppState> {
Router::new()
.route("/health/liveness", get(liveness))
.route("/health/readiness", get(readiness))
}
/// The process is up.
async fn liveness() -> StatusCode {
StatusCode::OK
}
/// The server is ready to accept traffic.
async fn readiness() -> StatusCode {
StatusCode::OK
}

View file

@ -1,532 +0,0 @@
//! `POST /v1/messages`, the Anthropic Messages HTTP surface.
mod service;
use axum::Router;
use axum::body::Body;
use axum::extract::{Json, State};
use axum::http::StatusCode;
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue};
use axum::response::{IntoResponse, Response};
use axum::routing::post;
use litellm_core::Error;
use serde_json::{Map, Value};
use crate::auth::RequireMasterKey;
use crate::constants::{MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH};
use crate::state::AppState;
/// This route's contribution to the app router.
pub fn router() -> Router<AppState> {
Router::new().route(MESSAGES_ROUTE_PATH, post(handle))
}
#[tracing::instrument(
name = "messages_gateway_route",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
async fn handle(
_auth: RequireMasterKey,
State(state): State<AppState>,
headers: HeaderMap,
Json(body): Json<Value>,
) -> Result<Response, MessagesRouteError> {
let extra_headers = forwarded_headers(&headers)?;
match service::run(&state.router, body, extra_headers)
.await
.map_err(MessagesRouteError::from)?
{
service::MessagesResponse::Json(body) => Ok(Json(body).into_response()),
service::MessagesResponse::Stream(upstream) => stream_response(upstream),
}
}
fn stream_response(upstream: reqwest::Response) -> Result<Response, MessagesRouteError> {
let content_type = upstream
.headers()
.get(CONTENT_TYPE)
.cloned()
.unwrap_or_else(|| HeaderValue::from_static("text/event-stream"));
let mut response = Response::builder()
.status(
StatusCode::from_u16(upstream.status().as_u16()).map_err(|error| {
MessagesRouteError(Error::InvalidResponse(format!(
"invalid upstream response status: {error}"
)))
})?,
)
.header(CONTENT_TYPE, content_type);
if let Some(value) = upstream.headers().get(CACHE_CONTROL) {
response = response.header(CACHE_CONTROL, value);
}
response
.body(Body::from_stream(upstream.bytes_stream()))
.map_err(|error| {
MessagesRouteError(Error::InvalidResponse(format!(
"failed to build streaming response: {error}"
)))
})
}
fn forwarded_headers(headers: &HeaderMap) -> Result<Option<Map<String, Value>>, Error> {
let forwarded = headers
.iter()
.filter(|(name, _)| {
!MESSAGES_HEADERS_NOT_FORWARDED
.iter()
.any(|excluded| name.as_str().eq_ignore_ascii_case(excluded))
})
.map(|(name, value)| {
let value = value.to_str().map_err(|_| {
Error::InvalidRequest(format!("invalid value for header {}", name.as_str()))
})?;
Ok((name.to_string(), Value::String(value.to_string())))
})
.collect::<Result<Map<_, _>, Error>>()?;
Ok((!forwarded.is_empty()).then_some(forwarded))
}
#[derive(Debug)]
struct MessagesRouteError(Error);
impl From<Error> for MessagesRouteError {
fn from(error: Error) -> Self {
Self(error)
}
}
impl IntoResponse for MessagesRouteError {
fn into_response(self) -> Response {
let (status, message) = match self.0 {
Error::InvalidRequest(message) => (StatusCode::BAD_REQUEST, message),
Error::InvalidProvider(_) | Error::Routing(_) => (
StatusCode::NOT_FOUND,
"no messages deployment is configured for this model".to_string(),
),
Error::Auth(_)
| Error::MissingApiKey { .. }
| Error::MissingAzureAiCredentials
| Error::MissingAzureDocumentIntelligenceCredentials
| Error::MissingReductoApiKey => (
StatusCode::BAD_GATEWAY,
"messages provider authentication failed".to_string(),
),
Error::Http { .. }
| Error::Network(_)
| Error::Connect(_)
| Error::InvalidResponse(_)
| Error::InvalidType { .. }
| Error::MissingField(_)
| Error::MissingDocumentUrl => (
StatusCode::BAD_GATEWAY,
"messages provider request failed".to_string(),
),
// The gateway has no Python implementation to decline to, so a
// request the core cannot serve is reported to the caller. The
// reason is a fixed internal string, never provider content.
Error::Unsupported(reason) => (
StatusCode::BAD_REQUEST,
format!("messages request is not supported: {reason}"),
),
};
(
status,
Json(serde_json::json!({"error": {"message": message}})),
)
.into_response()
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use axum::body::Body;
use axum::http::Request;
use axum::http::StatusCode;
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE};
use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter};
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use tower::ServiceExt;
use super::super::app;
use crate::io::realtime_pool::RealtimePool;
use crate::state::AppState;
fn state(model: &str, api_base: String, master_key: Option<&str>) -> AppState {
state_with_provider(model, model, api_base, master_key)
}
fn state_with_provider(
model_alias: &str,
provider_model: &str,
api_base: String,
master_key: Option<&str>,
) -> AppState {
AppState {
router: Arc::new(ModelRouter::new(vec![Deployment {
model_name: model_alias.to_string(),
litellm_params: LiteLLMParams {
model: format!("anthropic/{provider_model}"),
api_key: Some("upstream-key".to_string()),
api_base: Some(api_base),
},
}])),
master_key: master_key.map(Arc::from),
loggers: Arc::new(Vec::new()),
realtime_pool: RealtimePool::disabled(),
}
}
async fn upstream(listener: TcpListener) -> (String, tokio::task::JoinHandle<String>) {
let address = listener.local_addr().expect("listener has address");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let mut request = Vec::new();
let mut buffer = [0_u8; 4096];
loop {
let read = socket.read(&mut buffer).await.expect("reads request");
request.extend_from_slice(&buffer[..read]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
let request = String::from_utf8(request).expect("request is utf8");
let content_length = request
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
let header_end = request.find("\r\n\r\n").expect("request has headers") + 4;
let mut full_request = request.into_bytes();
while full_request.len().saturating_sub(header_end) < content_length {
let read = socket.read(&mut buffer).await.expect("reads body");
full_request.extend_from_slice(&buffer[..read]);
}
let request = String::from_utf8(full_request).expect("request is utf8");
let body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-test"}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
(format!("http://{address}"), server)
}
async fn streaming_upstream(
listener: TcpListener,
status: u16,
content_type: &'static str,
body: &'static str,
) -> (String, tokio::task::JoinHandle<String>) {
let address = listener.local_addr().expect("listener has address");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let mut request = Vec::new();
let mut buffer = [0_u8; 4096];
loop {
let read = socket.read(&mut buffer).await.expect("reads request");
request.extend_from_slice(&buffer[..read]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
let request_text = String::from_utf8(request).expect("request is utf8");
let content_length = request_text
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
let header_end = request_text.find("\r\n\r\n").expect("request has headers") + 4;
let mut full_request = request_text.into_bytes();
while full_request.len().saturating_sub(header_end) < content_length {
let read = socket.read(&mut buffer).await.expect("reads body");
full_request.extend_from_slice(&buffer[..read]);
}
let response = format!(
"HTTP/1.1 {status} OK\r\ncontent-type: {content_type}\r\ncache-control: no-cache\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len()
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
String::from_utf8(full_request).expect("request is utf8")
});
(format!("http://{address}"), server)
}
#[tokio::test]
async fn route_constructs_anthropic_upstream_request() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let (api_base, server) = upstream(listener).await;
let app = app(state("claude-test", api_base, Some("master-key")));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("x-api-key", "request-upstream-key")
.header("anthropic-beta", "beta-feature")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "claude-test",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hello"}]
})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body reads");
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&body).expect("json")["id"],
"msg_1"
);
let upstream_request = server.await.expect("upstream task completes");
let (head, body) = upstream_request
.split_once("\r\n\r\n")
.expect("upstream request has body");
let head = head.to_ascii_lowercase();
assert!(head.contains("x-api-key: request-upstream-key"));
assert!(head.contains("anthropic-beta: beta-feature"));
assert!(!head.contains("authorization: bearer master-key"));
let body: serde_json::Value = serde_json::from_str(body).expect("upstream body is json");
assert_eq!(body["model"], "claude-test");
assert_eq!(body["messages"][0]["content"], "hello");
}
#[tokio::test]
async fn route_substitutes_model_alias_with_provider_model_upstream() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let (api_base, server) = upstream(listener).await;
let app = app(state_with_provider(
"production",
"claude-sonnet-4-5",
api_base,
Some("master-key"),
));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "production",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hello"}]
})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::OK);
let upstream_request = server.await.expect("upstream task completes");
let (_, upstream_body) = upstream_request
.split_once("\r\n\r\n")
.expect("upstream request has body");
let upstream_body: serde_json::Value =
serde_json::from_str(upstream_body).expect("upstream body is json");
assert_eq!(upstream_body["model"], "claude-sonnet-4-5");
assert_ne!(upstream_body["model"], "production");
}
#[tokio::test]
async fn route_streams_anthropic_events_without_buffering_or_reordering() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let events = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: content_block_delta\ndata: {\"type\":\"content_block_delta\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
let (api_base, server) =
streaming_upstream(listener, 200, "text/event-stream", events).await;
let app = app(state("claude-test", api_base, Some("master-key")));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "claude-test",
"max_tokens": 16,
"stream": true,
"messages": [{"role": "user", "content": "hello"}]
})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get(CONTENT_TYPE)
.unwrap()
.to_str()
.unwrap(),
"text/event-stream"
);
assert_eq!(
response
.headers()
.get(CACHE_CONTROL)
.unwrap()
.to_str()
.unwrap(),
"no-cache"
);
let response_body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body reads");
assert_eq!(response_body, events.as_bytes());
let upstream_request = server.await.expect("upstream task completes");
let (_, upstream_body) = upstream_request
.split_once("\r\n\r\n")
.expect("upstream request has body");
assert_eq!(
serde_json::from_str::<serde_json::Value>(upstream_body)
.expect("upstream body is json")["stream"],
true
);
}
#[tokio::test]
async fn route_maps_streaming_upstream_errors_before_starting_response() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let (api_base, server) = streaming_upstream(
listener,
429,
"application/json",
r#"{"error":"rate limited"}"#,
)
.await;
let app = app(state("claude-test", api_base, Some("master-key")));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "claude-test",
"max_tokens": 16,
"stream": true,
"messages": [{"role": "user", "content": "hello"}]
})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
let response_body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body reads");
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&response_body).expect("error is json")["error"]
["message"],
"messages provider request failed"
);
server.await.expect("upstream task completes");
}
#[tokio::test]
async fn route_rejects_missing_master_key() {
let app = app(state(
"claude-test",
"http://127.0.0.1:1".to_string(),
Some("master-key"),
));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("content-type", "application/json")
.body(Body::from("{}"))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn route_rejects_invalid_master_key() {
let app = app(state(
"claude-test",
"http://127.0.0.1:1".to_string(),
Some("master-key"),
));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer wrong-key")
.header("content-type", "application/json")
.body(Body::from("{}"))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn route_rejects_malformed_json_without_panicking() {
let app = app(state(
"claude-test",
"http://127.0.0.1:1".to_string(),
Some("master-key"),
));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from("{not-json"))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
}
}

View file

@ -1,71 +0,0 @@
use std::sync::Arc;
use litellm_core::Error;
use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER;
use litellm_core::messages::types::MessagesRequest;
use litellm_core::messages::{messages, messages_stream};
use litellm_core::router::Router;
use serde_json::{Map, Value};
pub(crate) enum MessagesResponse {
Json(Value),
Stream(reqwest::Response),
}
#[tracing::instrument(
name = "messages_gateway_service",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub async fn run(
router: &Arc<Router>,
body: Value,
extra_headers: Option<Map<String, Value>>,
) -> Result<MessagesResponse, Error> {
let model = body
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|model| !model.is_empty())
.ok_or_else(|| Error::InvalidRequest("messages body requires a model".to_string()))?;
let deployment = router
.get_available_deployment(model)
.ok_or_else(|| Error::Routing(format!("no deployment available for model '{model}'")))?;
let provider_model = deployment.litellm_params.model.as_str();
let upstream_model = provider_model
.split_once('/')
.map_or(provider_model, |(_, model)| model);
let custom_llm_provider = if provider_model.contains('/') {
None
} else {
Some(ANTHROPIC_MESSAGES_PROVIDER)
};
let mut body = body;
body.as_object_mut()
.ok_or_else(|| Error::InvalidRequest("messages body must be an object".to_string()))?
.insert(
"model".to_string(),
Value::String(upstream_model.to_string()),
);
let request = MessagesRequest {
model: provider_model,
body,
api_key: deployment.litellm_params.api_key.as_deref(),
api_base: deployment.litellm_params.api_base.as_deref(),
custom_llm_provider,
extra_headers,
timeout: None,
};
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
return messages_stream(request).await.map(MessagesResponse::Stream);
}
let response = messages(request).await?;
serde_json::to_value(response)
.map(MessagesResponse::Json)
.map_err(|err| {
Error::InvalidResponse(format!("failed to serialize messages response: {err}"))
})
}

View file

@ -1,25 +0,0 @@
//! HTTP routes.
//!
//! **Template:** every route module exposes `pub fn router() -> Router<AppState>`
//! that mounts its own paths; [`app`] merges them. A trivial route is a single
//! file (`health.rs`); a non-trivial one is a folder (`realtime/`) with
//! `handler` (entry) + `service` (logic) + `transport` (adapters). See AGENTS.md.
pub mod health;
pub mod messages;
pub mod realtime;
pub mod responses;
use axum::Router;
use crate::state::AppState;
/// Assemble the application router by merging every route module's `router()`.
pub fn app(state: AppState) -> Router {
Router::new()
.merge(health::router())
.merge(messages::router())
.merge(realtime::router())
.merge(responses::router())
.with_state(state)
}

View file

@ -1,87 +0,0 @@
# Realtime route (`GET /v1/realtime`)
Proxies OpenAI's realtime WebSocket. `mod.rs` is the axum surface (handler +
socket↔events adapter); `service.rs` is the pure logic (select a deployment, then
splice client ↔ upstream). The pool itself lives in
`crates/providers/src/realtime_pool.rs`.
## Connection pooling
### The problem
The gateway's realtime overhead lives **entirely in session establishment**. On each
client connect it dials a *fresh* upstream WS to OpenAI and waits for
`session.created` before it can serve. Measured at 5000 calls / 500 concurrency, the
fresh-dial session phase is **~360 ms** vs **~7 ms** direct; dial, first-audio, and
streaming add ~0. So the one lever is removing that per-connect handshake from the
critical path.
### The idea
Keep a few upstream OpenAI sockets **already connected and already past
`session.created`** (buffered). On a client connect, hand off a warm socket — relay
its buffered `session.created` instantly (a local `Vec::pop`, sub-millisecond) and
splice exactly as a fresh dial would. A background task keeps the pool topped up. On
a miss or dead socket we fall back to fresh-dial: the pool is a latency optimization,
never a correctness dependency.
```
┌───────────────────────────────────────┐
client connect ──────► │ routes/realtime → service::run │
│ pool.take(key) │
│ hit → relay buffered │
│ session.created, then splice │
│ miss → fresh dial (original path) │
└───────────────┬───────────────────────┘
│ replenish (async, concurrent)
┌───────────────▼───────────────────────┐
background task ─────► │ RealtimePool: per-key warm sockets │
│ each = { ws, buffered session.created}│
│ liveness-checked before handoff │
└─────────────────────────────────────────┘
```
A warm session is indistinguishable from a fresh one: OpenAI sends `session.created`
unprompted on connect, we pre-read exactly that one frame and relay it on handoff,
and we send nothing else on the socket before a client exists — so the client's first
`session.update` behaves identically either way.
### Sizing
Each warm socket serves **exactly one** session (realtime isn't multiplexed), so the
pool is sized to the **peak concurrent connects per instance**, not total live
connections:
```
REALTIME_POOL_SIZE ≈ peak_concurrency / instance_count
```
e.g. 500 concurrency over 10 instances → ~5064 per instance. The replenisher dials
the missing sockets **concurrently**, so a drained pool refills in ~one handshake
window and keeps supply close to the connect rate. Over-provisioning just burns idle
upstream sockets, which is why warm sockets are short-lived
(`REALTIME_POOL_MAX_IDLE_SECS`).
### Config
| env | default | meaning |
| ----------------------------- | ------- | --------------------------------------------------------------- |
| `REALTIME_POOL_SIZE` | `4` | target warm sockets per key. `0` disables pooling (fresh-dial). |
| `REALTIME_POOL_MAX_IDLE_SECS` | `30` | max time a warm socket sits before it's closed and replaced. |
### Notes
- **Miss / dead socket → fresh dial.** Burst beyond warm supply, or a socket that
died, never blocks or fails — it falls back to the original path. The pool can only
make a connect faster, never slower or more fragile.
- **Auth scope.** The pool key includes `api_key`, so a warm socket is only handed to
a request resolving to the same key — no cross-tenant reuse.
- **Idle billing.** Warm sockets are liveness-checked at handoff and capped at
`REALTIME_POOL_MAX_IDLE_SECS` to bound idle billing and dodge OpenAI's idle timeout.
- **Replenish backoff.** If a key's warm-up dials all fail (invalid credentials, an
unreachable upstream), the replenisher puts that key into exponential backoff
(500 ms → 30 s cap) instead of re-dialing it every tick. This bounds connection
attempts against a broken key so it can't exhaust upstream rate limits and degrade
valid cold-path traffic; the backoff resets the moment a dial succeeds.
Benchmarks and repro: `../../benchmarks/realtime/README.md`.

View file

@ -1,166 +0,0 @@
//! `GET /v1/realtime` (WebSocket).
//!
//! This file is the **axum surface**: `router()`, the handler, and the small
//! socket↔events adapter. The pure logic (no axum) lives in [`service`]. Auth is
//! the `RequireMasterKey` extractor, so the handler stays thin.
mod service;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use crate::io::realtime_pool::RealtimePool;
use axum::Router;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Query, State};
use axum::http::StatusCode;
use axum::response::Response;
use axum::routing::get;
use futures_util::{SinkExt, StreamExt};
use litellm_core::realtime::types::RealtimeEvent;
use litellm_core::router::Router as ModelRouter;
use serde::Deserialize;
use crate::auth::RequireMasterKey;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
use crate::realtime::streaming::{RealTimeStreaming, SessionStatus};
use crate::state::AppState;
/// Process-local monotonic counter, mixed into the per-session call id so two
/// sessions opened in the same nanosecond still get distinct ids.
static CALL_SEQ: AtomicU64 = AtomicU64::new(0);
/// Generate a per-connection `litellm_call_id`. No external uuid dep: epoch
/// nanos + a process-local sequence is unique enough for log correlation.
fn new_call_id() -> String {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let seq = CALL_SEQ.fetch_add(1, Ordering::Relaxed);
format!("rt-{nanos:x}-{seq:x}")
}
/// This route's contribution to the app router.
pub fn router() -> Router<AppState> {
Router::new().route("/v1/realtime", get(handle))
}
#[derive(Debug, Deserialize)]
struct RealtimeQuery {
model: String,
}
/// Auth runs via the `RequireMasterKey` extractor. We validate the model BEFORE
/// the upgrade so failures are clean HTTP (400/404), not a socket that opens then
/// closes, then hand the socket to `bridge`.
async fn handle(
_auth: RequireMasterKey,
ws: WebSocketUpgrade,
State(state): State<AppState>,
Query(query): Query<RealtimeQuery>,
) -> Result<Response, (StatusCode, String)> {
if query.model.trim().is_empty() {
return Err((
StatusCode::BAD_REQUEST,
"missing 'model' query param".to_string(),
));
}
if !state.router.has_deployment(&query.model) {
return Err((
StatusCode::NOT_FOUND,
format!("no deployment for model '{}'", query.model),
));
}
let router = state.router.clone();
let pool = state.realtime_pool.clone();
let loggers = state.loggers.clone();
let master_key = state.master_key.clone();
let model = query.model;
Ok(ws.on_upgrade(move |socket| bridge(socket, router, pool, loggers, master_key, model)))
}
/// Adapt the axum socket (text frames) to the typed-event `Stream`/`Sink` the
/// service wants, keeping axum types out of `service`.
///
/// This is also the realtime-logging seam: every upstream→client event (the
/// direction carrying `session.created` and `response.done` with usage) is fed
/// to a [`RealTimeStreaming`] collector via the splice's `observe` callback. The
/// observe is O(1) and never buffers frames. When the splice returns (any of the
/// three break paths — client disconnect, upstream close, idle timeout), we flush
/// one logging payload to the registered callbacks.
async fn bridge(
socket: WebSocket,
router: Arc<ModelRouter>,
pool: Arc<RealtimePool>,
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
master_key: Option<Arc<str>>,
model: String,
) {
let (ws_sink, ws_stream) = socket.split();
// Attribute the spend log to the key that authenticated this session (the
// master key — the gateway is master-key auth). A non-null user_api_key_hash
// is required for the Python spend logger to write a SpendLogs row.
//
// SECURITY: hash the key — never send the raw credential. This field fans out
// to spend logs and every callback integration; the SHA-256 (matching the
// proxy's hash_token) keeps the plaintext master key out of all of them while
// still matching the key's hash in LiteLLM_SpendLogs.
let metadata = RequestMetadata {
user_api_key_hash: master_key.as_deref().map(crate::auth::hash_token),
..RequestMetadata::default()
};
// Owned by THIS task only. The splice observes it via a synchronous `&mut`
// callback (below), so there is no Arc/Mutex/atomic on the per-frame hot
// path — just a monomorphized FnMut mutating stack-local fields. This is
// what lets observe scale: 10K concurrent sessions = 10K independent
// collectors, zero cross-task synchronization.
let mut collector = RealTimeStreaming::new(
loggers.as_ref().clone(),
new_call_id(),
model.clone(),
metadata,
);
let client_in = ws_stream.filter_map(|message| async move {
match message {
Ok(Message::Text(text)) => serde_json::from_str::<RealtimeEvent>(&text).ok(),
_ => None,
}
});
// Plain forwarding sink — no observe here anymore.
let client_out = ws_sink.with(|event: RealtimeEvent| async move {
Ok::<Message, axum::Error>(Message::Text(
serde_json::to_string(&event).unwrap_or_default(),
))
});
futures_util::pin_mut!(client_in, client_out);
// The observe closure borrows `&mut collector` for the duration of the
// splice; the borrow ends when `run` returns, freeing the collector for the
// single post-session `log_messages` flush. `run` picks a pooled (warm) or
// fresh upstream — observe fires on the upstream arm either way.
let result = service::run(
&router,
&pool,
&model,
None,
|event: &RealtimeEvent| collector.observe(event),
client_in,
client_out,
)
.await;
let status = if result.is_ok() {
SessionStatus::Success
} else {
SessionStatus::Failure
};
collector.log_messages(status).await;
}

View file

@ -1,77 +0,0 @@
//! Business logic: select a deployment with the (pure) core router, then call the
//! provider splice. The seam between `core::router` (selection only) and
//! `io` (the actual WebSocket I/O).
//!
//! On connect we try a pre-warmed upstream from the pool (handshake already paid,
//! `session.created` buffered) and relay it instantly. On a pool miss or dead warm
//! socket we fresh-dial exactly as before — the pool is never on the critical path
//! for correctness, only latency.
use std::time::Duration;
use crate::io::realtime_pool::{RealtimePool, upstream_key};
use futures_util::{Sink, Stream};
use litellm_core::error::Error;
use litellm_core::realtime::types::RealtimeEvent;
use litellm_core::router::Router;
/// Select a deployment for `model` and splice the client stream to the provider.
///
/// `pool` supplies a pre-warmed upstream when one is available; otherwise we
/// fresh-dial. A disabled pool always misses, so this collapses to the original
/// fresh-dial behavior.
pub async fn run<In, Out>(
router: &Router,
pool: &RealtimePool,
model: &str,
idle_timeout: Option<Duration>,
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::Error: std::fmt::Display,
{
let deployment = router
.get_available_deployment(model)
.ok_or_else(|| Error::Routing(format!("no deployment available for model '{model}'")))?;
let params = &deployment.litellm_params;
// Strip a leading `openai/` so the OpenAI-only realtime fn gets the bare model.
let provider_model = params
.model
.strip_prefix("openai/")
.unwrap_or(&params.model);
// Warm path: take a pooled upstream (handshake already paid) and relay its
// buffered session.created immediately. On miss/dead socket fall through.
if let Some(key) = upstream_key(
provider_model,
params.api_key.as_deref(),
params.api_base.as_deref(),
) && let Some(handoff) = pool.take(&key)
{
return crate::io::realtime::realtime_warm(
provider_model,
handoff,
idle_timeout,
observe,
client_in,
client_out,
)
.await;
}
// Cold path: fresh dial (the original behavior).
crate::io::realtime::realtime(
provider_model,
params.api_key.as_deref(),
params.api_base.as_deref(),
idle_timeout,
observe,
client_in,
client_out,
)
.await
}

View file

@ -1,348 +0,0 @@
mod service;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use axum::Router;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Query, State};
use axum::http::StatusCode;
use axum::response::Response;
use axum::routing::get;
use futures_util::{Sink, SinkExt, StreamExt};
use litellm_core::responses::types::{ResponsesErrorFrame, ResponsesWsEvent, ResponsesWsEventType};
use litellm_core::router::Router as ModelRouter;
use serde::Deserialize;
use crate::auth::RequireMasterKey;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
use crate::state::AppState;
static CALL_SEQ: AtomicU64 = AtomicU64::new(0);
fn new_call_id() -> String {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_nanos())
.unwrap_or(0);
let sequence = CALL_SEQ.fetch_add(1, Ordering::Relaxed);
format!("respws-{nanos:x}-{sequence:x}")
}
pub fn router() -> Router<AppState> {
Router::new()
.route("/v1/responses", get(handle))
.route("/responses", get(handle))
}
#[derive(Debug, Deserialize)]
struct ResponsesQuery {
model: Option<String>,
}
async fn handle(
_auth: RequireMasterKey,
ws: WebSocketUpgrade,
State(state): State<AppState>,
Query(query): Query<ResponsesQuery>,
) -> Result<Response, (StatusCode, String)> {
if let Some(model) = query.model.as_deref() {
validate_model(&state.router, model)?;
}
let router = state.router.clone();
let loggers = state.loggers.clone();
let master_key = state.master_key.clone();
Ok(ws.on_upgrade(move |socket| bridge(socket, router, loggers, master_key, query.model)))
}
fn validate_model(router: &ModelRouter, model: &str) -> Result<(), (StatusCode, String)> {
if model.trim().is_empty() {
return Err((
StatusCode::BAD_REQUEST,
"missing 'model' query param".to_string(),
));
}
let Some(deployment) = router.get_available_deployment(model) else {
return Err((
StatusCode::NOT_FOUND,
format!("no deployment for model '{model}'"),
));
};
if deployment.litellm_params.model.contains('/')
&& !deployment.litellm_params.model.starts_with("openai/")
{
return Err((
StatusCode::BAD_REQUEST,
"Responses WebSocket route supports OpenAI deployments only".to_string(),
));
}
Ok(())
}
async fn send_error_and_close<S>(sink: &mut S, message: String)
where
S: futures_util::Sink<Message> + Unpin,
S::Error: std::fmt::Display,
{
if let Ok(payload) = serde_json::to_string(&ResponsesErrorFrame::invalid_request(message)) {
let _ = sink.send(Message::Text(payload)).await;
}
let _ = sink
.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code: 1008,
reason: "Pre-call error".into(),
})))
.await;
let _ = sink.close().await;
}
struct ResponseClientSink {
sink: futures_util::stream::SplitSink<WebSocket, Message>,
}
impl Sink<ResponsesWsEvent> for ResponseClientSink {
type Error = axum::Error;
fn poll_ready(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::pin::Pin::new(&mut self.sink).poll_ready(context)
}
fn start_send(
mut self: std::pin::Pin<&mut Self>,
item: ResponsesWsEvent,
) -> Result<(), Self::Error> {
let payload = serde_json::to_string(&item).map_err(axum::Error::new)?;
std::pin::Pin::new(&mut self.sink).start_send(Message::Text(payload))
}
fn poll_flush(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::pin::Pin::new(&mut self.sink).poll_flush(context)
}
fn poll_close(
mut self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
std::pin::Pin::new(&mut self.sink).poll_close(context)
}
}
impl ResponseClientSink {
async fn close_with_code(&mut self, code: u16, reason: &'static str) {
let _ = self
.sink
.send(Message::Close(Some(axum::extract::ws::CloseFrame {
code,
reason: reason.into(),
})))
.await;
let _ = self.sink.close().await;
}
}
async fn bridge(
socket: WebSocket,
router: Arc<ModelRouter>,
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
master_key: Option<Arc<str>>,
requested_model: Option<String>,
) {
let (mut ws_sink, ws_stream) = socket.split();
let (model, first_frame, stream) = if let Some(model) = requested_model {
(model, None, ws_stream)
} else {
let mut stream = ws_stream;
let first = match stream.next().await {
Some(Ok(Message::Text(text))) => {
match serde_json::from_str::<ResponsesWsEvent>(&text) {
Ok(event) => event,
Err(_) => {
send_error_and_close(
&mut ws_sink,
"Invalid JSON in response.create event".to_string(),
)
.await;
return;
}
}
}
_ => {
send_error_and_close(&mut ws_sink, "Missing response.create event".to_string())
.await;
return;
}
};
let Some(model) = first.model().filter(|value| !value.trim().is_empty()) else {
send_error_and_close(
&mut ws_sink,
"Missing model in response.create event".to_string(),
)
.await;
return;
};
if first.event_type != ResponsesWsEventType::ResponseCreate {
send_error_and_close(
&mut ws_sink,
"First frame must be a response.create event".to_string(),
)
.await;
return;
}
(model.to_string(), Some(first), stream)
};
if let Err((status, message)) = validate_model(&router, &model) {
let _ = status;
let _ = message;
send_error_and_close(&mut ws_sink, "Unknown model deployment".to_string()).await;
return;
}
let call_id = new_call_id();
let metadata = RequestMetadata {
user_api_key_hash: master_key.as_deref().map(crate::auth::hash_token),
..RequestMetadata::default()
};
let client_in = Box::pin(stream.filter_map(|message| async move {
match message {
Ok(Message::Text(text)) => serde_json::from_str::<ResponsesWsEvent>(&text).ok(),
_ => None,
}
}));
let mut client_out = ResponseClientSink { sink: ws_sink };
let result = service::run(
&router,
&model,
first_frame,
None,
loggers,
call_id,
metadata,
client_in,
&mut client_out,
)
.await;
if result.is_err() {
client_out
.close_with_code(1011, "Internal server error")
.await;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::io::realtime_pool::RealtimePool;
use crate::state::AppState;
use axum::body::Body;
use axum::http::Request;
use litellm_core::router::Router as ModelRouter;
use serde_json::json;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tower::ServiceExt;
struct RecordingSink {
messages: Vec<Message>,
}
impl Sink<Message> for RecordingSink {
type Error = std::convert::Infallible;
fn poll_ready(
self: Pin<&mut Self>,
_context: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
self.messages.push(item);
Ok(())
}
fn poll_flush(
self: Pin<&mut Self>,
_context: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_close(
self: Pin<&mut Self>,
_context: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn pre_call_error_matches_python_frame_and_close() {
let mut sink = RecordingSink {
messages: Vec::new(),
};
send_error_and_close(&mut sink, "missing model".to_string()).await;
let Message::Text(payload) = &sink.messages[0] else {
panic!("expected error text frame");
};
assert_eq!(
serde_json::from_str::<serde_json::Value>(payload).expect("error json"),
json!({
"type": "error",
"error": {
"type": "invalid_request_error",
"message": "missing model"
}
})
);
assert_eq!(
sink.messages[1],
Message::Close(Some(axum::extract::ws::CloseFrame {
code: 1008,
reason: "Pre-call error".into(),
}))
);
}
fn state() -> AppState {
AppState {
router: Arc::new(ModelRouter::default()),
master_key: Some(Arc::from("master-key")),
loggers: Arc::new(Vec::new()),
realtime_pool: RealtimePool::disabled(),
}
}
#[tokio::test]
async fn auth_rejects_responses_upgrade_before_handler() {
let request = Request::builder()
.uri("/responses?model=known")
.body(Body::empty())
.expect("request");
let response = router()
.with_state(state())
.oneshot(request)
.await
.expect("response");
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[test]
fn unknown_query_model_is_rejected_before_upgrade() {
assert_eq!(
validate_model(&ModelRouter::default(), "unknown").expect_err("unknown model"),
(
StatusCode::NOT_FOUND,
"no deployment for model 'unknown'".to_string()
)
);
}
}

View file

@ -1,156 +0,0 @@
use std::sync::Arc;
use std::time::Duration;
use futures_util::{Sink, Stream};
use litellm_core::Error;
use litellm_core::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use litellm_core::responses::instrumentation::{
ResponsesWsCallbackPayload, ResponsesWsInstrumentation, ResponsesWsLogOutcome,
ResponsesWsMetadata,
};
use litellm_core::responses::types::ResponsesWsEvent;
use crate::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::RequestMetadata;
#[allow(clippy::too_many_arguments)]
pub async fn run<In, Out>(
router: &litellm_core::router::Router,
model: &str,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
call_id: String,
metadata: RequestMetadata,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let deployment = router
.get_available_deployment(model)
.ok_or_else(|| Error::Routing(format!("no deployment available for model '{model}'")))?;
let params = &deployment.litellm_params;
let provider_model = params
.model
.strip_prefix("openai/")
.unwrap_or(&params.model);
if params.model.contains('/') && !params.model.starts_with("openai/") {
return Err(Error::InvalidProvider(
"Responses WebSocket route supports OpenAI deployments only".to_string(),
));
}
let instrumentation = Arc::new(ResponsesWsInstrumentation::new(
call_id.clone(),
model,
ResponsesWsMetadata {
user_api_key_hash: metadata.user_api_key_hash,
user_api_key_user_id: metadata.user_api_key_user_id,
user_api_key_team_id: metadata.user_api_key_team_id,
},
));
let observer_instrumentation = Arc::clone(&instrumentation);
let context = CallLifecycleContext::new("responses_websocket", model, "openai", call_id);
let result = CallLifecycle::default()
.run(context, (), instrumentation.as_ref(), |_| async move {
crate::io::responses_ws::async_responses_websocket(
provider_model,
params.api_key.as_deref(),
params.api_base.as_deref(),
first_frame,
idle_timeout,
move |event| {
observer_instrumentation.observe(event);
},
client_in,
client_out,
)
.await
})
.await;
let outcome = instrumentation.take_or_build_outcome(result.is_ok());
dispatch_outcome(loggers, outcome).await;
result
}
async fn dispatch_outcome(
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
outcome: ResponsesWsLogOutcome,
) {
let runner = CustomLoggerRunner::new(loggers.as_ref().clone());
match outcome {
ResponsesWsLogOutcome::Success { payload, callback } => {
let (details, response, start_time, end_time) = logging_values(payload, callback, None);
let _ = runner
.async_log_success_event(
&details,
&response,
CallbackTiming::new(start_time, end_time),
)
.await;
}
ResponsesWsLogOutcome::Failure {
payload,
callback,
error_message,
error_kind,
} => {
let error = LoggingError {
message: error_message,
kind: error_kind,
};
let (details, response, start_time, end_time) =
logging_values(payload, callback, Some(error));
let _ = runner
.async_log_failure_event(
&details,
Some(&response),
CallbackTiming::new(start_time, end_time),
)
.await;
}
}
}
fn logging_values(
payload: litellm_core::responses::instrumentation::ResponsesWsLogPayload,
callback: ResponsesWsCallbackPayload,
error: Option<LoggingError>,
) -> (ModelCallDetails, CallbackValue, f64, f64) {
let start_time = payload.start_time;
let end_time = payload.end_time;
let callback = CallbackValue::new(callback.object, callback.value);
let details = ModelCallDetails::from_standard_logging_payload(
crate::integrations::types::StandardLoggingPayload {
id: payload.id,
litellm_call_id: payload.litellm_call_id,
call_type: payload.call_type,
model: payload.model,
custom_llm_provider: payload.custom_llm_provider,
response_cost: payload.response_cost,
prompt_tokens: payload.usage.prompt_tokens,
completion_tokens: payload.usage.completion_tokens,
total_tokens: payload.usage.total_tokens,
start_time: payload.start_time,
end_time: payload.end_time,
stream: payload.stream,
metadata: crate::integrations::types::StandardLoggingMetadata {
user_api_key_hash: payload.metadata.user_api_key_hash,
user_api_key_user_id: payload.metadata.user_api_key_user_id,
user_api_key_team_id: payload.metadata.user_api_key_team_id,
..Default::default()
},
messages: None,
},
);
let details = match error {
Some(error) => details.with_failure_error(error),
None => details,
};
(details, callback, start_time, end_time)
}

View file

@ -1,21 +0,0 @@
use std::sync::Arc;
use crate::io::realtime_pool::RealtimePool;
use litellm_core::router::Router;
use crate::integrations::custom_logger::CustomLogger;
/// Shared application state handed to every route handler.
#[derive(Clone)]
pub struct AppState {
pub router: Arc<Router>,
/// The gateway master key. Any caller presenting it as a bearer token may
/// invoke the gateway. `None` → auth not configured (routes fail closed).
pub master_key: Option<Arc<str>>,
/// Logging callbacks fanned out at the end of each realtime session.
pub loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
/// 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.
pub realtime_pool: Arc<RealtimePool>,
}

View file

@ -1,100 +0,0 @@
//! Harness-only in-process adapters. Never mounted as production routes.
use std::sync::Arc;
use axum::body::{Body, to_bytes};
use axum::http::header::{AUTHORIZATION, CONTENT_TYPE};
use axum::http::{Request, StatusCode};
use litellm_core::Error;
use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter};
use serde::Serialize;
use serde_json::Value;
use tower::ServiceExt;
use tracing::instrument::WithSubscriber;
use crate::io::realtime_pool::RealtimePool;
use crate::routes;
use crate::state::AppState;
#[derive(Debug, Serialize)]
pub struct GatewayResponse {
pub status: u16,
pub body: Value,
}
#[derive(Debug, Serialize)]
pub struct TracedGatewayResponse {
pub response: Option<GatewayResponse>,
pub error: Option<String>,
pub trace: Vec<litellm_core::observability::FunctionTraceEvent>,
}
pub async fn traced_request(
path: String,
model_alias: String,
provider_model: String,
api_base: String,
body: Value,
) -> TracedGatewayResponse {
let trace = litellm_core::observability::FunctionTrace::default();
let result = request(path, model_alias, provider_model, api_base, body)
.with_subscriber(trace.dispatcher())
.await;
let events = trace.events();
match result {
Ok(response) => TracedGatewayResponse {
response: Some(response),
error: None,
trace: events,
},
Err(error) => TracedGatewayResponse {
response: None,
error: Some(error.to_string()),
trace: events,
},
}
}
pub async fn request(
path: String,
model_alias: String,
provider_model: String,
api_base: String,
body: Value,
) -> Result<GatewayResponse, Error> {
let state = AppState {
router: Arc::new(ModelRouter::new(vec![Deployment {
model_name: model_alias,
litellm_params: LiteLLMParams {
model: provider_model,
api_key: Some("trace-provider-key".to_string()),
api_base: Some(api_base),
},
}])),
master_key: Some(Arc::from("trace-master-key")),
loggers: Arc::new(Vec::new()),
realtime_pool: RealtimePool::disabled(),
};
let request = Request::builder()
.method("POST")
.uri(path)
.header(AUTHORIZATION, "Bearer trace-master-key")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(body.to_string()))
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
let response = match routes::app(state).oneshot(request).await {
Ok(response) => response,
Err(error) => match error {},
};
let status: StatusCode = response.status();
let bytes = to_bytes(response.into_body(), usize::MAX)
.await
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
let body = serde_json::from_slice(&bytes).map_err(|error| {
Error::InvalidResponse(format!("gateway returned invalid JSON: {error}"))
})?;
Ok(GatewayResponse {
status: status.as_u16(),
body,
})
}

View file

@ -1,53 +0,0 @@
//! Guards the wiring, not just the helper: a `wss://` dial through the public
//! API has to resolve its own crypto provider, in a test binary where nothing
//! has installed a process-wide one, and has to leave it uninstalled.
use std::time::Duration;
use futures_util::{sink, stream};
use litellm_ai_gateway::io::responses_ws::async_responses_websocket;
use tokio::net::TcpListener;
async fn dead_tls_server() -> u16 {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind a loopback port");
let port = listener
.local_addr()
.expect("read the bound address")
.port();
tokio::spawn(async move {
while let Ok((stream, _peer)) = listener.accept().await {
drop(stream);
}
});
port
}
#[tokio::test]
async fn dialing_wss_returns_an_error_instead_of_panicking() {
let port = dead_tls_server().await;
let result = async_responses_websocket(
"gpt-5",
Some("test-key"),
Some(&format!("wss://127.0.0.1:{port}/")),
None,
Some(Duration::from_secs(10)),
|_| {},
stream::empty(),
sink::drain(),
)
.await;
assert!(
result.is_err(),
"a plain TCP server cannot finish a TLS handshake"
);
assert!(
rustls::crypto::CryptoProvider::get_default().is_none(),
"the dial settles its provider on its own connector, not process-wide"
);
}

View file

@ -1,16 +0,0 @@
[package]
name = "litellm-config"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-core.workspace = true
pyo3 = { workspace = true, features = ["auto-initialize"], optional = true }
serde_json.workspace = true
thiserror.workspace = true
[features]
default = []
python = ["dep:pyo3"]

View file

@ -1,11 +0,0 @@
use thiserror::Error as ThisError;
#[derive(Debug, ThisError)]
pub enum Error {
#[error("read_model_list failed: {0}")]
PythonLoading(String),
#[error("serializing model_list failed: {0}")]
Serialization(String),
#[error("parsing model_list failed: {0}")]
ModelListParsing(#[source] serde_json::Error),
}

View file

@ -1,7 +0,0 @@
mod error;
#[cfg(feature = "python")]
mod python;
pub use error::Error;
#[cfg(feature = "python")]
pub use python::load_model_list;

View file

@ -1,76 +0,0 @@
use std::path::Path;
use litellm_core::router::Deployment;
use pyo3::prelude::*;
use crate::Error;
pub fn load_model_list(config_path: &Path) -> Result<Vec<Deployment>, Error> {
Python::attach(|python| {
let model_list = python
.import("litellm.proxy.read_model_list")
.and_then(|module| module.getattr("read_model_list"))
.and_then(|reader| reader.call1((config_path.to_string_lossy().as_ref(),)))
.map_err(|error| Error::PythonLoading(error.to_string()))?;
let model_list_json = python
.import("json")
.and_then(|json| json.getattr("dumps"))
.and_then(|dumps| dumps.call1((model_list,)))
.and_then(|encoded| encoded.extract::<String>())
.map_err(|error| Error::Serialization(error.to_string()))?;
parse_model_list(&model_list_json)
})
}
fn parse_model_list(model_list_json: &str) -> Result<Vec<Deployment>, Error> {
serde_json::from_str(model_list_json).map_err(Error::ModelListParsing)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_resolved_model_list() {
let deployments = parse_model_list(
r#"[
{
"model_name": "realtime",
"litellm_params": {
"model": "openai/gpt-realtime",
"api_key": "resolved-secret",
"api_base": "https://api.example.test/v1"
}
},
{
"model_name": "without-optional-values",
"litellm_params": {"model": "openai/gpt-4.1"}
}
]"#,
)
.expect("resolved model list should parse");
assert_eq!(deployments.len(), 2);
assert_eq!(deployments[0].model_name, "realtime");
assert_eq!(
deployments[0].litellm_params.api_key.as_deref(),
Some("resolved-secret")
);
assert_eq!(
deployments[0].litellm_params.api_base.as_deref(),
Some("https://api.example.test/v1")
);
assert_eq!(deployments[1].litellm_params.api_key, None);
assert_eq!(deployments[1].litellm_params.api_base, None);
}
#[test]
fn malformed_model_list_returns_parsing_error() {
let error = parse_model_list(r#"[{"model_name":"missing-params"}]"#)
.expect_err("missing litellm_params should fail");
assert!(matches!(error, Error::ModelListParsing(_)));
}
}

View file

@ -28,8 +28,6 @@ subtle.workspace = true
tokio = { workspace = true, features = ["sync"] }
tokio-tungstenite.workspace = true
thiserror.workspace = true
tracing.workspace = true
tracing-subscriber = { workspace = true, optional = true }
sha2.workspace = true
url.workspace = true
veil.workspace = true
@ -50,8 +48,6 @@ bedrock-auth = [
"dep:aws-types",
"dep:aws-smithy-runtime-api",
]
observability = ["dep:tracing-subscriber"]
[dev-dependencies]
rstest.workspace = true
tracing-subscriber.workspace = true

View file

@ -6,7 +6,6 @@ use crate::http_utils::{http_request, truncate_error_body};
use super::client::http_client;
use super::types::ProviderAudioTranscriptionRequest;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn execute_audio_transcription_provider_call(
request: ProviderAudioTranscriptionRequest,
) -> Result<Value, Error> {

View file

@ -11,7 +11,6 @@ pub use handler::execute_audio_transcription_provider_call;
pub use prepare::prepare_audio_transcription_provider_call;
pub use types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await

View file

@ -2,12 +2,11 @@ use crate::error::Error;
use crate::http_utils::{has_header, string_headers};
#[cfg(feature = "bedrock-auth")]
use crate::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig};
use super::types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
#[cfg(feature = "bedrock-auth")]
if provider == "bedrock" {
@ -17,7 +16,6 @@ fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProv
None
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub fn prepare_audio_transcription_provider_call(
request: AudioTranscriptionRequest<'_>,
) -> Result<ProviderAudioTranscriptionRequest, Error> {

View file

@ -15,7 +15,6 @@ pub enum AudioTranscriptionAuth {
pub trait AudioTranscriptionProviderConfig: Sync {
fn supported_transcription_params(&self) -> &'static [&'static str];
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn map_transcription_params(&self, params: &Map<String, Value>) -> Map<String, Value> {
params
.iter()

View file

@ -106,7 +106,6 @@ impl VertexAuth {
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) async fn validate_environment(
&self,
headers: Vec<(String, String)>,

View file

@ -7,7 +7,6 @@ use super::transformation::ChatCompletionsProviderConfig;
const HEADER_CONTEXT: &str = "chat completions";
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) fn chat_completions_provider_config(
provider: &str,
) -> Option<&'static dyn ChatCompletionsProviderConfig> {

View file

@ -11,7 +11,6 @@ use super::types::{
ResolvedChatCompletionsRequest,
};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) async fn execute_chat_completions_provider_call(
request: ResolvedChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -22,7 +22,6 @@ use handler::execute_chat_completions_provider_call;
use prepare::{parse_messages, resolve_provider_config, resolve_request};
use types::{ChatCompletionsRequest, ChatCompletionsResponse};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn chat_completions(
request: ChatCompletionsRequest<'_>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -2,7 +2,7 @@ use serde_json::Value;
use crate::error::Error;
use crate::http_utils::has_header;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{chat_completions_provider_config, string_headers};
use super::transformation::{ChatCompletionsAuth, ChatCompletionsProviderConfig};
@ -62,7 +62,6 @@ pub(super) fn resolve_request(
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn validate_environment(
request: &ResolvedChatCompletionsRequest<'_>,
model: &str,

View file

@ -42,8 +42,6 @@ pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion";
pub const EMPTY_TEXT_PLACEHOLDER: &str =
"[System: Empty message content sanitised to satisfy protocol]";
pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace";
pub(crate) const MEDIA_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const OCR_RESPONSE_MAX_BYTES: usize = 64 * 1024 * 1024;

View file

@ -38,7 +38,6 @@ pub(crate) fn with_headers(
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn http_request(
request: reqwest::RequestBuilder,
) -> Result<reqwest::Response, reqwest::Error> {

View file

@ -8,14 +8,9 @@ pub mod error;
pub mod http_utils;
mod media;
pub mod messages;
#[cfg(any(feature = "observability", test))]
pub mod observability;
pub mod ocr;
pub mod providers;
pub mod realtime;
pub mod responses;
pub mod router;
pub mod routing_utils;
mod url_utils;
pub use auth::AuthError;

View file

@ -10,7 +10,6 @@ pub(super) use crate::http_utils::{has_bearer_auth, has_header, truncate_error_b
const HEADER_CONTEXT: &str = "messages";
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) fn messages_provider_config(
provider: &str,
) -> Option<&'static dyn AnthropicMessagesProviderConfig> {

View file

@ -7,7 +7,6 @@ use super::common_utils::truncate_error_body;
use super::prepare::prepare_provider_request;
use super::types::{AnthropicMessagesResponse, MessagesRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) async fn execute_messages_provider_call(
request: MessagesRequest<'_>,
) -> Result<AnthropicMessagesResponse, Error> {

View file

@ -18,7 +18,6 @@ pub mod types;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use types::{AnthropicMessagesResponse, MessagesRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
execute_messages_provider_call(request).await
}

View file

@ -1,5 +1,5 @@
use crate::error::Error;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
use super::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy};
@ -56,7 +56,6 @@ pub(super) fn prepare_provider_request(
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn validate_environment(
config: &dyn AnthropicMessagesProviderConfig,
extra_headers: Option<Map<String, Value>>,

View file

@ -45,7 +45,6 @@ pub trait AnthropicMessagesProviderConfig: Sync {
]
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_request(
&self,
request: AnthropicMessagesRequest,
@ -53,7 +52,6 @@ pub trait AnthropicMessagesProviderConfig: Sync {
Ok(request)
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_response(
&self,
_model: &str,

View file

@ -1,215 +0,0 @@
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use serde::Serialize;
use tracing::span::{Attributes, Id};
use tracing::{Dispatch, Subscriber};
use tracing_subscriber::layer::Context;
use tracing_subscriber::prelude::*;
use tracing_subscriber::registry::LookupSpan;
use tracing_subscriber::{Layer, Registry};
use super::function_trace_filter;
#[derive(Clone, Debug, PartialEq, Serialize)]
pub struct FunctionTraceEvent {
pub id: usize,
pub parent_id: Option<usize>,
pub function: &'static str,
pub module_path: Option<&'static str>,
pub file: Option<&'static str>,
pub line: Option<u32>,
}
#[derive(Clone, Default)]
pub struct FunctionTrace {
events: Arc<Mutex<Vec<FunctionTraceEvent>>>,
span_events: Arc<Mutex<HashMap<Id, usize>>>,
}
impl FunctionTrace {
pub fn dispatcher(&self) -> Dispatch {
Dispatch::new(
Registry::default().with(
FunctionTraceLayer {
trace: self.clone(),
}
.with_filter(function_trace_filter()),
),
)
}
pub fn events(&self) -> Vec<FunctionTraceEvent> {
self.events
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone()
}
}
struct FunctionTraceLayer {
trace: FunctionTrace,
}
impl<S> Layer<S> for FunctionTraceLayer
where
S: Subscriber + for<'lookup> LookupSpan<'lookup>,
{
fn on_new_span(&self, attributes: &Attributes<'_>, id: &Id, context: Context<'_, S>) {
let parent_id = context.span(id).and_then(|span| {
let span_events = self
.trace
.span_events
.lock()
.unwrap_or_else(|error| error.into_inner());
span.scope()
.skip(1)
.find_map(|ancestor| span_events.get(&ancestor.id()).copied())
});
let mut events = self
.trace
.events
.lock()
.unwrap_or_else(|error| error.into_inner());
let event_id = events.len();
events.push(FunctionTraceEvent {
id: event_id,
parent_id,
function: attributes.metadata().name(),
module_path: attributes.metadata().module_path(),
file: attributes.metadata().file(),
line: attributes.metadata().line(),
});
self.trace
.span_events
.lock()
.unwrap_or_else(|error| error.into_inner())
.insert(id.clone(), event_id);
}
}
#[cfg(test)]
mod tests {
use crate::constants::FUNCTION_TRACE_TARGET;
use super::*;
fn event(
id: usize,
parent_id: Option<usize>,
function: &'static str,
) -> (usize, Option<usize>, &'static str) {
(id, parent_id, function)
}
fn structural_events(
events: &[FunctionTraceEvent],
) -> Vec<(usize, Option<usize>, &'static str)> {
events
.iter()
.map(|event| (event.id, event.parent_id, event.function))
.collect()
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn outer() {
tokio::task::yield_now().await;
inner().await;
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn inner() {
tokio::task::yield_now().await;
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn concurrent_parent() {
tokio::join!(inner(), inner());
}
#[tokio::test]
async fn concurrent_futures_keep_separate_traces_across_yields() {
use tracing::instrument::WithSubscriber;
let first = FunctionTrace::default();
let second = FunctionTrace::default();
let outside = FunctionTrace::default();
async {
tokio::join!(
outer().with_subscriber(first.dispatcher()),
inner().with_subscriber(second.dispatcher()),
);
inner().await;
}
.with_subscriber(outside.dispatcher())
.await;
assert_eq!(
structural_events(&first.events()),
vec![event(0, None, "outer"), event(1, Some(0), "inner")],
);
assert_eq!(
structural_events(&second.events()),
vec![event(0, None, "inner")],
);
assert_eq!(
structural_events(&outside.events()),
vec![event(0, None, "inner")],
);
}
#[tokio::test]
async fn concurrent_siblings_keep_the_same_parent() {
use tracing::instrument::WithSubscriber;
let trace = FunctionTrace::default();
concurrent_parent()
.with_subscriber(trace.dispatcher())
.await;
assert_eq!(
structural_events(&trace.events()),
vec![
event(0, None, "concurrent_parent"),
event(1, Some(0), "inner"),
event(2, Some(0), "inner"),
]
);
}
#[test]
fn records_matching_spans_in_creation_order() {
let trace = FunctionTrace::default();
let dispatch = trace.dispatcher();
tracing::dispatcher::with_default(&dispatch, || {
let _ignored = tracing::trace_span!(target: "other", "ignored");
let _first = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "same_name");
let _wrong_level = tracing::debug_span!(target: FUNCTION_TRACE_TARGET, "wrong_level");
let _second = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "same_name");
});
assert_eq!(
structural_events(&trace.events()),
vec![event(0, None, "same_name"), event(1, None, "same_name")]
);
}
#[test]
fn records_matching_span_nesting_depth() {
let trace = FunctionTrace::default();
let dispatch = trace.dispatcher();
tracing::dispatcher::with_default(&dispatch, || {
let outer = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "outer");
let _outer_guard = outer.enter();
let _inner = tracing::trace_span!(target: FUNCTION_TRACE_TARGET, "inner");
});
assert_eq!(
structural_events(&trace.events()),
vec![event(0, None, "outer"), event(1, Some(0), "inner")]
);
}
}

View file

@ -1,59 +0,0 @@
use tracing::span::Id;
use tracing::{Level, Metadata, Subscriber};
use tracing_subscriber::filter::{FilterFn, LevelFilter, filter_fn};
use tracing_subscriber::layer::Context;
use tracing_subscriber::registry::LookupSpan;
use crate::constants::FUNCTION_TRACE_TARGET;
pub mod function_trace;
pub use function_trace::{FunctionTrace, FunctionTraceEvent};
pub fn function_trace_filter() -> FilterFn<impl Fn(&Metadata<'_>) -> bool> {
filter_fn(|metadata| {
metadata.is_span()
&& metadata.target() == FUNCTION_TRACE_TARGET
&& *metadata.level() == Level::TRACE
})
.with_max_level_hint(LevelFilter::TRACE)
}
pub fn span_depth<S>(context: &Context<'_, S>, id: &Id) -> usize
where
S: Subscriber + for<'lookup> LookupSpan<'lookup>,
{
context
.span(id)
.map(|span| span.scope().skip(1).count())
.unwrap_or_default()
}
#[cfg(test)]
mod tests {
use tracing::instrument::WithSubscriber;
use super::*;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn instrumented_with_literal_target() {}
#[tokio::test]
async fn literal_instrument_target_matches_filter_constant() {
assert_eq!(FUNCTION_TRACE_TARGET, "litellm::function_trace");
let trace = FunctionTrace::default();
instrumented_with_literal_target()
.with_subscriber(trace.dispatcher())
.await;
let events = trace.events();
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, 0);
assert_eq!(events[0].parent_id, None);
assert_eq!(events[0].function, "instrumented_with_literal_target");
assert_eq!(events[0].module_path, Some(module_path!()));
assert_eq!(events[0].file, Some(file!()));
assert!(events[0].line.is_some());
}
}

View file

@ -75,7 +75,6 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn map_ocr_params(
request: &LiteLLMOcrRequest,
) -> Result<DocumentIntelligenceParams, OcrRequestError> {

View file

@ -36,12 +36,6 @@ impl OcrClient {
shared_client()
}
#[tracing::instrument(
name = "ocr",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
use super::{
NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost,

View file

@ -5,7 +5,6 @@ use super::types::*;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn transform_ocr_request(
provider_model: &str,
document: OcrDocument,

View file

@ -7,7 +7,6 @@ use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn transform_ocr_request(
document: OcrDocument,
) -> Result<DocumentIntelligenceRequest, OcrRequestError> {

View file

@ -2,7 +2,6 @@ use super::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse};
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn transform_ocr_request(
model: &str,
document: OcrDocument,

View file

@ -6,12 +6,6 @@ use super::types::*;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub(crate) fn transform_v3_ocr_request(
_model: &str,
document: OcrDocument,
@ -23,12 +17,6 @@ pub(crate) fn transform_v3_ocr_request(
})
}
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub(crate) fn transform_legacy_ocr_request(
_model: &str,
document: OcrDocument,

View file

@ -125,12 +125,6 @@ impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
Box::pin(async move { Ok(request) })
}
#[tracing::instrument(
name = "success_callback",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn async_log_success_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
@ -140,12 +134,6 @@ impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
self.hooks.success(context, response, timing)
}
#[tracing::instrument(
name = "failure_callback",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn async_log_failure_event<'a>(
&'a self,
context: &'a CallLifecycleContext,

View file

@ -14,7 +14,6 @@ pub(crate) struct ParsedProviderParams<T> {
pub extra_params: Map<String, Value>,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn _prepare_ocr_request<T: DeserializeOwned>(
request: &LiteLLMOcrRequest,
) -> Result<ParsedProviderParams<T>, OcrRequestError> {

View file

@ -1,6 +1,6 @@
use super::adapters::OcrAdapter;
use crate::Error;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
macro_rules! define_adapter_types {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {

View file

@ -117,7 +117,6 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)] {
SUPPORTED_PARAMS
}
@ -138,7 +137,6 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_request(
&self,
model: &str,
@ -150,7 +148,6 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_response(
&self,
_model: &str,

View file

@ -42,7 +42,6 @@ pub fn complete_anthropic_url(
}
impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn complete_url(
&self,
api_base: Option<&str>,

View file

@ -132,7 +132,6 @@ fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMess
}
impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn complete_url(
&self,
api_base: Option<&str>,

View file

@ -46,12 +46,10 @@ fn optional_string<'a>(params: &'a Map<String, Value>, key: &str) -> Option<&'a
}
impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn supported_transcription_params(&self) -> &'static [&'static str] {
SUPPORTED_PARAMS
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_transcription_request(
&self,
_model: &str,
@ -85,7 +83,6 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn transform_transcription_response(
&self,
_model: &str,

View file

@ -163,7 +163,6 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
&[("Content-Type", "application/json")]
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)] {
SUPPORTED_PARAMS
}

View file

@ -2,4 +2,5 @@ pub mod anthropic;
pub mod azure_ai;
#[cfg(feature = "bedrock-auth")]
pub mod bedrock;
pub mod custom_llm_provider;
pub mod openai;

View file

@ -1,2 +1 @@
pub mod realtime;
pub mod responses;

View file

@ -1 +0,0 @@
pub mod transformation;

View file

@ -1,189 +0,0 @@
use crate::Error;
use crate::realtime::transformation::RealtimeProviderConfig;
use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
/// Default OpenAI API base, used when the caller does not override `api_base`.
pub const OPENAI_REALTIME_DEFAULT_API_BASE: &str = "https://api.openai.com";
/// Path appended to the resolved host base to reach the realtime endpoint.
pub const OPENAI_REALTIME_PATH: &str = "/v1/realtime";
/// Percent-encode a query value, escaping any char outside the RFC 3986
/// unreserved set (`A-Za-z0-9-._~`). Keeps us dependency-free; common realtime
/// model slugs have no special chars, but this stays correct for the rest.
fn percent_encode(value: &str) -> String {
let mut encoded = String::with_capacity(value.len());
for byte in value.bytes() {
let unreserved = byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~');
if unreserved {
encoded.push(byte as char);
} else {
encoded.push('%');
encoded.push_str(&format!("{byte:02X}"));
}
}
encoded
}
/// Build the realtime WebSocket URL, porting Python's `OpenAIRealtime._construct_url`.
///
/// Blank/whitespace `api_base` is treated as absent (guard at resolution time),
/// falling back to the default. The scheme is swapped to its WebSocket
/// equivalent (`https://`→`wss://`, `http://`→`ws://`); bases already using
/// `ws`/`wss` are left untouched. A bare host or unrecognized scheme defaults to
/// secure `wss://` so we never hand a scheme-less URL to the connector (this is
/// a deliberate hardening over Python's `_construct_url`, which would emit a
/// scheme-less URL here). A trailing `/` is trimmed before the path and
/// `?model=<encoded>` are appended.
pub fn complete_url(api_base: Option<&str>, model: &str) -> String {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(OPENAI_REALTIME_DEFAULT_API_BASE);
let base = if let Some(rest) = base.strip_prefix("https://") {
format!("wss://{rest}")
} else if let Some(rest) = base.strip_prefix("http://") {
format!("ws://{rest}")
} else if base.starts_with("wss://") || base.starts_with("ws://") {
base.to_string()
} else {
format!("wss://{base}")
};
let base = base.trim_end_matches('/');
format!(
"{base}{OPENAI_REALTIME_PATH}?model={}",
percent_encode(model)
)
}
pub struct OpenAiRealtimeConfig;
pub const OPENAI_REALTIME_CONFIG: OpenAiRealtimeConfig = OpenAiRealtimeConfig;
impl RealtimeProviderConfig for OpenAiRealtimeConfig {
fn complete_url(&self, api_base: Option<&str>, model: &str) -> String {
complete_url(api_base, model)
}
fn transform_realtime_request(
&self,
event: &RealtimeEvent,
_model: &str,
) -> Result<RealtimeTransformResult, Error> {
Ok(RealtimeTransformResult::passthrough(event.clone()))
}
fn transform_realtime_response(
&self,
event: &RealtimeEvent,
_model: &str,
) -> Result<RealtimeTransformResult, Error> {
Ok(RealtimeTransformResult::passthrough(event.clone()))
}
}
pub fn transform_realtime_request(
event: &RealtimeEvent,
model: &str,
) -> Result<RealtimeTransformResult, Error> {
OPENAI_REALTIME_CONFIG.transform_realtime_request(event, model)
}
pub fn transform_realtime_response(
event: &RealtimeEvent,
model: &str,
) -> Result<RealtimeTransformResult, Error> {
OPENAI_REALTIME_CONFIG.transform_realtime_response(event, model)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn complete_url_defaults_to_openai_wss() {
assert_eq!(
complete_url(None, "gpt-4o-realtime-preview"),
"wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
);
}
#[test]
fn complete_url_blank_base_uses_default() {
assert_eq!(
complete_url(Some(" "), "gpt-4o-realtime-preview"),
"wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
);
}
#[test]
fn complete_url_swaps_http_to_ws() {
assert_eq!(
complete_url(Some("http://localhost:8080"), "gpt-4o-realtime-preview"),
"ws://localhost:8080/v1/realtime?model=gpt-4o-realtime-preview"
);
}
#[test]
fn complete_url_dedupes_trailing_slash() {
assert_eq!(
complete_url(Some("https://api.openai.com/"), "gpt-4o-realtime-preview"),
"wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview"
);
}
#[test]
fn complete_url_custom_base() {
assert_eq!(
complete_url(Some("https://oai.azure.example"), "gpt-4o-realtime-preview"),
"wss://oai.azure.example/v1/realtime?model=gpt-4o-realtime-preview"
);
}
#[test]
fn complete_url_preserves_existing_wss_scheme() {
assert_eq!(
complete_url(Some("wss://api.openai.com"), "gpt-realtime"),
"wss://api.openai.com/v1/realtime?model=gpt-realtime"
);
}
#[test]
fn complete_url_bare_host_defaults_to_wss() {
assert_eq!(
complete_url(Some("api.openai.com"), "gpt-realtime"),
"wss://api.openai.com/v1/realtime?model=gpt-realtime"
);
}
#[test]
fn complete_url_percent_encodes_model_space() {
assert_eq!(
complete_url(None, "gpt 4o"),
"wss://api.openai.com/v1/realtime?model=gpt%204o"
);
}
#[test]
fn transform_realtime_request_passthrough_preserves_event() {
let event: RealtimeEvent =
serde_json::from_str(r#"{"type":"session.update","session":{"voice":"alloy"}}"#)
.expect("valid event");
let result =
transform_realtime_request(&event, "gpt-realtime").expect("passthrough is infallible");
assert_eq!(result.events, vec![event]);
}
#[test]
fn transform_realtime_response_passthrough_preserves_event() {
let event: RealtimeEvent =
serde_json::from_str(r#"{"type":"response.output_audio.delta","delta":"abc=="}"#)
.expect("valid event");
let result =
transform_realtime_response(&event, "gpt-realtime").expect("passthrough is infallible");
assert_eq!(result.events, vec![event]);
}
}

View file

@ -1,2 +0,0 @@
pub mod transformation;
pub mod types;

View file

@ -1,22 +0,0 @@
use crate::Error;
use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
pub trait RealtimeProviderConfig {
/// Build the upstream WebSocket URL (e.g. `wss://api.openai.com/v1/realtime?model=…`).
/// Pure string construction only — no network, no env.
fn complete_url(&self, api_base: Option<&str>, model: &str) -> String;
/// Transform a client → backend event before it is forwarded upstream.
fn transform_realtime_request(
&self,
event: &RealtimeEvent,
model: &str,
) -> Result<RealtimeTransformResult, Error>;
/// Transform a backend → client event before it is forwarded downstream.
fn transform_realtime_response(
&self,
event: &RealtimeEvent,
model: &str,
) -> Result<RealtimeTransformResult, Error>;
}

View file

@ -1,60 +0,0 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
/// A single realtime event exchanged over the WebSocket.
///
/// The `type` discriminator is a typed field; the remaining fields are
/// preserved losslessly in `data` so a transform can pass an event through, or
/// inspect/modify specific fields, without enumerating every event variant.
/// Wire (de)serialization happens at the host edge — `core`/`providers` operate
/// only on this typed form.
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct RealtimeEvent {
#[serde(rename = "type")]
pub event_type: String,
#[serde(flatten)]
pub data: Map<String, Value>,
}
/// One or more typed events produced by a realtime transform.
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct RealtimeTransformResult {
pub events: Vec<RealtimeEvent>,
}
impl RealtimeTransformResult {
/// Forward a single event unchanged (the OpenAI baseline).
pub fn passthrough(event: RealtimeEvent) -> Self {
Self {
events: vec![event],
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn event(raw: &str) -> RealtimeEvent {
serde_json::from_str(raw).expect("valid event json")
}
#[test]
fn realtime_event_round_trips_type_and_extra_fields() {
let raw = r#"{"type":"response.output_text.delta","delta":"hi","response_id":"r1"}"#;
let parsed = event(raw);
assert_eq!(parsed.event_type, "response.output_text.delta");
assert_eq!(parsed.data.get("delta"), Some(&Value::String("hi".into())));
// Re-serializing yields a semantically-equal event (key order may differ).
let reparsed: RealtimeEvent =
serde_json::from_str(&serde_json::to_string(&parsed).unwrap()).unwrap();
assert_eq!(parsed, reparsed);
}
#[test]
fn passthrough_produces_single_element_vec() {
let parsed = event(r#"{"type":"session.update"}"#);
let result = RealtimeTransformResult::passthrough(parsed.clone());
assert_eq!(result.events, vec![parsed]);
}
}

Some files were not shown because too many files have changed in this diff Show more