mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge remote-tracking branch 'origin/main' into litellm_mcp_discovery_cache_fixes
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> # Conflicts: # litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
This commit is contained in:
commit
a157a696af
215 changed files with 12410 additions and 5620 deletions
|
|
@ -3419,7 +3419,7 @@ workflows:
|
|||
name: integration-<< matrix.suite >>
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, mcp, sdk, cost, browser]
|
||||
suite: [management, accounting, database, providers, mcp, sdk, cost, security, browser]
|
||||
- integration_contracts:
|
||||
name: integration-extensions
|
||||
suite: extensions
|
||||
|
|
|
|||
13
.github/workflows/create-rc-branch.yml
vendored
13
.github/workflows/create-rc-branch.yml
vendored
|
|
@ -15,6 +15,8 @@ jobs:
|
|||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
outputs:
|
||||
version: ${{ steps.version.outputs.version }}
|
||||
steps:
|
||||
- name: Require main
|
||||
env:
|
||||
|
|
@ -64,3 +66,14 @@ jobs:
|
|||
sha: context.sha,
|
||||
});
|
||||
core.info(`Created branch ${branchName} at ${context.sha}`);
|
||||
|
||||
linear-release:
|
||||
name: Move the Linear release to rc
|
||||
needs: create-rc-branch
|
||||
permissions:
|
||||
contents: read
|
||||
uses: ./.github/workflows/linear-release.yml
|
||||
with:
|
||||
rc_version: ${{ needs.create-rc-branch.outputs.version }}
|
||||
secrets:
|
||||
LINEAR_API_KEY: ${{ secrets.LINEAR_API_KEY }}
|
||||
|
|
|
|||
131
.github/workflows/linear-release.yml
vendored
Normal file
131
.github/workflows/linear-release.yml
vendored
Normal file
|
|
@ -0,0 +1,131 @@
|
|||
name: Linear Release
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- "rc/**"
|
||||
release:
|
||||
types: [published]
|
||||
workflow_call:
|
||||
inputs:
|
||||
rc_version:
|
||||
description: "X.Y.0 release whose rc branch was just cut"
|
||||
required: true
|
||||
type: string
|
||||
secrets:
|
||||
LINEAR_API_KEY:
|
||||
required: true
|
||||
|
||||
permissions: {}
|
||||
|
||||
jobs:
|
||||
linear-release:
|
||||
name: Linear Release
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Plan
|
||||
id: plan
|
||||
env:
|
||||
EVENT: ${{ github.event_name }}
|
||||
REF_NAME: ${{ github.ref_name }}
|
||||
BEFORE: ${{ github.event.before }}
|
||||
CREATED: ${{ github.event.created }}
|
||||
RC_VERSION: ${{ inputs.rc_version }}
|
||||
RELEASE_TAG: ${{ github.event.release.tag_name }}
|
||||
PRERELEASE: ${{ github.event.release.prerelease }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
sync_base="${BEFORE}"
|
||||
if [ "${CREATED}" = "true" ]; then
|
||||
sync_base=""
|
||||
fi
|
||||
if [ -n "${RC_VERSION}" ]; then
|
||||
echo "version=${RC_VERSION}" >> "$GITHUB_OUTPUT"
|
||||
echo "stage=rc" >> "$GITHUB_OUTPUT"
|
||||
elif [ "${EVENT}" = "release" ]; then
|
||||
if [ "${PRERELEASE}" = "true" ] || ! echo "${RELEASE_TAG}" | grep -qE '^v[0-9]+\.[0-9]+\.0$'; then
|
||||
echo "::notice::${RELEASE_TAG} is not an X.Y.0 stable release; nothing to complete"
|
||||
exit 0
|
||||
fi
|
||||
echo "version=${RELEASE_TAG#v}" >> "$GITHUB_OUTPUT"
|
||||
echo "complete=true" >> "$GITHUB_OUTPUT"
|
||||
elif [ "${REF_NAME}" = "main" ]; then
|
||||
version="$(python3 .github/scripts/read_rc_version.py | cut -d= -f2)"
|
||||
status=0
|
||||
git ls-remote --exit-code --heads origin "rc/${version}" > /dev/null || status=$?
|
||||
case "${status}" in
|
||||
0)
|
||||
IFS=. read -r major minor _ <<< "${version}"
|
||||
version="${major}.$((minor + 1)).0"
|
||||
;;
|
||||
2) ;;
|
||||
*)
|
||||
echo "::error::could not check whether rc/${version} exists (git ls-remote exit ${status})"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
echo "version=${version}" >> "$GITHUB_OUTPUT"
|
||||
echo "sync_base=${sync_base}" >> "$GITHUB_OUTPUT"
|
||||
echo "main=true" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "version=${REF_NAME#rc/}" >> "$GITHUB_OUTPUT"
|
||||
echo "sync_base=${sync_base}" >> "$GITHUB_OUTPUT"
|
||||
echo "stage=rc" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Sync commits into the release
|
||||
if: steps.plan.outputs.sync_base != ''
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: sync
|
||||
name: LiteLLM ${{ steps.plan.outputs.version }}
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
base_ref: ${{ steps.plan.outputs.sync_base }}
|
||||
cli_version: v0.18.0
|
||||
|
||||
- name: Keep the main stage unless the rc branch was cut during this run
|
||||
id: main_stage
|
||||
if: steps.plan.outputs.main == 'true'
|
||||
env:
|
||||
VERSION: ${{ steps.plan.outputs.version }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
status=0
|
||||
git ls-remote --exit-code --heads origin "rc/${VERSION}" > /dev/null || status=$?
|
||||
case "${status}" in
|
||||
0) echo "::notice::rc/${VERSION} was cut during this run; leaving the release in its rc stage" ;;
|
||||
2) echo "stage=main" >> "$GITHUB_OUTPUT" ;;
|
||||
*)
|
||||
echo "::error::could not check whether rc/${VERSION} exists (git ls-remote exit ${status})"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
- name: Move the release to its stage
|
||||
if: steps.plan.outputs.stage != '' || steps.main_stage.outputs.stage != ''
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: update
|
||||
stage: ${{ steps.plan.outputs.stage || steps.main_stage.outputs.stage }}
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
cli_version: v0.18.0
|
||||
|
||||
- name: Complete the release
|
||||
if: steps.plan.outputs.complete == 'true'
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: complete
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
cli_version: v0.18.0
|
||||
|
|
@ -1,37 +0,0 @@
|
|||
# Publish MCP servers in the AI Hub
|
||||
|
||||
Set `litellm_settings.public_mcp_servers` to the concrete IDs of the servers you want listed in the public AI Hub. Pin `server_id` in each configuration entry so the publication list stays stable across deployments
|
||||
|
||||
```yaml
|
||||
mcp_servers:
|
||||
documentation:
|
||||
server_id: documentation-mcp
|
||||
url: https://mcp.example.com/mcp
|
||||
transport: http
|
||||
available_on_public_internet: true
|
||||
|
||||
litellm_settings:
|
||||
public_mcp_hub_strict_whitelist: true
|
||||
public_mcp_servers:
|
||||
- documentation-mcp
|
||||
```
|
||||
|
||||
Use `documentation-mcp`, the `server_id`, in the publication list. The configuration key `documentation`, display names, and aliases are not publication IDs. Database-created servers use the ID returned by `/v1/mcp/server`
|
||||
|
||||
The dashboard's **AI Hub > MCP Hub > Manage MCP Hub Visibility** dialog edits this same list. Its YAML example includes the selected server IDs. With database-backed configuration (`store_model_in_db: true`), a value declared in YAML is owned by that file: edit the file and reload, or remove that key from YAML to let the dashboard manage it in the database. File-backed deployments can save the list directly to their configuration file
|
||||
|
||||
To remove all explicit entries, save an empty selection in the dialog or configure:
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
public_mcp_hub_strict_whitelist: true
|
||||
public_mcp_servers: []
|
||||
```
|
||||
|
||||
## Hub listing and network access
|
||||
|
||||
The **Hub listing** column in AI Hub identifies servers that appear in `/public/mcp_hub`. The dashboard derives this status from the current registry and publication settings. Setting `mcp_info.is_public` on a server does not publish it; that response field is derived metadata. `mcp_info.is_public_explicit` identifies registered servers included in the explicit publication list
|
||||
|
||||
Gateway cards and server details show **All Networks** when `available_on_public_internet` is enabled or the server is explicitly published in `public_mcp_servers`. They show **Internal Only** when both are false. The per-server flag defaults to `true`; explicit publication overrides a disabled flag for compatibility. Older proxies that omit the metadata needed to determine access show **Unknown**. These labels describe allowed client IPs; authentication and tool permissions still apply
|
||||
|
||||
The default `public_mcp_hub_strict_whitelist: true` lists only registered servers in `public_mcp_servers`. Legacy mode (`false`) additionally lists registered servers with `available_on_public_internet: true`. In legacy mode, clearing the explicit publication list leaves these automatically listed servers visible. Enable strict mode when the publication list should fully determine hub visibility
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "pinned_tools" JSONB DEFAULT '{}';
|
||||
|
|
@ -322,6 +322,7 @@ model LiteLLM_MCPServerTable {
|
|||
allowed_tools String[] @default([])
|
||||
tool_name_to_display_name Json? @default("{}")
|
||||
tool_name_to_description Json? @default("{}")
|
||||
pinned_tools Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
static_headers Json? @default("{}")
|
||||
// Admin-configured environment variables interpolated into static_headers
|
||||
|
|
|
|||
12
litellm-rust/Cargo.lock
generated
12
litellm-rust/Cargo.lock
generated
|
|
@ -3552,8 +3552,10 @@ name = "litellm-cache-response"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"litellm-cache",
|
||||
"litellm-cache-gcs",
|
||||
"litellm-cache-memory",
|
||||
"litellm-cache-redis",
|
||||
"litellm-http",
|
||||
"py_literal",
|
||||
"redis",
|
||||
"redis-test",
|
||||
|
|
@ -3562,6 +3564,7 @@ dependencies = [
|
|||
"serde_json",
|
||||
"sha2 0.10.9",
|
||||
"tokio",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3646,7 +3649,11 @@ dependencies = [
|
|||
"litellm-auth",
|
||||
"litellm-auth-aws",
|
||||
"litellm-auth-gcp",
|
||||
"litellm-cache",
|
||||
"litellm-cache-memory",
|
||||
"litellm-cache-response",
|
||||
"litellm-core-utils",
|
||||
"litellm-framing",
|
||||
"litellm-host",
|
||||
"litellm-host-native",
|
||||
"litellm-http",
|
||||
|
|
@ -3669,6 +3676,7 @@ dependencies = [
|
|||
"time",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"tokio-util",
|
||||
"tracing",
|
||||
"url",
|
||||
"veil",
|
||||
|
|
@ -3808,8 +3816,11 @@ dependencies = [
|
|||
"bytes",
|
||||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-cache-memory",
|
||||
"litellm-cache-response",
|
||||
"litellm-core",
|
||||
"litellm-gateway-auth",
|
||||
"litellm-host",
|
||||
"litellm-host-http",
|
||||
"litellm-http",
|
||||
"litellm-llms",
|
||||
|
|
@ -3817,6 +3828,7 @@ dependencies = [
|
|||
"litellm-secrets",
|
||||
"litellm-types",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ async fn set_writes_encoded_object_and_headers(#[future(awt)] server: MockServer
|
|||
)]
|
||||
#[case::missing("missing", ResponseTemplate::new(404), Ok(None))]
|
||||
#[case::server_error("server-error", ResponseTemplate::new(500), Err(Error::Unavailable))]
|
||||
#[case::unauthorized("unauthorized", ResponseTemplate::new(401), Err(Error::Unavailable))]
|
||||
#[case::invalid(
|
||||
"invalid",
|
||||
ResponseTemplate::new(200).set_body_string("not json"),
|
||||
|
|
|
|||
29
litellm-rust/crates/cache-response/AGENTS.md
Normal file
29
litellm-rust/crates/cache-response/AGENTS.md
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
# Response caching
|
||||
|
||||
Design this crate for shared Rust execution used by the Python SDK and the Rust gateway. The Python SDK will remain, with more core execution moving to Rust and Python callbacks staying in Python. The Rust gateway is still evolving and is intended to replace the Python proxy. Keep response-cache policy independent of Python, HTTP serving, and either proxy's configuration format
|
||||
|
||||
Separate what is cached, how a hit is matched, and where entries are stored. Chat Completions, Messages, Responses, and embeddings are API workloads. Exact and semantic matching are lookup behaviors. Memory, Redis, disk, and object stores are storage choices. Embeddings are inference too, so do not use an inference-cache name to imply a category that excludes embeddings. Consult the existing Python cache and caching handler for behavior and compatibility contracts without copying their class structure
|
||||
|
||||
Storage traits, codecs, and backend capabilities belong in `litellm-cache` and the storage crates. Keep storage reusable for value types beyond LLM responses. This crate owns response entries, matching and freshness semantics, the Python-compatible response codec, and deferred-write policy. Core owns route-specific request identity, response encoding and reconstruction, embedding partial-hit orchestration, and stream capture and replay. Boundaries own configuration translation, resource construction, and caller identity
|
||||
|
||||
Construct and inject the response-cache service at the Python bridge or gateway boundary, as with the HTTP client. Reuse it across calls. Core and provider code must not discover cache configuration through Python globals, process configuration, or backend-specific factories
|
||||
|
||||
Keep `ResponseCache<B>` generic over its storage backend. Preserve typed backend contexts and capability bounds internally. Inject an object-safe service into core for runtime backend selection, so storage types do not spread through route and host types. Keep API request and response types statically typed. Add a generic parameter only where it preserves a useful type relationship or capability
|
||||
|
||||
Keep the core service contract narrow. Lookup and store must not require connection testing, ping, flush, deletion, counters, queues, or scripts. Require batch operations where a consumer needs partial hits, and keep management capabilities on their own interfaces. An exact-only adapter must remain explicit about its matching restriction. Supporting semantic matching requires a defined lookup-context and embedding execution contract, not just a renamed trait
|
||||
|
||||
Separate reusable resources from per-call policy. Backend configuration, namespace, default expiry, and entry limits belong to the configured service or backend. Read/write controls, expiry and freshness overrides, and authenticated caller scope belong to the call. Passing call options must not replace or mutate the route's configured service
|
||||
|
||||
Keep cache misses and storage failures distinguishable in return values. Core owns the decision to continue with provider execution after a cache failure. A read can reject an entry for freshness while the backend still retains it. Preserve the timestamp at which a response was produced when writing it later
|
||||
|
||||
Define lookup placement explicitly relative to authorization, deployment and credential resolution, and request-transforming callbacks. Cache identity must account for every input that affects reuse, including API surface and caller scope, while preserving intentional Python caching groups. Preserve existing keys and response formats unless changing them is an explicit migration decision
|
||||
|
||||
Cache normalized provider results before caller-specific response transformations. Hits must still run the applicable response processing, success callbacks, and cache-hit accounting. Keep callback execution in the host. Python cache implementations and semantic embedders that require the caller's task must use the existing host-operation mechanism rather than Python calls from a Rust worker. Preserve legacy fallback until that contract is supported
|
||||
|
||||
Keep unary caching independent of stream-only methods. Store streams only after successful exhaustion and protocol completion. Errors, incomplete streams, cancellation, and oversized entries must not populate the cache. Embedding batches need ordered partial results and reconstruction around the uncached inputs
|
||||
|
||||
Test each contract in its owner: storage capabilities in backend tests, envelopes and freshness here, reuse and replay in core, Python callback and fallback behavior at the bridge, and HTTP behavior at the gateway. Run backend contract checks and Python response-codec fixtures before exposing a new backend
|
||||
|
||||
`ScopedCache` requires an explicit shared or isolated scope at construction. `CacheOptions` has no default sharing policy. Callers may override policy per invocation without replacing the attached service. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec
|
||||
|
||||
Response storage is not the source of budget or rate-limit coordination dependencies. Keep counters, reservations, and atomic admission operations out of `ResponseCacheService`, including when both services happen to use Redis
|
||||
|
|
@ -13,9 +13,12 @@ serde_json.workspace = true
|
|||
sha2.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-cache-gcs.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-memory.workspace = true
|
||||
litellm-cache-redis.workspace = true
|
||||
redis = "1.7.0"
|
||||
redis-test = "1.0.4"
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -1,51 +0,0 @@
|
|||
# Response cache
|
||||
|
||||
`ResponseCache<B>` adds request keys, independent read/write controls, response envelopes, and freshness checks to any `B: BaseCache<Value = CacheEntry>`
|
||||
|
||||
## Ownership
|
||||
|
||||
`litellm-cache` defines typed storage, codec, and capability traits. `BaseCache` is only get, set, TTL, and pipeline writes. Everything else is an optional capability a backend implements only where its Python class defines the method: `DisconnectCache`, `ConnectionCache` (`test_connection`), `PingCache`, `BatchCache`, `DeleteCache`, `FlushCache`, counters, queues, TTL, scan, and scripts. Memory, Redis, disk, S3, GCS, and Azure Blob implement those traits without depending on response policy, so other consumers can store their own value types in the same backends
|
||||
|
||||
Semantic backends (Redis, Valkey, Qdrant) are generic over their embedder and codec, and share one prompt and embedding contract from `litellm_cache::semantic`. They take a `SemanticCacheContext`, so `ResponseCache` drives them the same way it drives exact backends
|
||||
|
||||
`litellm-cache-response` owns response keys, controls, entries, the Python-compatible response codec, and `WriteBuffer`, the backend-neutral deferred-write policy. It has no runtime dependency on a specific cache backend or Python
|
||||
|
||||
`ExactResponseCache` is the object-safe view of a `ResponseCache` over an exact backend. `ConnectionProbe` is the object-safe `test_connection`, implemented only when the backend implements `ConnectionCache`, so a host holds one next to its `ExactResponseCache` and reports the operation as unsupported otherwise, as Python's `BaseCache` does. Lookup, store, batch, and flush never require it
|
||||
|
||||
## Native Rust use
|
||||
|
||||
```rust
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{CacheKeyInput, ResponseCache, ResponseCacheRequest};
|
||||
use serde_json::json;
|
||||
|
||||
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
|
||||
let request = ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some("example:key".into()),
|
||||
..Default::default()
|
||||
});
|
||||
let now = Duration::from_secs(100);
|
||||
cache.store(&request, json!({"answer": 7}), now)?;
|
||||
assert_eq!(cache.async_lookup(&request, now).await?, Some(json!({"answer": 7})));
|
||||
```
|
||||
|
||||
For Redis, inject `RedisCache::new(url, ttl, ResponseCacheCodec)` instead. Namespaces are optional and existing namespace prefixes are preserved
|
||||
|
||||
Callers supply Unix time for response freshness. Backend TTL uses its own clock. A read can reject an entry through `max_age` even while the backend still retains it
|
||||
|
||||
## Python integration
|
||||
|
||||
The bridge activates backends through the Rust catalog in `litellm/rust_bridge/catalog.py`. Every cache rule ships as `PYTHON_ONLY`, so SDK, Router, and proxy calls stay on Python and construct no native cache resources until a rule is changed
|
||||
|
||||
When a rule selects a backend, the Python `Cache` facade builds the native runtime from its own configuration and routes its storage calls (sync and async lookup and store, and pipelined batch store) to it. Stream replay, embedding partial-hit merging, response reconstruction, and callbacks stay in Python on top of that native store. The Python backend object remains for its direct API
|
||||
|
||||
Object responses are written as they are, and every other response shape is written as a serialized string, which is the pair of shapes Python reads. A string on the wire is therefore always a serialized response, so string-valued responses round trip. Typed backends such as memory never pass through the codec
|
||||
|
||||
Native cache handles must be recreated after fork. Native errors propagate to the host, which owns the existing fail-open and logging policy
|
||||
|
||||
## Adding another backend
|
||||
|
||||
Implement `BaseCache` for the backend with its associated value type and the capability traits its Python class supports, and accept a `CacheCodec` when wire serialization is needed. `ResponseCache<B>` then works without another response implementation
|
||||
|
||||
Run the `litellm-cache-testing` contract checks the backend's capabilities allow, and run response fixtures with `ResponseCacheCodec`, including both Python envelope encodings, before adding a catalog rule
|
||||
|
|
@ -4,6 +4,7 @@ mod codec;
|
|||
mod embedding;
|
||||
mod exact;
|
||||
mod response;
|
||||
mod service;
|
||||
|
||||
pub use buffer::WriteBuffer;
|
||||
pub use caching::{
|
||||
|
|
@ -14,3 +15,8 @@ pub use codec::ResponseCacheCodec;
|
|||
pub use embedding::PartialHits;
|
||||
pub use exact::{ConnectionProbe, ExactResponseCache};
|
||||
pub use response::{ResponseCache, ResponseCacheRequest};
|
||||
|
||||
pub use service::{
|
||||
CacheOptions, CacheScope, ResponseCacheConfig, ResponseCacheService, ResponseEnvelope,
|
||||
ScopedCache,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -7,7 +7,9 @@ use litellm_cache::{
|
|||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{CacheControls, CacheEntry, CacheKeyInput, PartialHits, cache_key};
|
||||
use crate::{
|
||||
CacheControls, CacheEntry, CacheKeyInput, PartialHits, ResponseCacheConfig, cache_key,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponseCacheRequest<C: CacheContext = litellm_cache::ExactCacheContext> {
|
||||
|
|
@ -50,6 +52,7 @@ where
|
|||
B::Context: Default + PartialEq,
|
||||
{
|
||||
backend: Arc<B>,
|
||||
config: ResponseCacheConfig,
|
||||
}
|
||||
|
||||
impl<B> ResponseCache<B>
|
||||
|
|
@ -58,7 +61,18 @@ where
|
|||
B::Context: Default + PartialEq,
|
||||
{
|
||||
pub fn new(backend: Arc<B>) -> Self {
|
||||
Self { backend }
|
||||
Self {
|
||||
backend,
|
||||
config: ResponseCacheConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_config(self, config: ResponseCacheConfig) -> Self {
|
||||
Self { config, ..self }
|
||||
}
|
||||
|
||||
pub fn config(&self) -> &ResponseCacheConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
pub fn backend(&self) -> &B {
|
||||
|
|
@ -221,7 +235,7 @@ where
|
|||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
if !request.controls.writes() {
|
||||
if !request.controls.writes() || !self.fits(&response) {
|
||||
return Ok(());
|
||||
}
|
||||
self.backend.set_cache(
|
||||
|
|
@ -240,7 +254,7 @@ where
|
|||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
if !request.controls.writes() {
|
||||
if !request.controls.writes() || !self.fits(&response) {
|
||||
return Ok(());
|
||||
}
|
||||
self.backend
|
||||
|
|
@ -277,7 +291,7 @@ where
|
|||
) -> Result<(), Error> {
|
||||
let writable = entries
|
||||
.into_iter()
|
||||
.filter(|(request, _, _)| request.controls.writes())
|
||||
.filter(|(request, response, _)| request.controls.writes() && self.fits(response))
|
||||
.map(|(request, response, now)| {
|
||||
(
|
||||
cache_key(&request.key),
|
||||
|
|
@ -312,6 +326,11 @@ where
|
|||
Ok(())
|
||||
}
|
||||
|
||||
fn fits(&self, response: &Value) -> bool {
|
||||
self.config.max_entry_bytes == usize::MAX
|
||||
|| response.to_string().len() <= self.config.max_entry_bytes
|
||||
}
|
||||
|
||||
fn partial_hits(
|
||||
requests: &[ResponseCacheRequest<B::Context>],
|
||||
readable: Vec<(usize, &ResponseCacheRequest<B::Context>)>,
|
||||
|
|
|
|||
177
litellm-rust/crates/cache-response/src/service.rs
Normal file
177
litellm-rust/crates/cache-response/src/service.rs
Normal file
|
|
@ -0,0 +1,177 @@
|
|||
use std::{future::Future, pin::Pin, time::Duration};
|
||||
|
||||
use litellm_cache::{BaseCache, Error, ExactCacheContext};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
CacheControls, CacheEntry, CacheKeyField, CacheKeyInput, ResponseCache, ResponseCacheRequest,
|
||||
};
|
||||
|
||||
type CacheFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponseCacheConfig {
|
||||
pub namespace: String,
|
||||
pub max_entry_bytes: usize,
|
||||
}
|
||||
|
||||
impl Default for ResponseCacheConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
namespace: String::new(),
|
||||
max_entry_bytes: usize::MAX,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait ResponseCacheService: Send + Sync {
|
||||
fn config(&self) -> &ResponseCacheConfig;
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, Option<Value>>;
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, ()>;
|
||||
}
|
||||
|
||||
impl<B> ResponseCacheService for ResponseCache<B>
|
||||
where
|
||||
B: BaseCache<Value = CacheEntry, Context = ExactCacheContext>,
|
||||
{
|
||||
fn config(&self) -> &ResponseCacheConfig {
|
||||
self.config()
|
||||
}
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, Option<Value>> {
|
||||
Box::pin(self.async_lookup(request, now))
|
||||
}
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, ()> {
|
||||
Box::pin(self.async_store(request, response, now))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum CacheScope {
|
||||
Shared,
|
||||
Isolated(String),
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CacheOptions {
|
||||
pub caching: Option<bool>,
|
||||
pub no_cache: bool,
|
||||
pub no_store: bool,
|
||||
pub ttl: Option<Duration>,
|
||||
pub max_age: Option<Duration>,
|
||||
pub scope: CacheScope,
|
||||
}
|
||||
|
||||
impl CacheOptions {
|
||||
pub fn new(scope: CacheScope) -> Self {
|
||||
Self {
|
||||
caching: None,
|
||||
no_cache: false,
|
||||
no_store: false,
|
||||
ttl: None,
|
||||
max_age: None,
|
||||
scope,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn enabled(&self) -> bool {
|
||||
self.caching != Some(false) && !(self.no_cache && self.no_store)
|
||||
}
|
||||
|
||||
pub fn request(self, namespace: &str, surface: &str, mut input: Value) -> ResponseCacheRequest {
|
||||
input.sort_all_objects();
|
||||
let scope = match self.scope {
|
||||
CacheScope::Shared => String::new(),
|
||||
CacheScope::Isolated(scope) => serde_json::json!(["isolated", scope]).to_string(),
|
||||
};
|
||||
ResponseCacheRequest {
|
||||
key: CacheKeyInput {
|
||||
namespace: Some(format!("{namespace}:inference-v2")),
|
||||
fields: [
|
||||
("surface", surface.to_owned()),
|
||||
("scope", scope),
|
||||
("request", input.to_string()),
|
||||
]
|
||||
.into_iter()
|
||||
.map(|(name, value)| CacheKeyField {
|
||||
name: name.into(),
|
||||
value: Some(value),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
})
|
||||
.collect(),
|
||||
..Default::default()
|
||||
},
|
||||
controls: CacheControls {
|
||||
configured: true,
|
||||
supported_call_type: true,
|
||||
native_backend: true,
|
||||
default_on: true,
|
||||
caching: self.caching,
|
||||
no_cache: self.no_cache,
|
||||
no_store: self.no_store,
|
||||
..Default::default()
|
||||
},
|
||||
context: ExactCacheContext { ttl: self.ttl },
|
||||
max_age: self.max_age,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize, serde::Deserialize)]
|
||||
pub struct ResponseEnvelope<T> {
|
||||
version: u32,
|
||||
surface: String,
|
||||
output: T,
|
||||
}
|
||||
|
||||
impl<T> ResponseEnvelope<T> {
|
||||
pub fn new(surface: &str, output: T) -> Self {
|
||||
Self {
|
||||
version: 1,
|
||||
surface: surface.into(),
|
||||
output,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decode(self, surface: &str) -> Option<T> {
|
||||
(self.version == 1 && self.surface == surface).then_some(self.output)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ScopedCache {
|
||||
pub service: std::sync::Arc<dyn ResponseCacheService>,
|
||||
pub scope: CacheScope,
|
||||
}
|
||||
|
||||
impl ScopedCache {
|
||||
pub fn new(service: std::sync::Arc<dyn ResponseCacheService>, scope: CacheScope) -> Self {
|
||||
Self { service, scope }
|
||||
}
|
||||
|
||||
pub fn options(&self, overrides: Option<CacheOptions>) -> CacheOptions {
|
||||
overrides.unwrap_or_else(|| CacheOptions::new(self.scope.clone()))
|
||||
}
|
||||
}
|
||||
|
|
@ -18,7 +18,7 @@ use litellm_cache_response::{
|
|||
WriteBuffer, cache_key,
|
||||
};
|
||||
use redis_test::MockCmd;
|
||||
use rstest::rstest;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
use support::{keyed, memory, redis, request};
|
||||
|
||||
|
|
@ -648,3 +648,129 @@ async fn write_buffer_clear_drops_pending_entries(memory: Memory, request: Respo
|
|||
assert_eq!(memory.lookup(&request, now).unwrap(), None);
|
||||
assert_eq!(memory.lookup(&other, now).unwrap(), None);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::python_sync("{'timestamp': 100.0, 'response': '{\"answer\": 7}'}")]
|
||||
#[case::python_async(r#"{"timestamp":100.0,"response":{"answer":7}}"#)]
|
||||
#[case::bare_response(r#"{"answer":7}"#)]
|
||||
#[tokio::test]
|
||||
async fn gcs_reads_python_entries_and_writes_python_compatible_envelopes(
|
||||
#[case] encoded: &str,
|
||||
#[values(false, true)] asynchronous: bool,
|
||||
#[future(awt)] gcs: (wiremock::MockServer, Gcs),
|
||||
) {
|
||||
use wiremock::{
|
||||
Mock, ResponseTemplate,
|
||||
matchers::{body_json, header, method, path, query_param},
|
||||
};
|
||||
|
||||
let (server, cache) = gcs;
|
||||
let response = json!({"answer": 7});
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fpython"))
|
||||
.and(query_param("alt", "media"))
|
||||
.and(header("authorization", "Bearer token"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string(encoded))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.and(query_param("uploadType", "media"))
|
||||
.and(query_param("name", "cache/native"))
|
||||
.and(header("authorization", "Bearer token"))
|
||||
.and(header("content-type", "application/json"))
|
||||
.and(body_json(json!({"timestamp": 102.0, "response": response})))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let lookup = if asynchronous {
|
||||
cache
|
||||
.async_lookup(&keyed("python"), Duration::from_secs(102))
|
||||
.await
|
||||
} else {
|
||||
cache.lookup(&keyed("python"), Duration::from_secs(102))
|
||||
};
|
||||
assert_eq!(lookup.unwrap(), Some(response.clone()));
|
||||
let request = ResponseCacheRequest {
|
||||
context: litellm_cache::ExactCacheContext {
|
||||
ttl: Some(Duration::from_secs(12)),
|
||||
},
|
||||
..keyed("native")
|
||||
};
|
||||
let stored = if asynchronous {
|
||||
cache
|
||||
.async_store(&request, response, Duration::from_secs(102))
|
||||
.await
|
||||
} else {
|
||||
cache.store(&request, response, Duration::from_secs(102))
|
||||
};
|
||||
assert_eq!(stored, Ok(()));
|
||||
let requests = server.received_requests().await.unwrap();
|
||||
let upload = requests
|
||||
.iter()
|
||||
.find(|request| request.method.as_str() == "POST")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
upload.url.query(),
|
||||
Some("uploadType=media&name=cache%2Fnative")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn gcs_batch_reads_preserve_order_and_treat_invalid_entries_as_misses(
|
||||
#[future(awt)] gcs: (wiremock::MockServer, Gcs),
|
||||
) {
|
||||
use wiremock::{
|
||||
Mock, ResponseTemplate,
|
||||
matchers::{method, path},
|
||||
};
|
||||
|
||||
let (server, cache) = gcs;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fhit"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(json!({"timestamp": 100.0, "response": {"answer":7}})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Finvalid"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string("not an entry"))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fmissing"))
|
||||
.respond_with(ResponseTemplate::new(404))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let requests = [keyed("hit"), keyed("missing"), keyed("invalid")];
|
||||
let partial = cache
|
||||
.async_lookup_batch(&requests, Duration::from_secs(102))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(partial.values, vec![Some(json!({"answer":7})), None, None]);
|
||||
assert_eq!(partial.missing_indices, vec![1, 2]);
|
||||
}
|
||||
|
||||
type Gcs = ResponseCache<litellm_cache_gcs::GcsCache<litellm_cache_response::ResponseCacheCodec>>;
|
||||
|
||||
#[fixture]
|
||||
async fn gcs() -> (wiremock::MockServer, Gcs) {
|
||||
let server = wiremock::MockServer::start().await;
|
||||
let cache = ResponseCache::new(Arc::new(litellm_cache_gcs::GcsCache::with_token_source(
|
||||
litellm_cache_gcs::GcsConfig {
|
||||
bucket_name: "bucket".into(),
|
||||
gcs_path: Some("cache".into()),
|
||||
path_service_account: None,
|
||||
endpoint: server.uri(),
|
||||
},
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_cache_response::ResponseCacheCodec,
|
||||
Arc::new(litellm_cache_gcs::StaticTokenSource("token".into())),
|
||||
)));
|
||||
(server, cache)
|
||||
}
|
||||
|
|
|
|||
176
litellm-rust/crates/cache-response/tests/service.rs
Normal file
176
litellm-rust/crates/cache-response/tests/service.rs
Normal file
|
|
@ -0,0 +1,176 @@
|
|||
use std::{
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicU64, Ordering},
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use litellm_cache::ExactCacheContext;
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{
|
||||
CacheEntry, CacheKeyInput, ResponseCache, ResponseCacheConfig, ResponseCacheRequest,
|
||||
ResponseCacheService,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn service_honors_per_call_expiry_and_freshness() {
|
||||
let clock = Arc::new(AtomicU64::new(0));
|
||||
let cache_clock = clock.clone();
|
||||
let cache: Arc<dyn ResponseCacheService> = Arc::new(ResponseCache::new(Arc::new(
|
||||
InMemoryCache::with_clock(Some(100), Some(Duration::from_secs(60)), move || {
|
||||
Duration::from_secs(cache_clock.load(Ordering::SeqCst))
|
||||
}),
|
||||
)));
|
||||
let request = ResponseCacheRequest {
|
||||
context: ExactCacheContext {
|
||||
ttl: Some(Duration::from_secs(5)),
|
||||
},
|
||||
..ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some("entry".into()),
|
||||
..Default::default()
|
||||
})
|
||||
};
|
||||
cache
|
||||
.store(&request, json!({"answer":7}), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
cache.lookup(&request, Duration::ZERO).await.unwrap(),
|
||||
Some(json!({"answer":7}))
|
||||
);
|
||||
let stale_request = ResponseCacheRequest {
|
||||
max_age: Some(Duration::from_secs(1)),
|
||||
..request.clone()
|
||||
};
|
||||
clock.store(2, Ordering::SeqCst);
|
||||
assert_eq!(
|
||||
cache
|
||||
.lookup(&stale_request, Duration::from_secs(2))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
assert!(
|
||||
cache
|
||||
.lookup(&request, Duration::from_secs(2))
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some()
|
||||
);
|
||||
clock.store(6, Ordering::SeqCst);
|
||||
assert_eq!(
|
||||
cache
|
||||
.lookup(&request, Duration::from_secs(6))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn entry_limit_applies_to_sync_async_and_batch_writes() {
|
||||
let storage = Arc::new(InMemoryCache::<CacheEntry>::default());
|
||||
let cache = ResponseCache::new(storage.clone()).with_config(ResponseCacheConfig {
|
||||
namespace: "service-test".into(),
|
||||
max_entry_bytes: json!({"answer":7}).to_string().len(),
|
||||
});
|
||||
let small = json!({"answer":7});
|
||||
let large = json!({"answer":"too large"});
|
||||
let request = |key: &str| {
|
||||
ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some(key.into()),
|
||||
..Default::default()
|
||||
})
|
||||
};
|
||||
cache
|
||||
.store(&request("sync"), large.clone(), Duration::ZERO)
|
||||
.unwrap();
|
||||
cache
|
||||
.async_store(&request("async"), large.clone(), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
cache
|
||||
.async_store_batch(
|
||||
vec![
|
||||
(request("batch-large"), large),
|
||||
(request("batch-small"), small.clone()),
|
||||
],
|
||||
Duration::ZERO,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let service: Arc<dyn ResponseCacheService> = Arc::new(cache);
|
||||
service
|
||||
.store(&request("service"), small.clone(), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
for key in ["sync", "async", "batch-large"] {
|
||||
assert!(storage.get_cache(key).unwrap().is_none());
|
||||
}
|
||||
for key in ["batch-small", "service"] {
|
||||
assert_eq!(
|
||||
service.lookup(&request(key), Duration::ZERO).await.unwrap(),
|
||||
Some(small.clone())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::same_scope("tenant-a", "tenant-a", true)]
|
||||
#[case::different_scope("tenant-a", "tenant-b", false)]
|
||||
#[case::empty_isolated_scope("", "", true)]
|
||||
#[tokio::test]
|
||||
async fn isolated_policy_controls_actual_entry_reuse(
|
||||
#[case] first: &str,
|
||||
#[case] second: &str,
|
||||
#[case] hit: bool,
|
||||
) {
|
||||
use litellm_cache_response::{CacheOptions, CacheScope};
|
||||
let service = ResponseCache::new(Arc::new(InMemoryCache::<CacheEntry>::default()));
|
||||
let request =
|
||||
|scope| CacheOptions::new(scope).request("test", "messages", json!({"prompt":"hello"}));
|
||||
service
|
||||
.async_store(
|
||||
&request(CacheScope::Isolated(first.into())),
|
||||
json!({"answer":7}),
|
||||
Duration::ZERO,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
service
|
||||
.async_lookup(
|
||||
&request(CacheScope::Isolated(second.into())),
|
||||
Duration::ZERO
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
hit.then(|| json!({"answer":7}))
|
||||
);
|
||||
assert_eq!(
|
||||
service
|
||||
.async_lookup(&request(CacheScope::Shared), Duration::ZERO)
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::valid(1, "messages", Some(7))]
|
||||
#[case::unknown_version(2, "messages", None)]
|
||||
#[case::another_surface(1, "responses", None)]
|
||||
fn envelopes_require_a_matching_surface_and_version(
|
||||
#[case] version: u32,
|
||||
#[case] surface: &str,
|
||||
#[case] expected: Option<u32>,
|
||||
) {
|
||||
let envelope: litellm_cache_response::ResponseEnvelope<u32> =
|
||||
serde_json::from_value(json!({"version":version,"surface":surface,"output":7})).unwrap();
|
||||
assert_eq!(envelope.decode("messages"), expected);
|
||||
}
|
||||
|
|
@ -56,6 +56,7 @@ pub struct LegacyLogging {
|
|||
stream: Option<DeliveredStream>,
|
||||
asynchronous: bool,
|
||||
internal: bool,
|
||||
cache_key: Option<String>,
|
||||
}
|
||||
|
||||
fn datetime(py: Python<'_>, epoch_seconds: f64) -> PyResult<Py<PyAny>> {
|
||||
|
|
@ -80,6 +81,7 @@ impl LegacyLogging {
|
|||
stream: None,
|
||||
asynchronous,
|
||||
internal: false,
|
||||
cache_key: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -207,10 +209,10 @@ impl LegacyLogging {
|
|||
logger.object(py),
|
||||
billing.url_route,
|
||||
billing.endpoint_type,
|
||||
&self
|
||||
.request
|
||||
.as_ref()
|
||||
.map(|request| request.body.clone_ref(py)),
|
||||
&self.request.as_ref().map_or_else(
|
||||
|| self.call.kwargs().clone_ref(py),
|
||||
|request| request.body.clone_ref(py),
|
||||
),
|
||||
&stream.chunks,
|
||||
&self.start,
|
||||
&self.end,
|
||||
|
|
@ -246,10 +248,10 @@ impl LegacyLogging {
|
|||
(
|
||||
logger.object(py),
|
||||
billing.endpoint_type,
|
||||
&self
|
||||
.request
|
||||
.as_ref()
|
||||
.map(|request| request.body.clone_ref(py)),
|
||||
&self.request.as_ref().map_or_else(
|
||||
|| self.call.kwargs().clone_ref(py),
|
||||
|request| request.body.clone_ref(py),
|
||||
),
|
||||
&stream.chunks,
|
||||
error,
|
||||
),
|
||||
|
|
@ -428,6 +430,40 @@ impl LegacyLogging {
|
|||
self.finalize(py)
|
||||
}
|
||||
|
||||
pub(crate) fn result_ready(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
facts: &litellm_host::interceptors::ExecutionFacts,
|
||||
) -> PyResult<HookStep<Self, ()>> {
|
||||
use litellm_host::interceptors::ResultSource;
|
||||
|
||||
let logger = self.logger()?.object(py);
|
||||
let params = logger
|
||||
.getattr("litellm_params")?
|
||||
.cast_into::<PyDict>()?
|
||||
.copy()?;
|
||||
params.set_item("custom_llm_provider", &facts.provider.provider)?;
|
||||
crate::python::Logging::Update.call(
|
||||
py,
|
||||
(
|
||||
&logger,
|
||||
self.call.kwargs(),
|
||||
&facts.provider.model,
|
||||
logger.getattr("optional_params")?,
|
||||
params,
|
||||
&facts.provider.provider,
|
||||
),
|
||||
)?;
|
||||
let details = logger.getattr("model_call_details")?;
|
||||
self.cache_key = match &facts.source {
|
||||
ResultSource::Provider => None,
|
||||
ResultSource::Cache { key } => Some(key.clone()),
|
||||
};
|
||||
details.set_item("cache_hit", self.cache_key.is_some())?;
|
||||
details.set_item("cache_key", self.cache_key.as_deref())?;
|
||||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
|
||||
pub(crate) fn post_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
|
|
@ -485,10 +521,14 @@ impl LegacyLogging {
|
|||
self.dispatch_failure(py)
|
||||
}
|
||||
|
||||
pub(crate) fn stream_opened(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
pub(crate) fn stream_opened(&mut self, py: Python<'_>, head: &Py<PyAny>) -> PyResult<()> {
|
||||
if self.stream_billing().is_none() {
|
||||
return Err(missing_state());
|
||||
}
|
||||
if let Some(key) = &self.cache_key {
|
||||
head.bind(py).set_item("cache_key", key)?;
|
||||
head.bind(py).set_item("cache_hit", true)?;
|
||||
}
|
||||
Streaming::Opened.call(py, (self.logger()?.object(py),))?;
|
||||
self.stream = Some(DeliveredStream {
|
||||
chunks: PyList::empty(py).unbind(),
|
||||
|
|
@ -1726,7 +1766,9 @@ assert logger.calls[1][1] is response
|
|||
operation: litellm_types::Operation::Messages,
|
||||
..logged(py, &locals, true)
|
||||
};
|
||||
logging.on_stream_open(py).unwrap();
|
||||
logging
|
||||
.on_stream_open(py, &pyo3::types::PyDict::new(py).into_any().unbind())
|
||||
.unwrap();
|
||||
logging
|
||||
.on_stream_chunk(py, &local(&locals, "first").unbind())
|
||||
.unwrap();
|
||||
|
|
|
|||
|
|
@ -56,7 +56,7 @@ type After = fn(&mut LegacyLogging, Python<'_>, &RawResponse) -> Step<()>;
|
|||
type Transform = fn(&mut LegacyLogging, Python<'_>, Py<PyAny>, Timing) -> Step<Py<PyAny>>;
|
||||
type Success = fn(&mut LegacyLogging, Python<'_>, Timing, &Py<PyAny>) -> Step<()>;
|
||||
type Failure = fn(&mut LegacyLogging, Python<'_>, Timing, FailureOrigin, &PyErr) -> Step<()>;
|
||||
type Open = fn(&mut LegacyLogging, Python<'_>) -> PyResult<()>;
|
||||
type Open = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
|
||||
type Chunk = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
|
||||
|
||||
const PREPARE: Binding<Prepare> = Binding {
|
||||
|
|
@ -173,6 +173,9 @@ impl CallHooks<PythonRuntime> for LegacyLogging {
|
|||
PythonCallEvent::Started { .. } | PythonCallEvent::Cancelled { .. } => {
|
||||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
PythonCallEvent::Execution(ExecutionEvent::ResultReady { facts }) => {
|
||||
self.result_ready(py, &facts)
|
||||
}
|
||||
PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
(AFTER.invoke)(self, py, raw)
|
||||
}
|
||||
|
|
@ -187,8 +190,8 @@ impl CallHooks<PythonRuntime> for LegacyLogging {
|
|||
}
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
(OPEN.invoke)(self, py)
|
||||
fn on_stream_open(&mut self, py: Python<'_>, head: &Py<PyAny>) -> PyResult<()> {
|
||||
(OPEN.invoke)(self, py, head)
|
||||
}
|
||||
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
|
|
|
|||
|
|
@ -33,3 +33,15 @@ Scope follows the concept, not the first caller. An error type under `litellm-ll
|
|||
`litellm_llms::Error` (`crates/llms/src/error.rs`) is the one transformation error for every provider and API. `base_llm/ocr/error.rs` is the recorded exception until OCR folds into it
|
||||
|
||||
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback execution of any kind. Core runs each route as a machine that yields host operations and call events; which integrations consume those events is the host's business.
|
||||
|
||||
## Response caching and accounting boundary
|
||||
|
||||
Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries per-call cache overrides and observation; attaching a service does not change the execution contract
|
||||
|
||||
Core owns request identity, typed response reconstruction and stream capture/replay. `cache-response` owns cache policy, namespacing, scope encoding, versioned envelopes and freshness. The SDK explicitly chooses shared scope. The gateway derives isolated scope from authenticated identity before attaching its service
|
||||
|
||||
Core delivers `ExecutionFacts` through the awaited `ResultReady` host operation for both provider and cached results, before public response processing or stream opening. Facts carry resolved model/provider and result source, including the hit key. Usage remains in the typed response or delivered stream, where completion and cancellation determine what was actually reported. Passive observation is not an accounting delivery mechanism
|
||||
|
||||
Core does not calculate prices, charge budgets, or update rate-limit counters. The legacy Python callback adapter translates execution facts into the existing Python logging contract; Python remains the accounting owner on that path. Native gateway accounting belongs to gateway dependencies, independently of `host-python`. Response-cache services expose no coordination counters or reservation APIs. A shared Redis deployment does not make response storage and accounting coordination the same dependency
|
||||
|
||||
Cache lookup follows provider preparation, credential resolution and the request interceptor. Keys describe the effective provider URL, authenticated headers and rewritten body. Signed requests bypass caching until the signing identity has a stable cache representation
|
||||
|
|
|
|||
|
|
@ -6,6 +6,10 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-cache.workspace = true
|
||||
litellm-cache-response.workspace = true
|
||||
litellm-framing.workspace = true
|
||||
tokio-util = { version = "0.7", features = ["codec"] }
|
||||
litellm-secrets.workspace = true
|
||||
litellm-types.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
|
|
@ -36,6 +40,7 @@ url.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-cache-memory.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-host-native.workspace = true
|
||||
|
|
|
|||
318
litellm-rust/crates/core/src/caching.rs
Normal file
318
litellm-rust/crates/core/src/caching.rs
Normal file
|
|
@ -0,0 +1,318 @@
|
|||
use std::{
|
||||
future::Future,
|
||||
sync::Arc,
|
||||
time::{Duration, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use litellm_cache_response::{
|
||||
CacheOptions, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope, cache_key,
|
||||
};
|
||||
use litellm_host::{
|
||||
call::{CallOutput, OutputOf},
|
||||
interceptors::{ExecutionFacts, Interceptors, ProviderIdentity, ResultSource, WireRequest},
|
||||
lifecycle::{CallEvent, ExecutionEvent},
|
||||
observation::ObservationSender,
|
||||
protocol::Protocol,
|
||||
};
|
||||
use serde::{Deserialize, Serialize, de::DeserializeOwned};
|
||||
use serde_json::Value;
|
||||
use tokio_util::codec::Decoder;
|
||||
|
||||
use crate::RouteError;
|
||||
|
||||
pub trait Cachable: Protocol<Error = RouteError> {
|
||||
const SURFACE: &'static str;
|
||||
|
||||
fn reusable(_response: &Self::Response) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
pub struct CacheRequest {
|
||||
pub identity: ProviderIdentity,
|
||||
pub input: Value,
|
||||
}
|
||||
|
||||
impl CacheRequest {
|
||||
pub fn from_wire(identity: ProviderIdentity, wire: Option<&WireRequest>) -> Self {
|
||||
Self {
|
||||
input: wire.map_or(Value::Null, |wire| {
|
||||
serde_json::json!({
|
||||
"provider": identity.provider,
|
||||
"model": identity.model,
|
||||
"url": wire.url,
|
||||
"headers": wire.headers,
|
||||
"body": wire.body,
|
||||
})
|
||||
}),
|
||||
identity,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait StreamCachable: Cachable {
|
||||
const TERMINAL_EVENT: &'static str;
|
||||
|
||||
fn replay(data: Bytes) -> Option<OutputOf<Self>>;
|
||||
fn bytes(chunk: &Self::Chunk) -> &[u8];
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", content = "value")]
|
||||
pub enum CachedOutput<R> {
|
||||
Response(R),
|
||||
Stream(String),
|
||||
}
|
||||
|
||||
struct CacheSession {
|
||||
service: Arc<dyn ResponseCacheService>,
|
||||
request: ResponseCacheRequest,
|
||||
}
|
||||
|
||||
impl CacheSession {
|
||||
fn prepare<P: Cachable>(
|
||||
service: Option<Arc<dyn ResponseCacheService>>,
|
||||
options: Option<CacheOptions>,
|
||||
request: &CacheRequest,
|
||||
) -> Option<Self> {
|
||||
let options = options.filter(CacheOptions::enabled)?;
|
||||
let service = service?;
|
||||
let input = request.input.clone();
|
||||
let request = options.request(&service.config().namespace, P::SURFACE, input);
|
||||
Some(Self { service, request })
|
||||
}
|
||||
|
||||
async fn lookup<P: Cachable>(&self) -> Option<CachedOutput<P::Response>>
|
||||
where
|
||||
P::Response: DeserializeOwned,
|
||||
{
|
||||
if !self.request.controls.reads() {
|
||||
return None;
|
||||
}
|
||||
match self.service.lookup(&self.request, now()).await {
|
||||
Ok(Some(value)) => {
|
||||
serde_json::from_value::<ResponseEnvelope<CachedOutput<P::Response>>>(value)
|
||||
.ok()
|
||||
.and_then(|entry| entry.decode(P::SURFACE))
|
||||
}
|
||||
Ok(None) => None,
|
||||
Err(_) => {
|
||||
tracing::warn!("response cache lookup failed");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn store(&self, entry: Value) {
|
||||
if !self.request.controls.writes() {
|
||||
return;
|
||||
}
|
||||
if self
|
||||
.service
|
||||
.store(&self.request, entry, now())
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
tracing::warn!("response cache write failed");
|
||||
}
|
||||
}
|
||||
|
||||
async fn store_response<P: Cachable>(&self, response: &P::Response)
|
||||
where
|
||||
P::Response: Serialize,
|
||||
{
|
||||
if !self.request.controls.writes() || !P::reusable(response) {
|
||||
return;
|
||||
}
|
||||
if let Ok(value) = serde_json::to_value(response)
|
||||
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
|
||||
P::SURFACE,
|
||||
CachedOutput::Response(value),
|
||||
))
|
||||
{
|
||||
self.store(entry).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute_unary<P, F, Fut>(
|
||||
request: CacheRequest,
|
||||
cache: Option<Arc<dyn ResponseCacheService>>,
|
||||
options: Option<CacheOptions>,
|
||||
interceptors: &impl Interceptors<RouteError>,
|
||||
observers: Option<&ObservationSender>,
|
||||
provider: F,
|
||||
) -> Result<P::Response, RouteError>
|
||||
where
|
||||
P: Cachable,
|
||||
P::Response: Serialize + DeserializeOwned,
|
||||
F: FnOnce() -> Fut,
|
||||
Fut: Future<Output = Result<P::Response, RouteError>>,
|
||||
{
|
||||
let identity = request.identity.clone();
|
||||
crate::diagnostic::provider(&identity.model, &identity.provider);
|
||||
let session = CacheSession::prepare::<P>(cache, options, &request);
|
||||
let hit = match &session {
|
||||
Some(session) => session.lookup::<P>().await.and_then(|entry| match entry {
|
||||
CachedOutput::Response(response) => Some((response, cache_key(&session.request.key))),
|
||||
CachedOutput::Stream(_) => None,
|
||||
}),
|
||||
None => None,
|
||||
};
|
||||
let (response, source) = match hit {
|
||||
Some((response, key)) => (response, ResultSource::Cache { key }),
|
||||
None => (provider().await?, ResultSource::Provider),
|
||||
};
|
||||
let from_provider = source == ResultSource::Provider;
|
||||
publish(
|
||||
ExecutionFacts {
|
||||
provider: identity,
|
||||
source,
|
||||
},
|
||||
interceptors,
|
||||
observers,
|
||||
)
|
||||
.await?;
|
||||
if from_provider && let Some(session) = session {
|
||||
session.store_response::<P>(&response).await;
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
pub async fn execute_streaming<P, F, Fut>(
|
||||
request: CacheRequest,
|
||||
cache: Option<Arc<dyn ResponseCacheService>>,
|
||||
options: Option<CacheOptions>,
|
||||
interceptors: &impl Interceptors<RouteError>,
|
||||
observers: Option<&ObservationSender>,
|
||||
provider: F,
|
||||
) -> Result<OutputOf<P>, RouteError>
|
||||
where
|
||||
P: StreamCachable,
|
||||
P::Response: Serialize + DeserializeOwned,
|
||||
F: FnOnce() -> Fut,
|
||||
Fut: Future<Output = Result<OutputOf<P>, RouteError>>,
|
||||
{
|
||||
let identity = request.identity.clone();
|
||||
crate::diagnostic::provider(&identity.model, &identity.provider);
|
||||
let session = CacheSession::prepare::<P>(cache, options, &request);
|
||||
let hit = match &session {
|
||||
Some(session) => session.lookup::<P>().await.and_then(|entry| {
|
||||
let output = match entry {
|
||||
CachedOutput::Response(response) => Some(CallOutput::Complete(response)),
|
||||
CachedOutput::Stream(data) => P::replay(Bytes::from(data)),
|
||||
};
|
||||
output.map(|output| (output, cache_key(&session.request.key)))
|
||||
}),
|
||||
None => None,
|
||||
};
|
||||
let (output, source) = match hit {
|
||||
Some((output, key)) => (output, ResultSource::Cache { key }),
|
||||
None => (provider().await?, ResultSource::Provider),
|
||||
};
|
||||
let from_provider = source == ResultSource::Provider;
|
||||
publish(
|
||||
ExecutionFacts {
|
||||
provider: identity,
|
||||
source,
|
||||
},
|
||||
interceptors,
|
||||
observers,
|
||||
)
|
||||
.await?;
|
||||
let Some(session) =
|
||||
session.filter(|session| from_provider && session.request.controls.writes())
|
||||
else {
|
||||
return Ok(output);
|
||||
};
|
||||
match output {
|
||||
CallOutput::Complete(response) => {
|
||||
session.store_response::<P>(&response).await;
|
||||
Ok(CallOutput::Complete(response))
|
||||
}
|
||||
CallOutput::Stream { head, chunks } => {
|
||||
let captured = stream::try_unfold(
|
||||
(chunks, Some(Vec::<u8>::new()), session),
|
||||
|(mut chunks, captured, session)| async move {
|
||||
match chunks.try_next().await? {
|
||||
Some(chunk) => {
|
||||
let captured = captured.and_then(|mut data| {
|
||||
let bytes = P::bytes(&chunk);
|
||||
if data.len().saturating_add(bytes.len())
|
||||
> session.service.config().max_entry_bytes
|
||||
{
|
||||
return None;
|
||||
}
|
||||
data.extend_from_slice(bytes);
|
||||
Some(data)
|
||||
});
|
||||
Ok(Some((chunk, (chunks, captured, session))))
|
||||
}
|
||||
None => {
|
||||
if let Some(data) = captured
|
||||
&& let Ok(text) = String::from_utf8(data)
|
||||
&& successful_stream(&text, P::TERMINAL_EVENT)
|
||||
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
|
||||
P::SURFACE,
|
||||
CachedOutput::<Value>::Stream(text),
|
||||
))
|
||||
{
|
||||
session.store(entry).await;
|
||||
}
|
||||
Ok::<_, RouteError>(None)
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
.boxed();
|
||||
Ok(CallOutput::Stream {
|
||||
head,
|
||||
chunks: captured,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn now() -> Duration {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn successful_stream(text: &str, terminal: &str) -> bool {
|
||||
let mut pending = BytesMut::from(text.as_bytes());
|
||||
let mut codec = litellm_framing::sse::SseCodec::default();
|
||||
let mut complete = false;
|
||||
loop {
|
||||
let event = match codec.decode(&mut pending) {
|
||||
Ok(Some(event)) => event,
|
||||
Ok(None) => return complete && pending.is_empty(),
|
||||
Err(_) => return false,
|
||||
};
|
||||
let Ok(value) = serde_json::from_str::<Value>(&event.data) else {
|
||||
return false;
|
||||
};
|
||||
let Some(kind) = value.get("type").and_then(Value::as_str) else {
|
||||
return false;
|
||||
};
|
||||
if matches!(kind, "error" | "response.failed" | "response.incomplete") {
|
||||
return false;
|
||||
}
|
||||
complete |= kind == terminal;
|
||||
}
|
||||
}
|
||||
|
||||
async fn publish(
|
||||
facts: ExecutionFacts,
|
||||
interceptors: &impl Interceptors<RouteError>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<(), RouteError> {
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: facts.clone(),
|
||||
}));
|
||||
}
|
||||
interceptors.result_ready(facts).await
|
||||
}
|
||||
|
|
@ -22,6 +22,8 @@ pub(super) async fn execute(
|
|||
http: &Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderChatCompletionsRequest,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
|
|
@ -45,6 +47,10 @@ pub(super) async fn execute(
|
|||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?;
|
||||
let identity = litellm_host::interceptors::ProviderIdentity {
|
||||
model: context.model.clone(),
|
||||
provider: context.custom_llm_provider.clone(),
|
||||
};
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
|
|
@ -55,59 +61,72 @@ pub(super) async fn execute(
|
|||
context,
|
||||
)
|
||||
.await?;
|
||||
let outbound = outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
let cache = cache.filter(|_| authenticated.signer.is_none());
|
||||
let cache_request =
|
||||
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
|
||||
crate::caching::execute_unary::<super::route::ChatCompletions, _, _>(
|
||||
cache_request,
|
||||
cache.as_ref().map(|cache| cache.service.clone()),
|
||||
cache.as_ref().map(|cache| cache.options(cache_options)),
|
||||
interceptors,
|
||||
observers,
|
||||
|| async move {
|
||||
let outbound = outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)?;
|
||||
|
||||
let response = crate::outbound::send(outbound, http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
if err.is_connect() || err.is_builder() {
|
||||
Error::Transport(litellm_http::transport::Error::Connect(err.to_string()))
|
||||
} else {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
}
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response.text().await.map_err(|err| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
})?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
let raw = RawResponse { body: text.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
.map_err(Error::from)
|
||||
.map_err(as_response_error)
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)?;
|
||||
|
||||
let response = crate::outbound::send(outbound, http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
if err.is_connect() || err.is_builder() {
|
||||
Error::Transport(litellm_http::transport::Error::Connect(err.to_string()))
|
||||
} else {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
}
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response.text().await.map_err(|err| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
})?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
let raw = RawResponse { body: text.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
.map_err(Error::from)
|
||||
.map_err(as_response_error)
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Re-tag an error raised while normalizing a response the provider already
|
||||
|
|
@ -232,6 +251,8 @@ mod tests {
|
|||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
None,
|
||||
None,
|
||||
&interceptors,
|
||||
None,
|
||||
)
|
||||
|
|
@ -271,6 +292,8 @@ mod tests {
|
|||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
None,
|
||||
None,
|
||||
&interceptors,
|
||||
None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ pub struct ChatCompletionsRoute {
|
|||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
}
|
||||
|
||||
impl ChatCompletionsRoute {
|
||||
|
|
@ -30,6 +31,14 @@ impl ChatCompletionsRoute {
|
|||
http,
|
||||
auth,
|
||||
secrets,
|
||||
cache: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -37,49 +46,48 @@ impl ChatCompletionsRoute {
|
|||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
litellm_host::lifecycle::observe_unary(
|
||||
observers.clone(),
|
||||
self.run(request, interceptors, observers.as_ref()),
|
||||
self.run_call(
|
||||
request.into(),
|
||||
cache_options,
|
||||
interceptors,
|
||||
observers.as_ref(),
|
||||
),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "chat_completions",
|
||||
model = %request.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let resolved = resolve_request(request)?;
|
||||
let snapshot = self
|
||||
.secrets
|
||||
.resolve(&resolved.config.secret_names())
|
||||
.await?;
|
||||
let prepared = prepare_provider_request(resolved, snapshot)?;
|
||||
crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<
|
||||
'_,
|
||||
Result<ChatCompletionsResponse, Error>,
|
||||
> = Box::pin(handler::execute(
|
||||
let resolved = resolve_request(request)?;
|
||||
let snapshot = self
|
||||
.secrets
|
||||
.resolve(&resolved.config.secret_names())
|
||||
.await?;
|
||||
let prepared = prepare_provider_request(resolved, snapshot)?;
|
||||
crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ChatCompletionsResponse, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
prepared,
|
||||
self.cache.clone(),
|
||||
cache_options,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
execute.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -27,26 +27,56 @@ impl ChatCompletionsRoute {
|
|||
pub fn machine(
|
||||
self,
|
||||
call: ChatCompletionsCall,
|
||||
observers: Option<ObservationSender>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> HostedMachine<ChatCompletions> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
hosted_call(
|
||||
call,
|
||||
observers,
|
||||
move |call: ChatCompletionsCall, _, interceptors, observers| async move {
|
||||
let request = ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers,
|
||||
timeout: call.timeout,
|
||||
};
|
||||
self.run(request, &interceptors, observers.as_ref())
|
||||
move |call, _, interceptors, observers| async move {
|
||||
self.run_call(call, cache_options, &interceptors, observers.as_ref())
|
||||
.await
|
||||
.map(CallOutput::Complete)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "chat_completions",
|
||||
model = %call.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub(super) async fn run_call(
|
||||
&self,
|
||||
call: ChatCompletionsCall,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let request = ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers,
|
||||
timeout: call.timeout,
|
||||
};
|
||||
self.run(request, cache_options, interceptors, observers)
|
||||
.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::Cachable for ChatCompletions {
|
||||
const SURFACE: &'static str = "chat_completions";
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
mod diagnostic;
|
||||
|
||||
pub mod audio_transcription;
|
||||
pub mod caching;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
pub mod error;
|
||||
|
|
@ -12,3 +13,27 @@ pub mod resources;
|
|||
pub mod responses;
|
||||
|
||||
pub use error::RouteError;
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct CallOptions {
|
||||
pub cache: Option<litellm_cache_response::CacheOptions>,
|
||||
pub observers: Option<litellm_host::observation::ObservationSender>,
|
||||
}
|
||||
|
||||
impl From<Option<litellm_host::observation::ObservationSender>> for CallOptions {
|
||||
fn from(observers: Option<litellm_host::observation::ObservationSender>) -> Self {
|
||||
Self {
|
||||
cache: None,
|
||||
observers,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<litellm_cache_response::CacheOptions> for CallOptions {
|
||||
fn from(cache: litellm_cache_response::CacheOptions) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
observers: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -27,6 +27,8 @@ pub(super) async fn execute(
|
|||
http: &litellm_http::Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderMessagesRequest,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
|
|
@ -47,6 +49,10 @@ pub(super) async fn execute(
|
|||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let identity = litellm_host::interceptors::ProviderIdentity {
|
||||
model: context.model.clone(),
|
||||
provider: context.custom_llm_provider.clone(),
|
||||
};
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
|
|
@ -57,44 +63,57 @@ pub(super) async fn execute(
|
|||
context,
|
||||
)
|
||||
.await?;
|
||||
let provider_name = provider.as_str();
|
||||
log_request_body(provider_name, stream, &wire.body);
|
||||
let response = send(
|
||||
http,
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
let cache = cache.filter(|_| authenticated.signer.is_none());
|
||||
let cache_request =
|
||||
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
|
||||
crate::caching::execute_streaming::<super::route::Messages, _, _>(
|
||||
cache_request,
|
||||
cache.as_ref().map(|cache| cache.service.clone()),
|
||||
cache.as_ref().map(|cache| cache.options(cache_options)),
|
||||
interceptors,
|
||||
observers,
|
||||
|| async move {
|
||||
let provider_name = provider.as_str();
|
||||
log_request_body(provider_name, stream, &wire.body);
|
||||
let response = send(
|
||||
http,
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
&wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)
|
||||
.await?;
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
let config = provider.config();
|
||||
if stream {
|
||||
return Ok(streaming_response(
|
||||
response,
|
||||
config.stream_decoder(),
|
||||
provider_name,
|
||||
));
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
log_response_body(&text);
|
||||
let raw = RawResponse { body: text.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
decode_response(config, &body.model, &text)
|
||||
.map(|message| MessagesResponse::Complete(Box::new(message)))
|
||||
},
|
||||
&wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)
|
||||
.await?;
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
let config = provider.config();
|
||||
if stream {
|
||||
return Ok(streaming_response(
|
||||
response,
|
||||
config.stream_decoder(),
|
||||
provider_name,
|
||||
));
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
log_response_body(&text);
|
||||
let raw = RawResponse { body: text.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
decode_response(config, &body.model, &text)
|
||||
.map(|message| MessagesResponse::Complete(Box::new(message)))
|
||||
.await
|
||||
}
|
||||
|
||||
fn serialize_failure(err: serde_json::Error) -> Error {
|
||||
|
|
|
|||
|
|
@ -17,18 +17,84 @@ pub struct MessagesRoute {
|
|||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
#[derive(Clone, Default)]
|
||||
pub struct MessagesRouteBuilder<Http = (), Auth = (), Secrets = ()> {
|
||||
http: Http,
|
||||
auth: Auth,
|
||||
secrets: Secrets,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
}
|
||||
|
||||
impl<Http, Auth, Secrets> MessagesRouteBuilder<Http, Auth, Secrets> {
|
||||
pub fn with_http(
|
||||
self,
|
||||
http: litellm_http::Client,
|
||||
) -> MessagesRouteBuilder<litellm_http::Client, Auth, Secrets> {
|
||||
MessagesRouteBuilder {
|
||||
http,
|
||||
auth: self.auth,
|
||||
secrets: self.secrets,
|
||||
cache: self.cache,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_auth(
|
||||
self,
|
||||
auth: Arc<AuthServices>,
|
||||
) -> MessagesRouteBuilder<Http, Arc<AuthServices>, Secrets> {
|
||||
MessagesRouteBuilder {
|
||||
http: self.http,
|
||||
auth,
|
||||
secrets: self.secrets,
|
||||
cache: self.cache,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_secrets(
|
||||
self,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> MessagesRouteBuilder<Http, Auth, Arc<dyn SecretSource>> {
|
||||
MessagesRouteBuilder {
|
||||
http: self.http,
|
||||
auth: self.auth,
|
||||
secrets,
|
||||
cache: self.cache,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl MessagesRouteBuilder<litellm_http::Client, Arc<AuthServices>, Arc<dyn SecretSource>> {
|
||||
pub fn build(self) -> MessagesRoute {
|
||||
MessagesRoute {
|
||||
http: self.http,
|
||||
auth: self.auth,
|
||||
secrets: self.secrets,
|
||||
cache: self.cache,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl MessagesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
pub fn builder() -> MessagesRouteBuilder {
|
||||
MessagesRouteBuilder::default()
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -36,11 +102,15 @@ impl MessagesRoute {
|
|||
&self,
|
||||
call: MessagesCall,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
litellm_host::lifecycle::observe_call(
|
||||
observers.clone(),
|
||||
self.run(call, interceptors, observers.as_ref()),
|
||||
self.run(call, cache_options, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
|
@ -56,22 +126,36 @@ impl MessagesRoute {
|
|||
async fn run(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.body.model, request.provider.as_str());
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<MessagesResponse, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
request,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
self.run_provider(call, cache_options, interceptors, observers)
|
||||
.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn run_provider(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.body.model, request.provider.as_str());
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<MessagesResponse, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
request,
|
||||
self.cache.clone(),
|
||||
cache_options,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
use litellm_host::observation::ObservationSender;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
|
|
@ -34,14 +33,40 @@ impl super::MessagesRoute {
|
|||
pub fn machine(
|
||||
self,
|
||||
request: super::MessagesCall,
|
||||
observers: Option<ObservationSender>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> MessagesMachine {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
hosted_call(
|
||||
request,
|
||||
observers,
|
||||
move |call, _, interceptors, observers| async move {
|
||||
self.run(call, &interceptors, observers.as_ref()).await
|
||||
self.run(call, cache_options, &interceptors, observers.as_ref())
|
||||
.await
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::Cachable for Messages {
|
||||
const SURFACE: &'static str = "messages";
|
||||
}
|
||||
|
||||
impl crate::caching::StreamCachable for Messages {
|
||||
const TERMINAL_EVENT: &'static str = "message_stop";
|
||||
|
||||
fn replay(data: bytes::Bytes) -> Option<litellm_host::call::OutputOf<Self>> {
|
||||
Some(litellm_host::call::CallOutput::Stream {
|
||||
head: MessagesStreamHead {
|
||||
headers: Vec::new(),
|
||||
},
|
||||
chunks: Box::pin(futures_util::stream::iter([Ok(data)])),
|
||||
})
|
||||
}
|
||||
|
||||
fn bytes(chunk: &Self::Chunk) -> &[u8] {
|
||||
chunk.as_ref()
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -15,10 +15,16 @@ pub(super) async fn execute(
|
|||
http: &litellm_http::Client,
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderResponsesRequest,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let authenticated = resolve_auth(auth, request.environment, &|_| None).await?;
|
||||
let identity = litellm_host::interceptors::ProviderIdentity {
|
||||
model: request.context.model.clone(),
|
||||
provider: request.context.custom_llm_provider.clone(),
|
||||
};
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
|
|
@ -29,65 +35,80 @@ pub(super) async fn execute(
|
|||
request.context,
|
||||
)
|
||||
.await?;
|
||||
let stream = match wire.body.get("stream") {
|
||||
None => false,
|
||||
Some(serde_json::Value::Bool(value)) => *value,
|
||||
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
|
||||
};
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
let cache = cache.filter(|_| authenticated.signer.is_none());
|
||||
let cache_request =
|
||||
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
|
||||
crate::caching::execute_streaming::<super::route::Responses, _, _>(
|
||||
cache_request,
|
||||
cache.as_ref().map(|cache| cache.service.clone()),
|
||||
cache.as_ref().map(|cache| cache.options(cache_options)),
|
||||
interceptors,
|
||||
observers,
|
||||
|| async move {
|
||||
let stream = match wire.body.get("stream") {
|
||||
None => false,
|
||||
Some(serde_json::Value::Bool(value)) => *value,
|
||||
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
|
||||
};
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(network)?;
|
||||
let status = response.status().as_u16();
|
||||
if !response.status().is_success() {
|
||||
let body = response.text().await.map_err(network)?;
|
||||
return Err(litellm_http::transport::Error::Http {
|
||||
status,
|
||||
body: litellm_http::request::truncate_error_body(&body),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
if stream {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| {
|
||||
Some((name.to_string(), value.to_str().ok()?.to_owned()))
|
||||
})
|
||||
.collect();
|
||||
let chunks = response
|
||||
.bytes_stream()
|
||||
.map(|chunk| chunk.map_err(network))
|
||||
.boxed();
|
||||
return Ok(ResponsesOutput::Stream {
|
||||
head: ResponsesStreamHead { headers },
|
||||
chunks,
|
||||
});
|
||||
}
|
||||
let body = response.text().await.map_err(network)?;
|
||||
let raw = RawResponse { body: body.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
let value = serde_json::from_str(&body)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into()))?;
|
||||
request
|
||||
.config
|
||||
.transform_response_api_response(value)
|
||||
.map(ResponsesOutput::Complete)
|
||||
.map_err(Error::from)
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(network)?;
|
||||
let status = response.status().as_u16();
|
||||
if !response.status().is_success() {
|
||||
let body = response.text().await.map_err(network)?;
|
||||
return Err(litellm_http::transport::Error::Http {
|
||||
status,
|
||||
body: litellm_http::request::truncate_error_body(&body),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
if stream {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_owned())))
|
||||
.collect();
|
||||
let chunks = response
|
||||
.bytes_stream()
|
||||
.map(|chunk| chunk.map_err(network))
|
||||
.boxed();
|
||||
return Ok(ResponsesOutput::Stream {
|
||||
head: ResponsesStreamHead { headers },
|
||||
chunks,
|
||||
});
|
||||
}
|
||||
let body = response.text().await.map_err(network)?;
|
||||
let raw = RawResponse { body: body.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
let value = serde_json::from_str(&body)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into()))?;
|
||||
request
|
||||
.config
|
||||
.transform_response_api_response(value)
|
||||
.map(ResponsesOutput::Complete)
|
||||
.map_err(Error::from)
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn network(error: reqwest::Error) -> Error {
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ pub struct ResponsesRoute {
|
|||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
|
|
@ -31,6 +32,14 @@ impl ResponsesRoute {
|
|||
http,
|
||||
auth,
|
||||
secrets,
|
||||
cache: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -38,11 +47,15 @@ impl ResponsesRoute {
|
|||
&self,
|
||||
call: ResponsesCall,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
litellm_host::lifecycle::observe_call(
|
||||
observers.clone(),
|
||||
self.run(call, interceptors, observers.as_ref()),
|
||||
self.run(call, cache_options, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
|
@ -58,25 +71,36 @@ impl ResponsesRoute {
|
|||
async fn run(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(
|
||||
&request.context.model,
|
||||
&request.context.custom_llm_provider,
|
||||
);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
request,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
self.run_provider(call, cache_options, interceptors, observers)
|
||||
.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn run_provider(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
cache_options: Option<litellm_cache_response::CacheOptions>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.context.model, &request.context.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
request,
|
||||
self.cache.clone(),
|
||||
cache_options,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -16,14 +16,9 @@ pub(super) async fn prepare(
|
|||
call: ResponsesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderResponsesRequest, Error> {
|
||||
let provider = call.custom_llm_provider.as_deref().unwrap_or("openai");
|
||||
if provider != "openai" {
|
||||
return Err(Error::Unsupported("native HTTP responses provider"));
|
||||
}
|
||||
let model = call.model.strip_prefix("openai/").unwrap_or(&call.model);
|
||||
if model.is_empty() || model.contains('/') {
|
||||
return Err(Error::InvalidProvider(call.model));
|
||||
}
|
||||
let identity = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
|
||||
let provider = identity.provider.as_str();
|
||||
let model = identity.model.as_str();
|
||||
let config: &'static dyn BaseResponsesApiConfig = &OpenAiResponsesApiConfig;
|
||||
let snapshot = secrets
|
||||
.resolve(config.secret_names(call.api_key.as_deref(), call.api_base.as_deref()))
|
||||
|
|
@ -56,3 +51,21 @@ pub(super) async fn prepare(
|
|||
timeout: call.timeout,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<litellm_host::interceptors::ProviderIdentity, Error> {
|
||||
let provider = custom_llm_provider.unwrap_or("openai");
|
||||
if provider != "openai" {
|
||||
return Err(Error::Unsupported("native HTTP responses provider"));
|
||||
}
|
||||
let resolved = model.strip_prefix("openai/").unwrap_or(model);
|
||||
if resolved.is_empty() || resolved.contains('/') {
|
||||
return Err(Error::InvalidProvider(model.into()));
|
||||
}
|
||||
Ok(litellm_host::interceptors::ProviderIdentity {
|
||||
model: resolved.into(),
|
||||
provider: provider.into(),
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
use litellm_host::observation::ObservationSender;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
|
|
@ -28,14 +27,48 @@ impl ResponsesRoute {
|
|||
pub fn machine(
|
||||
self,
|
||||
call: ResponsesCall,
|
||||
observers: Option<ObservationSender>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> HostedMachine<Responses> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
hosted_call(
|
||||
call,
|
||||
observers,
|
||||
move |call, _, interceptors, observers| async move {
|
||||
self.run(call, &interceptors, observers.as_ref()).await
|
||||
self.run(call, cache_options, &interceptors, observers.as_ref())
|
||||
.await
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::Cachable for Responses {
|
||||
const SURFACE: &'static str = "responses";
|
||||
|
||||
fn reusable(response: &Self::Response) -> bool {
|
||||
response
|
||||
.extra
|
||||
.get("status")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
== Some("completed")
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::StreamCachable for Responses {
|
||||
const TERMINAL_EVENT: &'static str = "response.completed";
|
||||
|
||||
fn replay(data: bytes::Bytes) -> Option<litellm_host::call::OutputOf<Self>> {
|
||||
Some(litellm_host::call::CallOutput::Stream {
|
||||
head: ResponsesStreamHead {
|
||||
headers: Vec::new(),
|
||||
},
|
||||
chunks: Box::pin(futures_util::stream::iter([Ok(data)])),
|
||||
})
|
||||
}
|
||||
|
||||
fn bytes(chunk: &Self::Chunk) -> &[u8] {
|
||||
chunk.as_ref()
|
||||
}
|
||||
}
|
||||
|
|
|
|||
1217
litellm-rust/crates/core/tests/caching.rs
Normal file
1217
litellm-rust/crates/core/tests/caching.rs
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -1,4 +1,8 @@
|
|||
use litellm_host::interceptors::RawResponse;
|
||||
use litellm_host::{
|
||||
interceptors::{ExecutionFacts, ResultSource},
|
||||
lifecycle::ExecutionEvent,
|
||||
};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest};
|
||||
|
|
@ -310,7 +314,13 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
|
|||
&events[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(_),
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
|
|
|
|||
|
|
@ -1,4 +1,8 @@
|
|||
use litellm_core::messages::{MessagesResponse, messages_body};
|
||||
use litellm_host::{
|
||||
interceptors::{ExecutionFacts, ResultSource},
|
||||
lifecycle::ExecutionEvent,
|
||||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use rstest::rstest;
|
||||
|
||||
|
|
@ -60,7 +64,13 @@ async fn calls_defer_execution_until_polled(
|
|||
true,
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(_),
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
)
|
||||
|
|
@ -69,7 +79,13 @@ async fn calls_defer_execution_until_polled(
|
|||
true,
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(_),
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
)
|
||||
|
|
@ -253,22 +269,25 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
|
|||
};
|
||||
|
||||
let resources = support::resources();
|
||||
let response = litellm_core::messages::MessagesRoute::new(
|
||||
provider_http(&resources, &Resolution::from(&settings).config),
|
||||
resources.auth,
|
||||
no_secrets(),
|
||||
)
|
||||
.execute(
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(base),
|
||||
..call
|
||||
},
|
||||
&(),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
let response = litellm_core::messages::MessagesRoute::builder()
|
||||
.with_http(provider_http(
|
||||
&resources,
|
||||
&Resolution::from(&settings).config,
|
||||
))
|
||||
.with_auth(resources.auth)
|
||||
.with_secrets(no_secrets())
|
||||
.build()
|
||||
.execute(
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(base),
|
||||
..call
|
||||
},
|
||||
&(),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
let MessagesResponse::Complete(message) = response else {
|
||||
panic!("a non-streaming request returns a message");
|
||||
|
|
@ -321,3 +340,56 @@ async fn message_route_summary_excludes_payload_diagnostics(
|
|||
assert!(summaries[0].get("body").is_none());
|
||||
assert!(!format!("{:?}", traces.records()).contains("private-key-sentinel"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::uncached(false, 2)]
|
||||
#[case::cached(true, 1)]
|
||||
#[tokio::test]
|
||||
async fn builder_preserves_dependencies_and_optional_cache(
|
||||
#[case] caching: bool,
|
||||
#[case] expected_requests: usize,
|
||||
) {
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{CacheScope, ResponseCache, ScopedCache};
|
||||
use litellm_core::messages::MessagesRoute;
|
||||
|
||||
let upstream = upstream([message_response(), message_response()]).await;
|
||||
let resources = resources();
|
||||
let builder = MessagesRoute::builder();
|
||||
let builder = if caching {
|
||||
builder.with_cache(ScopedCache::new(
|
||||
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
|
||||
Some(100),
|
||||
Some(Duration::from_secs(60)),
|
||||
)))),
|
||||
CacheScope::Shared,
|
||||
))
|
||||
} else {
|
||||
builder
|
||||
};
|
||||
let route = builder
|
||||
.with_secrets(Arc::new(RecordingSecrets::new([(
|
||||
"ANTHROPIC_API_KEY",
|
||||
"builder-key",
|
||||
)])))
|
||||
.with_auth(resources.auth.clone())
|
||||
.with_http(provider_http(&resources, &http_config()))
|
||||
.build();
|
||||
for _ in 0..2 {
|
||||
let request = MessagesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..super::call()
|
||||
};
|
||||
let MessagesResponse::Complete(response) = route.execute(request, &(), None).await.unwrap()
|
||||
else {
|
||||
panic!("expected a completed message");
|
||||
};
|
||||
assert_eq!(
|
||||
response.content,
|
||||
message_body()["content"].as_array().unwrap().as_slice()
|
||||
);
|
||||
}
|
||||
let requests = received(&upstream).await;
|
||||
assert_eq!(requests.len(), expected_requests);
|
||||
assert_eq!(requests[0].header("x-api-key"), Some("builder-key"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ use super::*;
|
|||
|
||||
pub(crate) fn event_name(event: &CallEvent) -> &'static str {
|
||||
match event {
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady { .. }) => "result_ready",
|
||||
CallEvent::Started { .. } => "started",
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }) => "response",
|
||||
CallEvent::Succeeded { .. } => "success",
|
||||
|
|
|
|||
|
|
@ -41,6 +41,11 @@ async fn drive_until(
|
|||
Err(error) => break Err(error),
|
||||
};
|
||||
let answer = match op {
|
||||
HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply }) => host
|
||||
.result_ready(facts)
|
||||
.await
|
||||
.map(|()| reply.send(()))
|
||||
.map_err(HostFailure::Error),
|
||||
HostRequest::Stream(stream) => match stream {
|
||||
litellm_host::protocol::StreamDelivery::Open(head, _) => match head {},
|
||||
litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {},
|
||||
|
|
@ -80,6 +85,7 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto
|
|||
_ = stop.notified() => break,
|
||||
step = machine.resume() => {
|
||||
match step.unwrap() {
|
||||
MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::ResultReady { reply, .. })) => reply.send(()),
|
||||
MachineStep::Suspended(HostRequest::HostCall(op)) => host.handle_host_call(op).await.unwrap(),
|
||||
MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::BeforeProviderRequest { wire, reply, .. })) => reply.send(*wire),
|
||||
MachineStep::Suspended(HostRequest::Intercept(InterceptRequest::AfterProviderResponse { reply, .. })) => reply.send(()),
|
||||
|
|
|
|||
|
|
@ -1,3 +1,7 @@
|
|||
use litellm_host::{
|
||||
interceptors::{ExecutionFacts, ResultSource},
|
||||
lifecycle::ExecutionEvent,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
|
|
@ -70,7 +74,13 @@ async fn http_responses_share_execution_and_hooks(call: ResponsesCall, #[case] h
|
|||
&host.events.0.lock().unwrap()[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(_),
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
|
|
@ -97,7 +107,8 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption(
|
|||
let (headers, bytes) = if hosted {
|
||||
assert_eq!(
|
||||
litellm_host_native::in_process::run_hosted(
|
||||
responses_route(no_secrets()).machine(host.request().unwrap(), None,),
|
||||
responses_route(no_secrets())
|
||||
.machine(host.request().unwrap(), Some(host.events.0.sender.clone())),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
|
|
@ -117,7 +128,18 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption(
|
|||
else {
|
||||
panic!()
|
||||
};
|
||||
assert_eq!(host.events.0.lock().unwrap().len(), 1);
|
||||
assert!(matches!(
|
||||
&host.events.0.lock().unwrap()[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
]
|
||||
));
|
||||
(
|
||||
head.headers,
|
||||
chunks.try_collect::<Vec<_>>().await.unwrap().concat(),
|
||||
|
|
@ -127,7 +149,16 @@ async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption(
|
|||
assert_eq!(bytes, body.as_bytes());
|
||||
assert!(matches!(
|
||||
&host.events.0.lock().unwrap()[..],
|
||||
[CallEvent::Started { .. }, CallEvent::Succeeded { .. }]
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -44,11 +44,11 @@ pub fn provider_http(
|
|||
|
||||
pub fn messages_route(secrets: Arc<dyn SecretSource>) -> litellm_core::messages::MessagesRoute {
|
||||
let resources = resources();
|
||||
litellm_core::messages::MessagesRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
secrets,
|
||||
)
|
||||
litellm_core::messages::MessagesRoute::builder()
|
||||
.with_http(provider_http(&resources, &http_config()))
|
||||
.with_auth(resources.auth)
|
||||
.with_secrets(secrets)
|
||||
.build()
|
||||
}
|
||||
|
||||
pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute {
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-cache-response.workspace = true
|
||||
axum = { workspace = true, features = ["json", "multipart", "original-uri"] }
|
||||
base64.workspace = true
|
||||
bytes.workspace = true
|
||||
|
|
@ -13,15 +14,18 @@ litellm-auth.workspace = true
|
|||
litellm-gateway-auth.workspace = true
|
||||
litellm-core.workspace = true
|
||||
litellm-host-http.workspace = true
|
||||
litellm-host.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-router.workspace = true
|
||||
litellm-secrets.workspace = true
|
||||
litellm-types.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-cache-memory.workspace = true
|
||||
futures-util.workspace = true
|
||||
tokio = { workspace = true, features = ["io-util"] }
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
109
litellm-rust/crates/gateway-inference/src/caching.rs
Normal file
109
litellm-rust/crates/gateway-inference/src/caching.rs
Normal file
|
|
@ -0,0 +1,109 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_cache_response::{CacheOptions, CacheScope};
|
||||
use litellm_gateway_auth::AuthenticatedRequest;
|
||||
use serde::Deserialize;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Default, Deserialize)]
|
||||
#[serde(default, deny_unknown_fields)]
|
||||
struct Controls {
|
||||
#[serde(rename = "no-cache")]
|
||||
no_cache: bool,
|
||||
#[serde(rename = "no-store")]
|
||||
no_store: bool,
|
||||
ttl: Option<f64>,
|
||||
#[serde(rename = "s-maxage", alias = "s-max-age")]
|
||||
max_age: Option<f64>,
|
||||
}
|
||||
|
||||
type Prepared = (Map<String, Value>, CacheOptions);
|
||||
|
||||
pub(crate) fn prepare(
|
||||
identity: &AuthenticatedRequest,
|
||||
body: Map<String, Value>,
|
||||
) -> Result<Prepared, Error> {
|
||||
let controls: Controls = match body.get("cache").filter(|value| !value.is_null()) {
|
||||
Some(value) => serde_json::from_value(value.clone())
|
||||
.map_err(|error| Error::InvalidBody(error.to_string()))?,
|
||||
None => Controls::default(),
|
||||
};
|
||||
let caching: Option<bool> = body
|
||||
.get("caching")
|
||||
.filter(|value| !value.is_null())
|
||||
.map(|value| serde_json::from_value(value.clone()))
|
||||
.transpose()
|
||||
.map_err(|error| Error::InvalidBody(error.to_string()))?;
|
||||
let caller = identity.caller();
|
||||
let options = CacheOptions {
|
||||
caching,
|
||||
no_cache: controls.no_cache,
|
||||
no_store: controls.no_store,
|
||||
ttl: controls.ttl.map(duration).transpose()?,
|
||||
max_age: controls.max_age.map(duration).transpose()?,
|
||||
scope: CacheScope::Isolated(
|
||||
serde_json::json!([
|
||||
caller.principal().authority(),
|
||||
caller.principal().subject(),
|
||||
caller.authentication().credential_id
|
||||
])
|
||||
.to_string(),
|
||||
),
|
||||
};
|
||||
Ok((
|
||||
body.into_iter()
|
||||
.filter(|(name, _)| !matches!(name.as_str(), "cache" | "caching"))
|
||||
.collect(),
|
||||
options,
|
||||
))
|
||||
}
|
||||
|
||||
fn duration(seconds: f64) -> Result<Duration, Error> {
|
||||
Duration::try_from_secs_f64(seconds)
|
||||
.ok()
|
||||
.filter(|duration| !duration.is_zero())
|
||||
.ok_or_else(|| Error::InvalidBody("cache durations must be finite and positive".into()))
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub(crate) struct CacheHeaders(std::sync::Arc<std::sync::OnceLock<String>>);
|
||||
|
||||
impl litellm_host::interceptors::Interceptors<litellm_core::RouteError> for CacheHeaders {
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: litellm_host::interceptors::WireRequest,
|
||||
_: litellm_host::interceptors::RequestContext,
|
||||
) -> Result<litellm_host::interceptors::WireRequest, litellm_core::RouteError> {
|
||||
Ok(wire)
|
||||
}
|
||||
|
||||
async fn after_provider_response(
|
||||
&self,
|
||||
_: litellm_host::interceptors::RawResponse,
|
||||
) -> Result<(), litellm_core::RouteError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn result_ready(
|
||||
&self,
|
||||
facts: litellm_host::interceptors::ExecutionFacts,
|
||||
) -> Result<(), litellm_core::RouteError> {
|
||||
if let litellm_host::interceptors::ResultSource::Cache { key } = facts.source {
|
||||
let _ = self.0.set(key);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl CacheHeaders {
|
||||
pub(crate) fn apply(&self, mut response: axum::response::Response) -> axum::response::Response {
|
||||
if let Some(key) = self.0.get()
|
||||
&& let Ok(value) = axum::http::HeaderValue::from_str(key)
|
||||
{
|
||||
response.headers_mut().insert("x-litellm-cache-key", value);
|
||||
}
|
||||
response
|
||||
}
|
||||
}
|
||||
|
|
@ -42,9 +42,20 @@ async fn handle(
|
|||
) -> Result<Response, Error> {
|
||||
let deployment = request::resolve_deployment(gateway, &body)?;
|
||||
request::authorize_model(identity, deployment, &body).await?;
|
||||
let (body, cache_options) = crate::caching::prepare(identity, body)?;
|
||||
let route = gateway.chat_completions.clone();
|
||||
let route = match &gateway.cache {
|
||||
Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new(
|
||||
cache.clone(),
|
||||
cache_options.scope.clone(),
|
||||
)),
|
||||
None => route,
|
||||
};
|
||||
|
||||
let messages = body.get("messages").cloned().unwrap_or_default();
|
||||
let headers = crate::caching::CacheHeaders::default();
|
||||
let response = litellm_host_http::serve_unary(
|
||||
gateway.chat_completions.clone().machine(
|
||||
route.machine(
|
||||
ChatCompletionsCall {
|
||||
model: deployment.model.clone(),
|
||||
messages,
|
||||
|
|
@ -58,13 +69,13 @@ async fn handle(
|
|||
extra_headers: None,
|
||||
timeout: deployment.timeout,
|
||||
},
|
||||
None,
|
||||
cache_options,
|
||||
),
|
||||
(),
|
||||
(),
|
||||
headers.clone(),
|
||||
litellm_host_http::Unary::new(Json),
|
||||
None,
|
||||
)
|
||||
.await?;
|
||||
Ok(response)
|
||||
Ok(headers.apply(response))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@
|
|||
//! maps a public model name to its deployment and runs the core route.
|
||||
|
||||
mod audio_transcription;
|
||||
mod caching;
|
||||
mod chat_completions;
|
||||
mod error;
|
||||
pub mod messages;
|
||||
|
|
@ -27,6 +28,7 @@ pub use litellm_router::{Deployment, Router as ModelRouter};
|
|||
pub use request::{JsonObject, RequestId};
|
||||
|
||||
pub struct Gateway {
|
||||
cache: Option<Arc<dyn litellm_cache_response::ResponseCacheService>>,
|
||||
pub audio_transcription: AudioTranscriptionRoute,
|
||||
pub chat_completions: ChatCompletionsRoute,
|
||||
pub messages: MessagesRoute,
|
||||
|
|
@ -39,6 +41,13 @@ pub struct Gateway {
|
|||
}
|
||||
|
||||
impl Gateway {
|
||||
pub fn with_cache(self, cache: Arc<dyn litellm_cache_response::ResponseCacheService>) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new(
|
||||
resources: CoreResources,
|
||||
http: HttpClientConfig,
|
||||
|
|
@ -48,6 +57,7 @@ impl Gateway {
|
|||
let provider = resources.pool.client(&http, ClientVariant::Provider)?;
|
||||
let auth = resources.auth.clone();
|
||||
Ok(Self {
|
||||
cache: None,
|
||||
audio_transcription: AudioTranscriptionRoute::new(
|
||||
provider.clone(),
|
||||
auth.clone(),
|
||||
|
|
@ -58,7 +68,11 @@ impl Gateway {
|
|||
auth.clone(),
|
||||
secrets.clone(),
|
||||
),
|
||||
messages: MessagesRoute::new(provider.clone(), auth.clone(), secrets.clone()),
|
||||
messages: MessagesRoute::builder()
|
||||
.with_http(provider.clone())
|
||||
.with_auth(auth.clone())
|
||||
.with_secrets(secrets.clone())
|
||||
.build(),
|
||||
responses: ResponsesRoute::new(provider, auth.clone(), secrets.clone()),
|
||||
ocr: OcrRoute::new(OcrClient::new(
|
||||
&resources.pool,
|
||||
|
|
|
|||
|
|
@ -43,11 +43,23 @@ async fn handle(
|
|||
) -> Result<Response, Error> {
|
||||
let deployment = request::resolve_deployment(gateway, &body)?;
|
||||
request::authorize_model(identity, deployment, &body).await?;
|
||||
let (body, cache_options) = crate::caching::prepare(identity, body)?;
|
||||
let route = gateway.messages.clone();
|
||||
let route = match &gateway.cache {
|
||||
Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new(
|
||||
cache.clone(),
|
||||
cache_options.scope.clone(),
|
||||
)),
|
||||
None => route,
|
||||
};
|
||||
|
||||
let call = project(deployment, body, headers)?;
|
||||
let machine = gateway.messages.clone().machine(call, None);
|
||||
let machine = route.machine(call, cache_options);
|
||||
let stream =
|
||||
Sse::<Messages, _, _>::new(Json, |error| Bytes::from(Error::from(error).sse_frame()));
|
||||
Ok(litellm_host_http::serve(machine, (), (), stream, None).await?)
|
||||
let headers = crate::caching::CacheHeaders::default();
|
||||
let response = litellm_host_http::serve(machine, (), headers.clone(), stream, None).await?;
|
||||
Ok(headers.apply(response))
|
||||
}
|
||||
|
||||
fn project(
|
||||
|
|
|
|||
|
|
@ -15,6 +15,16 @@ pub(crate) async fn create(
|
|||
) -> Result<Response, Error> {
|
||||
let deployment = request::resolve_deployment(&gateway, &body)?;
|
||||
request::authorize_model(&identity, deployment, &body).await?;
|
||||
let (body, cache_options) = crate::caching::prepare(&identity, body)?;
|
||||
let route = gateway.responses.clone();
|
||||
let route = match &gateway.cache {
|
||||
Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new(
|
||||
cache.clone(),
|
||||
cache_options.scope.clone(),
|
||||
)),
|
||||
None => route,
|
||||
};
|
||||
|
||||
let call = ResponsesCall {
|
||||
model: deployment.model.clone(),
|
||||
input: body.get("input").cloned().unwrap_or_default(),
|
||||
|
|
@ -28,7 +38,7 @@ pub(crate) async fn create(
|
|||
extra_headers: None,
|
||||
timeout: deployment.timeout,
|
||||
};
|
||||
let machine = gateway.responses.clone().machine(call, None);
|
||||
let machine = route.machine(call, cache_options);
|
||||
let stream = Sse::<Responses, _, _>::new(Json, |error| {
|
||||
let error = Error::from(error);
|
||||
Bytes::from(format!(
|
||||
|
|
@ -36,5 +46,7 @@ pub(crate) async fn create(
|
|||
json!({"type": "error", "code": error.status().as_u16().to_string(), "message": error.to_string(), "param": null})
|
||||
))
|
||||
});
|
||||
Ok(litellm_host_http::serve(machine, (), (), stream, None).await?)
|
||||
let headers = crate::caching::CacheHeaders::default();
|
||||
let response = litellm_host_http::serve(machine, (), headers.clone(), stream, None).await?;
|
||||
Ok(headers.apply(response))
|
||||
}
|
||||
|
|
|
|||
185
litellm-rust/crates/gateway-inference/tests/caching.rs
Normal file
185
litellm-rust/crates/gateway-inference/tests/caching.rs
Normal file
|
|
@ -0,0 +1,185 @@
|
|||
mod support;
|
||||
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use axum::body::to_bytes;
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{
|
||||
CacheKeyInput, ResponseCache, ResponseCacheConfig, ResponseCacheRequest, ResponseCacheService,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
|
||||
|
||||
#[rstest]
|
||||
#[case::chat("/v1/chat/completions", "anthropic/test-model", false)]
|
||||
#[case::messages("/v1/messages", "anthropic/test-model", false)]
|
||||
#[case::responses("/v1/responses", "openai/test-model", false)]
|
||||
#[case::messages_stream("/v1/messages", "anthropic/test-model", true)]
|
||||
#[case::responses_stream("/v1/responses", "openai/test-model", true)]
|
||||
#[tokio::test]
|
||||
async fn all_inference_endpoints_share_native_cache(
|
||||
#[case] path: &str,
|
||||
#[case] model: &str,
|
||||
#[case] stream: bool,
|
||||
#[values("s-maxage", "s-max-age")] max_age: &str,
|
||||
) {
|
||||
let upstream = MockServer::start().await;
|
||||
let is_responses = path.ends_with("responses");
|
||||
let provider_body = if is_responses {
|
||||
json!({"id":"response-1", "model":"test-model", "status":"completed", "output":[]})
|
||||
} else {
|
||||
json!({"id":"message-1", "model":"test-model", "type":"message", "role":"assistant", "content":[{"type":"text","text":"hello"}], "stop_reason":"end_turn", "usage":{"input_tokens":1,"output_tokens":1}})
|
||||
};
|
||||
let terminal = if is_responses {
|
||||
"response.completed"
|
||||
} else {
|
||||
"message_stop"
|
||||
};
|
||||
let events = format!("event: {terminal}\ndata: {{\"type\":\"{terminal}\"}}\n\n");
|
||||
let template = if stream {
|
||||
ResponseTemplate::new(200).set_body_raw(events.clone(), "text/event-stream")
|
||||
} else {
|
||||
ResponseTemplate::new(200).set_body_json(provider_body)
|
||||
};
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(template)
|
||||
.expect(2)
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let cache: Arc<dyn ResponseCacheService> = Arc::new(
|
||||
ResponseCache::new(Arc::new(InMemoryCache::new(
|
||||
Some(100),
|
||||
Some(Duration::from_secs(60)),
|
||||
)))
|
||||
.with_config(ResponseCacheConfig {
|
||||
namespace: "gateway-test".into(),
|
||||
max_entry_bytes: 4096,
|
||||
}),
|
||||
);
|
||||
let app = support::app_with_cache(model, &upstream.uri(), cache.clone());
|
||||
let request = if is_responses {
|
||||
json!({"model":"public/model", "input":"hello", "stream":stream, "cache":{(max_age):600}})
|
||||
} else {
|
||||
json!({"model":"public/model", "messages":[{"role":"user","content":"hello"}], "max_tokens":16, "stream":stream, "cache":{(max_age):600}})
|
||||
};
|
||||
let first = support::post(app.clone(), path, request.clone()).await;
|
||||
assert_eq!(first.status(), 200);
|
||||
assert!(!first.headers().contains_key("x-litellm-cache-key"));
|
||||
let first = to_bytes(first.into_body(), 4096).await.unwrap();
|
||||
let second = support::post(app.clone(), path, request.clone()).await;
|
||||
assert_eq!(second.status(), 200);
|
||||
let cache_key = second.headers().get("x-litellm-cache-key").unwrap().clone();
|
||||
assert!(!cache_key.as_bytes().is_empty());
|
||||
let stored = cache
|
||||
.lookup(
|
||||
&ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some(cache_key.to_str().unwrap().into()),
|
||||
..Default::default()
|
||||
}),
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(
|
||||
stored.is_some(),
|
||||
"the header must identify the stored entry"
|
||||
);
|
||||
let second = to_bytes(second.into_body(), 4096).await.unwrap();
|
||||
if stream {
|
||||
assert_eq!(first, events);
|
||||
assert_eq!(second, first);
|
||||
} else {
|
||||
assert_eq!(
|
||||
serde_json::from_slice::<Value>(&first).unwrap(),
|
||||
serde_json::from_slice::<Value>(&second).unwrap()
|
||||
);
|
||||
}
|
||||
let bypass_request = Value::Object(
|
||||
request
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.map(|(name, value)| {
|
||||
(
|
||||
name.clone(),
|
||||
if name == "cache" {
|
||||
json!({"no-cache": true, "no-store": true})
|
||||
} else {
|
||||
value.clone()
|
||||
},
|
||||
)
|
||||
})
|
||||
.collect(),
|
||||
);
|
||||
let bypassed = support::post(app.clone(), path, bypass_request).await;
|
||||
assert_eq!(bypassed.status(), 200);
|
||||
assert!(!bypassed.headers().contains_key("x-litellm-cache-key"));
|
||||
to_bytes(bypassed.into_body(), 4096).await.unwrap();
|
||||
let restored = support::post(app, path, request).await;
|
||||
assert_eq!(restored.status(), 200);
|
||||
assert_eq!(
|
||||
restored.headers().get("x-litellm-cache-key"),
|
||||
Some(&cache_key)
|
||||
);
|
||||
assert_eq!(to_bytes(restored.into_body(), 4096).await.unwrap(), second);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::different_subject("issuer", "tenant-b")]
|
||||
#[case::different_authority("other-issuer", "tenant-a")]
|
||||
#[tokio::test]
|
||||
async fn authenticated_callers_do_not_share_cached_responses(
|
||||
#[case] authority: &str,
|
||||
#[case] subject: &str,
|
||||
) {
|
||||
use litellm_gateway_auth::{Principal, PrincipalKind};
|
||||
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"id":"response-1", "model":"test-model", "status":"completed", "output":[]
|
||||
})))
|
||||
.expect(2)
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let cache: Arc<dyn ResponseCacheService> = Arc::new(ResponseCache::new(Arc::new(
|
||||
InMemoryCache::new(Some(100), Some(Duration::from_secs(60))),
|
||||
)));
|
||||
let first_caller = support::app_with_cache_for_principal(
|
||||
"openai/test-model",
|
||||
&upstream.uri(),
|
||||
cache.clone(),
|
||||
Principal::new("issuer".into(), "tenant-a".into(), PrincipalKind::Service),
|
||||
);
|
||||
let second_caller = support::app_with_cache_for_principal(
|
||||
"openai/test-model",
|
||||
&upstream.uri(),
|
||||
cache,
|
||||
Principal::new(authority.into(), subject.into(), PrincipalKind::Service),
|
||||
);
|
||||
let body = json!({"model":"public/model","input":"same prompt"});
|
||||
let first = support::post(first_caller.clone(), "/v1/responses", body.clone()).await;
|
||||
assert_eq!(first.status(), 200);
|
||||
assert!(!first.headers().contains_key("x-litellm-cache-key"));
|
||||
let first_hit = support::post(first_caller.clone(), "/v1/responses", body.clone()).await;
|
||||
assert_eq!(first_hit.status(), 200);
|
||||
let first_key = first_hit.headers().get("x-litellm-cache-key").unwrap();
|
||||
let second = support::post(second_caller.clone(), "/v1/responses", body.clone()).await;
|
||||
assert_eq!(second.status(), 200);
|
||||
assert!(!second.headers().contains_key("x-litellm-cache-key"));
|
||||
let second_hit = support::post(second_caller, "/v1/responses", body.clone()).await;
|
||||
assert_eq!(second_hit.status(), 200);
|
||||
assert_ne!(
|
||||
second_hit.headers().get("x-litellm-cache-key").unwrap(),
|
||||
first_key
|
||||
);
|
||||
let first_again = support::post(first_caller, "/v1/responses", body).await;
|
||||
assert_eq!(first_again.status(), 200);
|
||||
assert_eq!(
|
||||
first_again.headers().get("x-litellm-cache-key"),
|
||||
Some(first_key)
|
||||
);
|
||||
}
|
||||
|
|
@ -1,3 +1,6 @@
|
|||
// Shared across integration-test targets; each target uses a different subset.
|
||||
#![allow(dead_code)]
|
||||
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use axum::{
|
||||
|
|
@ -33,39 +36,83 @@ pub fn app_with_permissions(
|
|||
model: &str,
|
||||
api_base: &str,
|
||||
permissions: litellm_gateway_auth::Permissions,
|
||||
) -> Router {
|
||||
configured_app(model, api_base, permissions, None, None)
|
||||
}
|
||||
|
||||
pub fn app_with_cache(
|
||||
model: &str,
|
||||
api_base: &str,
|
||||
cache: Arc<dyn litellm_cache_response::ResponseCacheService>,
|
||||
) -> Router {
|
||||
configured_app(
|
||||
model,
|
||||
api_base,
|
||||
litellm_gateway_auth::Permissions::All,
|
||||
Some(cache),
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn app_with_cache_for_principal(
|
||||
model: &str,
|
||||
api_base: &str,
|
||||
cache: Arc<dyn litellm_cache_response::ResponseCacheService>,
|
||||
principal: litellm_gateway_auth::Principal,
|
||||
) -> Router {
|
||||
configured_app(
|
||||
model,
|
||||
api_base,
|
||||
litellm_gateway_auth::Permissions::All,
|
||||
Some(cache),
|
||||
Some(principal),
|
||||
)
|
||||
}
|
||||
|
||||
fn configured_app(
|
||||
model: &str,
|
||||
api_base: &str,
|
||||
permissions: litellm_gateway_auth::Permissions,
|
||||
cache: Option<Arc<dyn litellm_cache_response::ResponseCacheService>>,
|
||||
principal: Option<litellm_gateway_auth::Principal>,
|
||||
) -> Router {
|
||||
let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver)));
|
||||
let http = Resolution::from(&HttpSettings::default()).config;
|
||||
let secrets = Arc::new(NoSecrets);
|
||||
let resources = CoreResources::new(pool);
|
||||
router(Arc::new(
|
||||
Gateway::new(
|
||||
resources,
|
||||
http,
|
||||
secrets,
|
||||
[(
|
||||
"public/model".into(),
|
||||
Deployment {
|
||||
model: model.into(),
|
||||
api_base: Some(api_base.into()),
|
||||
api_key: Some("test-key".into()),
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
..Default::default()
|
||||
},
|
||||
)]
|
||||
.into_iter()
|
||||
.collect(),
|
||||
)
|
||||
.unwrap(),
|
||||
))
|
||||
.layer(axum::middleware::from_fn_with_state(
|
||||
permissions,
|
||||
let gateway = Gateway::new(
|
||||
resources,
|
||||
http,
|
||||
secrets,
|
||||
[(
|
||||
"public/model".into(),
|
||||
Deployment {
|
||||
model: model.into(),
|
||||
api_base: Some(api_base.into()),
|
||||
api_key: Some("test-key".into()),
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
..Default::default()
|
||||
},
|
||||
)]
|
||||
.into_iter()
|
||||
.collect(),
|
||||
)
|
||||
.unwrap();
|
||||
let gateway = match cache {
|
||||
Some(cache) => gateway.with_cache(cache),
|
||||
None => gateway,
|
||||
};
|
||||
router(Arc::new(gateway)).layer(axum::middleware::from_fn_with_state(
|
||||
(permissions, principal),
|
||||
test_identity,
|
||||
))
|
||||
}
|
||||
|
||||
async fn test_identity(
|
||||
axum::extract::State(permissions): axum::extract::State<litellm_gateway_auth::Permissions>,
|
||||
axum::extract::State((permissions, principal)): axum::extract::State<(
|
||||
litellm_gateway_auth::Permissions,
|
||||
Option<litellm_gateway_auth::Principal>,
|
||||
)>,
|
||||
mut request: axum::extract::Request,
|
||||
next: axum::middleware::Next,
|
||||
) -> Response {
|
||||
|
|
@ -74,7 +121,7 @@ async fn test_identity(
|
|||
Some(SecretValue::new("test-inbound-key")),
|
||||
Arc::new(NoSecrets),
|
||||
)),
|
||||
Arc::new(TestPermissions(permissions)),
|
||||
Arc::new(TestPermissions(permissions, principal)),
|
||||
Arc::new(litellm_gateway_auth::NoAdditionalPolicy),
|
||||
Arc::new(litellm_gateway_auth::SystemClock),
|
||||
);
|
||||
|
|
@ -101,7 +148,10 @@ pub async fn json(response: Response) -> Value {
|
|||
serde_json::from_slice(&to_bytes(response.into_body(), 1024 * 1024).await.unwrap()).unwrap()
|
||||
}
|
||||
|
||||
struct TestPermissions(litellm_gateway_auth::Permissions);
|
||||
struct TestPermissions(
|
||||
litellm_gateway_auth::Permissions,
|
||||
Option<litellm_gateway_auth::Principal>,
|
||||
);
|
||||
|
||||
impl litellm_gateway_auth::IdentityResolver for TestPermissions {
|
||||
fn resolve<'a>(
|
||||
|
|
@ -110,7 +160,7 @@ impl litellm_gateway_auth::IdentityResolver for TestPermissions {
|
|||
) -> litellm_gateway_auth::AuthFuture<'a, litellm_gateway_auth::ResolvedIdentity> {
|
||||
Box::pin(async move {
|
||||
Ok(litellm_gateway_auth::ResolvedIdentity {
|
||||
principal: identity.principal.clone(),
|
||||
principal: self.1.clone().unwrap_or_else(|| identity.principal.clone()),
|
||||
permissions: self.0.clone(),
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -66,6 +66,11 @@ where
|
|||
MachineStep::Suspended(request) => request,
|
||||
};
|
||||
let answered = match request {
|
||||
HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply }) => self
|
||||
.interceptors
|
||||
.result_ready(facts)
|
||||
.await
|
||||
.map(|()| reply.send(())),
|
||||
HostRequest::HostCall(call) => self.services.handle_host_call(call).await,
|
||||
HostRequest::Intercept(InterceptRequest::BeforeProviderRequest {
|
||||
wire,
|
||||
|
|
|
|||
|
|
@ -64,6 +64,7 @@ enum EventNext {
|
|||
}
|
||||
|
||||
enum Pending<L> {
|
||||
Host,
|
||||
Native,
|
||||
Arguments(HookResume<L, Py<PyDict>>),
|
||||
Wire(HookResume<L, Box<WireRequest>>, Reply<WireRequest>),
|
||||
|
|
@ -199,6 +200,19 @@ where
|
|||
Err(error) => self.hook_failed(py, error),
|
||||
}
|
||||
}
|
||||
(Some(Pending::Host), Some(result)) => {
|
||||
match self.binding.resume_host_call(py, result) {
|
||||
Ok(Some(awaitable)) => {
|
||||
self.pending = Some(Pending::Host);
|
||||
Ok(ExecutionStep::Await(awaitable))
|
||||
}
|
||||
Ok(None) => self.resume_machine(py, None),
|
||||
Err(InvokeError::Python(error)) => self.interrupt(py, error),
|
||||
Err(InvokeError::Native(error)) => {
|
||||
self.resume_machine(py, Some(HostFailure::Error(error)))
|
||||
}
|
||||
}
|
||||
}
|
||||
(Some(Pending::Native), Some(Ok(_))) => {
|
||||
let result = self.native.take_result()?;
|
||||
self.run_steps(py, NativePoll::Ready(result))
|
||||
|
|
@ -405,7 +419,13 @@ where
|
|||
Err(error) => return self.machine_failed(py, error).map(Next::Return),
|
||||
};
|
||||
let answered = match op {
|
||||
HostRequest::HostCall(op) => answered(self.binding.handle_host_call(py, op)),
|
||||
HostRequest::HostCall(op) => match self.binding.begin_host_call(py, op) {
|
||||
Ok(Some(awaitable)) => {
|
||||
self.pending = Some(Pending::Host);
|
||||
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
|
||||
}
|
||||
result => answered(result.map(|_| ())),
|
||||
},
|
||||
HostRequest::Intercept(InterceptRequest::BeforeProviderRequest {
|
||||
wire,
|
||||
context,
|
||||
|
|
@ -427,6 +447,21 @@ where
|
|||
HostRequest::Stream(StreamDelivery::Chunk(chunk, reply)) => {
|
||||
return self.delivered(py, chunk, reply).map(Next::Return);
|
||||
}
|
||||
HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply }) => {
|
||||
let event = PythonCallEvent::Execution(ExecutionEvent::ResultReady { facts });
|
||||
self.observe(&event);
|
||||
match self.hooks.on_event(py, event) {
|
||||
Ok(HookStep::Ready(())) => {
|
||||
reply.send(());
|
||||
Ok(Ok(()))
|
||||
}
|
||||
Ok(HookStep::Await(awaitable, resume)) => {
|
||||
self.pending = Some(Pending::Event(resume, EventNext::Emitted(reply)));
|
||||
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
HostRequest::Intercept(InterceptRequest::AfterProviderResponse { raw, reply }) => {
|
||||
let event = PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived {
|
||||
raw: &raw,
|
||||
|
|
@ -466,7 +501,7 @@ where
|
|||
Ok(head) => head,
|
||||
Err(error) => return self.interrupt(py, error),
|
||||
};
|
||||
match self.hooks.on_stream_open(py) {
|
||||
match self.hooks.on_stream_open(py, &head) {
|
||||
Ok(()) => {
|
||||
self.pending = Some(Pending::Consumer(reply));
|
||||
Ok(ExecutionStep::Open(head))
|
||||
|
|
@ -770,12 +805,15 @@ mod tests {
|
|||
RejectNatively,
|
||||
RejectRequestNatively,
|
||||
RaiseRequestPython,
|
||||
AwaitAnswer,
|
||||
AwaitFailure,
|
||||
}
|
||||
|
||||
struct SyntheticBinding {
|
||||
log: Log,
|
||||
op: OpScript,
|
||||
classifier_fails: bool,
|
||||
pending_reply: Option<Reply<String>>,
|
||||
}
|
||||
|
||||
/// The fake route's public exception, kept as a value so a test sees what `classify`
|
||||
|
|
@ -792,7 +830,7 @@ mod tests {
|
|||
impl SyntheticBinding {
|
||||
fn answer(&self, value: impl FnOnce() -> String) -> Result<String, InvokeError<Error>> {
|
||||
match self.op {
|
||||
OpScript::Answer => Ok(value()),
|
||||
OpScript::Answer | OpScript::AwaitAnswer | OpScript::AwaitFailure => Ok(value()),
|
||||
OpScript::RaisePython | OpScript::RaiseRequestPython => {
|
||||
Err(PyValueError::new_err("op failed").into())
|
||||
}
|
||||
|
|
@ -867,6 +905,41 @@ mod tests {
|
|||
self.answer(|| op.to_string())
|
||||
.map(|answer| reply.send(answer))
|
||||
}
|
||||
|
||||
fn begin_host_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
(op, reply): (&'static str, Reply<String>),
|
||||
) -> Result<Option<Py<PyAny>>, InvokeError<Error>> {
|
||||
if !matches!(self.op, OpScript::AwaitAnswer | OpScript::AwaitFailure) {
|
||||
return self.handle_host_call(py, (op, reply)).map(|()| None);
|
||||
}
|
||||
self.pending_reply = Some(reply);
|
||||
let module = PyModule::from_code(
|
||||
py,
|
||||
pyo3::ffi::c_str!(
|
||||
"async def answer(fail):\n if fail:\n raise LookupError('async host failed')\n return 'awaited'\n"
|
||||
),
|
||||
pyo3::ffi::c_str!("host_op.py"),
|
||||
pyo3::ffi::c_str!("host_op"),
|
||||
)?;
|
||||
Ok(Some(
|
||||
module
|
||||
.getattr("answer")?
|
||||
.call1((matches!(self.op, OpScript::AwaitFailure),))?
|
||||
.unbind(),
|
||||
))
|
||||
}
|
||||
|
||||
fn resume_host_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: PyResult<Py<PyAny>>,
|
||||
) -> Result<Option<Py<PyAny>>, InvokeError<Error>> {
|
||||
let answer = result?.extract::<String>(py)?;
|
||||
self.pending_reply.take().unwrap().send(answer);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
impl PythonOwned for SyntheticBinding {
|
||||
|
|
@ -988,6 +1061,9 @@ mod tests {
|
|||
return Err(PyValueError::new_err("callback failed"));
|
||||
}
|
||||
self.log.push(match event {
|
||||
PythonCallEvent::Execution(ExecutionEvent::ResultReady { .. }) => {
|
||||
"cache_hit".into()
|
||||
}
|
||||
PythonCallEvent::Started { .. } => "started".into(),
|
||||
PythonCallEvent::Cancelled { .. } => "cancelled".into(),
|
||||
PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
|
|
@ -1003,7 +1079,7 @@ mod tests {
|
|||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, _: Python<'_>) -> PyResult<()> {
|
||||
fn on_stream_open(&mut self, _: Python<'_>, _: &Py<PyAny>) -> PyResult<()> {
|
||||
self.log.push("opened");
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -1037,6 +1113,7 @@ mod tests {
|
|||
log: Log::default(),
|
||||
op,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
script,
|
||||
asynchronous,
|
||||
|
|
@ -1060,6 +1137,50 @@ mod tests {
|
|||
)
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::success(OpScript::AwaitAnswer)]
|
||||
#[case::failure(OpScript::AwaitFailure)]
|
||||
fn asynchronous_host_operations_resume_the_same_machine(#[case] op: OpScript) {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
.lock()
|
||||
.unwrap_or_else(std::sync::PoisonError::into_inner);
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
install_lifecycle_module(py);
|
||||
let (result, log) = run_scripted(
|
||||
py,
|
||||
|_| {
|
||||
CallMachine::<Synthetic>::new(None, |host| {
|
||||
Box::pin(async move { host.services.call(|reply| ("read", reply)).await })
|
||||
})
|
||||
},
|
||||
op,
|
||||
HookScript::Plain,
|
||||
true,
|
||||
);
|
||||
match op {
|
||||
OpScript::AwaitAnswer => {
|
||||
assert_eq!(result.unwrap().extract::<String>(py).unwrap(), "awaited")
|
||||
}
|
||||
OpScript::AwaitFailure => {
|
||||
assert!(
|
||||
result
|
||||
.unwrap_err()
|
||||
.is_instance_of::<pyo3::exceptions::PyLookupError>(py)
|
||||
);
|
||||
assert_eq!(
|
||||
log.iter()
|
||||
.filter(|entry| entry.starts_with("failed:"))
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
assert!(!log.iter().any(|entry| entry.starts_with("succeeded:")));
|
||||
}
|
||||
_ => unreachable!(),
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::synchronous(false)]
|
||||
#[case::asynchronous(true)]
|
||||
|
|
@ -1086,6 +1207,7 @@ mod tests {
|
|||
log: Log(log.0.clone()),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
hooks,
|
||||
PyDict::new(py).unbind(),
|
||||
|
|
@ -1223,6 +1345,7 @@ mod tests {
|
|||
log: Log::default(),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
HookScript::ReplaceResponse,
|
||||
std::convert::identity,
|
||||
|
|
@ -1282,6 +1405,7 @@ mod tests {
|
|||
log: Log::default(),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
script,
|
||||
std::convert::identity,
|
||||
|
|
@ -1787,6 +1911,7 @@ mod tests {
|
|||
log: Log::default(),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: true,
|
||||
pending_reply: None,
|
||||
},
|
||||
HookScript::Plain,
|
||||
false,
|
||||
|
|
@ -1902,6 +2027,7 @@ mod tests {
|
|||
log: Log(log.0.clone()),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
crate::HookChain::new()
|
||||
.with(SyntheticHooks {
|
||||
|
|
@ -1984,6 +2110,7 @@ mod tests {
|
|||
log: Log(log.0.clone()),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
SyntheticHooks {
|
||||
log: Log(log.0.clone()),
|
||||
|
|
@ -2020,6 +2147,7 @@ mod tests {
|
|||
log: Log::default(),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
HookScript::Plain,
|
||||
|hooks| {
|
||||
|
|
@ -2062,6 +2190,7 @@ mod tests {
|
|||
log: Log::default(),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
pending_reply: None,
|
||||
},
|
||||
HookScript::Plain,
|
||||
|hooks| {
|
||||
|
|
|
|||
|
|
@ -89,7 +89,7 @@ pub(super) trait ChainHooks: PythonOwned {
|
|||
result: PyResult<Py<PyAny>>,
|
||||
) -> PyResult<ChainStep<()>>;
|
||||
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()>;
|
||||
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()>;
|
||||
fn on_stream_open(&mut self, py: Python<'_>, head: &Py<PyAny>) -> PyResult<()>;
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()>;
|
||||
}
|
||||
|
||||
|
|
@ -181,8 +181,8 @@ impl<H: PythonCallHooks> ChainHooks for HookAdapter<H> {
|
|||
self.hooks.arguments_prepared(py, arguments)
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
self.hooks.on_stream_open(py)
|
||||
fn on_stream_open(&mut self, py: Python<'_>, head: &Py<PyAny>) -> PyResult<()> {
|
||||
self.hooks.on_stream_open(py, head)
|
||||
}
|
||||
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
|
|
|
|||
|
|
@ -236,6 +236,9 @@ fn notification_result(
|
|||
|
||||
fn retain_event(py: Python<'_>, event: PythonCallEvent<'_>) -> OwnedEvent {
|
||||
match event {
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady { facts }) => {
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady { facts })
|
||||
}
|
||||
CallEvent::Started { start_time } => CallEvent::Started { start_time },
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw: raw.clone() })
|
||||
|
|
@ -263,6 +266,12 @@ fn dispatch(
|
|||
event: &OwnedEvent,
|
||||
) -> PyResult<ChainStep<()>> {
|
||||
match event {
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady { facts }) => hooks.on_event(
|
||||
py,
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: facts.clone(),
|
||||
}),
|
||||
),
|
||||
CallEvent::Started { start_time } => hooks.on_event(
|
||||
py,
|
||||
CallEvent::Started {
|
||||
|
|
@ -350,10 +359,10 @@ impl CallHooks<PythonRuntime> for HookChain {
|
|||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
fn on_stream_open(&mut self, py: Python<'_>, head: &Py<PyAny>) -> PyResult<()> {
|
||||
self.hooks
|
||||
.iter_mut()
|
||||
.try_for_each(|hooks| hooks.on_stream_open(py))
|
||||
.try_for_each(|hooks| hooks.on_stream_open(py, head))
|
||||
}
|
||||
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
|
|
|
|||
|
|
@ -8,4 +8,20 @@ pub trait PythonHostCalls<P: Protocol>: PythonOwned {
|
|||
py: Python<'_>,
|
||||
call: P::HostCall,
|
||||
) -> Result<(), InvokeError<P::Error>>;
|
||||
|
||||
fn begin_host_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
call: P::HostCall,
|
||||
) -> Result<Option<Py<PyAny>>, InvokeError<P::Error>> {
|
||||
self.handle_host_call(py, call).map(|()| None)
|
||||
}
|
||||
|
||||
fn resume_host_call(
|
||||
&mut self,
|
||||
_: Python<'_>,
|
||||
result: PyResult<Py<PyAny>>,
|
||||
) -> Result<Option<Py<PyAny>>, InvokeError<P::Error>> {
|
||||
result.map(|_| None).map_err(InvokeError::Python)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -158,7 +158,7 @@ impl CallHooks<PythonRuntime> for ScriptHooks {
|
|||
}
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
fn on_stream_open(&mut self, py: Python<'_>, _head: &Py<PyAny>) -> PyResult<()> {
|
||||
self.object.call_method1(py, "stream", (py.None(),))?;
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -307,7 +307,7 @@ fn transformations_feed_each_other_and_notifications_share_final_values(
|
|||
)
|
||||
.unwrap();
|
||||
finish(py, &mut hooks, step).unwrap();
|
||||
hooks.on_stream_open(py).unwrap();
|
||||
hooks.on_stream_open(py, &py.None()).unwrap();
|
||||
hooks.on_stream_chunk(py, &response).unwrap();
|
||||
let locals = scripts.bind(py);
|
||||
locals.set_item("arguments", arguments).unwrap();
|
||||
|
|
|
|||
|
|
@ -25,3 +25,5 @@ Rust handlers answer suspensions through `litellm-host-native::Driver`, which `l
|
|||
Keep API policy in gateway-inference and python-bridge, and legacy callback policy in callbacks-legacy-python. Python bindings and hooks expose retained references through `PythonOwned`, with idempotent close and GC traversal. Runtime machinery stays in driver, native, handle and runtime modules
|
||||
|
||||
Interceptors run inline and can rewrite values or fail execution. Observers consume owned `CallEvent` snapshots from `observation_channel`; its bounded `ObservationSender` never waits for delivery and counts events dropped when the queue is full or closed. The host owns receiver processing and draining. Pass the same publisher to machine construction and the driver when one receiver should collect execution and lifecycle events. Legacy Python callbacks retain their existing awaited, fallible behavior through the Python adapter
|
||||
|
||||
`ExecutionFacts` and `ResultSource` describe execution without pricing or budget policy. `Interceptors::result_ready` delivers these facts through an awaited `InterceptRequest::ResultReady`; hosts receive them before response transformation or stream delivery. `ExecutionEvent::ResultReady` is the matching lifecycle event and can also be published as a passive snapshot. Accounting must consume the awaited path rather than a lossy observation queue
|
||||
|
|
|
|||
|
|
@ -61,7 +61,11 @@ pub trait CallHooks<R: HookRuntime>: Sized {
|
|||
Ok(R::ready(()))
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, _runtime: R::Context<'_>) -> Result<(), R::Error> {
|
||||
fn on_stream_open(
|
||||
&mut self,
|
||||
_runtime: R::Context<'_>,
|
||||
_head: &R::Response,
|
||||
) -> Result<(), R::Error> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -30,7 +30,29 @@ pub struct RawResponse {
|
|||
pub body: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct ProviderIdentity {
|
||||
pub model: String,
|
||||
pub provider: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum ResultSource {
|
||||
Provider,
|
||||
Cache { key: String },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct ExecutionFacts {
|
||||
pub provider: ProviderIdentity,
|
||||
pub source: ResultSource,
|
||||
}
|
||||
|
||||
pub trait Interceptors<E>: Send + Sync {
|
||||
fn result_ready(&self, _facts: ExecutionFacts) -> impl Future<Output = Result<(), E>> + Send {
|
||||
async { Ok(()) }
|
||||
}
|
||||
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
|
|
@ -44,6 +66,10 @@ pub trait Interceptors<E>: Send + Sync {
|
|||
}
|
||||
|
||||
impl<E, T: Interceptors<E> + ?Sized> Interceptors<E> for &T {
|
||||
fn result_ready(&self, facts: ExecutionFacts) -> impl Future<Output = Result<(), E>> + Send {
|
||||
(**self).result_ready(facts)
|
||||
}
|
||||
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
|
|
|
|||
|
|
@ -52,7 +52,12 @@ pub enum CallEvent<Response = (), Error = (), Raw = RawResponse> {
|
|||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum ExecutionEvent<Raw = RawResponse> {
|
||||
ProviderResponseReceived { raw: Raw },
|
||||
ResultReady {
|
||||
facts: crate::interceptors::ExecutionFacts,
|
||||
},
|
||||
ProviderResponseReceived {
|
||||
raw: Raw,
|
||||
},
|
||||
}
|
||||
|
||||
impl<Response, Error, Raw: std::borrow::Borrow<RawResponse>> CallEvent<Response, Error, Raw> {
|
||||
|
|
@ -66,6 +71,11 @@ impl<Response, Error, Raw: std::borrow::Borrow<RawResponse>> CallEvent<Response,
|
|||
raw: raw.borrow().clone(),
|
||||
})
|
||||
}
|
||||
Self::Execution(ExecutionEvent::ResultReady { facts }) => {
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: facts.clone(),
|
||||
})
|
||||
}
|
||||
Self::Succeeded { timing, .. } => CallEvent::Succeeded {
|
||||
timing: *timing,
|
||||
response: (),
|
||||
|
|
|
|||
|
|
@ -101,6 +101,17 @@ where
|
|||
.await
|
||||
}
|
||||
|
||||
async fn result_ready(
|
||||
&self,
|
||||
facts: crate::interceptors::ExecutionFacts,
|
||||
) -> Result<(), P::Error> {
|
||||
self.0
|
||||
.request_reply(|reply| {
|
||||
HostRequest::Intercept(InterceptRequest::ResultReady { facts, reply })
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), P::Error> {
|
||||
self.0
|
||||
.request_reply(|reply| {
|
||||
|
|
|
|||
|
|
@ -20,6 +20,10 @@ pub enum HostRequest<P: Protocol> {
|
|||
}
|
||||
|
||||
pub enum InterceptRequest {
|
||||
ResultReady {
|
||||
facts: crate::interceptors::ExecutionFacts,
|
||||
reply: Reply<()>,
|
||||
},
|
||||
BeforeProviderRequest {
|
||||
wire: Box<WireRequest>,
|
||||
context: Box<RequestContext>,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,10 @@
|
|||
|
||||
This folder owns Python cache API compatibility: argument projection, facade identity, public result construction, Python embedding calls and per-operation composition of native backends. Cache algorithms, storage protocols and response-cache semantics belong to their cache crates
|
||||
|
||||
`mod.rs` exposes the cache boundary to routes and module registration; adapter directories remain private. `selection.rs` owns global cache selection, route admission and inference protocol composition for both adapters. `runtime.rs` exposes the Python-facing runtime that can wrap either native storage or a Python callback. `future.rs` converts cache results into ready Futures
|
||||
|
||||
`python/` delegates operations to the selected Python cache without discovering configuration. `native/` owns native backend construction, configuration projection, facade validation, embedding and storage bindings, including experimental V2 handles. Neither adapter depends on shared selection or the other adapter. Shared composition depends on the adapters, and routes use only the parent module's exports
|
||||
|
||||
`SemanticExecution` belongs here because its steps select cache operations and invoke the Python embedder. Use the shared `Execution` handle and inline lifecycle driver; do not duplicate coroutine state validation, runtime waiting or GIL machinery. Python embedding awaits stay in the caller's task, and cancellation must prevent later backend or batch operations from starting
|
||||
|
||||
Resolved asyncio Future construction is generic host machinery. Use `litellm-host-python::ready_future` with an already constructed Python value. Keep cache-specific conversion and disabled-cache return values here. Preserve the Future-returning API and running-loop requirement
|
||||
|
|
|
|||
|
|
@ -1,363 +0,0 @@
|
|||
use crate::http::host_client;
|
||||
use crate::logger::run_sync_value;
|
||||
use litellm_auth_aws::AwsAuthConfig;
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
use litellm_cache_qdrant_semantic::{OpenAiEmbedderConfig, Quantization};
|
||||
use litellm_cache_redis::{RedisNode, RedisTopology};
|
||||
use litellm_cache_redis_semantic::RedisSemanticConfig;
|
||||
use litellm_cache_s3::{S3CacheConfig, S3Endpoint};
|
||||
use litellm_host_python::release_gil;
|
||||
use litellm_http::ClientVariant;
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
exceptions::{PyRuntimeError, PyTypeError},
|
||||
prelude::*,
|
||||
};
|
||||
use url::Url;
|
||||
|
||||
use super::{
|
||||
cache_error,
|
||||
config::{QdrantSemanticCacheConfig, project_redis_semantic},
|
||||
embedder::PythonEmbedder,
|
||||
facade::FacadeGuard,
|
||||
native::NativeResponseCache,
|
||||
request::duration,
|
||||
};
|
||||
|
||||
#[pyclass(frozen, name = "_CacheTestHandle")]
|
||||
pub(crate) struct CacheTestHandle {
|
||||
service: NativeResponseCache,
|
||||
pub(super) guard: Option<FacadeGuard>,
|
||||
pid: u32,
|
||||
}
|
||||
|
||||
impl CacheTestHandle {
|
||||
pub(super) fn service(&self) -> PyResult<NativeResponseCache> {
|
||||
if self.pid != std::process::id() {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"native cache handles must be recreated after fork",
|
||||
));
|
||||
}
|
||||
Ok(self.service.clone())
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl CacheTestHandle {
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (*, capacity=200, ttl_seconds=600.0, max_entry_bytes=1048576))]
|
||||
fn memory(capacity: usize, ttl_seconds: f64, max_entry_bytes: usize) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
service: NativeResponseCache::memory(capacity, duration(ttl_seconds)?, max_entry_bytes),
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (url, *, ttl_seconds=60.0, namespace=None, startup_nodes=None))]
|
||||
fn redis(
|
||||
py: Python<'_>,
|
||||
url: String,
|
||||
ttl_seconds: f64,
|
||||
namespace: Option<String>,
|
||||
startup_nodes: Option<Vec<(String, u16)>>,
|
||||
) -> PyResult<Self> {
|
||||
let ttl = Some(duration(ttl_seconds)?);
|
||||
let topology = match startup_nodes {
|
||||
None => RedisTopology::Standalone,
|
||||
Some(nodes) => RedisTopology::Cluster {
|
||||
startup_nodes: nodes
|
||||
.into_iter()
|
||||
.map(|(host, port)| RedisNode { host, port })
|
||||
.collect(),
|
||||
},
|
||||
};
|
||||
let service = release_gil(py, move || {
|
||||
NativeResponseCache::redis(&url, &topology, ttl, namespace)
|
||||
})
|
||||
.map_err(cache_error)?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
#[pyo3(signature = (bucket, *, region, endpoint_url=None, key_prefix="", access_key_id=None, secret_access_key=None, session_token=None))]
|
||||
fn s3(
|
||||
py: Python<'_>,
|
||||
bucket: String,
|
||||
region: String,
|
||||
endpoint_url: Option<String>,
|
||||
key_prefix: &str,
|
||||
access_key_id: Option<String>,
|
||||
secret_access_key: Option<String>,
|
||||
session_token: Option<String>,
|
||||
) -> PyResult<Self> {
|
||||
let config = S3CacheConfig {
|
||||
bucket,
|
||||
key_prefix: key_prefix.to_string(),
|
||||
region: region.clone(),
|
||||
endpoint: endpoint_url.map(|url| S3Endpoint { url }),
|
||||
auth: AwsAuthConfig {
|
||||
access_key_id,
|
||||
secret_access_key,
|
||||
session_token,
|
||||
region_name: Some(region),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let http = host_client(py, ClientVariant::NoRedirect)?;
|
||||
let service = run_sync_value(py, async move {
|
||||
Ok(NativeResponseCache::s3(config, http).await)
|
||||
})?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (bucket_name, *, gcs_path=None, path_service_account=None, endpoint=None, token=None))]
|
||||
fn gcs(
|
||||
py: Python<'_>,
|
||||
bucket_name: String,
|
||||
gcs_path: Option<String>,
|
||||
path_service_account: Option<String>,
|
||||
endpoint: Option<String>,
|
||||
token: Option<String>,
|
||||
) -> PyResult<Self> {
|
||||
let config = GcsConfig {
|
||||
bucket_name,
|
||||
gcs_path,
|
||||
path_service_account,
|
||||
endpoint: endpoint.unwrap_or_else(|| DEFAULT_ENDPOINT.to_string()),
|
||||
};
|
||||
let client = host_client(py, ClientVariant::NoRedirect)?;
|
||||
let service = NativeResponseCache::gcs(config, client, token);
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (directory))]
|
||||
fn disk(py: Python<'_>, directory: String) -> PyResult<Self> {
|
||||
let service =
|
||||
release_gil(py, move || NativeResponseCache::disk(&directory)).map_err(cache_error)?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (url, *, collection_name, similarity_threshold, vector_size, embedding_model="text-embedding-3-small", api_key=None, embedding_api_key=None, embedding_api_base=None, embedding_timeout_seconds=None, quantization="binary"))]
|
||||
#[expect(
|
||||
clippy::too_many_arguments,
|
||||
reason = "the test handle exposes the complete Qdrant constructor"
|
||||
)]
|
||||
fn qdrant_semantic(
|
||||
py: Python<'_>,
|
||||
url: String,
|
||||
collection_name: String,
|
||||
similarity_threshold: f64,
|
||||
vector_size: u64,
|
||||
embedding_model: &str,
|
||||
api_key: Option<String>,
|
||||
embedding_api_key: Option<String>,
|
||||
embedding_api_base: Option<String>,
|
||||
embedding_timeout_seconds: Option<f64>,
|
||||
quantization: &str,
|
||||
) -> PyResult<Self> {
|
||||
let parsed = Url::parse(&url).map_err(|_| {
|
||||
pyo3::exceptions::PyValueError::new_err(
|
||||
"native Qdrant requires the default REST port so the gRPC port can be derived",
|
||||
)
|
||||
})?;
|
||||
if !matches!(parsed.scheme(), "http" | "https")
|
||||
|| (!parsed.path().is_empty() && parsed.path() != "/")
|
||||
|| parsed.query().is_some()
|
||||
|| parsed.host_str().is_none()
|
||||
|| parsed.port() != Some(6333)
|
||||
{
|
||||
return Err(pyo3::exceptions::PyValueError::new_err(
|
||||
"native Qdrant requires the default REST port so the gRPC port can be derived",
|
||||
));
|
||||
}
|
||||
let mut grpc_url = parsed;
|
||||
grpc_url.set_port(Some(6334)).map_err(|_| {
|
||||
pyo3::exceptions::PyValueError::new_err(
|
||||
"native Qdrant requires the default REST port so the gRPC port can be derived",
|
||||
)
|
||||
})?;
|
||||
grpc_url.set_path("");
|
||||
grpc_url.set_query(None);
|
||||
let embedding_api_key = embedding_api_key
|
||||
.or_else(|| {
|
||||
std::env::var("OPENAI_API_KEY")
|
||||
.ok()
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
pyo3::exceptions::PyValueError::new_err(
|
||||
"native semantic embedding requires an OpenAI API key",
|
||||
)
|
||||
})?;
|
||||
let embedding_api_base = embedding_api_base.unwrap_or_else(|| {
|
||||
std::env::var("OPENAI_BASE_URL")
|
||||
.or_else(|_| std::env::var("OPENAI_API_BASE"))
|
||||
.unwrap_or_else(|_| "https://api.openai.com/v1".to_owned())
|
||||
});
|
||||
let quantization = match quantization {
|
||||
"binary" => Quantization::Binary,
|
||||
"scalar" => Quantization::Scalar,
|
||||
"product" => Quantization::Product,
|
||||
_ => {
|
||||
return Err(pyo3::exceptions::PyValueError::new_err(
|
||||
"unsupported Qdrant quantization",
|
||||
));
|
||||
}
|
||||
};
|
||||
let config = QdrantSemanticCacheConfig {
|
||||
grpc_url: grpc_url.to_string().trim_end_matches('/').to_owned(),
|
||||
api_key,
|
||||
collection_name,
|
||||
similarity_threshold,
|
||||
vector_size,
|
||||
embedding: OpenAiEmbedderConfig {
|
||||
api_base: embedding_api_base,
|
||||
api_key: embedding_api_key,
|
||||
model: embedding_model.to_owned(),
|
||||
timeout: embedding_timeout_seconds.map(duration).transpose()?,
|
||||
},
|
||||
quantization,
|
||||
};
|
||||
let client = host_client(py, ClientVariant::Provider)?;
|
||||
let service = run_sync_value(py, async move {
|
||||
let handle = tokio::runtime::Handle::current();
|
||||
NativeResponseCache::qdrant_semantic(config, client, handle)
|
||||
.await
|
||||
.map_err(cache_error)
|
||||
})?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (url, similarity_threshold, index_name, embedder))]
|
||||
fn valkey_semantic(
|
||||
url: String,
|
||||
similarity_threshold: f64,
|
||||
index_name: String,
|
||||
embedder: &Bound<'_, PyAny>,
|
||||
) -> PyResult<Self> {
|
||||
let python_embedder = PythonEmbedder::new(embedder.clone().unbind());
|
||||
let service = NativeResponseCache::valkey_semantic(
|
||||
&url,
|
||||
similarity_threshold,
|
||||
index_name,
|
||||
python_embedder,
|
||||
)
|
||||
.map_err(cache_error)?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (account_url, container))]
|
||||
fn azure_blob(py: Python<'_>, account_url: String, container: String) -> PyResult<Self> {
|
||||
let http = host_client(py, ClientVariant::NoRedirect)?;
|
||||
let service = run_sync_value(py, async move {
|
||||
NativeResponseCache::azure_blob(&account_url, &container, http)
|
||||
.await
|
||||
.map_err(cache_error)
|
||||
})?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
fn redis_semantic(py: Python<'_>, backend: Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
let class = py
|
||||
.import("litellm.caching.redis_semantic_cache")?
|
||||
.getattr("RedisSemanticCache")?;
|
||||
if !backend.get_type().is(&class) {
|
||||
return Err(PyTypeError::new_err(
|
||||
"native redis-semantic handles require the built-in RedisSemanticCache",
|
||||
));
|
||||
}
|
||||
let config = project_redis_semantic(&backend)?;
|
||||
let embedder = PythonEmbedder::new(backend.unbind());
|
||||
let service = release_gil(py, move || {
|
||||
NativeResponseCache::redis_semantic(
|
||||
&config.redis_url,
|
||||
embedder,
|
||||
RedisSemanticConfig {
|
||||
index_name: config.index_name,
|
||||
similarity_threshold: config.similarity_threshold as f32,
|
||||
},
|
||||
)
|
||||
})
|
||||
.map_err(cache_error)?;
|
||||
Ok(Self {
|
||||
service,
|
||||
guard: None,
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[getter]
|
||||
fn backend(&self) -> &'static str {
|
||||
self.service.kind()
|
||||
}
|
||||
|
||||
fn _bind_facade(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<()> {
|
||||
let service = self.service()?;
|
||||
let guard = FacadeGuard::capture(py, facade, &service)?;
|
||||
let service = service
|
||||
.with_scope(
|
||||
facade
|
||||
.getattr("semantic_cache_scope")?
|
||||
.extract::<String>()?,
|
||||
)
|
||||
.with_redis_flush_size(
|
||||
facade
|
||||
.getattr("redis_flush_size")?
|
||||
.extract::<Option<usize>>()?,
|
||||
);
|
||||
let handle = Py::new(
|
||||
py,
|
||||
Self {
|
||||
service,
|
||||
guard: Some(guard),
|
||||
pid: self.pid,
|
||||
},
|
||||
)?;
|
||||
facade.setattr("_native_cache_handle", handle)
|
||||
}
|
||||
|
||||
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
self.service.traverse(&visit)?;
|
||||
if let Some(guard) = &self.guard {
|
||||
guard.traverse(visit)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
@ -1,16 +1,13 @@
|
|||
mod activation;
|
||||
mod binding;
|
||||
mod callback;
|
||||
mod config;
|
||||
mod embedder;
|
||||
mod facade;
|
||||
mod future;
|
||||
mod handle;
|
||||
mod identity;
|
||||
mod native;
|
||||
mod request;
|
||||
mod resolver;
|
||||
mod semantic;
|
||||
mod python;
|
||||
mod runtime;
|
||||
mod selection;
|
||||
|
||||
pub(crate) use native::NativeCacheHandle;
|
||||
pub(crate) use python::{CacheCall, PythonCache};
|
||||
pub(crate) use runtime::ResolvedCache;
|
||||
pub(crate) use selection::{Cached, Selection, admit_native, configure, configured_native};
|
||||
|
||||
use litellm_cache::Error;
|
||||
use pyo3::{
|
||||
|
|
@ -18,8 +15,6 @@ use pyo3::{
|
|||
prelude::*,
|
||||
};
|
||||
|
||||
pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver};
|
||||
|
||||
fn cache_error(error: Error) -> PyErr {
|
||||
match error {
|
||||
Error::InvalidEntry => PyValueError::new_err(error.to_string()),
|
||||
|
|
|
|||
9
litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md
vendored
Normal file
9
litellm-rust/crates/python-bridge/src/cache/native/AGENTS.md
vendored
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
# Native cache bindings
|
||||
|
||||
This directory constructs and exposes Rust cache backends to Python. It owns backend configuration projection, facade validation, native request conversion, semantic embedding integration and experimental V2 handles. Cache algorithms and storage protocols remain in their cache crates
|
||||
|
||||
Accept the cache object or projected configuration selected by the parent module. Do not read global `litellm.cache`, decide route admission, or select the Python cache adapter here
|
||||
|
||||
Keep Python-facing cache classes and method signatures stable when reorganizing modules. Native internals stay private to this directory unless the shared cache boundary or Python module registration needs them. Python embedding awaits use the existing inline lifecycle driver, preserving caller task identity and cancellation
|
||||
|
||||
Verify changes with the existing backend and facade tests using a freshly built extension. Test cache behavior, not module paths or file structure
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::cache::cache_error;
|
||||
use crate::logger::run_sync_value;
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
use litellm_cache_redis_semantic::RedisSemanticConfig;
|
||||
|
|
@ -6,10 +7,9 @@ use litellm_http::ClientVariant;
|
|||
use pyo3::prelude::*;
|
||||
|
||||
use super::{
|
||||
cache_error,
|
||||
backend::NativeResponseCache,
|
||||
config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig},
|
||||
embedder::PythonEmbedder,
|
||||
native::NativeResponseCache,
|
||||
};
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use crate::http::host_client;
|
||||
|
|
@ -20,7 +20,7 @@ fn declined(reason: UnsupportedCacheConfig) -> PyErr {
|
|||
|
||||
/// Builds the native backend a `Cache` facade's projected configuration describes. `backend` is
|
||||
/// the facade's `.cache` object, which owns embedding for the Python-embedded semantic caches.
|
||||
pub(super) fn activate(
|
||||
pub(in crate::cache) fn activate(
|
||||
py: Python<'_>,
|
||||
backend: &Bound<'_, PyAny>,
|
||||
config: NativeCacheConfig,
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::cache::cache_error;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use litellm_cache::{CacheCodec, CacheConnectionResult, Error, semantic::SemanticLookup};
|
||||
|
|
@ -14,7 +15,7 @@ use litellm_cache_response::{
|
|||
};
|
||||
use litellm_cache_s3::{S3Cache, S3CacheConfig};
|
||||
use litellm_cache_valkey_semantic::{ValkeySemanticCache, ValkeySemanticConfig};
|
||||
use pyo3::{PyTraverseError, PyVisit, prelude::*};
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
|
|
@ -26,13 +27,13 @@ use super::{
|
|||
};
|
||||
|
||||
/// What the Python embedder receives for one semantic request.
|
||||
pub(super) struct EmbeddingInput {
|
||||
pub(super) prompt: String,
|
||||
pub(super) metadata: Option<Value>,
|
||||
pub(in crate::cache) struct EmbeddingInput {
|
||||
pub(in crate::cache) prompt: String,
|
||||
pub(in crate::cache) metadata: Option<Value>,
|
||||
}
|
||||
|
||||
/// An exact-match backend behind one pointer, with the identity its facade must reproduce.
|
||||
pub(super) struct ExactService {
|
||||
pub(in crate::cache) struct ExactService {
|
||||
cache: Arc<dyn ExactResponseCache>,
|
||||
probe: Option<Arc<dyn ConnectionProbe>>,
|
||||
buffer: Option<WriteBuffer>,
|
||||
|
|
@ -40,7 +41,7 @@ pub(super) struct ExactService {
|
|||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(super) enum NativeResponseCache {
|
||||
pub(in crate::cache) enum NativeResponseCache {
|
||||
Exact(Arc<ExactService>),
|
||||
ValkeySemantic {
|
||||
cache: Arc<ResponseCache<ValkeySemanticCache<PythonEmbedder, ResponseCacheCodec>>>,
|
||||
|
|
@ -240,7 +241,7 @@ impl NativeResponseCache {
|
|||
})
|
||||
}
|
||||
|
||||
pub async fn qdrant_semantic(
|
||||
pub(super) async fn qdrant_semantic(
|
||||
config: QdrantSemanticCacheConfig,
|
||||
client: litellm_http::Client,
|
||||
runtime: tokio::runtime::Handle,
|
||||
|
|
@ -285,10 +286,6 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn kind(&self) -> &'static str {
|
||||
self.identity().kind()
|
||||
}
|
||||
|
||||
pub fn with_redis_flush_size(self, flush_size: Option<usize>) -> Self {
|
||||
match self {
|
||||
Self::Exact(service) if matches!(service.identity, BackendIdentity::Redis { .. }) => {
|
||||
|
|
@ -324,7 +321,10 @@ impl NativeResponseCache {
|
|||
}
|
||||
|
||||
/// The prompt and metadata this backend would embed for `request`, if it has a prompt.
|
||||
pub(super) fn embedding_input(&self, request: &NativeRequest) -> Option<EmbeddingInput> {
|
||||
pub(in crate::cache) fn embedding_input(
|
||||
&self,
|
||||
request: &NativeRequest,
|
||||
) -> Option<EmbeddingInput> {
|
||||
let context = match self {
|
||||
Self::ValkeySemantic { scope, .. } => request.scoped_semantic(scope).context,
|
||||
Self::RedisSemantic { .. } => request.semantic().context,
|
||||
|
|
@ -462,7 +462,7 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) fn async_lookup_semantic_py<'py>(
|
||||
pub(in crate::cache) fn async_lookup_semantic_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: NativeRequest,
|
||||
|
|
@ -478,7 +478,7 @@ impl NativeResponseCache {
|
|||
.await
|
||||
.map(SemanticReply::from)
|
||||
},
|
||||
super::cache_error,
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
|
|
@ -487,7 +487,7 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) fn async_lookup_py<'py>(
|
||||
pub(in crate::cache) fn async_lookup_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: NativeRequest,
|
||||
|
|
@ -498,7 +498,7 @@ impl NativeResponseCache {
|
|||
crate::logger::run_async(
|
||||
py,
|
||||
async move { service.async_lookup(&request, now()).await },
|
||||
super::cache_error,
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
|
|
@ -541,7 +541,7 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) fn async_store_py<'py>(
|
||||
pub(in crate::cache) fn async_store_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
request: NativeRequest,
|
||||
|
|
@ -553,7 +553,7 @@ impl NativeResponseCache {
|
|||
crate::logger::run_async(
|
||||
py,
|
||||
async move { service.async_store(&request, response, now()).await },
|
||||
super::cache_error,
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
|
|
@ -611,7 +611,7 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) fn async_store_batch_py<'py>(
|
||||
pub(in crate::cache) fn async_store_batch_py<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
entries: Vec<(NativeRequest, Value)>,
|
||||
|
|
@ -622,7 +622,7 @@ impl NativeResponseCache {
|
|||
crate::logger::run_async(
|
||||
py,
|
||||
async move { service.async_store_batch(entries, now()).await },
|
||||
super::cache_error,
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
Self::ValkeySemantic { .. } | Self::RedisSemantic { .. } => {
|
||||
|
|
@ -656,15 +656,6 @@ impl NativeResponseCache {
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
match self {
|
||||
Self::ValkeySemantic { embedder, .. } | Self::RedisSemantic { embedder, .. } => {
|
||||
embedder.traverse(visit)
|
||||
}
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => Ok(()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn exact_requests(requests: &[NativeRequest]) -> Vec<litellm_cache_response::ResponseCacheRequest> {
|
||||
|
|
@ -673,7 +664,10 @@ fn exact_requests(requests: &[NativeRequest]) -> Vec<litellm_cache_response::Res
|
|||
|
||||
/// What `lookup_semantic` hands Python: the response and the similarity to stamp, if any.
|
||||
#[derive(serde::Serialize)]
|
||||
pub(super) struct SemanticReply(pub(super) Option<Value>, pub(super) Option<f64>);
|
||||
pub(in crate::cache) struct SemanticReply(
|
||||
pub(in crate::cache) Option<Value>,
|
||||
pub(in crate::cache) Option<f64>,
|
||||
);
|
||||
|
||||
impl From<SemanticLookup<Value>> for SemanticReply {
|
||||
fn from(lookup: SemanticLookup<Value>) -> Self {
|
||||
|
|
@ -11,7 +11,7 @@ use pyo3::{
|
|||
types::{PyAny, PyBool, PyDict, PyList, PyString},
|
||||
};
|
||||
|
||||
use super::{identity::BackendIdentity, native::NativeResponseCache, request::duration};
|
||||
use super::{backend::NativeResponseCache, identity::BackendIdentity, request::duration};
|
||||
|
||||
pub(super) struct CachePolicy {
|
||||
pub(super) redis_flush_size: Option<usize>,
|
||||
|
|
@ -233,12 +233,12 @@ pub(super) enum CacheBackendConfig {
|
|||
QdrantSemantic(Box<QdrantSemanticCacheConfig>),
|
||||
}
|
||||
|
||||
pub(super) struct NativeCacheConfig {
|
||||
pub(in crate::cache) struct NativeCacheConfig {
|
||||
pub(super) policy: CachePolicy,
|
||||
pub(super) backend: CacheBackendConfig,
|
||||
}
|
||||
|
||||
pub(super) enum UnsupportedCacheConfig {
|
||||
pub(in crate::cache) enum UnsupportedCacheConfig {
|
||||
Backend,
|
||||
RedisTopology,
|
||||
RedisCredentials,
|
||||
|
|
@ -263,7 +263,7 @@ pub(super) enum UnsupportedCacheConfig {
|
|||
}
|
||||
|
||||
impl UnsupportedCacheConfig {
|
||||
pub(super) fn message(&self) -> &'static str {
|
||||
pub(in crate::cache) fn message(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Backend => "native cache backend is not implemented",
|
||||
Self::RedisTopology => "native Redis topology is not implemented",
|
||||
|
|
@ -305,14 +305,14 @@ impl UnsupportedCacheConfig {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) enum CacheConfigProjection {
|
||||
pub(in crate::cache) enum CacheConfigProjection {
|
||||
Native(Box<NativeCacheConfig>),
|
||||
Unsupported(UnsupportedCacheConfig),
|
||||
}
|
||||
|
||||
impl NativeCacheConfig {
|
||||
#[inline(never)]
|
||||
pub(super) fn project(facade: &Bound<'_, PyAny>) -> PyResult<CacheConfigProjection> {
|
||||
pub(in crate::cache) fn project(facade: &Bound<'_, PyAny>) -> PyResult<CacheConfigProjection> {
|
||||
let backend_name = facade.getattr("type")?.extract::<String>()?;
|
||||
let policy = CachePolicy {
|
||||
redis_flush_size: facade
|
||||
|
|
@ -1170,7 +1170,7 @@ mod tests {
|
|||
GcsCacheConfig, NativeCacheConfig, REDIS_PY_DEFAULT_MAX_CONNECTIONS, RedisConnectionConfig,
|
||||
RedisProtocol, RedisSemanticCacheConfig, RedisTlsConfig, UnsupportedCacheConfig,
|
||||
};
|
||||
use crate::cache::{embedder::PythonEmbedder, native::NativeResponseCache};
|
||||
use crate::cache::native::{backend::NativeResponseCache, embedder::PythonEmbedder};
|
||||
|
||||
fn cluster_facade<'py>(py: Python<'py>, startup_nodes: &str, hook: &str) -> Bound<'py, PyAny> {
|
||||
facade(
|
||||
|
|
@ -11,7 +11,7 @@ tokio::task_local! {
|
|||
|
||||
/// Runs `future` with the vector the Python embedder already produced, so the backend's
|
||||
/// `async_embed` never has to call back into Python from the runtime.
|
||||
pub(super) fn with_prepared_embedding<F: Future>(
|
||||
pub(in crate::cache) fn with_prepared_embedding<F: Future>(
|
||||
vector: Result<Vec<f32>, Error>,
|
||||
future: F,
|
||||
) -> impl Future<Output = F::Output> {
|
||||
|
|
@ -19,7 +19,7 @@ pub(super) fn with_prepared_embedding<F: Future>(
|
|||
}
|
||||
|
||||
/// The Python object that owns embedding for a semantic backend.
|
||||
pub(super) struct PythonEmbedder(Py<PyAny>);
|
||||
pub(in crate::cache) struct PythonEmbedder(Py<PyAny>);
|
||||
|
||||
impl Clone for PythonEmbedder {
|
||||
fn clone(&self) -> Self {
|
||||
|
|
@ -28,15 +28,15 @@ impl Clone for PythonEmbedder {
|
|||
}
|
||||
|
||||
impl PythonEmbedder {
|
||||
pub(super) fn new(object: Py<PyAny>) -> Self {
|
||||
pub(in crate::cache) fn new(object: Py<PyAny>) -> Self {
|
||||
Self(object)
|
||||
}
|
||||
|
||||
pub(super) fn object(&self) -> &Py<PyAny> {
|
||||
pub(in crate::cache) fn object(&self) -> &Py<PyAny> {
|
||||
&self.0
|
||||
}
|
||||
|
||||
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.0)
|
||||
}
|
||||
|
||||
|
|
@ -50,7 +50,7 @@ impl PythonEmbedder {
|
|||
}
|
||||
|
||||
/// The awaitable of `_get_async_embedding(prompt, metadata=...)`, to run in the caller's loop.
|
||||
pub(super) fn async_embedding(
|
||||
pub(in crate::cache) fn async_embedding(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
prompt: &str,
|
||||
|
|
@ -63,7 +63,7 @@ impl PythonEmbedder {
|
|||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub(super) fn extract(vector: Bound<'_, PyAny>) -> PyResult<Vec<f32>> {
|
||||
pub(in crate::cache) fn extract(vector: Bound<'_, PyAny>) -> PyResult<Vec<f32>> {
|
||||
Ok(vector
|
||||
.extract::<Vec<f64>>()?
|
||||
.into_iter()
|
||||
|
|
@ -9,10 +9,9 @@ use pyo3::{
|
|||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
backend::NativeResponseCache,
|
||||
config::{CacheConfigProjection, NativeCacheConfig},
|
||||
handle::CacheTestHandle,
|
||||
identity::BackendIdentity,
|
||||
native::NativeResponseCache,
|
||||
};
|
||||
|
||||
struct ClassGuard {
|
||||
|
|
@ -84,7 +83,7 @@ const VALKEY_POOL: RedisPoolAttributes = STANDALONE_POOL;
|
|||
/// `Cache._native_cache` holds the runtime `Cache.__init__` resolved.
|
||||
const INSTANCE_STATE: &[&str] = &["_native_cache"];
|
||||
|
||||
pub(super) struct FacadeGuard {
|
||||
pub(in crate::cache) struct FacadeGuard {
|
||||
outer: ObjectGuard,
|
||||
backend: ObjectGuard,
|
||||
disk_store: Option<DiskStoreGuard>,
|
||||
|
|
@ -354,7 +353,7 @@ impl ConnectionGuard {
|
|||
}
|
||||
|
||||
impl FacadeGuard {
|
||||
pub(super) fn capture(
|
||||
pub(in crate::cache) fn capture(
|
||||
py: Python<'_>,
|
||||
facade: &Bound<'_, PyAny>,
|
||||
service: &NativeResponseCache,
|
||||
|
|
@ -472,7 +471,11 @@ impl FacadeGuard {
|
|||
})
|
||||
}
|
||||
|
||||
pub(super) fn matches(&self, py: Python<'_>, facade: &Bound<'_, PyAny>) -> PyResult<bool> {
|
||||
pub(in crate::cache) fn matches(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
facade: &Bound<'_, PyAny>,
|
||||
) -> PyResult<bool> {
|
||||
if !self.outer.matches(py, facade)? {
|
||||
return Ok(false);
|
||||
}
|
||||
|
|
@ -488,7 +491,7 @@ impl FacadeGuard {
|
|||
self.connection.matches(py, &backend)
|
||||
}
|
||||
|
||||
pub(super) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
pub(in crate::cache) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
self.outer.traverse(&visit)?;
|
||||
self.backend.traverse(&visit)?;
|
||||
if let Some(guard) = &self.disk_store {
|
||||
|
|
@ -497,28 +500,3 @@ impl FacadeGuard {
|
|||
self.connection.traverse(&visit)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn resolve(
|
||||
py: Python<'_>,
|
||||
facade: &Bound<'_, PyAny>,
|
||||
) -> PyResult<Option<NativeResponseCache>> {
|
||||
let Ok(dict) = facade
|
||||
.getattr("__dict__")
|
||||
.and_then(|dict| dict.cast_into::<PyDict>().map_err(Into::into))
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(handle) = dict.get_item("_native_cache_handle")? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Ok(handle) = handle.extract::<PyRef<'_, CacheTestHandle>>() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(guard) = &handle.guard else {
|
||||
return Ok(None);
|
||||
};
|
||||
if !guard.matches(py, facade).unwrap_or(false) {
|
||||
return Ok(None);
|
||||
}
|
||||
handle.service().map(Some)
|
||||
}
|
||||
|
|
@ -6,7 +6,7 @@ use litellm_cache_redis::RedisTopology;
|
|||
/// observe on the Python object, captured once so facade projection and native construction
|
||||
/// compare plain data instead of reaching into each backend type.
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub(super) enum BackendIdentity {
|
||||
pub(in crate::cache) enum BackendIdentity {
|
||||
Memory {
|
||||
capacity: usize,
|
||||
max_entry_bytes: Option<usize>,
|
||||
|
|
@ -55,8 +55,7 @@ pub(super) enum BackendIdentity {
|
|||
const TYPES: &str = "facade and native backend types must match";
|
||||
|
||||
impl BackendIdentity {
|
||||
/// The native backend name reported to Python through `_CacheTestHandle.backend`.
|
||||
pub(super) fn kind(&self) -> &'static str {
|
||||
pub(in crate::cache) fn kind(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Memory { .. } => "memory",
|
||||
Self::Redis { .. } => "redis",
|
||||
|
|
@ -71,7 +70,7 @@ impl BackendIdentity {
|
|||
}
|
||||
|
||||
/// The `LiteLLMCacheType` value a facade of this backend carries in `Cache.type`.
|
||||
pub(super) fn cache_type(&self) -> &'static str {
|
||||
pub(in crate::cache) fn cache_type(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Memory { .. } => "local",
|
||||
Self::Redis { .. } => "redis",
|
||||
|
|
@ -87,7 +86,7 @@ impl BackendIdentity {
|
|||
|
||||
/// The first difference between the facade's configuration (`self`) and the native
|
||||
/// backend (`native`), in the order Python users see the attributes.
|
||||
pub(super) fn mismatch(&self, native: &Self) -> Option<&'static str> {
|
||||
pub(in crate::cache) fn mismatch(&self, native: &Self) -> Option<&'static str> {
|
||||
let mut differences: Vec<(bool, &'static str)> = Vec::new();
|
||||
let mut differs = |condition: bool, message: &'static str| {
|
||||
differences.push((condition, message));
|
||||
11
litellm-rust/crates/python-bridge/src/cache/native/mod.rs
vendored
Normal file
11
litellm-rust/crates/python-bridge/src/cache/native/mod.rs
vendored
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
pub(super) mod activation;
|
||||
pub(super) mod backend;
|
||||
pub(super) mod config;
|
||||
mod embedder;
|
||||
pub(super) mod facade;
|
||||
mod identity;
|
||||
pub(super) mod request;
|
||||
mod semantic;
|
||||
pub(super) mod v2;
|
||||
|
||||
pub(crate) use v2::NativeCacheHandle;
|
||||
|
|
@ -23,7 +23,7 @@ struct RequestInput {
|
|||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(super) struct NativeRequest {
|
||||
pub(in crate::cache) struct NativeRequest {
|
||||
pub(super) key: CacheKeyInput,
|
||||
pub(super) controls: CacheControls,
|
||||
pub(super) ttl: Option<Duration>,
|
||||
|
|
@ -129,7 +129,7 @@ fn semantic_key(request: &NativeRequest, scope: &str) -> CacheKeyInput {
|
|||
key
|
||||
}
|
||||
|
||||
pub(super) fn request(value: &Bound<'_, PyAny>) -> PyResult<NativeRequest> {
|
||||
pub(in crate::cache) fn request(value: &Bound<'_, PyAny>) -> PyResult<NativeRequest> {
|
||||
let input: RequestInput = from_py(value)?;
|
||||
request_input(input)
|
||||
}
|
||||
|
|
@ -152,7 +152,7 @@ fn request_input(input: RequestInput) -> PyResult<NativeRequest> {
|
|||
})
|
||||
}
|
||||
|
||||
pub(super) fn requests(value: &Bound<'_, PyAny>) -> PyResult<Vec<NativeRequest>> {
|
||||
pub(in crate::cache) fn requests(value: &Bound<'_, PyAny>) -> PyResult<Vec<NativeRequest>> {
|
||||
from_py::<Vec<RequestInput>>(value)?
|
||||
.into_iter()
|
||||
.map(request_input)
|
||||
|
|
@ -164,7 +164,7 @@ pub(super) fn duration(seconds: f64) -> PyResult<Duration> {
|
|||
.map_err(|_| PyValueError::new_err("cache durations must be finite and nonnegative"))
|
||||
}
|
||||
|
||||
pub(super) fn now() -> Duration {
|
||||
pub(in crate::cache) fn now() -> Duration {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::cache::cache_error;
|
||||
use crate::logger::run_async;
|
||||
use std::{collections::VecDeque, time::Duration};
|
||||
|
||||
|
|
@ -11,9 +12,8 @@ use pyo3::{
|
|||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
cache_error,
|
||||
backend::{NativeResponseCache, SemanticReply},
|
||||
embedder::{PythonEmbedder, with_prepared_embedding},
|
||||
native::{NativeResponseCache, SemanticReply},
|
||||
request::{NativeRequest, now},
|
||||
};
|
||||
|
||||
345
litellm-rust/crates/python-bridge/src/cache/native/v2.rs
vendored
Normal file
345
litellm-rust/crates/python-bridge/src/cache/native/v2.rs
vendored
Normal file
|
|
@ -0,0 +1,345 @@
|
|||
use crate::cache::cache_error;
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use litellm_cache::{DeleteCache, DisconnectCache, PingCache};
|
||||
use litellm_host_python::{from_py, release_gil, to_py};
|
||||
use serde_json::Value;
|
||||
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_redis::{RedisCache, RedisTopology};
|
||||
use litellm_cache_response::{
|
||||
CacheEntry, CacheKeyInput, ExactResponseCache, ResponseCache, ResponseCacheCodec,
|
||||
ResponseCacheConfig, ResponseCacheRequest, ResponseCacheService,
|
||||
};
|
||||
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
|
||||
|
||||
#[pyclass(
|
||||
frozen,
|
||||
name = "NativeCacheHandle",
|
||||
module = "litellm.rust_bridge._native"
|
||||
)]
|
||||
pub(crate) struct NativeCacheHandle {
|
||||
service: Arc<dyn ResponseCacheService>,
|
||||
backend: Arc<dyn ExactResponseCache>,
|
||||
storage: Storage,
|
||||
pid: u32,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
enum Storage {
|
||||
Memory(Arc<InMemoryCache<CacheEntry>>),
|
||||
Redis(Arc<RedisCache<ResponseCacheCodec>>),
|
||||
}
|
||||
|
||||
impl NativeCacheHandle {
|
||||
fn check_process(&self) -> PyResult<()> {
|
||||
if self.pid != std::process::id() {
|
||||
return Err(pyo3::exceptions::PyRuntimeError::new_err(
|
||||
"recreate the v2 cache after fork",
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn request(key: String, ttl: Option<f64>) -> PyResult<ResponseCacheRequest> {
|
||||
let mut request: ResponseCacheRequest = ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some(key),
|
||||
..Default::default()
|
||||
});
|
||||
request.context.ttl = ttl.map(duration).transpose()?;
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl NativeCacheHandle {
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (*, ttl=600.0, capacity=200, max_entry_bytes=4194304))]
|
||||
fn memory(ttl: f64, capacity: usize, max_entry_bytes: usize) -> PyResult<Self> {
|
||||
let ttl = duration(ttl)?;
|
||||
if capacity == 0 || max_entry_bytes == 0 {
|
||||
return Err(PyValueError::new_err("cache limits must be positive"));
|
||||
}
|
||||
let storage = Arc::new(InMemoryCache::new(Some(capacity), Some(ttl)));
|
||||
let backend = Arc::new(ResponseCache::new(storage.clone()).with_config(
|
||||
ResponseCacheConfig {
|
||||
namespace: "sdk".into(),
|
||||
max_entry_bytes,
|
||||
},
|
||||
));
|
||||
Ok(Self {
|
||||
service: backend.clone(),
|
||||
backend,
|
||||
storage: Storage::Memory(storage),
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
|
||||
#[staticmethod]
|
||||
#[pyo3(signature = (url, *, namespace, ttl=600.0, max_entry_bytes=4194304))]
|
||||
fn redis(
|
||||
py: Python<'_>,
|
||||
url: &str,
|
||||
namespace: String,
|
||||
ttl: f64,
|
||||
max_entry_bytes: usize,
|
||||
) -> PyResult<Self> {
|
||||
let ttl = duration(ttl)?;
|
||||
if namespace.is_empty() || max_entry_bytes == 0 {
|
||||
return Err(PyValueError::new_err(
|
||||
"namespace and a positive cache limit are required",
|
||||
));
|
||||
}
|
||||
let storage = Arc::new(
|
||||
release_gil(py, || {
|
||||
RedisCache::connect(
|
||||
url,
|
||||
&RedisTopology::Standalone,
|
||||
Some(ttl),
|
||||
ResponseCacheCodec,
|
||||
)
|
||||
})
|
||||
.map_err(cache_error)?
|
||||
.with_namespace(Some(namespace.clone())),
|
||||
);
|
||||
let backend = Arc::new(ResponseCache::new(storage.clone()).with_config(
|
||||
ResponseCacheConfig {
|
||||
namespace,
|
||||
max_entry_bytes,
|
||||
},
|
||||
));
|
||||
Ok(Self {
|
||||
service: backend.clone(),
|
||||
backend,
|
||||
storage: Storage::Redis(storage),
|
||||
pid: std::process::id(),
|
||||
})
|
||||
}
|
||||
fn get(&self, py: Python<'_>, key: String) -> PyResult<Py<PyAny>> {
|
||||
self.check_process()?;
|
||||
let request = request(key, None)?;
|
||||
let value = release_gil(py, || self.backend.lookup(&request, super::request::now()))
|
||||
.map_err(cache_error)?;
|
||||
to_py(py, &value)
|
||||
}
|
||||
|
||||
#[pyo3(signature = (key, value, *, ttl=None))]
|
||||
fn set(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
key: String,
|
||||
value: &Bound<'_, PyAny>,
|
||||
ttl: Option<f64>,
|
||||
) -> PyResult<()> {
|
||||
self.check_process()?;
|
||||
let request = request(key, ttl)?;
|
||||
let value: Value = from_py(value)?;
|
||||
release_gil(py, || {
|
||||
self.backend.store(&request, value, super::request::now())
|
||||
})
|
||||
.map_err(cache_error)
|
||||
}
|
||||
|
||||
fn async_get<'py>(&self, py: Python<'py>, key: String) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let request = request(key, None)?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move { backend.async_lookup(&request, super::request::now()).await },
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyo3(signature = (key, value, *, ttl=None))]
|
||||
fn async_set<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
key: String,
|
||||
value: &Bound<'_, PyAny>,
|
||||
ttl: Option<f64>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let request = request(key, ttl)?;
|
||||
let value: Value = from_py(value)?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move {
|
||||
backend
|
||||
.async_store(&request, value, super::request::now())
|
||||
.await
|
||||
},
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyo3(signature = (entries, *, ttl=None))]
|
||||
fn async_set_many<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
entries: &Bound<'_, PyAny>,
|
||||
ttl: Option<f64>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let entries: Vec<(String, Value)> = from_py(entries)?;
|
||||
let entries = entries
|
||||
.into_iter()
|
||||
.map(|(key, value)| Ok((request(key, ttl)?, value)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move {
|
||||
backend
|
||||
.async_store_batch(entries, super::request::now())
|
||||
.await
|
||||
},
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
||||
fn flush(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
|
||||
self.check_process()?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_sync(py, async move { backend.async_flush().await }, cache_error)
|
||||
}
|
||||
|
||||
fn async_flush<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let backend = self.backend.clone();
|
||||
crate::logger::run_async(py, async move { backend.async_flush().await }, cache_error)
|
||||
}
|
||||
|
||||
fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let storage = self.storage.clone();
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move {
|
||||
match storage {
|
||||
Storage::Memory(_) => Ok(true),
|
||||
Storage::Redis(cache) => cache.ping().await,
|
||||
}
|
||||
},
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
||||
fn disconnect<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let storage = self.storage.clone();
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move {
|
||||
match storage {
|
||||
Storage::Memory(cache) => cache.disconnect().await,
|
||||
Storage::Redis(cache) => cache.disconnect().await,
|
||||
}
|
||||
},
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
|
||||
fn delete<'py>(&self, py: Python<'py>, keys: Vec<String>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.check_process()?;
|
||||
let storage = self.storage.clone();
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move {
|
||||
for key in keys {
|
||||
match &storage {
|
||||
Storage::Memory(cache) => cache.async_delete_cache(&key).await?,
|
||||
Storage::Redis(cache) => cache.async_delete_cache(&key).await?,
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
},
|
||||
cache_error,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn duration(seconds: f64) -> PyResult<Duration> {
|
||||
Duration::try_from_secs_f64(seconds)
|
||||
.ok()
|
||||
.filter(|value| !value.is_zero())
|
||||
.ok_or_else(|| PyValueError::new_err("cache durations must be finite and positive"))
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn native_handle<'py>(
|
||||
configured: &Bound<'py, PyAny>,
|
||||
) -> PyResult<Option<Bound<'py, PyAny>>> {
|
||||
Ok(configured
|
||||
.getattr_opt("cache")?
|
||||
.map(|backend| backend.getattr_opt("native_handle"))
|
||||
.transpose()?
|
||||
.flatten()
|
||||
.filter(|handle| handle.is_instance_of::<NativeCacheHandle>()))
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn configured(
|
||||
configured: &Bound<'_, PyAny>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
) -> PyResult<(
|
||||
Option<Arc<dyn ResponseCacheService>>,
|
||||
litellm_cache_response::CacheOptions,
|
||||
)> {
|
||||
let handle = native_handle(configured)?.ok_or_else(|| {
|
||||
pyo3::exceptions::PyRuntimeError::new_err(
|
||||
"the configured cache changed to a Python cache after native admission",
|
||||
)
|
||||
})?;
|
||||
let cache = handle.extract::<PyRef<'_, NativeCacheHandle>>()?;
|
||||
cache.check_process()?;
|
||||
let controls = kwargs.get_item("cache")?.filter(|value| !value.is_none());
|
||||
let controls = controls
|
||||
.as_ref()
|
||||
.map(|value| value.cast::<PyDict>())
|
||||
.transpose()?;
|
||||
if let Some(controls) = controls {
|
||||
for name in controls.keys() {
|
||||
let name = name.extract::<String>()?;
|
||||
if !matches!(
|
||||
name.as_str(),
|
||||
"no-cache" | "no-store" | "ttl" | "s-maxage" | "s-max-age" | "use-cache"
|
||||
) {
|
||||
return Err(PyValueError::new_err(format!(
|
||||
"unsupported v2 cache control: {name}"
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
let boolean = |name: &str| -> PyResult<bool> {
|
||||
controls
|
||||
.map(|values| values.get_item(name))
|
||||
.transpose()?
|
||||
.flatten()
|
||||
.map(|value| value.extract())
|
||||
.transpose()
|
||||
.map(|value| value.unwrap_or(false))
|
||||
};
|
||||
let seconds = |name: &str| -> PyResult<Option<Duration>> {
|
||||
controls
|
||||
.map(|values| values.get_item(name))
|
||||
.transpose()?
|
||||
.flatten()
|
||||
.map(|value| duration(value.extract()?))
|
||||
.transpose()
|
||||
};
|
||||
Ok((
|
||||
Some(cache.service.clone()),
|
||||
litellm_cache_response::CacheOptions {
|
||||
caching: kwargs
|
||||
.get_item("caching")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| value.extract())
|
||||
.transpose()?,
|
||||
no_cache: boolean("no-cache")?,
|
||||
no_store: boolean("no-store")?,
|
||||
ttl: seconds("ttl")?,
|
||||
max_age: seconds("s-max-age")?.or(seconds("s-maxage")?),
|
||||
scope: litellm_cache_response::CacheScope::Shared,
|
||||
},
|
||||
))
|
||||
}
|
||||
9
litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md
vendored
Normal file
9
litellm-rust/crates/python-bridge/src/cache/python/AGENTS.md
vendored
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
# Python cache delegation
|
||||
|
||||
This directory lets Rust inference use a selected Python cache. `service.rs` implements the injected Rust response-cache service and yields typed cache operations. `host.rs` calls the Python cache's sync or async API and delivers the result back to Rust. `callback.rs` provides Python cache delegation for the Python-facing cache runtime
|
||||
|
||||
Keep the shared inference protocol wrapper and adapter selection in the parent module. Receive the configured cache and prepared arguments from the parent module. Do not discover global configuration, choose native backends, or move inference to Python. Core remains independent of Python objects and cache implementation details
|
||||
|
||||
Await asynchronous cache operations through the existing host driver in the caller's task. Do not create another asyncio task or event loop. Cancellation must prevent subsequent provider requests and cache writes. Preserve ordinary cache failure handling without swallowing cancellation or other Python base exceptions
|
||||
|
||||
Traverse retained Python references for GC and release pending replies when the call closes. Regression tests must require native inference with Python fallback disabled and assert observable hits, provider request counts, cache-key headers, task identity and cancellation
|
||||
|
|
@ -5,16 +5,16 @@ use pyo3::{
|
|||
types::{PyDict, PyList, PyTuple},
|
||||
};
|
||||
|
||||
use super::future::ready_none;
|
||||
use crate::cache::future::ready_none;
|
||||
|
||||
pub(super) struct PythonCallback(Py<PyAny>);
|
||||
pub(in crate::cache) struct PythonCallback(Py<PyAny>);
|
||||
|
||||
impl PythonCallback {
|
||||
pub(super) fn new(object: Py<PyAny>) -> Self {
|
||||
pub(in crate::cache) fn new(object: Py<PyAny>) -> Self {
|
||||
Self(object)
|
||||
}
|
||||
|
||||
pub(super) fn lookup<'py>(
|
||||
pub(in crate::cache) fn lookup<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
kwargs: Option<&Bound<'py, PyDict>>,
|
||||
|
|
@ -24,7 +24,7 @@ impl PythonCallback {
|
|||
.call_method("get_cache", (), Some(callback_kwargs(kwargs)?))
|
||||
}
|
||||
|
||||
pub(super) fn async_lookup<'py>(
|
||||
pub(in crate::cache) fn async_lookup<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
kwargs: Option<&Bound<'py, PyDict>>,
|
||||
|
|
@ -34,7 +34,7 @@ impl PythonCallback {
|
|||
.call_method("async_get_cache", (), Some(callback_kwargs(kwargs)?))
|
||||
}
|
||||
|
||||
pub(super) fn store(
|
||||
pub(in crate::cache) fn store(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
response: &Bound<'_, PyAny>,
|
||||
|
|
@ -46,7 +46,7 @@ impl PythonCallback {
|
|||
.map(|_| ())
|
||||
}
|
||||
|
||||
pub(super) fn async_store<'py>(
|
||||
pub(in crate::cache) fn async_store<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
response: &Bound<'py, PyAny>,
|
||||
|
|
@ -59,7 +59,7 @@ impl PythonCallback {
|
|||
)
|
||||
}
|
||||
|
||||
pub(super) fn lookup_batch<'py>(
|
||||
pub(in crate::cache) fn lookup_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
requests: &Bound<'py, PyAny>,
|
||||
|
|
@ -76,7 +76,7 @@ impl PythonCallback {
|
|||
Ok(results.into_any())
|
||||
}
|
||||
|
||||
pub(super) fn async_lookup_batch<'py>(
|
||||
pub(in crate::cache) fn async_lookup_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
requests: &Bound<'py, PyAny>,
|
||||
|
|
@ -94,7 +94,7 @@ impl PythonCallback {
|
|||
.call_method1("gather", PyTuple::new(py, awaitables)?)
|
||||
}
|
||||
|
||||
pub(super) fn async_store_batch<'py>(
|
||||
pub(in crate::cache) fn async_store_batch<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
result: Option<&Bound<'py, PyAny>>,
|
||||
|
|
@ -110,7 +110,10 @@ impl PythonCallback {
|
|||
)
|
||||
}
|
||||
|
||||
pub(super) fn async_flush<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
pub(in crate::cache) fn async_flush<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let object = self.0.bind(py);
|
||||
let backend = match object.getattr_opt("cache")? {
|
||||
Some(backend) if !backend.is_none() => backend,
|
||||
|
|
@ -123,11 +126,11 @@ impl PythonCallback {
|
|||
ready_none(py)
|
||||
}
|
||||
|
||||
pub(super) fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
pub(in crate::cache) fn ping<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
|
||||
self.0.bind(py).call_method0("ping")
|
||||
}
|
||||
|
||||
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
pub(in crate::cache) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.0)
|
||||
}
|
||||
}
|
||||
141
litellm-rust/crates/python-bridge/src/cache/python/host.rs
vendored
Normal file
141
litellm-rust/crates/python-bridge/src/cache/python/host.rs
vendored
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
use litellm_cache::Error;
|
||||
use litellm_host::protocol::Reply;
|
||||
use litellm_host_python::{from_py, to_py};
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::service::CacheCall;
|
||||
|
||||
enum Pending {
|
||||
Lookup(Reply<Result<Option<Value>, Error>>),
|
||||
Store(Reply<Result<(), Error>>),
|
||||
}
|
||||
|
||||
pub(crate) struct PythonCache {
|
||||
cache: Option<Py<PyAny>>,
|
||||
arguments: Option<Py<PyDict>>,
|
||||
pending: Option<Pending>,
|
||||
asynchronous: bool,
|
||||
}
|
||||
|
||||
impl PythonCache {
|
||||
pub fn new(asynchronous: bool) -> Self {
|
||||
Self {
|
||||
cache: None,
|
||||
arguments: None,
|
||||
pending: None,
|
||||
asynchronous,
|
||||
}
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn bind(
|
||||
&mut self,
|
||||
cache: Bound<'_, PyAny>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) {
|
||||
self.cache = Some(cache.unbind());
|
||||
self.arguments = Some(arguments.clone().unbind());
|
||||
}
|
||||
|
||||
pub fn begin(&mut self, py: Python<'_>, call: CacheCall) -> PyResult<Option<Py<PyAny>>> {
|
||||
let Some(cache) = self.cache.as_ref() else {
|
||||
return Err(pyo3::exceptions::PyRuntimeError::new_err(
|
||||
"cache operation without configured cache",
|
||||
));
|
||||
};
|
||||
let arguments = self
|
||||
.arguments
|
||||
.as_ref()
|
||||
.ok_or_else(|| {
|
||||
pyo3::exceptions::PyRuntimeError::new_err("cache arguments unavailable")
|
||||
})?
|
||||
.bind(py)
|
||||
.copy()?;
|
||||
let (method, result) = match call {
|
||||
CacheCall::Lookup { reply } => {
|
||||
self.pending = Some(Pending::Lookup(reply));
|
||||
(
|
||||
if self.asynchronous {
|
||||
"async_get_cache"
|
||||
} else {
|
||||
"get_cache"
|
||||
},
|
||||
None,
|
||||
)
|
||||
}
|
||||
CacheCall::Store { value, reply } => {
|
||||
self.pending = Some(Pending::Store(reply));
|
||||
(
|
||||
if self.asynchronous {
|
||||
"async_add_cache"
|
||||
} else {
|
||||
"add_cache"
|
||||
},
|
||||
Some(to_py(py, &value)?),
|
||||
)
|
||||
}
|
||||
};
|
||||
let result = match result {
|
||||
Some(value) => cache
|
||||
.bind(py)
|
||||
.call_method(method, (value,), Some(&arguments)),
|
||||
None => cache.bind(py).call_method(method, (), Some(&arguments)),
|
||||
}
|
||||
.map(Bound::unbind);
|
||||
if self.asynchronous && result.is_ok() {
|
||||
return result.map(Some);
|
||||
}
|
||||
self.resume(py, result)
|
||||
}
|
||||
|
||||
pub fn resume(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: PyResult<Py<PyAny>>,
|
||||
) -> PyResult<Option<Py<PyAny>>> {
|
||||
if let Err(error) = &result
|
||||
&& !error.is_instance_of::<pyo3::exceptions::PyException>(py)
|
||||
{
|
||||
self.pending = None;
|
||||
return Err(result.err().unwrap());
|
||||
}
|
||||
match self.pending.take() {
|
||||
Some(Pending::Lookup(reply)) => {
|
||||
let value = result.map_err(|_| Error::Unavailable).and_then(|value| {
|
||||
if value.bind(py).is_none() {
|
||||
Ok(None)
|
||||
} else {
|
||||
from_py(value.bind(py))
|
||||
.map(Some)
|
||||
.map_err(|_| Error::InvalidEntry)
|
||||
}
|
||||
});
|
||||
reply.send(value);
|
||||
}
|
||||
Some(Pending::Store(reply)) => {
|
||||
reply.send(result.map(|_| ()).map_err(|_| Error::Unavailable));
|
||||
}
|
||||
None => {
|
||||
return Err(pyo3::exceptions::PyRuntimeError::new_err(
|
||||
"cache reply without pending operation",
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
pub fn close(&mut self) {
|
||||
self.pending = None;
|
||||
self.cache = None;
|
||||
self.arguments = None;
|
||||
}
|
||||
|
||||
pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.cache)?;
|
||||
visit.call(&self.arguments)
|
||||
}
|
||||
}
|
||||
8
litellm-rust/crates/python-bridge/src/cache/python/mod.rs
vendored
Normal file
8
litellm-rust/crates/python-bridge/src/cache/python/mod.rs
vendored
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
mod callback;
|
||||
mod host;
|
||||
mod service;
|
||||
|
||||
pub(super) use callback::PythonCallback;
|
||||
pub(crate) use host::PythonCache;
|
||||
pub(crate) use service::CacheCall;
|
||||
pub(super) use service::service;
|
||||
107
litellm-rust/crates/python-bridge/src/cache/python/service.rs
vendored
Normal file
107
litellm-rust/crates/python-bridge/src/cache/python/service.rs
vendored
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
use std::{future::Future, pin::Pin, time::Duration};
|
||||
|
||||
use litellm_cache::Error;
|
||||
use litellm_cache_response::{
|
||||
ResponseCacheConfig, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope,
|
||||
};
|
||||
use litellm_core::caching::CachedOutput;
|
||||
use litellm_host::{
|
||||
machine::{HostServices, MachineFault},
|
||||
protocol::{Protocol, Reply},
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
const STREAM_EVENTS_KEY: &str = "litellm_cached_anthropic_sse_events";
|
||||
|
||||
fn from_python(value: Value) -> Result<Value, Error> {
|
||||
let output = match value.get(STREAM_EVENTS_KEY) {
|
||||
Some(events) => {
|
||||
let events: Vec<String> =
|
||||
serde_json::from_value(events.clone()).map_err(|_| Error::InvalidEntry)?;
|
||||
CachedOutput::Stream(events.concat())
|
||||
}
|
||||
None => CachedOutput::Response(value),
|
||||
};
|
||||
serde_json::to_value(ResponseEnvelope::new("messages", output)).map_err(|_| Error::InvalidEntry)
|
||||
}
|
||||
|
||||
fn to_python(value: Value) -> Result<Value, Error> {
|
||||
let envelope: ResponseEnvelope<CachedOutput<Value>> =
|
||||
serde_json::from_value(value).map_err(|_| Error::InvalidEntry)?;
|
||||
match envelope.decode("messages").ok_or(Error::InvalidEntry)? {
|
||||
CachedOutput::Response(response) => Ok(response),
|
||||
CachedOutput::Stream(text) => Ok(serde_json::json!({
|
||||
STREAM_EVENTS_KEY: text.split_inclusive("\n\n").collect::<Vec<_>>()
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) enum CacheCall {
|
||||
Lookup {
|
||||
reply: Reply<Result<Option<Value>, Error>>,
|
||||
},
|
||||
Store {
|
||||
value: Value,
|
||||
reply: Reply<Result<(), Error>>,
|
||||
},
|
||||
}
|
||||
|
||||
struct PythonCacheService<P: Protocol> {
|
||||
services: HostServices<P>,
|
||||
config: ResponseCacheConfig,
|
||||
}
|
||||
|
||||
pub(in crate::cache) fn service<P: Protocol<HostCall = CacheCall>>(
|
||||
services: HostServices<P>,
|
||||
namespace: String,
|
||||
) -> std::sync::Arc<dyn ResponseCacheService>
|
||||
where
|
||||
P::Error: From<MachineFault>,
|
||||
{
|
||||
std::sync::Arc::new(PythonCacheService {
|
||||
services,
|
||||
config: ResponseCacheConfig {
|
||||
namespace,
|
||||
..Default::default()
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
impl<P: Protocol<HostCall = CacheCall>> ResponseCacheService for PythonCacheService<P>
|
||||
where
|
||||
P::Error: From<MachineFault>,
|
||||
{
|
||||
fn config(&self) -> &ResponseCacheConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
_: &'a ResponseCacheRequest,
|
||||
_: Duration,
|
||||
) -> Pin<Box<dyn Future<Output = Result<Option<Value>, Error>> + Send + 'a>> {
|
||||
Box::pin(async move {
|
||||
self.services
|
||||
.call(|reply| CacheCall::Lookup { reply })
|
||||
.await
|
||||
.map_err(|_| Error::Unavailable)??
|
||||
.map(from_python)
|
||||
.transpose()
|
||||
})
|
||||
}
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
_: &'a ResponseCacheRequest,
|
||||
value: Value,
|
||||
_: Duration,
|
||||
) -> Pin<Box<dyn Future<Output = Result<(), Error>> + Send + 'a>> {
|
||||
Box::pin(async move {
|
||||
let value = to_python(value)?;
|
||||
self.services
|
||||
.call(|reply| CacheCall::Store { value, reply })
|
||||
.await
|
||||
.map_err(|_| Error::Unavailable)?
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -1,25 +0,0 @@
|
|||
use pyo3::{PyTraverseError, PyVisit, prelude::*};
|
||||
|
||||
use super::binding::ResolvedCache;
|
||||
|
||||
#[pyclass(frozen, name = "_CacheResolver")]
|
||||
pub(crate) struct CacheResolver {
|
||||
namespace: Py<PyAny>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl CacheResolver {
|
||||
#[new]
|
||||
fn new(namespace: Py<PyAny>) -> Self {
|
||||
Self { namespace }
|
||||
}
|
||||
|
||||
pub(crate) fn resolve(&self, py: Python<'_>) -> PyResult<ResolvedCache> {
|
||||
let object = self.namespace.bind(py).getattr("cache")?;
|
||||
ResolvedCache::from_selected(&object)
|
||||
}
|
||||
|
||||
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.namespace)
|
||||
}
|
||||
}
|
||||
|
|
@ -10,13 +10,13 @@ use pyo3::{
|
|||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
activation::activate,
|
||||
cache_error,
|
||||
callback::PythonCallback,
|
||||
config::{CacheConfigProjection, NativeCacheConfig},
|
||||
future::{ready_none, ready_value},
|
||||
native::{NativeResponseCache, SemanticReply},
|
||||
request::{now, request, requests},
|
||||
native::activation::activate,
|
||||
native::backend::{NativeResponseCache, SemanticReply},
|
||||
native::config::{CacheConfigProjection, NativeCacheConfig},
|
||||
native::request::{now, request, requests},
|
||||
python::PythonCallback,
|
||||
};
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
|
||||
|
|
@ -29,7 +29,7 @@ pub(super) enum CacheBinding {
|
|||
#[pyclass(frozen, name = "_ResponseCacheRuntime")]
|
||||
pub(crate) struct ResolvedCache {
|
||||
binding: CacheBinding,
|
||||
guard: Option<super::facade::FacadeGuard>,
|
||||
guard: Option<super::native::facade::FacadeGuard>,
|
||||
pid: u32,
|
||||
}
|
||||
|
||||
|
|
@ -42,7 +42,7 @@ impl ResolvedCache {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) fn with_guard(mut self, guard: super::facade::FacadeGuard) -> Self {
|
||||
pub(super) fn with_guard(mut self, guard: super::native::facade::FacadeGuard) -> Self {
|
||||
self.guard = Some(guard);
|
||||
self
|
||||
}
|
||||
|
|
@ -90,10 +90,6 @@ impl ResolvedCache {
|
|||
let py = cache.py();
|
||||
let binding = if cache.is_none() {
|
||||
CacheBinding::Disabled
|
||||
} else if let Ok(handle) = cache.extract::<PyRef<'_, super::handle::CacheTestHandle>>() {
|
||||
CacheBinding::Native(handle.service()?)
|
||||
} else if let Some(service) = super::facade::resolve(py, cache)? {
|
||||
CacheBinding::Native(service)
|
||||
} else if let Some(runtime) = cache
|
||||
.getattr_opt("_native_cache")?
|
||||
.filter(|value| !value.is_none())
|
||||
|
|
@ -134,7 +130,7 @@ impl ResolvedCache {
|
|||
let service = activate(cache.py(), &backend, config)?;
|
||||
let resolved = Self::new(CacheBinding::Native(service.clone()));
|
||||
Ok(
|
||||
match super::facade::FacadeGuard::capture(cache.py(), cache, &service) {
|
||||
match super::native::facade::FacadeGuard::capture(cache.py(), cache, &service) {
|
||||
Ok(guard) => resolved.with_guard(guard),
|
||||
Err(_) => resolved,
|
||||
},
|
||||
176
litellm-rust/crates/python-bridge/src/cache/selection.rs
vendored
Normal file
176
litellm-rust/crates/python-bridge/src/cache/selection.rs
vendored
Normal file
|
|
@ -0,0 +1,176 @@
|
|||
use super::{native, python};
|
||||
use litellm_cache_response::{CacheOptions, CacheScope, ResponseCacheService, ScopedCache};
|
||||
use litellm_host::{
|
||||
machine::{HostServices, MachineFault},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
use std::sync::Arc;
|
||||
|
||||
pub(crate) struct Cached<P>(std::marker::PhantomData<P>);
|
||||
|
||||
impl<P: Protocol> Protocol for Cached<P> {
|
||||
type Request = (P::Request, Selection);
|
||||
type Response = P::Response;
|
||||
type Error = P::Error;
|
||||
type HostCall = python::CacheCall;
|
||||
type Chunk = P::Chunk;
|
||||
type StreamHead = P::StreamHead;
|
||||
}
|
||||
|
||||
enum Backend {
|
||||
Disabled,
|
||||
Native(Arc<dyn ResponseCacheService>),
|
||||
Python { namespace: String },
|
||||
}
|
||||
|
||||
pub(crate) struct Selection {
|
||||
backend: Backend,
|
||||
options: CacheOptions,
|
||||
}
|
||||
|
||||
impl Selection {
|
||||
pub(crate) fn attach<P: Protocol<HostCall = python::CacheCall>>(
|
||||
self,
|
||||
services: HostServices<P>,
|
||||
) -> (Option<ScopedCache>, CacheOptions)
|
||||
where
|
||||
P::Error: From<MachineFault>,
|
||||
{
|
||||
let service = match self.backend {
|
||||
Backend::Disabled => None,
|
||||
Backend::Native(service) => Some(service),
|
||||
Backend::Python { namespace } => Some(python::service(services, namespace)),
|
||||
};
|
||||
(
|
||||
service.map(|service| ScopedCache::new(service, CacheScope::Shared)),
|
||||
self.options,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn selected_cache<'py>(
|
||||
py: Python<'py>,
|
||||
kwargs: &Bound<'py, PyDict>,
|
||||
call_type: &str,
|
||||
) -> PyResult<Option<Bound<'py, PyAny>>> {
|
||||
let configured = py.import("litellm")?.getattr("cache")?;
|
||||
if configured.is_none()
|
||||
|| kwargs
|
||||
.get_item("caching")?
|
||||
.is_some_and(|value| value.is(pyo3::types::PyBool::new(py, false)))
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
let supported = configured.getattr("supported_call_types")?;
|
||||
if supported.is_none() || !supported.contains(call_type)? {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(Some(configured))
|
||||
}
|
||||
|
||||
pub(crate) fn admit_native(
|
||||
py: Python<'_>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
call_type: &str,
|
||||
) -> PyResult<()> {
|
||||
if let Some(configured) = selected_cache(py, kwargs, call_type)?
|
||||
&& native::v2::native_handle(&configured)?.is_none()
|
||||
{
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(
|
||||
"the configured cache requires Python inference",
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn configured_native(
|
||||
py: Python<'_>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
call_type: &str,
|
||||
) -> PyResult<(
|
||||
Option<Arc<dyn ResponseCacheService>>,
|
||||
litellm_cache_response::CacheOptions,
|
||||
)> {
|
||||
let Some(configured) = selected_cache(py, kwargs, call_type)? else {
|
||||
return Ok((
|
||||
None,
|
||||
litellm_cache_response::CacheOptions::new(litellm_cache_response::CacheScope::Shared),
|
||||
));
|
||||
};
|
||||
native_configuration(&configured, kwargs)
|
||||
}
|
||||
|
||||
fn native_configuration(
|
||||
configured: &Bound<'_, PyAny>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
) -> PyResult<(Option<Arc<dyn ResponseCacheService>>, CacheOptions)> {
|
||||
if !configured
|
||||
.call_method("should_use_cache", (), Some(kwargs))?
|
||||
.extract::<bool>()?
|
||||
{
|
||||
return Ok((
|
||||
None,
|
||||
litellm_cache_response::CacheOptions::new(litellm_cache_response::CacheScope::Shared),
|
||||
));
|
||||
}
|
||||
native::v2::configured(configured, kwargs)
|
||||
}
|
||||
|
||||
pub(crate) fn configure(
|
||||
python: &mut python::PythonCache,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
call_type: &str,
|
||||
) -> PyResult<Selection> {
|
||||
let selected = selected_cache(py, arguments, call_type)?;
|
||||
let Some(cache) = selected else {
|
||||
return Ok(Selection {
|
||||
backend: Backend::Disabled,
|
||||
options: CacheOptions::new(CacheScope::Shared),
|
||||
});
|
||||
};
|
||||
if native::v2::native_handle(&cache)?.is_some() {
|
||||
let (native, options) = native_configuration(&cache, arguments)?;
|
||||
return Ok(Selection {
|
||||
backend: native.map_or(Backend::Disabled, Backend::Native),
|
||||
options,
|
||||
});
|
||||
}
|
||||
let enabled = cache
|
||||
.call_method("should_use_cache", (), Some(arguments))?
|
||||
.extract::<bool>()?;
|
||||
let controls = arguments
|
||||
.get_item("cache")?
|
||||
.filter(|value| !value.is_none());
|
||||
let boolean = |name: &str| -> PyResult<bool> {
|
||||
controls
|
||||
.as_ref()
|
||||
.map(|value| value.cast::<PyDict>()?.get_item(name))
|
||||
.transpose()?
|
||||
.flatten()
|
||||
.map(|value| value.extract())
|
||||
.transpose()
|
||||
.map(|value| value.unwrap_or(false))
|
||||
};
|
||||
let options = CacheOptions {
|
||||
no_cache: boolean("no-cache")?,
|
||||
no_store: boolean("no-store")?,
|
||||
..CacheOptions::new(CacheScope::Shared)
|
||||
};
|
||||
let namespace = cache
|
||||
.getattr_opt("namespace")?
|
||||
.filter(|value| !value.is_none())
|
||||
.map(|value| value.extract())
|
||||
.transpose()?
|
||||
.unwrap_or_default();
|
||||
python.bind(cache, arguments);
|
||||
Ok(Selection {
|
||||
backend: if enabled {
|
||||
Backend::Python { namespace }
|
||||
} else {
|
||||
Backend::Disabled
|
||||
},
|
||||
options,
|
||||
})
|
||||
}
|
||||
|
|
@ -16,7 +16,7 @@ mod tokenizer;
|
|||
|
||||
#[pymodule(gil_used = true)]
|
||||
mod _native {
|
||||
use crate::cache::{CacheResolver, CacheTestHandle, ResolvedCache};
|
||||
use crate::cache::ResolvedCache;
|
||||
#[cfg(feature = "panic-test")]
|
||||
#[pymodule_export]
|
||||
use crate::diagnostics::_panic_for_test;
|
||||
|
|
@ -55,9 +55,10 @@ mod _native {
|
|||
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
let py = module.py();
|
||||
let dict = module.dict();
|
||||
dict.set_item("_CacheTestHandle", py.get_type::<CacheTestHandle>())?;
|
||||
dict.set_item("_CacheResolver", py.get_type::<CacheResolver>())?;
|
||||
dict.set_item("_CacheTestResolver", py.get_type::<CacheResolver>())?;
|
||||
dict.set_item(
|
||||
"NativeCacheHandle",
|
||||
py.get_type::<crate::cache::NativeCacheHandle>(),
|
||||
)?;
|
||||
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
|
||||
dict.set_item(
|
||||
"_SecretManagerRuntime",
|
||||
|
|
@ -77,11 +78,12 @@ pub(crate) fn native_module(py: Python<'_>) -> Bound<'_, PyModule> {
|
|||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn module_registration_preserves_the_public_surface() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let mut expected = vec![
|
||||
"NativeCacheHandle",
|
||||
"RustBridgeDeclined",
|
||||
"RustUpstreamError",
|
||||
"ForkedAfterNativeRuntimeStarted",
|
||||
|
|
|
|||
|
|
@ -142,6 +142,12 @@ fn run_public(
|
|||
request.clone().unbind(),
|
||||
"litellm.rust_bridge.chat_completions.route_host",
|
||||
);
|
||||
let cache_call_type = if asynchronous {
|
||||
"acompletion"
|
||||
} else {
|
||||
"completion"
|
||||
};
|
||||
crate::cache::admit_native(py, &kwargs, cache_call_type)?;
|
||||
let (arguments, hooks) = crate::routes::call_hooks(
|
||||
py,
|
||||
Operation::Completion,
|
||||
|
|
@ -160,7 +166,16 @@ fn run_public(
|
|||
crate::http::resources().auth.clone(),
|
||||
crate::secrets::source(py)?,
|
||||
);
|
||||
Ok(route.machine(request, None))
|
||||
let (cache, cache_options) =
|
||||
crate::cache::configured_native(py, arguments, cache_call_type)?;
|
||||
let route = match cache {
|
||||
Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new(
|
||||
cache,
|
||||
litellm_cache_response::CacheScope::Shared,
|
||||
)),
|
||||
None => route,
|
||||
};
|
||||
Ok(route.machine(request, cache_options))
|
||||
},
|
||||
host::ChatCompletionsPythonHost(host),
|
||||
hooks,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
use crate::cache::{CacheCall, Cached, PythonCache, Selection};
|
||||
use litellm_host_python::{PythonHostCalls, PythonOwned};
|
||||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_core::messages::{
|
||||
|
|
@ -89,11 +89,15 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
|||
/// public response, chunks and exceptions.
|
||||
pub(super) struct MessagesPythonHost {
|
||||
request: Py<PyAny>,
|
||||
cache: PythonCache,
|
||||
}
|
||||
|
||||
impl MessagesPythonHost {
|
||||
pub(super) fn new(request: Py<PyAny>) -> Self {
|
||||
Self { request }
|
||||
pub(super) fn new(request: Py<PyAny>, asynchronous: bool) -> Self {
|
||||
Self {
|
||||
request,
|
||||
cache: PythonCache::new(asynchronous),
|
||||
}
|
||||
}
|
||||
|
||||
fn projection(
|
||||
|
|
@ -223,17 +227,21 @@ impl MessagesPythonHost {
|
|||
}
|
||||
|
||||
impl PythonBinding for MessagesPythonHost {
|
||||
type Protocol = Messages;
|
||||
type Protocol = Cached<Messages>;
|
||||
type Failure = PyErr;
|
||||
|
||||
fn decode_request(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> Result<MessagesCall, InvokeError<Error>> {
|
||||
) -> Result<(MessagesCall, Selection), InvokeError<Error>> {
|
||||
let selection =
|
||||
crate::cache::configure(&mut self.cache, py, arguments, "anthropic_messages")
|
||||
.map_err(InvokeError::Python)?;
|
||||
self.projection(py, arguments)
|
||||
.map_err(|error| InvokeError::Python(self.map_failure(py, error)))?
|
||||
.map_err(InvokeError::Native)
|
||||
.map(|request| (request, selection))
|
||||
}
|
||||
|
||||
fn encode_response(
|
||||
|
|
@ -278,20 +286,42 @@ impl PythonBinding for MessagesPythonHost {
|
|||
}
|
||||
}
|
||||
|
||||
impl PythonHostCalls<Messages> for MessagesPythonHost {
|
||||
impl PythonHostCalls<Cached<Messages>> for MessagesPythonHost {
|
||||
fn handle_host_call(
|
||||
&mut self,
|
||||
_: Python<'_>,
|
||||
op: Infallible,
|
||||
py: Python<'_>,
|
||||
op: CacheCall,
|
||||
) -> Result<(), InvokeError<Error>> {
|
||||
match op {}
|
||||
self.cache
|
||||
.begin(py, op)
|
||||
.map(|_| ())
|
||||
.map_err(InvokeError::Python)
|
||||
}
|
||||
|
||||
fn begin_host_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
op: CacheCall,
|
||||
) -> Result<Option<Py<PyAny>>, InvokeError<Error>> {
|
||||
self.cache.begin(py, op).map_err(InvokeError::Python)
|
||||
}
|
||||
|
||||
fn resume_host_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: PyResult<Py<PyAny>>,
|
||||
) -> Result<Option<Py<PyAny>>, InvokeError<Error>> {
|
||||
self.cache.resume(py, result).map_err(InvokeError::Python)
|
||||
}
|
||||
}
|
||||
|
||||
impl PythonOwned for MessagesPythonHost {
|
||||
fn close(&mut self, _: Python<'_>) {}
|
||||
fn close(&mut self, _: Python<'_>) {
|
||||
self.cache.close();
|
||||
}
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.request)
|
||||
visit.call(&self.request)?;
|
||||
self.cache.traverse(visit)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -26,15 +26,40 @@ fn run_messages(
|
|||
py,
|
||||
arguments,
|
||||
move |py, arguments, request| {
|
||||
let route = litellm_core::messages::MessagesRoute::new(
|
||||
crate::http::provider_client(py, arguments, asynchronous)?
|
||||
.map_err(crate::http::client_error)?,
|
||||
crate::http::resources().auth.clone(),
|
||||
crate::secrets::source(py)?,
|
||||
);
|
||||
Ok(route.machine(request, None))
|
||||
let builder = litellm_core::messages::MessagesRoute::builder()
|
||||
.with_http(
|
||||
crate::http::provider_client(py, arguments, asynchronous)?
|
||||
.map_err(crate::http::client_error)?,
|
||||
)
|
||||
.with_auth(crate::http::resources().auth.clone())
|
||||
.with_secrets(crate::secrets::source(py)?);
|
||||
let route = builder.build();
|
||||
Ok(litellm_host::call::hosted_call(
|
||||
request,
|
||||
None,
|
||||
move |(call, selection): (_, crate::cache::Selection),
|
||||
services,
|
||||
interceptors,
|
||||
observers| async move {
|
||||
let (cache, options) = selection.attach(services);
|
||||
let route = match cache {
|
||||
Some(cache) => route.with_cache(cache),
|
||||
None => route,
|
||||
};
|
||||
route
|
||||
.execute(
|
||||
call,
|
||||
&interceptors,
|
||||
litellm_core::CallOptions {
|
||||
cache: Some(options),
|
||||
observers,
|
||||
},
|
||||
)
|
||||
.await
|
||||
},
|
||||
))
|
||||
},
|
||||
MessagesPythonHost::new(request.unbind()),
|
||||
MessagesPythonHost::new(request.unbind(), asynchronous),
|
||||
hooks,
|
||||
asynchronous,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -63,6 +63,12 @@ fn run_public(
|
|||
"native Python responses streaming",
|
||||
));
|
||||
}
|
||||
let cache_call_type = if asynchronous {
|
||||
"aresponses"
|
||||
} else {
|
||||
"responses"
|
||||
};
|
||||
crate::cache::admit_native(py, &kwargs, cache_call_type)?;
|
||||
let (arguments, hooks) = crate::routes::call_hooks(
|
||||
py,
|
||||
Operation::Responses,
|
||||
|
|
@ -81,7 +87,16 @@ fn run_public(
|
|||
crate::http::resources().auth.clone(),
|
||||
crate::secrets::source(py)?,
|
||||
);
|
||||
Ok(route.machine(request, None))
|
||||
let (cache, cache_options) =
|
||||
crate::cache::configured_native(py, arguments, cache_call_type)?;
|
||||
let route = match cache {
|
||||
Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new(
|
||||
cache,
|
||||
litellm_cache_response::CacheScope::Shared,
|
||||
)),
|
||||
None => route,
|
||||
};
|
||||
Ok(route.machine(request, cache_options))
|
||||
},
|
||||
host::ResponsesPythonHost(host),
|
||||
hooks,
|
||||
|
|
|
|||
3
litellm/_v2/AGENTS.md
Normal file
3
litellm/_v2/AGENTS.md
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
Everything here is experimental and should not be documented
|
||||
|
||||
Use this directory to explore alternative APIs where the Rust migration makes backward compatibility difficult. The gateway can use them for performance, but keep them behind the v2 flag to avoid breaking SDK users
|
||||
3
litellm/_v2/__init__.py
Normal file
3
litellm/_v2/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm._v2.cache import Cache
|
||||
|
||||
__all__ = ("Cache",)
|
||||
11
litellm/_v2/cache/AGENTS.md
vendored
Normal file
11
litellm/_v2/cache/AGENTS.md
vendored
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
# Python v2 cache
|
||||
|
||||
Keep `litellm._v2.cache.Cache` import-compatible when reorganizing this package. This package owns Python cache factories and the adapter between the existing `BaseCache` interface and `NativeCacheHandle`
|
||||
|
||||
Construct native handles at this boundary and inject the adapter through the existing cache facade's `_backend` parameter. Keep the facade's `type`, namespace, and TTL consistent with the configured native backend
|
||||
|
||||
Keep storage implementation in the Rust storage crates and response-cache policy in `litellm-cache-response` and core. Do not duplicate cache-key generation, freshness rules, response encoding, or inference orchestration here
|
||||
|
||||
Preserve synchronous and asynchronous cache operations, including TTL forwarding and lifecycle methods. Validate Python values before passing them to typed native interfaces. Keep native extension imports lazy so importing the package does not require loading the extension
|
||||
|
||||
Extend the existing v2 cache tests in `tests/test_litellm_rust/test_v2.py` for behavioral changes, following that directory's `AGENTS.md`. Test observable cache behavior rather than package layout or implementation structure
|
||||
72
litellm/_v2/cache/__init__.py
vendored
Normal file
72
litellm/_v2/cache/__init__.py
vendored
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.caching.base_cache import BaseCache
|
||||
from litellm.caching.caching import Cache as CacheFacade
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.rust_bridge._native import NativeCacheHandle
|
||||
|
||||
_DURATION: Final[TypeAdapter[float | None]] = TypeAdapter(float | None)
|
||||
|
||||
|
||||
class NativeBackend(BaseCache):
|
||||
def __init__(self, handle: NativeCacheHandle) -> None:
|
||||
self.native_handle = handle
|
||||
|
||||
def get_cache(self, key: str, **kwargs: object) -> object:
|
||||
return self.native_handle.get(key)
|
||||
|
||||
async def async_get_cache(self, key: str, **kwargs: object) -> object:
|
||||
return await self.native_handle.async_get(key)
|
||||
|
||||
def set_cache(self, key: str, value: object, **kwargs: object) -> None:
|
||||
self.native_handle.set(key, value, ttl=_DURATION.validate_python(kwargs.get("ttl")))
|
||||
|
||||
async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None:
|
||||
await self.native_handle.async_set(key, value, ttl=_DURATION.validate_python(kwargs.get("ttl")))
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list: Sequence[tuple[str, object]], **kwargs: object) -> None:
|
||||
await self.native_handle.async_set_many(cache_list, ttl=_DURATION.validate_python(kwargs.get("ttl")))
|
||||
|
||||
async def batch_cache_write(self, key: str, value: object, **kwargs: object) -> None:
|
||||
await self.async_set_cache(key, value, **kwargs)
|
||||
|
||||
def flush_cache(self) -> None:
|
||||
self.native_handle.flush()
|
||||
|
||||
async def async_flush_cache(self) -> None:
|
||||
await self.native_handle.async_flush()
|
||||
|
||||
async def ping(self) -> bool:
|
||||
return await self.native_handle.ping()
|
||||
|
||||
async def disconnect(self) -> None:
|
||||
await self.native_handle.disconnect()
|
||||
|
||||
async def delete_cache_keys(self, keys: Sequence[str]) -> None:
|
||||
await self.native_handle.delete(keys)
|
||||
|
||||
async def test_connection(self) -> dict[str, str]:
|
||||
return {"status": "success" if await self.ping() else "failed"}
|
||||
|
||||
|
||||
class Cache:
|
||||
@staticmethod
|
||||
def memory(*, ttl: float = 600, capacity: int = 200, max_entry_bytes: int = 4194304) -> CacheFacade:
|
||||
from litellm.rust_bridge._native import NativeCacheHandle
|
||||
|
||||
handle: Final = NativeCacheHandle.memory(ttl=ttl, capacity=capacity, max_entry_bytes=max_entry_bytes)
|
||||
return CacheFacade(type=LiteLLMCacheType.LOCAL, ttl=ttl, _backend=NativeBackend(handle))
|
||||
|
||||
@staticmethod
|
||||
def redis(url: str, *, namespace: str, ttl: float = 600, max_entry_bytes: int = 4194304) -> CacheFacade:
|
||||
from litellm.rust_bridge._native import NativeCacheHandle
|
||||
|
||||
handle: Final = NativeCacheHandle.redis(url, namespace=namespace, ttl=ttl, max_entry_bytes=max_entry_bytes)
|
||||
return CacheFacade(type=LiteLLMCacheType.REDIS, namespace=namespace, ttl=ttl, _backend=NativeBackend(handle))
|
||||
|
|
@ -127,6 +127,7 @@ class Cache:
|
|||
# GCP IAM authentication parameters
|
||||
gcp_service_account: str | None = None,
|
||||
gcp_ssl_ca_certs: str | None = None,
|
||||
_backend: BaseCache | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -183,7 +184,9 @@ class Cache:
|
|||
Returns:
|
||||
None. Cache is set as a litellm param
|
||||
"""
|
||||
if type == LiteLLMCacheType.REDIS:
|
||||
if _backend is not None:
|
||||
self.cache: BaseCache = _backend
|
||||
elif type == LiteLLMCacheType.REDIS:
|
||||
# Check REDIS_CLUSTER_NODES env var if no explicit startup nodes
|
||||
if not redis_startup_nodes:
|
||||
_env_cluster_nodes: Final = litellm.get_secret("REDIS_CLUSTER_NODES")
|
||||
|
|
@ -205,7 +208,7 @@ class Cache:
|
|||
if gcp_ssl_ca_certs is not None:
|
||||
cluster_kwargs["gcp_ssl_ca_certs"] = gcp_ssl_ca_certs
|
||||
|
||||
self.cache: BaseCache = RedisClusterCache(**cluster_kwargs)
|
||||
self.cache = RedisClusterCache(**cluster_kwargs)
|
||||
else:
|
||||
self.cache = RedisCache(
|
||||
host=host,
|
||||
|
|
@ -314,12 +317,6 @@ class Cache:
|
|||
if self.namespace is not None and isinstance(self.cache, RedisCache):
|
||||
self.cache.namespace = self.namespace
|
||||
|
||||
from litellm.rust_bridge.response_cache import resolve_response_cache
|
||||
|
||||
# The Rust catalog picks the store per backend. When it selects Rust, the storage calls
|
||||
# below go to the native runtime and the Python backend stays only for its direct API.
|
||||
self._native_cache = resolve_response_cache(self)
|
||||
|
||||
# Params whose values carry prompt content. Excluded from semantic-cache
|
||||
# scope keys so differently worded prompts share a bucket and match via
|
||||
# vector similarity rather than being split into per-wording buckets.
|
||||
|
|
|
|||
|
|
@ -949,6 +949,7 @@ openai_compatible_endpoints: Final[list] = [
|
|||
"https://api.sailresearch.com/v1",
|
||||
"https://api.cognition.ai/v1",
|
||||
"https://api.scx.ai/v1",
|
||||
"https://api.prisminference.com/v1",
|
||||
"https://gigachat.devices.sberbank.ru/api/v1",
|
||||
]
|
||||
|
||||
|
|
@ -1022,6 +1023,7 @@ openai_compatible_providers: Final[list] = [
|
|||
"meta", # Meta Model API (Muse Spark) - JSON-configured provider
|
||||
"cognition",
|
||||
"scx-ai",
|
||||
"prism",
|
||||
"sail",
|
||||
]
|
||||
|
||||
|
|
@ -2140,16 +2142,12 @@ MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: "
|
|||
PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__"
|
||||
PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job"
|
||||
PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900
|
||||
USAGE_TOP_API_KEYS_LIMIT: Final[int] = int(os.getenv("USAGE_TOP_API_KEYS_LIMIT", "100"))
|
||||
# Furthest back the catch-up pass looks for unpriced PTU days when a deployment
|
||||
# declares no ptu_effective_from, bounding the scan for an open-ended window.
|
||||
PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90
|
||||
# Deployments named in the lapsed-window alert before it is truncated, so a fleet-wide
|
||||
# expiry cannot produce an alert too large for the channel delivering it.
|
||||
PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
|
||||
DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID: Final[str] = "daily_global_spend_reconcile_job"
|
||||
DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS: Final[int] = 3600
|
||||
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM: Final[str] = "daily_global_spend_reconciled_through"
|
||||
SPEND_CAPTURE_RATE_CHECK_JOB_ID: Final[str] = "spend_capture_rate_check_job"
|
||||
SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS: Final[int] = 900
|
||||
SPEND_CAPTURE_RATE_MAX_RANGE_DAYS: Final[int] = 180
|
||||
|
|
|
|||
|
|
@ -1480,6 +1480,7 @@ class CustomGuardrail(CustomLogger):
|
|||
or call_type == CallTypes.acompletion.value
|
||||
or call_type == CallTypes.anthropic_messages.value
|
||||
or call_type == CallTypes.call_mcp_tool.value
|
||||
or call_type == CallTypes.list_mcp_tools.value
|
||||
):
|
||||
return data.get("messages")
|
||||
|
||||
|
|
|
|||
|
|
@ -201,6 +201,12 @@
|
|||
},
|
||||
"supported_endpoints": ["/v1/chat/completions"]
|
||||
},
|
||||
"prism": {
|
||||
"base_url": "https://api.prisminference.com/v1",
|
||||
"api_key_env": "PRISM_API_KEY",
|
||||
"api_base_env": "PRISM_API_BASE",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/messages"]
|
||||
},
|
||||
"sail": {
|
||||
"base_url": "https://api.sailresearch.com/v1",
|
||||
"api_key_env": "SAIL_API_KEY",
|
||||
|
|
|
|||
|
|
@ -12913,13 +12913,13 @@
|
|||
"source": "https://aws.amazon.com/bedrock/pricing/"
|
||||
},
|
||||
"bedrock/ap-southeast-2/minimax.minimax-m2.5": {
|
||||
"input_cost_per_token": 3.09e-07,
|
||||
"input_cost_per_token": 3.1e-07,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"source": "https://aws.amazon.com/bedrock/pricing/",
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
@ -12927,7 +12927,7 @@
|
|||
"supports_audio_input": false,
|
||||
"supports_response_schema": true,
|
||||
"supports_vision": false,
|
||||
"output_cost_per_token": 1.236e-06
|
||||
"output_cost_per_token": 1.24e-06
|
||||
},
|
||||
"bedrock/ap-southeast-3/deepseek.v3.2": {
|
||||
"input_cost_per_token": 7.4e-07,
|
||||
|
|
@ -37373,24 +37373,26 @@
|
|||
"source": "https://aws.amazon.com/bedrock/pricing/"
|
||||
},
|
||||
"meta.llama3-1-405b-instruct-v1:0": {
|
||||
"input_cost_per_token": 5.32e-06,
|
||||
"input_cost_per_token": 2.4e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"meta.llama3-1-70b-instruct-v1:0": {
|
||||
"input_cost_per_token": 9.9e-07,
|
||||
"input_cost_per_token": 7.2e-07,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 2048,
|
||||
"max_tokens": 2048,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 9.9e-07,
|
||||
"output_cost_per_token": 7.2e-07,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
|
|
@ -37406,13 +37408,14 @@
|
|||
"supports_tool_choice": false
|
||||
},
|
||||
"meta.llama3-2-11b-instruct-v1:0": {
|
||||
"input_cost_per_token": 3.5e-07,
|
||||
"input_cost_per_token": 1.6e-07,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.5e-07,
|
||||
"output_cost_per_token": 1.6e-07,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true
|
||||
|
|
@ -37440,13 +37443,14 @@
|
|||
"supports_tool_choice": false
|
||||
},
|
||||
"meta.llama3-2-90b-instruct-v1:0": {
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token": 7.2e-07,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-06,
|
||||
"output_cost_per_token": 7.2e-07,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true
|
||||
|
|
@ -38085,13 +38089,14 @@
|
|||
"supports_function_calling": true
|
||||
},
|
||||
"mistral.mistral-large-2407-v1:0": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 8191,
|
||||
"max_tokens": 8191,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 9e-06,
|
||||
"output_cost_per_token": 6e-06,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
|
|
@ -47427,24 +47432,26 @@
|
|||
"supports_vision": false
|
||||
},
|
||||
"us.meta.llama3-1-405b-instruct-v1:0": {
|
||||
"input_cost_per_token": 5.32e-06,
|
||||
"input_cost_per_token": 2.4e-06,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"us.meta.llama3-1-70b-instruct-v1:0": {
|
||||
"input_cost_per_token": 9.9e-07,
|
||||
"input_cost_per_token": 7.2e-07,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 2048,
|
||||
"max_tokens": 2048,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 9.9e-07,
|
||||
"output_cost_per_token": 7.2e-07,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
|
|
@ -47460,13 +47467,14 @@
|
|||
"supports_tool_choice": false
|
||||
},
|
||||
"us.meta.llama3-2-11b-instruct-v1:0": {
|
||||
"input_cost_per_token": 3.5e-07,
|
||||
"input_cost_per_token": 1.6e-07,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.5e-07,
|
||||
"output_cost_per_token": 1.6e-07,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true
|
||||
|
|
@ -47494,13 +47502,14 @@
|
|||
"supports_tool_choice": false
|
||||
},
|
||||
"us.meta.llama3-2-90b-instruct-v1:0": {
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token": 7.2e-07,
|
||||
"litellm_provider": "bedrock",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-06,
|
||||
"output_cost_per_token": 7.2e-07,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true
|
||||
|
|
@ -78411,7 +78420,9 @@
|
|||
"supports_mid_conversation_system": true,
|
||||
"cache_creation_input_token_cost": 2.5e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 4e-06,
|
||||
"cache_creation_input_token_cost_batches": 1.25e-06,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
|
|
@ -78449,7 +78460,9 @@
|
|||
"supports_mid_conversation_system": true,
|
||||
"cache_creation_input_token_cost": 2.5e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 4e-06,
|
||||
"cache_creation_input_token_cost_batches": 1.25e-06,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-anthropic_models",
|
||||
|
|
@ -78481,5 +78494,103 @@
|
|||
"thinking_always_on": true,
|
||||
"prompt_cache_min_tokens": 512,
|
||||
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
|
||||
},
|
||||
"prism/deepseek-v4.1-flash": {
|
||||
"cache_read_input_token_cost": 7e-08,
|
||||
"input_cost_per_token": 1.7e-07,
|
||||
"litellm_provider": "prism",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 384000,
|
||||
"max_tokens": 384000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6.3e-07,
|
||||
"source": "https://prisminference.com/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses",
|
||||
"/v1/messages"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"prism/deepseek-v4-flash": {
|
||||
"cache_read_input_token_cost": 7e-08,
|
||||
"input_cost_per_token": 1.7e-07,
|
||||
"litellm_provider": "prism",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 384000,
|
||||
"max_tokens": 384000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.1e-07,
|
||||
"source": "https://prisminference.com/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses",
|
||||
"/v1/messages"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"global.xai.grok-4.7": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 500000,
|
||||
"max_output_tokens": 500000,
|
||||
"max_tokens": 500000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6e-06,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": false,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"us.xai.grok-4.7": {
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 500000,
|
||||
"max_output_tokens": 500000,
|
||||
"max_tokens": 500000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6.6e-06,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": false,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"xai.grok-4.7": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "bedrock_converse",
|
||||
"max_input_tokens": 500000,
|
||||
"max_output_tokens": 500000,
|
||||
"max_tokens": 500000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6e-06,
|
||||
"source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": false,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ from pydantic import Field, ValidationInfo, field_validator
|
|||
|
||||
from litellm.types.llms.base import LiteLLMPydanticObjectBase
|
||||
from litellm.types.mcp import MCPAuthType, MCPCredentials, MCPTransportType
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, PinnedMCPTool, parse_pinned_tools
|
||||
|
||||
|
||||
class MCPEnvVarScope(str, enum.Enum):
|
||||
|
|
@ -69,6 +69,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
allowed_tools: list[str] = Field(default_factory=list)
|
||||
tool_name_to_display_name: dict[str, str] | None = None
|
||||
tool_name_to_description: dict[str, str] | None = None
|
||||
pinned_tools: dict[str, PinnedMCPTool] | None = None
|
||||
extra_headers: list[str] = Field(default_factory=list)
|
||||
mcp_info: MCPInfo | None = None
|
||||
static_headers: dict[str, str] | None = None
|
||||
|
|
@ -119,6 +120,11 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
reviewed_at: datetime | None = None
|
||||
review_notes: str | None = None
|
||||
|
||||
@field_validator("pinned_tools", mode="before")
|
||||
@classmethod
|
||||
def decode_stored_pinned_tools(cls, value: object) -> dict[str, PinnedMCPTool] | None:
|
||||
return parse_pinned_tools(value)
|
||||
|
||||
@field_validator("static_headers", "env", mode="before")
|
||||
@classmethod
|
||||
def decode_stored_secret_map(cls, value: object, info: ValidationInfo) -> Mapping[str, str] | None:
|
||||
|
|
|
|||
|
|
@ -1943,6 +1943,23 @@
|
|||
"interactions": true
|
||||
}
|
||||
},
|
||||
"prism": {
|
||||
"display_name": "Prism (`prism`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/prism",
|
||||
"endpoints": {
|
||||
"chat_completions": true,
|
||||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": false,
|
||||
"image_generations": false,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": false
|
||||
}
|
||||
},
|
||||
"recraft": {
|
||||
"display_name": "Recraft (`recraft`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/recraft",
|
||||
|
|
|
|||
|
|
@ -53,6 +53,7 @@ from litellm.repositories.verification_token_repository import (
|
|||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import MCPCredentials
|
||||
from litellm.types.mcp_server.mcp_server_manager import PinnedMCPTool
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_db_models
|
||||
|
|
@ -412,7 +413,6 @@ def _prepare_mcp_server_data(
|
|||
data_dict["tool_name_to_display_name"] = safe_dumps(data_dict["tool_name_to_display_name"] or {})
|
||||
if "tool_name_to_description" in data_dict:
|
||||
data_dict["tool_name_to_description"] = safe_dumps(data_dict["tool_name_to_description"] or {})
|
||||
|
||||
# mcp_access_groups is already List[str], no serialization needed
|
||||
|
||||
# On create, force is_byok so a False value is always written to the DB. On
|
||||
|
|
@ -2143,6 +2143,28 @@ async def approve_mcp_server(
|
|||
return table
|
||||
|
||||
|
||||
async def set_mcp_server_pinned_tools(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
pinned_tools: Mapping[str, PinnedMCPTool] | None,
|
||||
touched_by: str,
|
||||
) -> LiteLLM_MCPServerTable | None:
|
||||
"""Replace the server's pinned catalog; ``None`` unpins. Only this write path sets the pin."""
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
if await _db_find_mcp_server_row(prisma_client, server_id) is None:
|
||||
return None
|
||||
snapshot: Final = {name: tool.model_dump() for name, tool in (pinned_tools or {}).items()}
|
||||
updated: Final = await _db_update_mcp_server_row(
|
||||
prisma_client,
|
||||
server_id,
|
||||
{"pinned_tools": safe_dumps(snapshot), "updated_by": touched_by},
|
||||
)
|
||||
table: Final = LiteLLM_MCPServerTable.model_validate(updated.model_dump())
|
||||
decrypt_global_env_var_values(table.env_vars)
|
||||
return table
|
||||
|
||||
|
||||
async def reject_mcp_server(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.types.utils import CallTypes
|
|||
|
||||
guardrail_translation_mappings: Final = {
|
||||
CallTypes.call_mcp_tool: MCPGuardrailTranslationHandler,
|
||||
CallTypes.list_mcp_tools: MCPGuardrailTranslationHandler,
|
||||
}
|
||||
|
||||
__all__ = ["MCPGuardrailTranslationHandler", "guardrail_translation_mappings"]
|
||||
|
|
|
|||
|
|
@ -7,6 +7,11 @@ every string leaf of the call arguments as ``texts`` so text guardrails can
|
|||
detect and mask sensitive values in the payload. Works with the synthetic
|
||||
request from ProxyLogging._convert_mcp_to_llm_format.
|
||||
|
||||
A discovery scan (``list_mcp_tools``) hands the same handler the tool's
|
||||
description and input schema instead of call arguments: the description and
|
||||
every ``description`` string in the schema lead ``texts``, so a guardrail that
|
||||
blocks or masks them decides what the client gets to see in ``tools/list``.
|
||||
|
||||
Note: For MCP tool definitions (schema) -> OpenAI tools=[], see
|
||||
litellm.experimental_mcp_client.tools.transform_mcp_tool_to_openai_tool
|
||||
when you have a full MCP Tool from list_tools. Here we only have the call
|
||||
|
|
@ -58,23 +63,39 @@ def _too_deeply_nested() -> HTTPException:
|
|||
)
|
||||
|
||||
|
||||
def _argument_replacements(
|
||||
argument_leaves: tuple[tuple[JSONLeafPath, str], ...],
|
||||
masked_texts: Sequence[str] | None,
|
||||
) -> Mapping[JSONLeafPath, str]:
|
||||
"""Positionally pair the guardrail's returned texts with the leaves they came from.
|
||||
def _masked_texts(guarded: Mapping[str, object] | None, scanned: int) -> Sequence[str] | None:
|
||||
"""The guardrail's returned texts, or None when it returned nothing to write back.
|
||||
|
||||
Only leaves the guardrail actually rewrote are returned, so a guardrail that
|
||||
detects nothing leaves the outbound tool call byte-identical. A guardrail that
|
||||
returns the wrong number of texts fails closed, because a positional write-back
|
||||
would scramble the arguments rather than mask them.
|
||||
A guardrail that returns the wrong number of texts fails closed, because the
|
||||
positional write-back would scramble the payload rather than mask it.
|
||||
"""
|
||||
if masked_texts is not None and len(masked_texts) != len(argument_leaves):
|
||||
masked: Final[object] = guarded.get("texts") if guarded else None
|
||||
if masked is None:
|
||||
return None
|
||||
if not isinstance(masked, Sequence) or isinstance(masked, str) or len(masked) != scanned:
|
||||
raise _blocked(
|
||||
f"guardrail returned {len(masked_texts)} texts for {len(argument_leaves)} MCP tool call argument strings, "
|
||||
"so the redaction cannot be mapped back to the arguments"
|
||||
f"guardrail returned {len(masked) if isinstance(masked, Sequence) else 'no'} texts for {scanned} "
|
||||
"MCP tool strings, so the redaction cannot be mapped back"
|
||||
)
|
||||
return {path: masked for (path, original), masked in zip(argument_leaves, masked_texts or ()) if masked != original}
|
||||
return tuple(str(text) for text in masked)
|
||||
|
||||
|
||||
def _leaf_replacements(
|
||||
leaves: tuple[tuple[JSONLeafPath, str], ...],
|
||||
masked_texts: Sequence[str],
|
||||
) -> Mapping[JSONLeafPath, str]:
|
||||
"""Only the leaves the guardrail actually rewrote, so a guardrail that detects nothing leaves the payload byte-identical."""
|
||||
return {path: masked for (path, original), masked in zip(leaves, masked_texts) if masked != original}
|
||||
|
||||
|
||||
def _schema_description_leaves(input_schema: object) -> tuple[tuple[JSONLeafPath, str], ...]:
|
||||
leaves: Final = json_string_leaves(input_schema) if isinstance(input_schema, Mapping) else ()
|
||||
if leaves is None:
|
||||
raise _blocked(
|
||||
f"MCP tool input schema exceeds the maximum nesting depth of {MAX_STRUCTURED_CONTENT_SCAN_DEPTH} "
|
||||
"and cannot be scanned by the configured guardrail"
|
||||
)
|
||||
return tuple((path, text) for path, text in leaves if path and path[-1] == "description")
|
||||
|
||||
|
||||
def _conflicting_rewrite_paths(
|
||||
|
|
@ -125,6 +146,7 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
|
|||
mcp_tool_name: Final = data.get("mcp_tool_name") or data.get("name")
|
||||
mcp_arguments: Final[object] = data.get("mcp_arguments") or data.get("arguments")
|
||||
mcp_tool_description: Final = data.get("mcp_tool_description") or data.get("description")
|
||||
mcp_input_schema: Final[object] = data.get("mcp_input_schema")
|
||||
|
||||
if not mcp_tool_name:
|
||||
verbose_proxy_logger.debug("MCP Guardrail: mcp_tool_name missing")
|
||||
|
|
@ -135,7 +157,9 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
|
|||
mcp_tool: Final = MCPTool(
|
||||
name=mcp_tool_name,
|
||||
description=mcp_tool_description or "",
|
||||
input_schema={}, # mutable-ok: call payload has no schema; guardrail gets args from request_data
|
||||
input_schema=dict(mcp_input_schema)
|
||||
if isinstance(mcp_input_schema, Mapping)
|
||||
else {}, # mutable-ok: SDK dict field
|
||||
)
|
||||
openai_tool: Final = transform_mcp_tool_to_openai_tool(mcp_tool)
|
||||
fn: Final = openai_tool["function"]
|
||||
|
|
@ -153,12 +177,19 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
|
|||
strict=fn.get("strict", False) or False, # Default to False if None
|
||||
),
|
||||
}
|
||||
description_texts: Final = (str(mcp_tool_description),) if mcp_tool_description else ()
|
||||
schema_leaves: Final = _schema_description_leaves(mcp_input_schema)
|
||||
argument_leaves: Final = json_string_leaves(mcp_arguments)
|
||||
if argument_leaves is None:
|
||||
raise _too_deeply_nested()
|
||||
scanned_texts: Final = (
|
||||
*description_texts,
|
||||
*(text for _, text in schema_leaves),
|
||||
*(text for _, text in argument_leaves),
|
||||
)
|
||||
inputs: Final[GenericGuardrailAPIInputs] = GenericGuardrailAPIInputs(
|
||||
tools=[tool_def],
|
||||
texts=[text for _, text in argument_leaves],
|
||||
texts=list(scanned_texts),
|
||||
)
|
||||
|
||||
guarded: Final = await guardrail_to_apply.apply_guardrail(
|
||||
|
|
@ -167,10 +198,18 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
|
|||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
replacements: Final = _argument_replacements(
|
||||
argument_leaves=argument_leaves,
|
||||
masked_texts=guarded.get("texts") if guarded else None,
|
||||
)
|
||||
masked_texts: Final = _masked_texts(guarded, len(scanned_texts))
|
||||
if masked_texts is None:
|
||||
return data
|
||||
schema_start: Final = len(description_texts)
|
||||
argument_start: Final = schema_start + len(schema_leaves)
|
||||
if description_texts and masked_texts[0] != description_texts[0]:
|
||||
data["mcp_tool_description"] = masked_texts[0] # rebind-ok: serve the masked description
|
||||
schema_replacements: Final = _leaf_replacements(schema_leaves, masked_texts[schema_start:argument_start])
|
||||
if schema_replacements:
|
||||
masked_schema: Final = with_json_string_leaves(mcp_input_schema, schema_replacements)
|
||||
data["mcp_input_schema"] = masked_schema # rebind-ok: serve the masked schema
|
||||
replacements: Final = _leaf_replacements(argument_leaves, masked_texts[argument_start:])
|
||||
if not replacements:
|
||||
return data
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue