Merge branch 'main' into fix/responses-stream-empty-output-recovery

This commit is contained in:
Rain. 2026-09-23 17:57:14 +08:00 • committed by GitHub
commit 69f6634f3d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
1708 changed files with 163446 additions and 33232 deletions

View file

@ -6,6 +6,9 @@ parameters:
migration_candidate_image:
type: string
default: ""
migration_baseline_image:
type: string
default: "ghcr.io/berriai/litellm-database:v1.102.0"
migration_source_sha:
type: string
default: ""
@ -1508,7 +1511,7 @@ jobs:
- run:
name: Run tests
command: |
uv run --no-sync python -m pytest -vv tests/local_testing/test_basic_python_version.py -k "not v2_resolver"
uv run --no-sync python -m pytest -vv tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
installing_litellm_on_python_3_13:
docker:
@ -1532,7 +1535,7 @@ jobs:
- run:
name: Run tests
command: |
uv run --no-sync python -m pytest -v tests/local_testing/test_basic_python_version.py -k "not v2_resolver"
uv run --no-sync python -m pytest -v tests/local_testing/test_basic_python_version.py -k "not legacy_resolver"
installing_litellm_on_python_v2_migration_resolver:
docker:
@ -1561,10 +1564,11 @@ jobs:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Run v2 migration resolver proxy smoke test
name: Run both migration resolvers against Postgres
command: |
uv run --no-sync python -m pytest -vv \
tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings_v2_resolver
tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings \
tests/local_testing/test_basic_python_version.py::test_litellm_proxy_server_config_no_general_settings_legacy_resolver
helm_chart_testing:
machine:
@ -2945,7 +2949,10 @@ jobs:
parameters:
suite:
type: enum
enum: [startup, recovery, legacy]
enum: [startup, recovery, legacy, upgrade, shaped]
baseline:
type: boolean
default: false
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
@ -2953,6 +2960,7 @@ jobs:
environment:
LITELLM_MIGRATION_TESTS: "1"
LITELLM_MIGRATION_TEST_IMAGE: litellm-docker-database:ci
LITELLM_MIGRATION_BASELINE_IMAGE: << pipeline.parameters.migration_baseline_image >>
MIGRATION_TEST_ADMIN_URL: postgresql://postgres:postgres@127.0.0.1:5432/postgres
MIGRATION_TEST_CONTAINER_ADMIN_URL: postgresql://postgres:postgres@host.docker.internal:5432/postgres
MIGRATION_TEST_OUTPUT: /tmp/migration-results
@ -2980,6 +2988,16 @@ jobs:
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- when:
condition: << parameters.baseline >>
steps:
- run:
name: Pull the baseline release the upgrade starts from
environment:
BASELINE_IMAGE: << pipeline.parameters.migration_baseline_image >>
command: |
[[ "$BASELINE_IMAGE" =~ ^ghcr.io/berriai/[a-z0-9._/-]+(@sha256:[0-9a-f]{64}|:v[0-9][0-9a-z.-]*)$ ]] || exit 1
docker pull "$BASELINE_IMAGE"
- run:
name: Run migration startup regressions
environment:
@ -3032,28 +3050,29 @@ jobs:
- run:
name: Run Docker container with bad DATABASE_URL
command: |
set +e
docker run --name my-app \
-p 4000:4000 \
-e LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true \
-e DEFAULT_NUM_WORKERS_LITELLM_PROXY=1 \
-e DATABASE_URL="postgresql://wrong:wrong@wrong:5432/wrong" \
myapp:latest \
--port 4000 > docker_output.log 2>&1 || true
--port 4000 > docker_output.log 2>&1
echo "$?" > docker_exit_code
set -e
- run:
name: Display Docker logs
command: cat docker_output.log
- run:
name: Check for expected error
name: Proxy must refuse to serve on an unreachable database
command: |
if grep -q "Error: P1001: Can't reach database server at" docker_output.log && \
(grep -q "Database setup failed after multiple retries" docker_output.log || \
grep -q "ERROR: Application startup failed. Exiting." docker_output.log); then
echo "Expected error found. Test passed."
else
echo "Expected error not found. Test failed."
cat docker_output.log
exit 1
fi
fail() { echo "FAILED: $1"; cat docker_output.log; exit 1; }
exit_code="$(cat docker_exit_code)"
[ "$exit_code" -ne 0 ] || fail "proxy exited 0 with an unreachable database"
grep -q "P1001" docker_output.log || fail "log does not name the unreachable database server"
! grep -q "Application startup complete" docker_output.log || fail "proxy reached serving state"
! docker exec my-app true 2>/dev/null || fail "container is still running"
echo "Proxy refused to serve (exit $exit_code) and never reached startup. Test passed."
provider_replay_harness:
docker:
@ -3187,6 +3206,16 @@ workflows:
name: migration-legacy-and-pooling
suite: legacy
requires: [build_docker_database_image]
- migration_startup_tests:
name: migration-upgrade
suite: upgrade
baseline: true
requires: [build_docker_database_image]
- migration_startup_tests:
name: migration-upgrade-shaped
suite: shaped
baseline: true
requires: [build_docker_database_image]
migration_startup_scheduled:
triggers:
- schedule:

View file

@ -9,7 +9,6 @@ fi
suite="${1:?integration suite required}"
results="test-results/integration-${suite}"
mkdir -p "$results"
shard_timeout=11m
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
upstream_pid=""
proxy_pid=""
@ -121,12 +120,17 @@ start_proxy() {
"LITELLM_MODEL_COST_MAP_URL=$INTEGRATION_UPSTREAM_URL/_cost_map"
"MODEL_COST_MAP_MIN_MODEL_COUNT=1"
"MODEL_COST_MAP_MAX_SHRINK_RATIO=0"
"GEMINI_API_BASE=$INTEGRATION_UPSTREAM_URL"
"ANTHROPIC_API_BASE=$INTEGRATION_UPSTREAM_URL"
"GEMINI_API_KEY=sk-scripted-provider"
"ANTHROPIC_API_KEY=sk-scripted-provider"
)
else
cost_map_env=("LITELLM_LOCAL_MODEL_COST_MAP=True")
fi
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
@ -172,7 +176,7 @@ if [ "$suite" = browser ]; then
exit 0
fi
timeout --signal=TERM --kill-after=20s "$shard_timeout" env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
INTEGRATION_PROXY_URL="$INTEGRATION_PROXY_URL" INTEGRATION_PEER_URL="$INTEGRATION_PEER_URL" \

View file

@ -13,6 +13,8 @@ SUITES: Final = {
"startup": (("test_startup.py",), 12),
"recovery": (("test_recovery.py",), 15),
"legacy": (("test_legacy.py", "test_pooling.py"), 11),
"upgrade": (("test_upgrade.py", "test_rolling_upgrade.py"), 5),
"shaped": (("test_shaped_database.py",), 1),
}
@ -93,6 +95,7 @@ def main() -> int:
{
**metadata,
"suite": suite,
"baseline_image": os.environ.get("LITELLM_MIGRATION_BASELINE_IMAGE", ""),
"expected_cases": expected,
"passed": passed,
"pytest_exit_code": result.returncode,

View file

@ -15,6 +15,12 @@ description: >-
cache the same directory for different workloads, and a shared key would let
whichever ran first deny the others a save.
inputs:
profile:
description: "Cargo profile the build uses (dev or release)"
required: false
default: "dev"
runs:
using: composite
steps:
@ -25,6 +31,6 @@ runs:
~/.cargo/registry
~/.cargo/git
litellm-rust/target
key: ${{ runner.os }}-maturin-dev-${{ hashFiles('litellm-rust/Cargo.lock') }}
key: ${{ runner.os }}-maturin-${{ inputs.profile }}-${{ hashFiles('litellm-rust/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-maturin-dev-
${{ runner.os }}-maturin-${{ inputs.profile }}-

83
.github/e2e-stack/redact_output.py vendored Normal file
View file

@ -0,0 +1,83 @@
import argparse
import os
import sys
from functools import reduce
from pathlib import Path
from typing import Final
from xml.sax.saxutils import escape
from pydantic import JsonValue, TypeAdapter, ValidationError
from secrets_to_env import MIN_MASKED_LENGTH
REDACTED: Final = "***"
json_adapter: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
def string_leaves(node: JsonValue) -> tuple[str, ...]:
match node:
case str():
return (node,)
case list():
return tuple(leaf for child in node for leaf in string_leaves(child))
case dict():
return tuple(leaf for child in node.values() for leaf in string_leaves(child))
return ()
def field_lines(value: str) -> tuple[str, ...]:
try:
return tuple(line for leaf in string_leaves(json_adapter.validate_json(value)) for line in leaf.splitlines())
except ValidationError:
return ()
def masked_values(values_files: tuple[Path, ...]) -> tuple[str, ...]:
values: Final = frozenset(
line.split("=", 1)[1].strip().strip("'")
for path in values_files
for line in path.read_text().splitlines()
if "=" in line
)
texts: Final = frozenset(text for value in values for text in (value, *field_lines(value)))
renderings: Final = frozenset(
rendering
for text in texts
if len(text) >= MIN_MASKED_LENGTH
for rendering in (text, escape(text), escape(text, {'"': "&quot;"}))
)
return tuple(sorted(renderings, key=lambda rendering: (-len(rendering), rendering)))
def redact(text: str, values: tuple[str, ...]) -> str:
return reduce(lambda redacted, value: redacted.replace(value, REDACTED), values, text)
def write_redacted(source: Path, out_dir: Path, values: tuple[str, ...]) -> None:
target: Final = out_dir / source.name
with os.fdopen(os.open(target, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW, 0o600), "w") as handle:
_ = handle.write(redact(source.read_text(errors="replace"), values))
def main() -> int:
parser: Final = argparse.ArgumentParser()
_ = parser.add_argument("--values", action="append", type=Path, required=True)
_ = parser.add_argument("--out", type=Path, required=True)
_ = parser.add_argument("files", nargs="*", type=Path)
args: Final = parser.parse_args()
values_files: Final = tuple(args.values)
out_dir: Final[Path] = args.out
sources: Final = tuple(args.files)
try:
values: Final = masked_values(values_files)
out_dir.mkdir(mode=0o700, exist_ok=True)
for source in sources:
write_redacted(source, out_dir, values)
except OSError as error:
_ = sys.stderr.write(f"could not redact {error.filename}\n")
return 1
_ = sys.stdout.write(f"redacted {len(sources)} file(s) into {out_dir}\n")
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -9,6 +9,7 @@ from pydantic import TypeAdapter, ValidationError
secrets_adapter: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict[str, str])
ENV_NAME: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
MIN_MASKED_LENGTH: Final = 8
ACTIONS_RUNNER_FLAG: Final = "GITHUB_ACTIONS"
def main() -> int:
@ -30,10 +31,15 @@ def main() -> int:
f"these names or values cannot be represented in both bash and dotenv: {' '.join(sorted(unusable))}\n"
)
return 1
for value in secrets.values():
if len(value) >= MIN_MASKED_LENGTH:
_ = sys.stdout.write(f"::add-mask::{value.replace('%', '%25')}\n")
sys.stdout.flush()
if os.environ.get(ACTIONS_RUNNER_FLAG) == "true":
_ = sys.stdout.write(
"".join(
f"::add-mask::{value.replace('%', '%25')}\n"
for value in secrets.values()
if len(value) >= MIN_MASKED_LENGTH
)
)
sys.stdout.flush()
lines: Final = tuple(f"{key}='{value}'" for key, value in secrets.items() if value)
try:
with os.fdopen(os.open(env_path, os.O_WRONLY | os.O_APPEND | os.O_CREAT | os.O_NOFOLLOW, 0o600), "w") as handle:

View file

@ -9,6 +9,9 @@ UNSUPPORTED: Final = re.compile(
r"|^tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e\.py$"
r"|^tests/e2e/batches/test_managed_files_enforcement_e2e\.py$"
r"|^tests/e2e/guardrails/test_presidio_masking_e2e\.py$"
r"|^tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e\.py$"
r"|^tests/e2e/logging/test_langsmith_batch_serialization_e2e\.py$"
r"|^tests/e2e/secret_manager/"
)
HARNESS: Final = re.compile(
r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$"

View file

@ -24,6 +24,7 @@ DATABASE_USER="${E2E_DATABASE_USER:-litellm}"
DATABASE_PASSWORD="${E2E_DATABASE_PASSWORD:-dbpassword9090}"
DATABASE_NAME="${E2E_DATABASE_NAME:-litellm}"
JAEGER_OTLP_PORT="${E2E_JAEGER_OTLP_PORT:-4318}"
JAEGER_OTLP_TLS_PORT="${E2E_JAEGER_OTLP_TLS_PORT:-4319}"
JAEGER_QUERY_PORT="${E2E_JAEGER_QUERY_PORT:-16686}"
KEYCLOAK_PORT="${E2E_KEYCLOAK_PORT:-8081}"
@ -122,7 +123,7 @@ SERVER_ENV=(
"CONFIG_FILE_PATH=${CONFIG_PATH}"
"STORE_MODEL_IN_DB=True"
"OTEL_EXPORTER_OTLP_PROTOCOL=http/protobuf"
"OTEL_EXPORTER_OTLP_ENDPOINT=http://127.0.0.1:${JAEGER_OTLP_PORT}"
"OTEL_EXPORTER_OTLP_ENDPOINT=https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}"
"SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem"
"PYTHONPATH=${REPO_ROOT}"
"JWT_PUBLIC_KEY_URL=http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e/protocol/openid-connect/certs"
@ -143,20 +144,16 @@ env "${SERVER_ENV[@]}" uv run --no-sync python migrations/run.py >"${LOGS_DIR}/m
start_server() {
local name="$1"; shift
env "${SERVER_ENV[@]}" "$@" >"${LOGS_DIR}/${name}.log" 2>&1 &
env -u AWS_ROLE_NAME "${SERVER_ENV[@]}" "$@" >"${LOGS_DIR}/${name}.log" 2>&1 &
echo $! > "${PIDS_DIR}/${name}.pid"
}
start_server backend uv run --no-sync uvicorn backend.main:app --host 0.0.0.0 --port "${BACKEND_PORT}"
start_server gateway-1 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_1}"
start_server gateway-2 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_2}"
if [[ "$(uname)" == "Linux" ]]; then
NGINX_UPSTREAM_HOST=127.0.0.1
NGINX_DOCKER_ARGS=(--network host)
else
NGINX_UPSTREAM_HOST=host.docker.internal
NGINX_DOCKER_ARGS=(-p "${LB_PORT}:${LB_PORT}")
NGINX_DOCKER_ARGS=(-p "${LB_PORT}:${LB_PORT}" -p "${JAEGER_OTLP_TLS_PORT}:${JAEGER_OTLP_TLS_PORT}")
fi
cat > "${STACK_DIR}/nginx.conf" <<EOF
@ -186,12 +183,29 @@ http {
proxy_send_timeout 600s;
}
}
server {
listen ${JAEGER_OTLP_TLS_PORT} ssl;
ssl_certificate /certs/server.crt;
ssl_certificate_key /certs/server.key;
client_max_body_size 100m;
location / {
proxy_pass http://${NGINX_UPSTREAM_HOST}:${JAEGER_OTLP_PORT};
}
}
}
EOF
docker rm -f e2e-nginx >/dev/null 2>&1 || true
docker run -d --name e2e-nginx "${NGINX_DOCKER_ARGS[@]}" \
-v "${STACK_DIR}/nginx.conf:/etc/nginx/nginx.conf:ro" "${NGINX_IMAGE}" >/dev/null
-v "${STACK_DIR}/nginx.conf:/etc/nginx/nginx.conf:ro" \
-v "${CERTS_DIR}:/certs:ro" "${NGINX_IMAGE}" >/dev/null
wait_for "Jaeger OTLP TLS listener" \
"curl -sS --cacert ${CERTS_DIR}/ca.crt https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}/ -o /dev/null -w '%{http_code}' | grep -qE '^[2345]'"
start_server backend uv run --no-sync uvicorn backend.main:app --host 0.0.0.0 --port "${BACKEND_PORT}"
start_server gateway-1 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_1}"
start_server gateway-2 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_2}"
wait_for "backend" "curl -fs http://127.0.0.1:${BACKEND_PORT}/health/liveliness >/dev/null" 300
wait_for "gateway-1" "curl -fs http://127.0.0.1:${GATEWAY_PORT_1}/health/liveliness >/dev/null" 300
@ -206,6 +220,7 @@ LITELLM_MASTER_KEY=${MASTER_KEY}
REDIS_HOST=127.0.0.1
REDIS_PORT=${REDIS_PORT}
E2E_OTEL_QUERY_URL=http://127.0.0.1:${JAEGER_QUERY_PORT}
E2E_OTEL_EXPORTER_ENDPOINT=https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}
E2E_KEYCLOAK_URL=http://127.0.0.1:${KEYCLOAK_PORT}
E2E_KEYCLOAK_ADMIN_USER=admin
E2E_KEYCLOAK_ADMIN_PASSWORD=e2e-ephemeral-idp-not-a-secret

View file

@ -134,7 +134,16 @@ def main(
uncompressed_wheel_size: Final = sum(member.file_size for member in wheel_members)
native_path: Final = wheel.parent / "native" / Path(native_member.filename).name
native_path.parent.mkdir(parents=True, exist_ok=True)
native_path.write_bytes(archive.read(native_member))
native_bytes: Final = archive.read(native_member)
native_path.write_bytes(native_bytes)
duplicated_vocabularies: Final = tuple(
member.filename
for member in wheel_members
if member.filename.startswith("litellm/litellm_core_utils/tokenizers/")
and re.fullmatch(r"[0-9a-f]{40}", PurePosixPath(member.filename).name)
and member.file_size > 0
and archive.read(member) in native_bytes
)
wheel_metadata_tags_match: Final = (
len(wheel_metadata_tags) == len(expanded_filename_tags)
@ -205,7 +214,7 @@ def main(
native_module: Final = load_native_module(native_path)
native_module_loads: Final = native_module is not None
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
native_size_limit: Final = 25_000_000
native_size_limit: Final = 35_000_000
native_size_within_limit: Final = native_member.file_size <= native_size_limit
validations: Final = (
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
@ -222,7 +231,8 @@ def main(
("Python extension entry point is present", extension_entry_point_present),
("Native module loads", native_module_loads),
("Production module omits the panic test hook", panic_test_hook_absent),
("Native extension does not exceed 25 MB", native_size_within_limit),
(f"Native extension does not exceed {native_size_limit / 1_000_000:.0f} MB", native_size_within_limit),
("Tokenizer vocabularies are not duplicated in the native extension", not duplicated_vocabularies),
("Wheel contents are valid", not unexpected_members),
)
@ -267,7 +277,8 @@ def main(
),
(
not native_size_within_limit,
f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB",
f"native extension exceeds {native_size_limit / 1_000_000:.0f} MB: "
f"{native_member.file_size / 1_000_000:.2f} MB",
),
(bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"),
)

View file

@ -4,7 +4,13 @@ on:
workflow_call:
inputs:
test-path:
description: "Pytest path(s) to run"
description: >-
Space-separated pytest paths to run. A path that no longer exists is
dropped with a warning instead of being passed to pytest, because one
missing path makes pytest-xdist collect nothing and report exit 5, which
the step treats as a drained shard. Options are passed through as
written, so use the `--flag=value` form: a bare `--ignore path` would
have its path existence-checked like any other token.
required: true
type: string
workers:
@ -165,14 +171,22 @@ jobs:
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
run: |
found_path=false
for path in ${TEST_PATH}; do
if [ -e "${path%%::*}" ]; then
found_path=true
break
fi
pytest_args=()
existing_paths=0
for token in ${TEST_PATH:?}; do
case "${token}" in
-*) pytest_args+=("${token}") ;;
*)
if [ -e "${token%%::*}" ]; then
pytest_args+=("${token}")
existing_paths=$((existing_paths + 1))
else
echo "::warning::${token} does not exist; drop it from this shard's test-path"
fi
;;
esac
done
if [ "$found_path" = false ]; then
if [ "${existing_paths}" -eq 0 ]; then
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
exit 0
fi
@ -181,7 +195,7 @@ jobs:
xdist_args=(-n "${WORKERS}" --dist="${DIST}")
fi
set +e
uv run --no-sync pytest ${TEST_PATH:?} \
uv run --no-sync pytest "${pytest_args[@]}" \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
"${xdist_args[@]}" \

View file

@ -13,6 +13,7 @@ on:
- ".github/workflows/codspeed.yml"
- ".github/actions/setup-uv-with-retries/**"
- ".github/actions/cache-cargo-build/**"
- ".github/scripts/uv_sync_with_retries.sh"
pull_request:
branches:
- main
@ -25,6 +26,7 @@ on:
- ".github/workflows/codspeed.yml"
- ".github/actions/setup-uv-with-retries/**"
- ".github/actions/cache-cargo-build/**"
- ".github/scripts/uv_sync_with_retries.sh"
# Allow CodSpeed to trigger backtest performance analysis
# in order to generate initial data
workflow_dispatch:
@ -59,19 +61,27 @@ jobs:
- name: Cache the Rust build
uses: ./.github/actions/cache-cargo-build
with:
profile: release
# Build the wheel and resolve every dependency outside the CodSpeed
# runner: the same maturin build took 42 minutes inside `codspeed run`
# versus under 3 minutes as a plain step (LIT-6183)
- name: Build environment
- name: Build the release wheel
run: uv build --wheel --out-dir dist
- name: Install the wheel into the benchmark environment
run: |
UV_PROJECT_ENVIRONMENT="${RUNNER_TEMP}/benchmark-venv" .github/scripts/uv_sync_with_retries.sh --frozen --no-default-groups --group benchmarks --no-install-project --python 3.12
uv pip install --python "${RUNNER_TEMP}/benchmark-venv/bin/python" --no-deps dist/*.whl
- name: Collect benchmarks
env:
PYTEST_DISABLE_PLUGIN_AUTOLOAD: "1"
LITELLM_REQUIRE_INSTALLED_WHEEL: "1"
run: >
env PYTEST_DISABLE_PLUGIN_AUTOLOAD=1
uv run --frozen --no-default-groups
--with pytest==8.3.5
--with pytest-codspeed==4.3.0
--with "mcp>=2.2.0,<3.0"
--with "a2a-sdk>=1.1.0,<2.0"
pytest
"${RUNNER_TEMP}/benchmark-venv/bin/python" -I -m pytest
--import-mode=importlib
-p pytest_codspeed.plugin
tests/benchmarks/
--codspeed
@ -82,13 +92,9 @@ jobs:
with:
mode: simulation
run: >
env PYTEST_DISABLE_PLUGIN_AUTOLOAD=1
uv run --frozen --no-default-groups
--with pytest==8.3.5
--with pytest-codspeed==4.3.0
--with "mcp>=2.2.0,<3.0"
--with "a2a-sdk>=1.1.0,<2.0"
pytest
env PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 LITELLM_REQUIRE_INSTALLED_WHEEL=1
"${RUNNER_TEMP}/benchmark-venv/bin/python" -I -m pytest
--import-mode=importlib
-p pytest_codspeed.plugin
tests/benchmarks/
--codspeed

View file

@ -0,0 +1,33 @@
name: Compat Matrix Image
on:
pull_request:
paths:
- tests/e2e/claude_code/cron_vm/**
- .github/workflows/compat-matrix-image.yml
workflow_dispatch:
permissions: {}
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
compat-matrix-image:
name: compat-matrix-image
runs-on: ubuntu-latest
timeout-minutes: 15
permissions:
contents: read
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Build the Render cron image
run: docker build -f tests/e2e/claude_code/cron_vm/Dockerfile -t compat-matrix:${{ github.sha }} tests/e2e
- name: Run the pinned binaries as the cron user
run: |
docker run --rm compat-matrix:${{ github.sha }} bash -c 'set -e; whoami; claude --version; gh --version; uv --version'

View file

@ -1,186 +0,0 @@
name: Create Release
on:
workflow_dispatch:
inputs:
tag:
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0-dev.2, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
required: true
type: string
commit_hash:
description: "Full 40-char commit SHA to target"
required: true
type: string
permissions: {}
jobs:
release:
name: Create Release
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Validate inputs
env:
TAG: ${{ inputs.tag }}
COMMIT_HASH: ${{ inputs.commit_hash }}
run: |
if ! echo "${COMMIT_HASH}" | grep -qE '^[0-9a-f]{40}$'; then
echo "::error::commit_hash must be a full 40-character commit SHA"
exit 1
fi
if ! echo "${TAG}" | grep -qE '^v?[0-9]+\.[0-9]+\.[0-9]+'; then
echo "::error::tag must start with X.Y.Z (optional leading v), e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, or v1.83.10-stable"
exit 1
fi
- name: Create release
env:
TAG: ${{ inputs.tag }}
COMMIT_HASH: ${{ inputs.commit_hash }}
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
with:
script: |
const tag = process.env.TAG;
const commitHash = process.env.COMMIT_HASH;
// Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases.
// Accept both PEP 440 (`.dev`) and SemVer (`-dev`) separators so tags
// like `1.84.0.dev2` and `1.84.0-dev.2` are both detected.
// PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]`
// are stable maintenance releases, not pre-releases.
const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag);
// A stable release should only claim the repo "latest" badge when its
// version is >= the current latest. Otherwise a backport (e.g. 1.84.6)
// would steal "latest" from a newer line (e.g. 1.88.1).
const versionKey = (rawTag) => {
const m = String(rawTag).match(/^v?(\d+)\.(\d+)\.(\d+)/);
if (!m) return null;
const maintenance = String(rawTag).match(/(?:\.post|\.patch\.)(\d+)/i);
return [Number(m[1]), Number(m[2]), Number(m[3]), maintenance ? Number(maintenance[1]) : 0];
};
const isAtLeast = (a, b) => {
for (let i = 0; i < a.length; i++) {
if (a[i] !== b[i]) return a[i] > b[i];
}
return true;
};
const cosignSection = [
`## Verify Docker Image Signature`,
``,
`All LiteLLM Docker images are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit \`0112e53\`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).`,
``,
`**Verify using the pinned commit hash (recommended):**`,
``,
`A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:`,
``,
'```bash',
`cosign verify \\`,
` --key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \\`,
` ghcr.io/berriai/litellm:${tag}`,
'```',
``,
`**Verify using the release tag (convenience):**`,
``,
`Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:`,
``,
'```bash',
`cosign verify \\`,
` --key https://raw.githubusercontent.com/BerriAI/litellm/${tag}/cosign.pub \\`,
` ghcr.io/berriai/litellm:${tag}`,
'```',
``,
`Expected output:`,
``,
'```',
`The following checks were performed on each of these signatures:`,
` - The cosign claims were validated`,
` - The signatures were verified against the specified public key`,
'```',
``,
`---`,
``,
].join('\n');
try {
let makeLatest = "false";
const newVersion = versionKey(tag);
if (!isPrerelease && newVersion) {
let latestVersion = null;
try {
const latest = await github.rest.repos.getLatestRelease({
owner: context.repo.owner,
repo: context.repo.repo,
});
latestVersion = versionKey(latest.data.tag_name);
} catch (error) {
if (error.status !== 404) throw error;
}
makeLatest = (!latestVersion || isAtLeast(newVersion, latestVersion)) ? "true" : "false";
}
try {
await github.rest.git.createRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `refs/tags/${tag}`,
sha: commitHash,
});
} catch (error) {
if (error.status !== 422) throw error;
const existing = await github.rest.git.getRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `tags/${tag}`,
});
if (existing.data.object.sha !== commitHash) {
throw new Error(`Tag ${tag} already exists at ${existing.data.object.sha}, expected ${commitHash}`);
}
}
const response = await github.rest.repos.createRelease({
draft: true,
generate_release_notes: true,
name: tag,
owner: context.repo.owner,
prerelease: isPrerelease,
repo: context.repo.repo,
tag_name: tag,
});
const updatedBody = cosignSection + (response.data.body ?? '');
await github.rest.repos.updateRelease({
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
body: updatedBody,
draft: false,
});
if (!isPrerelease) {
await github.rest.repos.updateRelease({
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
make_latest: makeLatest,
});
}
} catch (error) {
core.setFailed(error.message);
}
create-branch:
name: Create Release Branch
needs: release
permissions:
contents: write
uses: ./.github/workflows/create-release-branch.yml
with:
tag: ${{ inputs.tag }}
commit_hash: ${{ inputs.commit_hash }}

View file

@ -6,8 +6,12 @@ on:
workflow_dispatch:
inputs:
issue_number:
description: "Closed issue number to comment on manually."
required: true
description: "Closed issue number to comment on and close the superseded pull requests of. Ignored by a sweep."
required: false
sweep:
description: "Close every open pull request whose linked issues were all fixed on the default branch. Reads every open pull request, so run it at most once an hour."
type: boolean
default: false
pull_request:
paths:
- .github/workflows/issue_fixed_comment.yml
@ -39,16 +43,17 @@ jobs:
with:
bun-version: "1.4.0"
- name: Test the closer lookup, the release placement and the comment
- name: Test the closer lookup, the release placement, the comment and the superseded pull request close
run: bun test scripts/comment-fixed-issue.test.ts
comment-fixed-issue:
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
timeout-minutes: 5
timeout-minutes: 15
permissions:
contents: read
issues: write
pull-requests: write
steps:
- name: Checkout scripts
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
@ -59,13 +64,16 @@ jobs:
- name: Setup Bun
uses: oven-sh/setup-bun@0c5077e51419868618aeaa5fe8019c62421857d6 # v2.2.0
with:
# Exact version, never latest: the next step holds an issues: write token
# Exact version, never latest: the next step holds issues: write and pull-requests: write tokens
bun-version: "1.4.0"
- name: Name the release that carries the fix
- name: Name the release that carries the fix and close the pull requests it supersedes
shell: bash
run: bun run scripts/comment-fixed-issue.ts | tee -a "${GITHUB_STEP_SUMMARY}"
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
ISSUE_NUMBER: ${{ github.event.issue.number || github.event.inputs.issue_number }}
SWEEP: ${{ github.event.inputs.sweep }}
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
DRY_RUN: ${{ vars.ISSUE_FIXED_COMMENT_ENABLED != 'true' }}
CLOSE_PRS_DRY_RUN: ${{ vars.ISSUE_FIXED_CLOSE_PRS_ENABLED != 'true' }}

View file

@ -175,6 +175,8 @@ jobs:
env:
TESTS: ${{ needs.detect.outputs.tests }}
E2E_FIXTURE_MODE: live
E2E_PROVIDER_EDGE_HOST_REACHABLE: '1'
COLUMNS: '400'
run: |
umask 077
read -r -a test_files <<< "${TESTS}"
@ -189,6 +191,7 @@ jobs:
uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}"
verified=$?
set -e
grep -E '^(FAILED|ERROR) ' "${log}" || true
grep -E '^=+ .* in [0-9.]+s( \([0-9:]+\))? =+$' "${log}" | tail -n 1
echo "::endgroup::"
if [ "${status}" = "5" ]; then
@ -206,6 +209,24 @@ jobs:
echo "pass ${pass} of 3 passed"
done
- name: Redact the pytest output
if: always() && steps.boot.outcome == 'success'
run: |
umask 077
shopt -s nullglob
uv run --no-sync python .github/e2e-stack/redact_output.py \
--values tests/e2e/.env --values "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" \
--out "${RUNNER_TEMP}/e2e-redacted" "${RUNNER_TEMP}"/e2e-pass-*.log "${RUNNER_TEMP}"/e2e-pass-*.xml
- name: Keep the redacted pytest output
if: always() && steps.boot.outcome == 'success'
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: e2e-changed-pytest-output-${{ github.run_attempt }}
path: ${{ runner.temp }}/e2e-redacted
retention-days: 14
if-no-files-found: ignore
- name: Stop the stack
if: always() && steps.boot.outcome != 'skipped'
run: bash .github/e2e-stack/down.sh
@ -214,7 +235,7 @@ jobs:
if: always()
run: |
rm -f tests/e2e/.env "${RUNNER_TEMP}/e2e-boot.log" "${RUNNER_TEMP}"/e2e-pass-*.log "${RUNNER_TEMP}"/e2e-pass-*.xml
rm -rf "${RUNNER_TEMP}/litellm-e2e-stack"
rm -rf "${RUNNER_TEMP}/litellm-e2e-stack" "${RUNNER_TEMP}/e2e-redacted"
gate:
name: e2e-changed-tests

View file

@ -130,6 +130,10 @@ jobs:
echo "File content around line 43:"
head -50 litellm/litellm_core_utils/custom_logger_registry.py | tail -10
- name: Check MCP operation boundary
if: steps.changes.outputs.decision != 'skip'
run: uv run --no-sync python scripts/check_mcp_operation_boundary.py
- name: Run Ruff linting
if: steps.changes.outputs.decision != 'skip'
run: |

View file

@ -105,6 +105,16 @@ jobs:
with:
python-version: "3.12"
- uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install Python dependencies for the bridge tests
working-directory: .
run: |
uv sync --frozen --no-install-project
echo "PYTHONPATH=$PWD/.venv/lib/$(ls .venv/lib)/site-packages" >> "$GITHUB_ENV"
- run: rustup toolchain install --no-self-update
- uses: taiki-e/install-action@d438492cf8a250514fa2d34b30bc3c0dc37c65ff # v2.87.8
@ -130,7 +140,7 @@ jobs:
- name: Test secret manager feature combinations
run: |
cargo test -p litellm-auth-gcp --locked --no-default-features
for features in '' aws google aws,google; do
for features in '' aws google hashicorp azure cyberark aws,google aws,azure google,azure aws,google,azure aws,google,cyberark aws,google,azure,cyberark aws,google,hashicorp,azure,cyberark; do
cargo test -p litellm-secrets --locked --no-default-features --features "$features"
done

View file

@ -107,26 +107,18 @@ jobs:
tests/test_litellm/batches
tests/test_litellm/secret_managers
tests/test_litellm/a2a_protocol
tests/test_litellm/anthropic_interface
tests/test_litellm/chat_completions
tests/test_litellm/completion_extras
tests/test_litellm/compression
tests/test_litellm/containers
tests/test_litellm/endpoints
tests/test_litellm/models
tests/test_litellm/repositories
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/messages
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rag
tests/test_litellm/realtime_api
tests/test_litellm/rerank_api
tests/test_litellm/rust_bridge
tests/test_litellm/sandbox
tests/test_litellm/skills
tests/test_litellm/test_router
tests/test_litellm/vector_stores
tests/test_litellm/videos
tests/test_litellm/test_*.py

View file

@ -1,10 +1,10 @@
# syntax=docker/dockerfile:1.7
# Base image for building
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43

View file

@ -164,6 +164,7 @@ lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
# Linting targets
lint-ruff: $(LINT_DEP_INSTALL)
$(UV_RUN) python scripts/check_mcp_operation_boundary.py
cd litellm && $(UV_RUN) ruff check . && cd ..
$(UV_RUN) ruff check --config ruff-tests.toml tests

View file

@ -307,6 +307,7 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
| [Deepgram (`deepgram`)](https://docs.litellm.ai/docs/providers/deepgram) | ✅ | ✅ | ✅ | | | ✅ | | | | |
| [DeepInfra (`deepinfra`)](https://docs.litellm.ai/docs/providers/deepinfra) | ✅ | ✅ | ✅ | | | | | | | |
| [Deepseek (`deepseek`)](https://docs.litellm.ai/docs/providers/deepseek) | ✅ | ✅ | ✅ | | | | | | | |
| [Eden AI (`edenai`)](https://docs.litellm.ai/docs/providers/edenai) | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | | | |
| [ElevenLabs (`elevenlabs`)](https://docs.litellm.ai/docs/providers/elevenlabs) | ✅ | ✅ | ✅ | | | ✅ | ✅ | | | |
| [Empower (`empower`)](https://docs.litellm.ai/docs/providers/empower) | ✅ | ✅ | ✅ | | | | | | | |
| [Fal AI (`fal_ai`)](https://docs.litellm.ai/docs/providers/fal_ai) | ✅ | ✅ | ✅ | | ✅ | | | | | |
@ -356,7 +357,7 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
| [Petals (`petals`)](https://docs.litellm.ai/docs/providers/petals) | ✅ | ✅ | ✅ | | | | | | | |
| [Pinstripes (`pinstripes`)](https://docs.litellm.ai/docs/providers/pinstripes) | ✅ | ✅ | ✅ | | | | | | | |
| [Predibase (`predibase`)](https://docs.litellm.ai/docs/providers/predibase) | ✅ | ✅ | ✅ | | | | | | | |
| [Qwen AI Platform (`qwen_ai_platform`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
| [Qianwen AI Platform (`qwen_ai_platform`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
| [QwenCloud (`qwencloud`)](https://docs.litellm.ai/docs/providers/qwencloud) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | ✅ |
| [Recraft (`recraft`)](https://docs.litellm.ai/docs/providers/recraft) | | | | | ✅ | | | | | |
| [Replicate (`replicate`)](https://docs.litellm.ai/docs/providers/replicate) | ✅ | ✅ | ✅ | | | | | | | |

View file

@ -1,5 +1,5 @@
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
@ -61,6 +61,8 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra saml \
--python python3.13
RUN cp "$(python -c 'import sysconfig; print(sysconfig.get_paths()["purelib"])')"/litellm/rust_bridge/_native*.so litellm/rust_bridge/
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
npm_config_cache=/root/.npm \
prisma generate --schema=./schema.prisma

View file

@ -26,6 +26,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/v2/login",
"/v3/login",
"/logout",
"/session/logout",
"/token",
"/onboarding/",
"/audit",
@ -51,6 +52,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/cache_settings",
"/coordination_redis/",
"/cost_tracking",
"/cost_optimization/",
"/cost/",
"/credentials",
"/credential",

View file

@ -1,8 +1,11 @@
"""Guard the cost map on pull requests.
Every pull request gets the file checks: the three cost map files parse, the backup copy matches the root file,
and the JSON schema is in sync and validates the map. Pull requests from the cost map sync bot (branches named
litellm_cost_map_sync_*) additionally may only touch those three files and may only add or update models.
Every pull request whose diff against its merge base touches one of the three cost map files gets the file
checks: the files parse, the backup copy matches the root file, and the JSON schema is in sync and validates the
map. A pull request that leaves all three untouched skips them, since merging it keeps the base branch's copies
and its head tree only carries whatever state the branch was cut from. Pull requests from the cost map sync bot
(branches named litellm_cost_map_sync_*) always get the file checks and additionally may only touch those three
files and may only add or update models.
"""
from __future__ import annotations
@ -108,20 +111,37 @@ def _bot_failures(base: Snapshot, head_map: CostMap, changed_files: Sequence[str
)
def touches_cost_map(changed_files: Sequence[str]) -> bool:
return any(path in GUARDED_PATHS for path in changed_files)
def contract_for(bot: bool, changed_files: Sequence[str]) -> str:
if bot:
return "bot contract enforced"
return "human PR, file checks only" if touches_cost_map(changed_files) else "human PR, cost map untouched"
def guard_failures(base: Snapshot, head: Snapshot, changed_files: Sequence[str], bot: bool) -> tuple[str, ...]:
if not bot and not touches_cost_map(changed_files):
return ()
head_map: Final = _parse_object(head.cost_map, COST_MAP_PATH)
if isinstance(head_map, str):
return (head_map,)
return (*_file_failures(head, head_map), *(_bot_failures(base, head_map, changed_files) if bot else ()))
def _git(*args: str) -> str:
def _git(*args: str) -> str | None:
result: Final = subprocess.run(("git", *args), check=False, capture_output=True, text=True)
return result.stdout if result.returncode == 0 else ""
return result.stdout if result.returncode == 0 else None
def snapshot(revision: str) -> Snapshot:
return Snapshot(*(_git("show", f"{revision}:{path}") for path in GUARDED_PATHS))
return Snapshot(*(_git("show", f"{revision}:{path}") or "" for path in GUARDED_PATHS))
def changed_files(base: str, head: str) -> tuple[str, ...] | None:
diff: Final = _git("diff", "--name-only", "--no-renames", base, head)
return None if diff is None else tuple(diff.splitlines())
def main(argv: Sequence[str]) -> int:
@ -131,9 +151,12 @@ def main(argv: Sequence[str]) -> int:
parser.add_argument("--head-ref", required=True, help="head branch name of the pull request")
args: Final = parser.parse_args(argv)
bot: Final = args.head_ref.startswith(BOT_BRANCH_PREFIX)
changed_files: Final = tuple(_git("diff", "--name-only", args.base, args.head).splitlines())
failures: Final = guard_failures(snapshot(args.base), snapshot(args.head), changed_files, bot)
contract: Final = "bot contract enforced" if bot else "human PR, file checks only"
changed: Final = changed_files(args.base, args.head)
if changed is None:
print(f"cost map guard failed: git diff {args.base} {args.head} failed, so the changed files are unknown")
return 1
failures: Final = guard_failures(snapshot(args.base), snapshot(args.head), changed, bot)
contract: Final = contract_for(bot, changed)
if failures:
print(f"cost map guard failed ({contract}):")
print("\n".join(f"- {failure}" for failure in failures))

View file

@ -6267,6 +6267,63 @@
],
"title": "Spend update queue sizes (litellm_<queue>_size)",
"type": "timeseries"
},
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"description": "Requests that carried usage but were logged at $0 on a model whose pricing entry has a non-zero rate, by requested model and reason",
"fieldConfig": {
"defaults": {
"color": {
"mode": "palette-classic"
},
"custom": {
"drawStyle": "line",
"fillOpacity": 10,
"lineWidth": 1,
"showPoints": "never",
"spanNulls": false
},
"unit": "short"
},
"overrides": []
},
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 430
},
"id": 110,
"options": {
"legend": {
"calcs": [],
"displayMode": "list",
"placement": "bottom",
"showLegend": true
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "sum(rate(litellm_zero_cost_requests_total[$__rate_interval])) by (requested_model, reason)",
"legendFormat": "{{requested_model}} / {{reason}}",
"range": true,
"refId": "A"
}
],
"title": "litellm_zero_cost_requests rate",
"type": "timeseries"
}
],
"preload": false,

View file

@ -1,10 +1,10 @@
# syntax=docker/dockerfile:1.7
# Base image for building
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
# Runtime image
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
# Pinned by digest like the other base images; bump explicitly on Node upgrades.
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43

View file

@ -1,8 +1,8 @@
# syntax=docker/dockerfile:1.7
# Base images
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG PROXY_EXTRAS_SOURCE=published
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
# Pinned by digest like the other base images; bump explicitly on Node upgrades.

View file

@ -2,6 +2,17 @@
This guide provides instructions for building and running the LiteLLM application using Docker and Docker Compose.
> **Just want to run LiteLLM?** This guide builds from source. To run the published
> image instead, use `docker-compose.quickstart.yml` in this directory — the
> two-service stack (gateway + Postgres) that the
> [Docker quickstart](https://docs.litellm.ai/docs/proxy/docker_quick_start) documents:
>
> ```bash
> curl -sSLO https://github.com/BerriAI/litellm/raw/main/docker/docker-compose.quickstart.yml
> printf 'LITELLM_MASTER_KEY=sk-%s\nLITELLM_SALT_KEY=sk-%s\n' "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" > .env
> docker compose -f docker-compose.quickstart.yml up -d
> ```
## Prerequisites
- Docker

View file

@ -0,0 +1,41 @@
# LiteLLM quickstart stack: the gateway plus a Postgres database that stores
# models, virtual keys, and spend logs. Used by
# https://docs.litellm.ai/docs/proxy/docker_quick_start
#
# curl -sSLO https://github.com/BerriAI/litellm/raw/main/docker/docker-compose.quickstart.yml
# printf 'LITELLM_MASTER_KEY=sk-%s\nLITELLM_SALT_KEY=sk-%s\n' "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" > .env
# docker compose -f docker-compose.quickstart.yml up -d
#
# Compose reads .env from this directory. Keep it: regenerating LITELLM_SALT_KEY
# makes credentials already stored in the database unreadable. For anything
# beyond local evaluation, pin the image to a specific release tag.
services:
litellm:
image: docker.litellm.ai/berriai/litellm:main-stable
ports:
- "4000:4000"
environment:
LITELLM_MASTER_KEY: ${LITELLM_MASTER_KEY:?set it in .env - see the header of this file}
LITELLM_SALT_KEY: ${LITELLM_SALT_KEY:?set it in .env - see the header of this file}
DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm
STORE_MODEL_IN_DB: "True"
depends_on:
db:
condition: service_healthy
db:
image: postgres:16
environment:
POSTGRES_USER: litellm
POSTGRES_PASSWORD: litellm
POSTGRES_DB: litellm
healthcheck:
test: ["CMD-SHELL", "pg_isready -U litellm"]
interval: 5s
timeout: 5s
retries: 10
volumes:
- postgres_data:/var/lib/postgresql/data
volumes:
postgres_data:

View file

@ -2,10 +2,11 @@
Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if the cost has been tracked.
"""
from collections.abc import Sequence
from dataclasses import replace as dataclasses_replace
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast
from typing import TYPE_CHECKING, Final, List, Literal, Optional, Protocol, Tuple, cast
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -18,8 +19,8 @@ if TYPE_CHECKING:
from prisma import models as prisma_models
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy._types import LiteLLM_ManagedObjectTable
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.prisma_protocols import TableActions
from litellm.router import Router
from litellm.types.router import Deployment
from litellm.types.utils import LiteLLMBatch
@ -41,6 +42,42 @@ TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = (
)
class _ManagedObjectRow(Protocol):
@property
def id(self) -> str: ...
@property
def unified_object_id(self) -> str: ...
@property
def created_by(self) -> str | None: ...
@property
def file_object(self) -> object: ...
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]":
table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable
return table
def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]":
table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = (
prisma_client.db.litellm_verificationtoken
)
return table
def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]":
table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable
return table
class CheckBatchCost:
def __init__(
self,
@ -73,7 +110,7 @@ class CheckBatchCost:
inline for a batch the first poll cycle then accounts again.
"""
try:
await self.prisma_client.db.litellm_managedobjecttable.find_first(
await _managed_object_table(self.prisma_client).find_first(
where={"file_purpose": "batch", "batch_processed": False}
)
except Exception as probe_err:
@ -97,10 +134,8 @@ class CheckBatchCost:
if not user_id:
return {}
try:
user_row: prisma_models.LiteLLM_UserTable | None = (
await self.prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_id}
)
user_row: prisma_models.LiteLLM_UserTable | None = await _user_table(self.prisma_client).find_unique(
where={"user_id": user_id}
)
if user_row is None:
return {}
@ -117,11 +152,9 @@ class CheckBatchCost:
if not api_key:
return None
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = (
await self.prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": api_key}
)
)
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
self.prisma_client
).find_unique(where={"token": api_key})
return getattr(key_row, "key_alias", None) if key_row is not None else None
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up key alias for batch {batch_id}: {e}")
@ -132,17 +165,15 @@ class CheckBatchCost:
if not team_id:
return None
try:
team_row: prisma_models.LiteLLM_TeamTable | None = (
await self.prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique(
where={"team_id": team_id}
)
return getattr(team_row, "team_alias", None) if team_row is not None else None
except Exception as e:
verbose_proxy_logger.error(f"CheckBatchCost: could not look up team alias for team {team_id}: {e}")
return None
async def _get_org_id(self, job: "LiteLLM_ManagedObjectTable", batch_id: str) -> str | None:
async def _get_org_id(self, job: "_ManagedObjectRow", batch_id: str) -> str | None:
org_id = getattr(job, "org_id", None)
if org_id:
return org_id
@ -150,11 +181,9 @@ class CheckBatchCost:
team_id = getattr(job, "team_id", None)
if api_key:
try:
key_row: prisma_models.LiteLLM_VerificationToken | None = (
await self.prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": api_key}
)
)
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
self.prisma_client
).find_unique(where={"token": api_key})
key_org_id = getattr(key_row, "organization_id", None) if key_row is not None else None
if key_org_id:
return key_org_id
@ -166,10 +195,8 @@ class CheckBatchCost:
if not team_id:
return None
try:
team_row: prisma_models.LiteLLM_TeamTable | None = (
await self.prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique(
where={"team_id": team_id}
)
return getattr(team_row, "organization_id", None) if team_row is not None else None
except Exception as e:
@ -177,7 +204,7 @@ class CheckBatchCost:
return None
async def _build_creator_attribution_metadata(
self, job: "LiteLLM_ManagedObjectTable", batch_id: str
self, job: "_ManagedObjectRow", batch_id: str
) -> dict[str, object]:
"""
Rebuild the spend-tracking metadata for the key, team, and tags that created the
@ -225,7 +252,7 @@ class CheckBatchCost:
should not be polled.
"""
cutoff: Final = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS)
result: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
result: Final = await _managed_object_table(self.prisma_client).update_many(
where={
"file_purpose": "batch",
"status": {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)},
@ -244,7 +271,7 @@ class CheckBatchCost:
# A row already in a terminal status is never rewritten by the sweep above, so
# without this it keeps a poll-page slot forever and starves newer batches.
retired: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
retired: Final = await _managed_object_table(self.prisma_client).update_many(
where={
"file_purpose": "batch",
"batch_processed": False,
@ -259,9 +286,9 @@ class CheckBatchCost:
f"{MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days that were never costed"
)
async def _fallback_find_jobs(self) -> list:
async def _fallback_find_jobs(self) -> "Sequence[_ManagedObjectRow]":
"""Query batch jobs without the batch_processed filter (for older schemas)."""
return await self.prisma_client.db.litellm_managedobjecttable.find_many(
return await _managed_object_table(self.prisma_client).find_many(
where={
"file_purpose": "batch",
"status": {
@ -279,7 +306,7 @@ class CheckBatchCost:
order={"created_at": "asc"},
)
async def _retire_job(self, job: "LiteLLM_ManagedObjectTable", reason: str) -> None:
async def _retire_job(self, job: "_ManagedObjectRow", reason: str) -> None:
"""
Take a row that can never be costed out of the poll page. Leaving it selectable
would burn one of the MAX_OBJECTS_PER_POLL_CYCLE slots on every future cycle, and
@ -292,7 +319,7 @@ class CheckBatchCost:
else {"status": "stale_expired"}
)
try:
await self.prisma_client.db.litellm_managedobjecttable.update(
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=data,
)
@ -306,7 +333,7 @@ class CheckBatchCost:
"so it will no longer be polled"
)
async def _claim_job_for_costing(self, job: "LiteLLM_ManagedObjectTable") -> bool:
async def _claim_job_for_costing(self, job: "_ManagedObjectRow") -> bool:
"""
Atomically flip batch_processed from false to true, returning whether this pod won
the row. Every pod and uvicorn worker schedules its own poller against the shared
@ -321,7 +348,7 @@ class CheckBatchCost:
if not self._has_batch_processed_column:
return True
try:
claimed: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many(
claimed: Final = await _managed_object_table(self.prisma_client).update_many(
where={"id": job.id, "batch_processed": False},
data={"batch_processed": True},
)
@ -332,7 +359,7 @@ class CheckBatchCost:
return False
return claimed > 0
async def _release_job_claim(self, job: "LiteLLM_ManagedObjectTable") -> None:
async def _release_job_claim(self, job: "_ManagedObjectRow") -> None:
"""Give a claimed row back once billing it failed, so a later poll cycle retries it.
Safe to match on batch_processed=True: while this poller is active the retrieve
@ -342,7 +369,7 @@ class CheckBatchCost:
if not self._has_batch_processed_column:
return
try:
await self.prisma_client.db.litellm_managedobjecttable.update_many(
await _managed_object_table(self.prisma_client).update_many(
where={"id": job.id, "batch_processed": True},
data={"batch_processed": False},
)
@ -353,7 +380,7 @@ class CheckBatchCost:
)
@staticmethod
def _has_unified_id_without_model(job: "LiteLLM_ManagedObjectTable") -> bool:
def _has_unified_id_without_model(job: "_ManagedObjectRow") -> bool:
"""A unified id that decodes but carries no model_id can never be routed."""
from litellm.proxy.openai_files_endpoints.common_utils import (
convert_b64_uid_to_unified_uid,
@ -402,7 +429,7 @@ class CheckBatchCost:
return isinstance(error, (NotFoundError, openai.NotFoundError)) and output_file_id in str(error)
async def _finalize_unbilled_terminal_job(
self, job: "prisma_models.LiteLLM_ManagedObjectTable", response: "LiteLLMBatch"
self, job: "_ManagedObjectRow", response: "LiteLLMBatch"
) -> None:
"""Persist a terminal batch that has nothing billable, converting any raw
provider file ids to managed ids, and take it out of the poll page."""
@ -426,7 +453,7 @@ class CheckBatchCost:
"file_object": response.model_dump_json(),
**({"batch_processed": True} if self._has_batch_processed_column else {}),
}
await self.prisma_client.db.litellm_managedobjecttable.update(
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=update_data,
)
@ -447,7 +474,7 @@ class CheckBatchCost:
def _resolve_job_routing(
self,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
prom_logger: Optional["PrometheusLogger"],
) -> Optional[Tuple[str, str]]:
"""
@ -524,7 +551,7 @@ class CheckBatchCost:
def _resolve_unmanaged_provider_routing(
self,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
prom_logger: Optional["PrometheusLogger"],
llm_provider: str,
bare_model_name: str,
@ -620,7 +647,7 @@ class CheckBatchCost:
@classmethod
def _get_managed_file_model_name(
cls,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
deployment_info: "Deployment",
) -> Optional[str]:
"""
@ -640,7 +667,7 @@ class CheckBatchCost:
)
@staticmethod
def _get_input_file_id(job: "LiteLLM_ManagedObjectTable") -> Optional[str]:
def _get_input_file_id(job: "_ManagedObjectRow") -> Optional[str]:
import json
from litellm.types.utils import LiteLLMBatch
@ -660,7 +687,7 @@ class CheckBatchCost:
async def _track_completed_batch_cost(
self,
job: "LiteLLM_ManagedObjectTable",
job: "_ManagedObjectRow",
response: "LiteLLMBatch",
model_id: str,
batch_id: str,
@ -936,7 +963,7 @@ class CheckBatchCost:
# endpoint may transition a batch to "complete" before
# CheckBatchCost runs. The batch_processed=False filter
# already prevents reprocessing finished batches.
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
jobs = await _managed_object_table(self.prisma_client).find_many(
where={
"file_purpose": "batch",
"batch_processed": False,
@ -1038,7 +1065,7 @@ class CheckBatchCost:
}
if self._has_batch_processed_column:
update_data["batch_processed"] = True
await self.prisma_client.db.litellm_managedobjecttable.update(
await _managed_object_table(self.prisma_client).update(
where={"id": job.id},
data=update_data,
)

View file

@ -6,7 +6,7 @@ same route are non-inference and free.
"""
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Dict, Optional, cast
from typing import TYPE_CHECKING, Dict, Final, Optional, Protocol, cast
import litellm
from litellm._logging import verbose_proxy_logger
@ -22,11 +22,31 @@ from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.prisma_protocols import TableActions
from litellm.router import Router
TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete"})
class _ManagedObjectRow(Protocol):
@property
def id(self) -> str: ...
@property
def unified_object_id(self) -> str: ...
@property
def created_by(self) -> str | None: ...
@property
def file_object(self) -> object: ...
def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]":
table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable
return table
class CheckResponsesCost:
def __init__(
self,
@ -128,7 +148,7 @@ class CheckResponsesCost:
f"CheckResponsesCost: stale cleanup failed (poll will continue): {cleanup_err}"
)
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
jobs = await _managed_object_table(self.prisma_client).find_many(
where={
"status": {"in": ["queued", "in_progress"]},
"file_purpose": "response",
@ -138,7 +158,7 @@ class CheckResponsesCost:
)
verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check")
completed_jobs = []
completed_jobs: Final[list[_ManagedObjectRow]] = []
for job in jobs:
unified_object_id = job.unified_object_id
@ -189,7 +209,7 @@ class CheckResponsesCost:
# Mark completed jobs in the database
if len(completed_jobs) > 0:
await self.prisma_client.db.litellm_managedobjecttable.update_many(
await _managed_object_table(self.prisma_client).update_many(
where={"id": {"in": [job.id for job in completed_jobs]}},
data={"status": "completed"},
)

View file

@ -481,10 +481,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"""
if self.prisma_client is None:
return
managed_object = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]}
)
managed_object = await _managed_object_table(self.prisma_client).find_first(
where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]}
)
if managed_object is None:
return
@ -509,10 +507,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"""
if self.prisma_client is None:
return
managed_file = (
await self.prisma_client.db.litellm_managedfiletable.find_first(
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
)
managed_file = await _managed_file_table(self.prisma_client).find_first(
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
)
if managed_file is None:
return
@ -535,8 +531,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
provider_file_ids = tuple(
file_id
for file_id in (
getattr(response, "output_file_id", None),
getattr(response, "error_file_id", None),
response.output_file_id,
response.error_file_id,
)
if file_id and not _is_base64_encoded_unified_file_id(file_id)
)
@ -544,10 +540,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return
if self.prisma_client is None:
return
batch_row = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"unified_object_id": response.id}
)
batch_row = await _managed_object_table(self.prisma_client).find_first(
where={"unified_object_id": response.id}
)
if batch_row is None or (
batch_row.created_by is None and batch_row.team_id is None

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.69"
version = "0.1.70"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.69"
version = "0.1.70"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -1,5 +1,5 @@
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x)
ARG PGBOUNCER_VERSION=1.25.2

View file

@ -96,16 +96,19 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/assemblyai/",
"/eu.assemblyai/",
"/deepgram/",
"/fal_ai/",
"/langfuse/",
"/vllm/",
"/mistral/",
"/typesafe/",
"/openrouter/",
"/nvidia_nim/",
"/groq/",
"/voyage/",
"/cursor/",
"/milvus/",
"/openai_passthrough/",
"/tinyfish/",
# Dynamic provider / toolset passthrough (path templates)
"/{provider}/",
"/toolset/",
@ -128,6 +131,7 @@ GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
"/redoc",
"/test",
"/debug/memory/summary",
"/api/event_logging/batch",
}
)

View file

@ -7,6 +7,9 @@ metadata:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: backend
spec:
{{- if and (not .Values.backend.hpa.enabled) (not (kindIs "invalid" .Values.backend.replicaCount)) }}
replicas: {{ .Values.backend.replicaCount }}
{{- end }}
{{- with .Values.backend.strategy }}
strategy:
{{- toYaml . | nindent 4 }}

View file

@ -7,6 +7,9 @@ metadata:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: gateway
spec:
{{- if and (not .Values.gateway.hpa.enabled) (not (kindIs "invalid" .Values.gateway.replicaCount)) }}
replicas: {{ .Values.gateway.replicaCount }}
{{- end }}
{{- with .Values.gateway.strategy }}
strategy:
{{- toYaml . | nindent 4 }}

View file

@ -61,7 +61,7 @@
"/v1/fine-tuning" "/fine-tuning" "/v1/responses" "/responses" "/v1/threads" "/threads"
"/v1/assistants" "/assistants" "/v1/vector_stores" "/vector_stores" "/v1/indexes"
"/v1/models" "/models" "/openai" "/engines"
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a" "/a2a"
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a" "/a2a" "/api/event_logging"
"/v1/rerank" "/v2/rerank" "/rerank" "/v1/ocr" "/ocr" "/v1/rag" "/rag"
"/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search"
"/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat"

View file

@ -7,6 +7,9 @@ metadata:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: ui
spec:
{{- if and (not .Values.ui.hpa.enabled) (not (kindIs "invalid" .Values.ui.replicaCount)) }}
replicas: {{ .Values.ui.replicaCount }}
{{- end }}
{{- with .Values.ui.strategy }}
strategy:
{{- toYaml . | nindent 4 }}

View file

@ -0,0 +1,100 @@
suite: test fixed replica count when HPA is disabled
templates:
- gateway/deployment.yaml
- gateway/configmap.yaml
- backend/deployment.yaml
- ui/deployment.yaml
values:
- ./values/required.yaml
tests:
- it: gateway renders replicaCount into spec.replicas when its HPA is disabled
template: gateway/deployment.yaml
set:
gateway.hpa.enabled: false
gateway.replicaCount: 3
asserts:
- isKind:
of: Deployment
- equal:
path: spec.replicas
value: 3
- it: backend renders replicaCount into spec.replicas when its HPA is disabled
template: backend/deployment.yaml
set:
backend.hpa.enabled: false
backend.replicaCount: 2
asserts:
- equal:
path: spec.replicas
value: 2
- it: ui renders replicaCount into spec.replicas when its HPA is disabled
template: ui/deployment.yaml
set:
ui.hpa.enabled: false
ui.replicaCount: 2
asserts:
- equal:
path: spec.replicas
value: 2
- it: replicaCount 0 scales the gateway to zero instead of being treated as unset
template: gateway/deployment.yaml
set:
gateway.hpa.enabled: false
gateway.replicaCount: 0
asserts:
- equal:
path: spec.replicas
value: 0
- it: a component with HPA disabled but no replicaCount set keeps omitting spec.replicas, so upgrades do not reset a hand-scaled Deployment
set:
gateway.hpa.enabled: false
backend.hpa.enabled: false
ui.hpa.enabled: false
asserts:
- notExists:
path: spec.replicas
template: gateway/deployment.yaml
- notExists:
path: spec.replicas
template: backend/deployment.yaml
- notExists:
path: spec.replicas
template: ui/deployment.yaml
- it: every component omits spec.replicas when its HPA is enabled, so the autoscaler owns the count
set:
gateway.hpa.enabled: true
gateway.replicaCount: 3
backend.hpa.enabled: true
backend.replicaCount: 3
ui.hpa.enabled: true
ui.replicaCount: 3
asserts:
- notExists:
path: spec.replicas
template: gateway/deployment.yaml
- notExists:
path: spec.replicas
template: backend/deployment.yaml
- notExists:
path: spec.replicas
template: ui/deployment.yaml
- it: a component with HPA disabled renders replicas while a sibling with HPA enabled does not
set:
gateway.hpa.enabled: false
gateway.replicaCount: 4
backend.hpa.enabled: true
backend.replicaCount: 4
asserts:
- equal:
path: spec.replicas
value: 4
template: gateway/deployment.yaml
- notExists:
path: spec.replicas
template: backend/deployment.yaml

View file

@ -397,6 +397,11 @@ gateway:
# failureThreshold: 30
# periodSeconds: 10
startupProbe: {}
# Optional fixed pod count, rendered into the Deployment's spec.replicas only
# when hpa.enabled is false. Unset by default so an existing Deployment keeps
# its current count; with the HPA on, the autoscaler owns the count, e.g.:
# replicaCount: 3
replicaCount:
hpa:
enabled: true
minReplicas: 1
@ -524,6 +529,8 @@ backend:
strategy: {}
# Optional startupProbe; same shape as gateway.startupProbe. Empty by default.
startupProbe: {}
# Same semantics as gateway.replicaCount.
replicaCount:
hpa:
enabled: true
minReplicas: 1
@ -590,6 +597,8 @@ ui:
strategy: {}
# Optional startupProbe; same shape as gateway.startupProbe. Empty by default.
startupProbe: {}
# Same semantics as gateway.replicaCount.
replicaCount:
hpa:
enabled: false
minReplicas: 1

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "access_group_ids" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -0,0 +1,42 @@
CREATE TABLE IF NOT EXISTS "LiteLLM_AutoRouterUserSession" (
"user_id" TEXT NOT NULL,
"api_key" TEXT NOT NULL,
"session_id" TEXT NOT NULL,
"router_name" TEXT NOT NULL,
"router_type" TEXT NOT NULL,
"first_turn_at" TIMESTAMP(3) NOT NULL,
"last_turn_at" TIMESTAMP(3) NOT NULL,
"last_model" TEXT NOT NULL,
"models" JSONB NOT NULL DEFAULT '{}',
"turns" INTEGER NOT NULL DEFAULT 0,
"unordered_turns" INTEGER NOT NULL DEFAULT 0,
"covered_turns" INTEGER NOT NULL DEFAULT 0,
"cache_hits" INTEGER NOT NULL DEFAULT 0,
"same_model_turns" INTEGER NOT NULL DEFAULT 0,
"same_model_hits" INTEGER NOT NULL DEFAULT 0,
"first_visit_turns" INTEGER NOT NULL DEFAULT 0,
"first_visit_hits" INTEGER NOT NULL DEFAULT 0,
"return_turns" INTEGER NOT NULL DEFAULT 0,
"return_hits" INTEGER NOT NULL DEFAULT 0,
"return_expired_misses" INTEGER NOT NULL DEFAULT 0,
"return_within_ttl_misses" INTEGER NOT NULL DEFAULT 0,
"ttl_5m_turns" INTEGER NOT NULL DEFAULT 0,
"ttl_1h_turns" INTEGER NOT NULL DEFAULT 0,
"total_tokens" BIGINT NOT NULL DEFAULT 0,
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"savings_estimated_turns" INTEGER NOT NULL DEFAULT 0,
"savings_estimated_actual_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"savings_estimated_saved_spend" DOUBLE PRECISION NOT NULL DEFAULT 0,
"savings_estimated_baseline_models" JSONB NOT NULL DEFAULT '{}',
"classifier_cost" DOUBLE PRECISION NOT NULL DEFAULT 0,
"classifier_cost_recorded_turns" INTEGER NOT NULL DEFAULT 0,
"tier_turns" JSONB NOT NULL DEFAULT '{}',
"baseline_models" JSONB NOT NULL DEFAULT '{}',
CONSTRAINT "LiteLLM_AutoRouterUserSession_pkey" PRIMARY KEY ("user_id", "api_key", "session_id", "router_name")
);
CREATE INDEX IF NOT EXISTS "idx_autorouter_user_session_last_turn" ON "LiteLLM_AutoRouterUserSession"("last_turn_at");
CREATE INDEX IF NOT EXISTS "idx_autorouter_user_session_user_last_turn" ON "LiteLLM_AutoRouterUserSession"("user_id", "last_turn_at");

View file

@ -0,0 +1 @@
ALTER TABLE "LiteLLM_PolicyAttachmentTable" ADD COLUMN IF NOT EXISTS "is_default" BOOLEAN NOT NULL DEFAULT false;

View file

@ -0,0 +1,3 @@
ALTER TABLE "LiteLLM_UserTable" ADD COLUMN IF NOT EXISTS "password_reset_required" BOOLEAN;
ALTER TABLE "LiteLLM_UserTable" ADD COLUMN IF NOT EXISTS "last_breach_check_at" TIMESTAMP(3);

View file

@ -73,6 +73,7 @@ model LiteLLM_AgentsTable {
static_headers Json? @default("{}")
extra_headers String[] @default([])
agent_access_groups String[] @default([])
access_group_ids String[] @default([])
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
spend Float @default(0.0)
@ -246,6 +247,8 @@ model LiteLLM_UserTable {
organization_id String?
object_permission_id String?
password String?
password_reset_required Boolean?
last_breach_check_at DateTime?
teams String[] @default([])
user_role String?
max_budget Float?
@ -1419,6 +1422,7 @@ model LiteLLM_PolicyAttachmentTable {
models String[] @default([]) // Model names or patterns
tags String[] @default([]) // Tag patterns (e.g., ["healthcare", "prod-*"])
priority Int? // Explicit execution order
is_default Boolean @default(false) // Applied only when no non-default attachment matches
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
@ -1620,6 +1624,47 @@ model LiteLLM_AutoRouterSession {
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
}
model LiteLLM_AutoRouterUserSession {
user_id String
api_key String
session_id String
router_name String
router_type String
first_turn_at DateTime
last_turn_at DateTime
last_model String
models Json @default("{}")
turns Int @default(0)
unordered_turns Int @default(0)
covered_turns Int @default(0)
cache_hits Int @default(0)
same_model_turns Int @default(0)
same_model_hits Int @default(0)
first_visit_turns Int @default(0)
first_visit_hits Int @default(0)
return_turns Int @default(0)
return_hits Int @default(0)
return_expired_misses Int @default(0)
return_within_ttl_misses Int @default(0)
ttl_5m_turns Int @default(0)
ttl_1h_turns Int @default(0)
total_tokens BigInt @default(0)
spend Float @default(0)
saved_spend Float @default(0)
savings_estimated_turns Int @default(0)
savings_estimated_actual_spend Float @default(0)
savings_estimated_saved_spend Float @default(0)
savings_estimated_baseline_models Json @default("{}")
classifier_cost Float @default(0)
classifier_cost_recorded_turns Int @default(0)
tier_turns Json @default("{}")
baseline_models Json @default("{}")
@@id([user_id, api_key, session_id, router_name])
@@index([last_turn_at], map: "idx_autorouter_user_session_last_turn")
@@index([user_id, last_turn_at], map: "idx_autorouter_user_session_user_last_turn")
}
// Shadow eval: evaluation of an auto-router against one or more keys' live traffic, in
// either direction. forward duplicates the requests the keys did not route through the
// router through it, answering whether they should adopt it; reverse duplicates the

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.100"
version = "0.4.101"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.100"
version = "0.4.101"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

921
litellm-rust/Cargo.lock generated

File diff suppressed because it is too large Load diff

View file

@ -9,6 +9,8 @@ license = "MIT"
repository = "https://github.com/BerriAI/litellm"
[workspace.dependencies]
litellm-tracing = { path = "crates/tracing" }
tracing = "0.1"
litellm-core = { path = "crates/core" }
litellm-host = { path = "crates/host" }
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
@ -22,12 +24,24 @@ litellm-secrets = { path = "crates/secrets" }
litellm-secrets-types = { path = "crates/secrets-types" }
litellm-secrets-aws = { path = "crates/secrets-aws" }
litellm-secrets-google = { path = "crates/secrets-google" }
litellm-secrets-hashicorp = { path = "crates/secrets-hashicorp" }
litellm-secrets-azure = { path = "crates/secrets-azure" }
litellm-secrets-cyberark = { path = "crates/secrets-cyberark" }
litellm-http = { path = "crates/http" }
litellm-llms = { path = "crates/llms" }
litellm-types = { path = "crates/types" }
litellm-core-utils = { path = "crates/core-utils" }
litellm-cache = { path = "crates/cache" }
litellm-cache-azure-blob = { path = "crates/cache-azure-blob" }
litellm-cache-memory = { path = "crates/cache-memory" }
litellm-cache-redis = { path = "crates/cache-redis" }
litellm-cache-s3 = { path = "crates/cache-s3" }
litellm-cache-gcs = { path = "crates/cache-gcs" }
litellm-cache-disk = { path = "crates/cache-disk" }
litellm-cache-redis-semantic = { path = "crates/cache-redis-semantic" }
litellm-cache-response = { path = "crates/cache-response" }
litellm-cache-qdrant-semantic = { path = "crates/cache-qdrant-semantic" }
litellm-cache-testing = { path = "crates/cache-testing" }
litellm-token-counter = { path = "crates/token-counter" }
litellm-token-counter-fast = { path = "crates/token-counter-fast" }
litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" }
@ -45,9 +59,14 @@ pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
pythonize = "0.29.0"
rand = "0.8"
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] }
qdrant-client = { version = "1.19.0", default-features = false }
uuid = { version = "1", features = ["v4"] }
rstest = "0.26.1"
rstest_reuse = "0.7.0"
rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] }
rustify = "=0.7.0"
rustify_derive = "=0.5.5"
vaultrs = { version = "=0.8.0", default-features = false, features = ["rustls"] }
rustls-native-certs = "0.8"
serde = { version = "1.0", features = ["derive"] }
serde_json = { version = "1.0", features = ["float_roundtrip"] }
@ -64,6 +83,7 @@ base64 = "0.22"
moka = { version = "0.12.16", features = ["future"] }
strum = { version = "0.28.0", features = ["derive"] }
url = "2.5.8"
percent-encoding = "2.3"
webpki-roots = "1"
time = { version = "0.3.53", features = ["parsing"] }
criterion = "0.8.2"
@ -72,7 +92,7 @@ veil = "0.3.0"
[profile.release]
opt-level = 3
lto = "thin"
lto = "fat"
codegen-units = 1
panic = "unwind"
debug = false

View file

@ -621,6 +621,26 @@ mod tests {
None
}
#[test]
fn secret_names_cover_environment_reads() {
let seen = std::sync::Arc::new(std::sync::Mutex::new(
std::collections::BTreeSet::<String>::new(),
));
let recorded = seen.clone();
let env = |name: &str| {
recorded.lock().unwrap().insert(name.to_string());
None
};
resolve_aws_region(None, &Map::new(), &env);
aws_auth_config(&Map::new(), &env);
assert!(
seen.lock()
.unwrap()
.iter()
.all(|name| crate::constants::SECRET_NAMES.contains(&name.as_str()))
);
}
#[test]
fn a_region_comes_from_the_call_then_the_model_then_the_environment() {
let params = Map::from_iter([("aws_region_name".to_string(), Value::from("eu-west-1"))]);

View file

@ -14,6 +14,19 @@ pub const AWS_WEB_IDENTITY_TOKEN_FILE: &str = "AWS_WEB_IDENTITY_TOKEN_FILE";
pub const AWS_STS_ENDPOINT: &str = "AWS_STS_ENDPOINT";
pub const AWS_EXTERNAL_ID: &str = "AWS_EXTERNAL_ID";
pub const AWS_BEARER_TOKEN_BEDROCK: &str = "AWS_BEARER_TOKEN_BEDROCK";
pub const SECRET_NAMES: &[&str] = &[
AWS_ACCESS_KEY_ID,
AWS_SECRET_ACCESS_KEY,
AWS_SESSION_TOKEN,
AWS_REGION_NAME,
AWS_REGION,
AWS_SESSION_NAME,
AWS_PROFILE_NAME,
AWS_ROLE_NAME,
AWS_WEB_IDENTITY_TOKEN,
AWS_STS_ENDPOINT,
AWS_EXTERNAL_ID,
];
/// Headers SigV4 covers, beyond the `x-amz-` / `x-amzn-` prefixes. Mirrors
/// Python's `_filter_headers_for_aws_signature` allowlist.

View file

@ -3,5 +3,5 @@ mod native;
mod resolve;
mod types;
pub use resolve::AzureAuthService;
pub use types::AzureAuthInputs;
pub use resolve::{AzureAuthService, SECRET_NAMES};
pub use types::{AzureAuthInputs, ConfigValue};

View file

@ -19,6 +19,17 @@ const AZURE_AUTHORITY_HOST_ENV: &str = "AZURE_AUTHORITY_HOST";
const AZURE_CREDENTIAL_ENV: &str = "AZURE_CREDENTIAL";
const AZURE_FEDERATED_TOKEN_FILE_ENV: &str = "AZURE_FEDERATED_TOKEN_FILE";
pub const SECRET_NAMES: &[&str] = &[
AZURE_AD_TOKEN_ENV,
AZURE_TENANT_ID_ENV,
AZURE_CLIENT_ID_ENV,
AZURE_CLIENT_SECRET_ENV,
AZURE_SCOPE_ENV,
AZURE_AUTHORITY_HOST_ENV,
AZURE_CREDENTIAL_ENV,
AZURE_FEDERATED_TOKEN_FILE_ENV,
];
#[derive(Clone, Debug)]
pub(crate) enum AzureCredentialPlan {
Supplied(Sourced<ResolvedCredential>),
@ -440,13 +451,14 @@ fn non_empty_reference(value: &str, kind: &str) -> Result<String, Error> {
#[cfg(test)]
mod tests {
use std::collections::BTreeSet;
use std::future::Future;
use std::sync::{Arc, Mutex};
use serde_json::json;
use super::{
AzureAuthService, AzureCredentialPlan, AzureTokenAcquirer, oidc_reference,
AzureAuthService, AzureCredentialPlan, AzureTokenAcquirer, SECRET_NAMES, oidc_reference,
resolve_reference, select_auth_plan,
};
use crate::native::ValidatedAzureRequest;
@ -517,6 +529,24 @@ mod tests {
assert!(matches!(plan, AzureCredentialPlan::Native(_)));
}
#[test]
fn secret_names_cover_environment_reads() {
let seen = std::sync::Arc::new(std::sync::Mutex::new(BTreeSet::<String>::new()));
let recorded = seen.clone();
let inputs = AzureAuthInputs::default();
select_auth_plan(&inputs, &|name| {
recorded.lock().unwrap().insert(name.to_string());
None
})
.unwrap();
assert!(
seen.lock()
.unwrap()
.iter()
.all(|name| SECRET_NAMES.contains(&name.as_str()))
);
}
#[test]
fn supplied_token_does_not_require_refresh() {
let params = json!({"azure_ad_token": "token"});

View file

@ -51,6 +51,21 @@ pub struct AzureAuthInputs {
}
impl AzureAuthInputs {
pub fn default_credential_for_scope(scope: &str) -> Self {
Self {
azure_scope: ConfigValue::Value(Sourced::new(
scope.to_string(),
InputSource::Deployment,
)),
azure_credential: ConfigValue::Value(Sourced::new(
"DefaultAzureCredential".to_string(),
InputSource::Deployment,
)),
enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment),
..Self::default()
}
}
pub fn or_configured_token_refresh(self, enabled: bool) -> Self {
if *self.enable_azure_ad_token_refresh.value() || !enabled {
return self;

View file

@ -23,6 +23,16 @@ const VERTEXAI_PROJECT_ENV: &str = "VERTEXAI_PROJECT";
const VERTEXAI_LOCATION_ENV: &str = "VERTEXAI_LOCATION";
const VERTEX_LOCATION_ENV: &str = "VERTEX_LOCATION";
pub const SECRET_NAMES: &[&str] = &[
VERTEX_AI_API_KEY_ENV,
VERTEXAI_API_KEY_ENV,
VERTEXAI_CREDENTIALS_ENV,
GOOGLE_APPLICATION_CREDENTIALS_ENV,
VERTEXAI_PROJECT_ENV,
VERTEXAI_LOCATION_ENV,
VERTEX_LOCATION_ENV,
];
#[derive(Clone, Debug, Default)]
pub struct VertexConfig {
credentials: Option<Sourced<SecretValue>>,
@ -128,6 +138,14 @@ impl VertexAuth {
}
}
pub async fn access_token(
&self,
config: &VertexConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<String, Error> {
self.load_provider(config, env_lookup).await?.token().await
}
pub async fn validate_environment(
&self,
headers: Vec<(String, String)>,
@ -398,6 +416,7 @@ fn auth_acquisition_error(error: gcp_auth::Error) -> Error {
#[cfg(test)]
mod tests {
use std::collections::BTreeSet;
use std::sync::atomic::{AtomicUsize, Ordering};
use serde_json::json;
@ -468,6 +487,27 @@ mod tests {
);
}
#[tokio::test]
async fn secret_names_cover_environment_reads() {
let seen = Arc::new(std::sync::Mutex::new(BTreeSet::<String>::new()));
let recorded = seen.clone();
let env = |name: &str| {
recorded.lock().unwrap().insert(name.to_string());
None
};
let auth = auth(Arc::new(AtomicUsize::new(0)), Arc::new(AtomicUsize::new(0)));
auth.validate_environment(Vec::new(), None, &VertexConfig::default(), &env)
.await
.unwrap();
get_vertex_ai_location(&VertexConfig::default(), &env);
assert!(
seen.lock()
.unwrap()
.iter()
.all(|name| SECRET_NAMES.contains(&name.as_str()))
);
}
#[test]
fn empty_primary_values_fall_back_to_python_aliases() {
let config = config(json!({

View file

@ -1,4 +1,5 @@
use serde::Deserialize;
use std::hash::{Hash, Hasher};
use veil::Redact;
#[derive(Redact, Clone, Deserialize)]
@ -23,6 +24,12 @@ impl PartialEq for SecretValue {
impl Eq for SecretValue {}
impl Hash for SecretValue {
fn hash<H: Hasher>(&self, state: &mut H) {
self.0.hash(state);
}
}
#[cfg(test)]
mod tests {
use super::SecretValue;

View file

@ -0,0 +1,27 @@
[package]
name = "litellm-cache-azure-blob"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-auth-azure.workspace = true
litellm-auth-types.workspace = true
litellm-cache.workspace = true
async-trait = "0.1"
azure_core = "1.1.0"
azure_storage_blob = "1.1.0"
futures-util.workspace = true
reqwest.workspace = true
tokio.workspace = true
url.workspace = true
[dev-dependencies]
litellm-cache-response.workspace = true
litellm-cache-testing.workspace = true
rstest.workspace = true
serde_json.workspace = true
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
wiremock = "0.6.5"

View file

@ -0,0 +1,248 @@
use std::{sync::Arc, time::Duration};
use azure_core::{
credentials::TokenCredential,
error::ErrorKind,
http::{ClientOptions, RequestContent, Transport},
};
use azure_storage_blob::{
BlobContainerClient, BlobContainerClientOptions,
models::{BlobClientUploadOptions, StorageErrorCode},
};
use futures_util::{TryStreamExt, future::try_join_all};
use litellm_cache::{
BaseCache, BatchCache, CacheCodec, DisconnectCache, Error, ExactCacheContext, FlushCache,
};
use tokio::runtime::Handle;
use url::Url;
use crate::{credential::AzureBlobCredential, transport::ReqwestTransport};
pub struct AzureBlobCache<C> {
container: BlobContainerClient,
codec: C,
runtime: Handle,
account_url: String,
container_name: String,
}
impl<C: CacheCodec> AzureBlobCache<C> {
/// `http` is the host's pooled client; the SDK sends every request through it.
pub async fn connect(
account_url: &str,
container: &str,
http: reqwest::Client,
codec: C,
runtime: Handle,
) -> Result<Self, Error> {
Self::connect_with_options(
account_url,
container,
Some(Arc::new(AzureBlobCredential::default())),
ClientOptions {
transport: Some(Transport::new(Arc::new(ReqwestTransport(http)))),
..ClientOptions::default()
},
codec,
runtime,
)
.await
}
pub async fn connect_with_options(
account_url: &str,
container: &str,
credential: Option<Arc<dyn TokenCredential>>,
client_options: ClientOptions,
codec: C,
runtime: Handle,
) -> Result<Self, Error> {
let parsed = Url::parse(account_url).map_err(|_| Error::Unavailable)?;
let account_url = parsed.as_str().trim_end_matches('/').to_string();
let container_url = {
let mut url = parsed;
url.path_segments_mut()
.map_err(|()| Error::Unavailable)?
.pop_if_empty()
.push(container);
url
};
let client = BlobContainerClient::new(
container_url,
credential,
Some(BlobContainerClientOptions {
client_options,
..BlobContainerClientOptions::default()
}),
)
.map_err(|_| Error::Unavailable)?;
let cache = Self {
container: client,
codec,
runtime,
account_url,
container_name: container.to_string(),
};
cache.create_container().await?;
Ok(cache)
}
pub fn account_url(&self) -> &str {
&self.account_url
}
pub fn container_name(&self) -> &str {
&self.container_name
}
async fn create_container(&self) -> Result<(), Error> {
match self.container.create(None).await {
Ok(_) => Ok(()),
Err(error) if is_storage_error(&error, StorageErrorCode::ContainerAlreadyExists) => {
Ok(())
}
Err(_) => Err(Error::Unavailable),
}
}
async fn upload(&self, key: &str, value: &C::Value, overwrite: bool) -> Result<(), Error> {
let payload = self.codec.encode(value)?;
let options = (!overwrite).then(|| BlobClientUploadOptions::default().if_not_exists());
match self
.container
.blob_client(key)
.upload(RequestContent::from(payload), options)
.await
{
Ok(_) => Ok(()),
Err(error) if !overwrite && is_already_present(&error) => Ok(()),
Err(_) => Err(Error::Unavailable),
}
}
async fn download(&self, key: &str) -> Result<Option<C::Value>, Error> {
let response = match self.container.blob_client(key).download(None).await {
Ok(response) => response,
Err(error) if is_storage_error(&error, StorageErrorCode::BlobNotFound) => {
return Ok(None);
}
Err(_) => return Err(Error::Unavailable),
};
let bytes = response
.body
.collect()
.await
.map_err(|_| Error::Unavailable)?;
self.codec.decode(&bytes).map(Some)
}
async fn delete_all_blobs(&self) -> Result<(), Error> {
let mut pages = self
.container
.list_blobs(None)
.map_err(|_| Error::Unavailable)?
.into_pages();
while let Some(page) = pages.try_next().await.map_err(|_| Error::Unavailable)? {
let page = page.into_model().map_err(|_| Error::Unavailable)?;
for name in page.blob_items.into_iter().filter_map(|item| item.name) {
self.container
.blob_client(&name)
.delete(None)
.await
.map_err(|_| Error::Unavailable)?;
}
}
Ok(())
}
fn block_on<T>(&self, future: impl Future<Output = T>) -> T {
if Handle::try_current().is_ok() {
tokio::task::block_in_place(|| self.runtime.block_on(future))
} else {
self.runtime.block_on(future)
}
}
}
fn is_already_present(error: &azure_core::Error) -> bool {
is_storage_error(error, StorageErrorCode::BlobAlreadyExists)
|| is_storage_error(error, StorageErrorCode::ConditionNotMet)
}
fn is_storage_error(error: &azure_core::Error, code: StorageErrorCode) -> bool {
matches!(
error.kind(),
ErrorKind::HttpResponse {
error_code: Some(error_code),
..
} if error_code == code.as_ref()
)
}
impl<C: CacheCodec> BaseCache for AzureBlobCache<C> {
type Value = C::Value;
type Context = ExactCacheContext;
fn get_ttl(&self, _: &ExactCacheContext) -> Option<Duration> {
None
}
fn set_cache(&self, key: &str, value: C::Value, _: &ExactCacheContext) -> Result<(), Error> {
self.block_on(self.upload(key, &value, false))
}
fn get_cache(&self, key: &str, _: &ExactCacheContext) -> Result<Option<C::Value>, Error> {
self.block_on(self.download(key))
}
async fn async_set_cache(
&self,
key: &str,
value: C::Value,
_: ExactCacheContext,
) -> Result<(), Error> {
self.upload(key, &value, true).await
}
async fn async_get_cache(
&self,
key: &str,
_: &ExactCacheContext,
) -> Result<Option<C::Value>, Error> {
self.download(key).await
}
async fn async_set_cache_pipeline(
&self,
entries: Vec<(String, C::Value)>,
_: ExactCacheContext,
) -> Result<(), Error> {
try_join_all(
entries
.iter()
.map(|(key, value)| self.upload(key, value, true)),
)
.await
.map(drop)
}
}
impl<C: CacheCodec> BatchCache for AzureBlobCache<C> {}
impl<C: CacheCodec> FlushCache for AzureBlobCache<C> {
fn flush_cache(&self) -> Result<(), Error> {
self.block_on(self.delete_all_blobs())
}
async fn async_flush_cache(&self) -> Result<(), Error> {
self.delete_all_blobs().await
}
}
impl<C: CacheCodec> DisconnectCache for AzureBlobCache<C> {
/// Python closes its two SDK clients; the Rust clients hold no connection of their own
/// (the pooled transport belongs to the host), so there is nothing to release.
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
}

View file

@ -0,0 +1,84 @@
use std::{
fmt,
sync::Arc,
time::{Duration, SystemTime},
};
use azure_core::{
credentials::{AccessToken, TokenCredential, TokenRequestOptions},
error::ErrorKind,
time::OffsetDateTime,
};
use litellm_auth_azure::{AzureAuthInputs, AzureAuthService};
use litellm_auth_types::ResolvedCredential;
const STATIC_TOKEN_LIFETIME: Duration = Duration::from_secs(300);
const LLM_TOKEN_ENV: &str = "AZURE_AD_TOKEN";
type EnvLookup = Arc<dyn Fn(&str) -> Option<String> + Send + Sync>;
pub struct AzureBlobCredential {
service: AzureAuthService,
env_lookup: EnvLookup,
}
impl fmt::Debug for AzureBlobCredential {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("AzureBlobCredential")
}
}
impl Default for AzureBlobCredential {
fn default() -> Self {
Self::new(
AzureAuthService::default(),
Arc::new(|name| std::env::var(name).ok()),
)
}
}
impl AzureBlobCredential {
pub fn new(service: AzureAuthService, env_lookup: EnvLookup) -> Self {
Self {
service,
env_lookup,
}
}
}
#[async_trait::async_trait]
impl TokenCredential for AzureBlobCredential {
async fn get_token(
&self,
scopes: &[&str],
_options: Option<TokenRequestOptions<'_>>,
) -> azure_core::Result<AccessToken> {
let env_lookup = &self.env_lookup;
let lookup = move |name: &str| (name != LLM_TOKEN_ENV).then(|| env_lookup(name)).flatten();
let credential = self
.service
.get_azure_ad_token(
&AzureAuthInputs::default_credential_for_scope(&scopes.join(" ")),
&lookup,
)
.await
.map_err(|error| {
azure_core::Error::with_message(ErrorKind::Credential, error.to_string())
})?
.ok_or_else(|| {
azure_core::Error::with_message(
ErrorKind::Credential,
"no Azure credential is available for blob storage",
)
})?;
let (token, expires_on) = match credential.into_value() {
ResolvedCredential::AccessToken { token, expires_on } => (token, expires_on),
ResolvedCredential::Static(token) => (token, None),
};
let expires_on = expires_on.unwrap_or_else(|| SystemTime::now() + STATIC_TOKEN_LIFETIME);
Ok(AccessToken::new(
token.expose().to_string(),
OffsetDateTime::from(expires_on),
))
}
}

View file

@ -0,0 +1,7 @@
mod cache;
mod credential;
mod transport;
pub use cache::AzureBlobCache;
pub use credential::AzureBlobCredential;
pub use transport::ReqwestTransport;

View file

@ -0,0 +1,49 @@
use azure_core::{
error::ErrorKind,
http::{
AsyncRawResponse, Body, HttpClient, Request,
headers::{HeaderName, HeaderValue, Headers},
},
};
use futures_util::TryStreamExt;
#[derive(Debug)]
pub struct ReqwestTransport(pub reqwest::Client);
#[async_trait::async_trait]
impl HttpClient for ReqwestTransport {
async fn execute_request(&self, request: &Request) -> azure_core::Result<AsyncRawResponse> {
let method = reqwest::Method::from_bytes(request.method().as_ref().as_bytes())
.map_err(|error| azure_core::Error::new(ErrorKind::Other, error))?;
let mut outgoing = self.0.request(method, request.url().as_str());
for (name, value) in request.headers().iter() {
outgoing = outgoing.header(name.as_str(), value.as_str());
}
let outgoing = match request.body().clone() {
Body::Bytes(bytes) => outgoing.body(bytes),
Body::SeekableStream(stream) => outgoing.body(reqwest::Body::wrap_stream(stream)),
};
let response = outgoing.send().await.map_err(|error| {
let kind = if error.is_connect() {
ErrorKind::Connection
} else {
ErrorKind::Io
};
azure_core::Error::new(kind, error)
})?;
let status = response.status().as_u16().into();
let mut headers = Headers::new();
for (name, value) in response.headers() {
if let Ok(value) = value.to_str() {
headers.insert(
HeaderName::from(name.as_str().to_owned()),
HeaderValue::from(value.to_owned()),
);
}
}
let body = response
.bytes_stream()
.map_err(|error| azure_core::Error::new(ErrorKind::Io, error));
Ok(AsyncRawResponse::new(status, headers, Box::pin(body)))
}
}

View file

@ -0,0 +1,494 @@
mod support;
use std::{sync::Arc, time::Duration};
use azure_core::http::Method;
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, DisconnectCache, Error, ExactCacheContext, FlushCache,
};
use litellm_cache_azure_blob::AzureBlobCache;
use litellm_cache_response::{
CacheEntry, CacheKeyField, CacheKeyInput, ResponseCache, ResponseCacheCodec,
ResponseCacheRequest, cache_key,
};
use rstest::{fixture, rstest};
use serde_json::json;
use support::{ACCOUNT_URL, CONTAINER, FakeBlobService, RecordedRequest};
use tokio::runtime::Runtime;
type Fixture = support::Fixture<ResponseCacheCodec>;
#[fixture]
fn fixture() -> Fixture {
Fixture::new(FakeBlobService::default(), ResponseCacheCodec)
}
fn response_cache(fixture: &Fixture) -> ResponseCache<AzureBlobCache<ResponseCacheCodec>> {
ResponseCache::new(fixture.cache.clone())
}
fn request(model: &str) -> ResponseCacheRequest {
ResponseCacheRequest::new(CacheKeyInput {
fields: vec![CacheKeyField {
name: "model".into(),
value: Some(model.into()),
api_parameter: true,
internal_parameter: false,
}],
preset: None,
namespace: None,
include_provider_parameters: false,
})
}
fn now() -> Duration {
Duration::from_secs(1_700_000_000)
}
fn entry(value: serde_json::Value) -> CacheEntry {
CacheEntry {
timestamp: Some(1_700_000_000.5),
response: value,
}
}
fn no_ttl() -> ExactCacheContext {
ExactCacheContext::default()
}
fn with_ttl(seconds: u64) -> ExactCacheContext {
ExactCacheContext {
ttl: Some(Duration::from_secs(seconds)),
}
}
fn connect_to(account_url: &str) -> (FakeBlobService, AzureBlobCache<ResponseCacheCodec>) {
let runtime = Runtime::new().unwrap();
let service = FakeBlobService::default();
let cache = runtime
.block_on(support::connect(
&service,
account_url,
ResponseCacheCodec,
runtime.handle().clone(),
))
.unwrap();
(service, cache)
}
#[rstest]
fn connect_creates_the_container_once(fixture: Fixture) {
assert!(fixture.service.container_exists());
assert_eq!(
fixture.service.requests(),
vec![RecordedRequest {
method: Method::Put,
path: format!("/{CONTAINER}"),
query: "restype=container".into(),
if_none_match: None,
}]
);
assert_eq!(fixture.cache.account_url(), ACCOUNT_URL);
assert_eq!(fixture.cache.container_name(), CONTAINER);
}
#[rstest]
fn connect_accepts_an_existing_container() {
let fixture = Fixture::new(
FakeBlobService::with_existing_container(),
ResponseCacheCodec,
);
assert!(fixture.service.container_exists());
assert_eq!(fixture.service.requests().len(), 1);
}
#[rstest]
fn connect_accepts_account_urls_with_trailing_slash() {
let (service, cache) = connect_to("https://example.blob.core.windows.net/");
assert_eq!(service.requests()[0].path, format!("/{CONTAINER}"));
assert_eq!(cache.account_url(), "https://example.blob.core.windows.net");
}
#[rstest]
fn connect_keeps_account_url_query_parameters_on_the_container_path() {
let (service, _) = connect_to("https://example.blob.core.windows.net/?sv=2024-01-01&sig=abc");
let create = &service.requests()[0];
assert_eq!(create.path, format!("/{CONTAINER}"));
assert!(create.query.contains("sig=abc"));
}
#[rstest]
fn connect_surfaces_service_failures() {
let runtime = Runtime::new().unwrap();
let service = FakeBlobService::default();
service.set_failing(true);
let result = runtime.block_on(support::connect(
&service,
ACCOUNT_URL,
ResponseCacheCodec,
runtime.handle().clone(),
));
assert!(matches!(result, Err(Error::Unavailable)));
}
#[rstest]
fn sync_set_and_get_round_trip_python_json_shape(fixture: Fixture) {
let value = entry(json!({"choices": [{"message": {"content": "héllo 🌍"}}]}));
fixture
.cache
.set_cache("key-1", value.clone(), &no_ttl())
.unwrap();
assert_eq!(
fixture.stored_json("key-1"),
json!({
"timestamp": 1_700_000_000.5,
"response": {"choices": [{"message": {"content": "héllo 🌍"}}]}
})
);
assert_eq!(
fixture.cache.get_cache("key-1", &no_ttl()).unwrap(),
Some(value)
);
}
#[rstest]
#[case::blob_already_exists(false)]
#[case::precondition_conflict(true)]
fn sync_set_does_not_overwrite_an_existing_blob(fixture: Fixture, #[case] precondition: bool) {
fixture.service.set_precondition_conflicts(precondition);
fixture
.cache
.set_cache("key", entry(json!({"v": "first"})), &no_ttl())
.unwrap();
fixture
.cache
.set_cache("key", entry(json!({"v": "second"})), &no_ttl())
.unwrap();
assert_eq!(
fixture.stored_json("key")["response"],
json!({"v": "first"})
);
let uploads: Vec<_> = fixture
.service
.requests()
.into_iter()
.filter(|request| request.method == Method::Put && request.path.ends_with("/key"))
.collect();
assert_eq!(uploads.len(), 2);
assert!(
uploads
.iter()
.all(|request| request.if_none_match.as_deref() == Some("*"))
);
}
#[rstest]
fn async_set_overwrites_an_existing_blob(fixture: Fixture) {
fixture.runtime.block_on(async {
fixture
.cache
.async_set_cache("key", entry(json!({"v": "first"})), no_ttl())
.await
.unwrap();
fixture
.cache
.async_set_cache("key", entry(json!({"v": "second"})), no_ttl())
.await
.unwrap();
assert_eq!(
fixture
.cache
.async_get_cache("key", &no_ttl())
.await
.unwrap(),
Some(entry(json!({"v": "second"})))
);
});
assert_eq!(
fixture.stored_json("key")["response"],
json!({"v": "second"})
);
assert!(
fixture
.service
.requests()
.iter()
.filter(|request| request.method == Method::Put && request.path.ends_with("/key"))
.all(|request| request.if_none_match.is_none())
);
}
#[rstest]
fn missing_blobs_are_misses(fixture: Fixture) {
assert_eq!(fixture.cache.get_cache("absent", &no_ttl()).unwrap(), None);
assert_eq!(
fixture
.runtime
.block_on(fixture.cache.async_get_cache("absent", &no_ttl()))
.unwrap(),
None
);
}
#[rstest]
fn ttl_is_ignored_and_entries_never_expire(fixture: Fixture) {
assert_eq!(fixture.cache.get_ttl(&with_ttl(1)), None);
assert_eq!(fixture.cache.get_ttl(&no_ttl()), None);
fixture
.cache
.set_cache("key", entry(json!("value")), &with_ttl(1))
.unwrap();
std::thread::sleep(Duration::from_millis(1100));
assert_eq!(
fixture.cache.get_cache("key", &with_ttl(1)).unwrap(),
Some(entry(json!("value")))
);
assert!(
fixture
.service
.requests()
.iter()
.all(|request| !request.query.contains("expiry"))
);
}
#[rstest]
#[case::broken_json("broken-json", b"{not json".as_slice())]
#[case::broken_utf8("broken-utf8", &[0xff, 0xfe, 0x22])]
#[case::wrong_shape("wrong-shape", br#"{"timestamp": "yesterday"}"#.as_slice())]
fn malformed_blobs_are_invalid_entries(fixture: Fixture, #[case] key: &str, #[case] bytes: &[u8]) {
fixture.service.seed_blob(key, bytes);
assert!(matches!(
fixture.cache.get_cache(key, &no_ttl()),
Err(Error::InvalidEntry)
));
}
#[rstest]
fn malformed_blobs_are_response_cache_misses(fixture: Fixture) {
let response_cache = response_cache(&fixture);
let broken = request("broken");
fixture
.service
.seed_blob(&cache_key(&broken.key), b"{not json");
assert_eq!(response_cache.lookup(&broken, now()).unwrap(), None);
assert_eq!(
fixture
.runtime
.block_on(response_cache.async_lookup(&broken, now()))
.unwrap(),
None
);
}
#[rstest]
fn batch_get_preserves_order_and_marks_misses_and_invalid_entries(fixture: Fixture) {
fixture
.cache
.set_cache("a", entry(json!("A")), &no_ttl())
.unwrap();
fixture
.cache
.set_cache("c", entry(json!("C")), &no_ttl())
.unwrap();
fixture.service.seed_blob("bad", b"nope");
let keys = ["c", "missing", "a", "bad"].map(String::from);
let sync = fixture.cache.batch_get_cache(&keys, &no_ttl()).unwrap();
assert_eq!(
sync,
vec![
BatchEntry::Hit(entry(json!("C"))),
BatchEntry::Miss,
BatchEntry::Hit(entry(json!("A"))),
BatchEntry::Invalid,
]
);
let asynchronous = fixture
.runtime
.block_on(fixture.cache.async_batch_get_cache(keys.to_vec(), no_ttl()))
.unwrap();
assert_eq!(asynchronous, sync);
let response_cache = response_cache(&fixture);
let requests = [request("hit"), request("missing"), request("bad")];
response_cache
.store(&requests[0], json!("HIT"), now())
.unwrap();
fixture
.service
.seed_blob(&cache_key(&requests[2].key), b"nope");
let hits = response_cache.lookup_batch(&requests, now()).unwrap();
assert_eq!(hits.values, vec![Some(json!("HIT")), None, None]);
assert_eq!(hits.missing_indices, vec![1, 2]);
let async_hits = fixture
.runtime
.block_on(response_cache.async_lookup_batch(&requests, now()))
.unwrap();
assert_eq!(async_hits.values, hits.values);
}
#[rstest]
fn async_pipeline_writes_every_entry_with_overwrite(fixture: Fixture) {
fixture.service.seed_blob("k2", b"stale");
fixture
.runtime
.block_on(fixture.cache.async_set_cache_pipeline(
vec![
("k1".into(), entry(json!({"n": 1}))),
("k2".into(), entry(json!({"n": 2}))),
("k3".into(), entry(json!({"n": 3}))),
],
with_ttl(30),
))
.unwrap();
assert_eq!(fixture.service.blob_names(), ["k1", "k2", "k3"]);
assert_eq!(fixture.stored_json("k2")["response"], json!({"n": 2}));
}
#[rstest]
fn flush_deletes_every_blob_in_the_container(fixture: Fixture) {
for key in ["x", "y", "z"] {
fixture
.cache
.set_cache(key, entry(json!(key)), &no_ttl())
.unwrap();
}
fixture.cache.flush_cache().unwrap();
assert!(fixture.service.blob_names().is_empty());
assert!(fixture.service.container_exists());
fixture
.cache
.set_cache("again", entry(json!(1)), &no_ttl())
.unwrap();
fixture
.runtime
.block_on(fixture.cache.async_flush_cache())
.unwrap();
assert!(fixture.service.blob_names().is_empty());
}
#[rstest]
fn service_failures_map_to_unavailable(fixture: Fixture) {
fixture.service.set_failing(true);
assert!(matches!(
fixture.cache.get_cache("key", &no_ttl()),
Err(Error::Unavailable)
));
assert!(matches!(
fixture.cache.set_cache("key", entry(json!(1)), &no_ttl()),
Err(Error::Unavailable)
));
assert!(matches!(
fixture.cache.flush_cache(),
Err(Error::Unavailable)
));
assert!(matches!(
fixture.runtime.block_on(
fixture
.cache
.async_set_cache_pipeline(vec![("k".into(), entry(json!(1)))], no_ttl())
),
Err(Error::Unavailable)
));
}
#[rstest]
fn disconnect_is_idempotent_and_keeps_data(fixture: Fixture) {
fixture
.cache
.set_cache("key", entry(json!(1)), &no_ttl())
.unwrap();
fixture.runtime.block_on(async {
fixture.cache.disconnect().await.unwrap();
fixture.cache.disconnect().await.unwrap();
});
assert_eq!(
fixture.cache.get_cache("key", &no_ttl()).unwrap(),
Some(entry(json!(1)))
);
}
#[rstest]
fn response_cache_stores_and_reads_through_the_backend(fixture: Fixture) {
let response_cache = response_cache(&fixture);
let mut request = request("gpt");
request.context = with_ttl(60);
let response = json!({"id": "chatcmpl-1"});
response_cache
.store(&request, response.clone(), now())
.unwrap();
assert_eq!(
fixture.stored_json(&cache_key(&request.key)),
json!({"timestamp": 1_700_000_000.0, "response": {"id": "chatcmpl-1"}})
);
assert_eq!(
response_cache
.lookup(&request, now() + Duration::from_secs(3600))
.unwrap(),
Some(response.clone())
);
assert_eq!(
fixture
.runtime
.block_on(response_cache.async_lookup(&request, now() + Duration::from_secs(3600)))
.unwrap(),
Some(response.clone())
);
fixture.runtime.block_on(async {
response_cache
.async_store(&request, json!("replaced"), now())
.await
.unwrap();
assert_eq!(
response_cache.async_lookup(&request, now()).await.unwrap(),
Some(json!("replaced"))
);
response_cache.async_flush().await.unwrap();
assert_eq!(
response_cache.async_lookup(&request, now()).await.unwrap(),
None
);
});
}
#[rstest]
fn non_object_responses_are_written_serialized_like_python(fixture: Fixture) {
fixture
.cache
.set_cache("s", entry(json!("plain")), &no_ttl())
.unwrap();
assert_eq!(
fixture.stored_json("s"),
json!({"timestamp": 1_700_000_000.5, "response": "\"plain\""})
);
assert_eq!(
fixture.cache.get_cache("s", &no_ttl()).unwrap(),
Some(entry(json!("plain")))
);
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn sync_methods_block_inside_a_multi_thread_runtime() {
let service = FakeBlobService::default();
let cache = support::connect(
&service,
ACCOUNT_URL,
ResponseCacheCodec,
tokio::runtime::Handle::current(),
)
.await
.map(Arc::new)
.unwrap();
cache.set_cache("key", entry(json!(1)), &no_ttl()).unwrap();
assert_eq!(
cache.get_cache("key", &no_ttl()).unwrap(),
Some(entry(json!(1)))
);
}

View file

@ -0,0 +1,81 @@
mod support;
use litellm_cache::{ExactCacheContext, JsonCodec};
use litellm_cache_azure_blob::AzureBlobCache;
use litellm_cache_testing as contract;
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use support::{ACCOUNT_URL, FakeBlobService};
use tokio::runtime::Handle;
#[fixture]
async fn azure() -> AzureBlobCache<JsonCodec<Value>> {
support::connect(
&FakeBlobService::default(),
ACCOUNT_URL,
JsonCodec::new(),
Handle::current(),
)
.await
.unwrap()
}
#[fixture]
fn context() -> ExactCacheContext {
ExactCacheContext::default()
}
const PREFIX: &str = "contract:";
// `overwrite_replaces` does not apply: sync `set_cache` never overwrites a blob, as in Python.
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn hit_and_miss(
#[future(awt)] azure: AzureBlobCache<JsonCodec<Value>>,
context: ExactCacheContext,
) {
contract::hit_and_miss(&azure, context, PREFIX, json!({"answer": 42})).await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn sync_async_equivalence(
#[future(awt)] azure: AzureBlobCache<JsonCodec<Value>>,
context: ExactCacheContext,
) {
contract::sync_async_equivalence(&azure, context, PREFIX, json!("first"), json!([2])).await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn pipeline_writes_every_entry(
#[future(awt)] azure: AzureBlobCache<JsonCodec<Value>>,
context: ExactCacheContext,
) {
contract::pipeline_writes_every_entry(
&azure,
context,
PREFIX,
vec![json!("a"), json!(2), json!({"c": true})],
)
.await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn batch_preserves_order(
#[future(awt)] azure: AzureBlobCache<JsonCodec<Value>>,
context: ExactCacheContext,
) {
contract::batch_preserves_order(&azure, context, PREFIX, json!("first"), json!(2)).await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn flush_clears(
#[future(awt)] azure: AzureBlobCache<JsonCodec<Value>>,
context: ExactCacheContext,
) {
contract::flush_clears(&azure, context, PREFIX, json!("value")).await;
}

View file

@ -0,0 +1,239 @@
#![allow(dead_code)]
use std::{
collections::BTreeMap,
sync::{Arc, Mutex},
};
use azure_core::http::{
AsyncRawResponse, Body, ClientOptions, HttpClient, Method, Request, StatusCode, Transport,
headers::{HeaderName, Headers},
};
use litellm_cache::{CacheCodec, Error};
use litellm_cache_azure_blob::AzureBlobCache;
use tokio::runtime::{Handle, Runtime};
pub const ACCOUNT_URL: &str = "https://example.blob.core.windows.net";
pub const CONTAINER: &str = "litellm-cache";
const IF_NONE_MATCH: HeaderName = HeaderName::from_static("if-none-match");
const ERROR_CODE: HeaderName = HeaderName::from_static("x-ms-error-code");
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RecordedRequest {
pub method: Method,
pub path: String,
pub query: String,
pub if_none_match: Option<String>,
}
#[derive(Default)]
struct FakeState {
container_exists: bool,
blobs: BTreeMap<String, Vec<u8>>,
requests: Vec<RecordedRequest>,
failing: bool,
precondition_conflicts: bool,
}
#[derive(Clone, Default)]
pub struct FakeBlobService {
state: Arc<Mutex<FakeState>>,
}
impl std::fmt::Debug for FakeBlobService {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("FakeBlobService")
}
}
impl FakeBlobService {
pub fn with_existing_container() -> Self {
let service = Self::default();
service.state.lock().unwrap().container_exists = true;
service
}
pub fn blob(&self, name: &str) -> Option<Vec<u8>> {
self.state.lock().unwrap().blobs.get(name).cloned()
}
pub fn blob_names(&self) -> Vec<String> {
self.state.lock().unwrap().blobs.keys().cloned().collect()
}
pub fn seed_blob(&self, name: &str, bytes: &[u8]) {
self.state
.lock()
.unwrap()
.blobs
.insert(name.to_string(), bytes.to_vec());
}
pub fn set_failing(&self, failing: bool) {
self.state.lock().unwrap().failing = failing;
}
pub fn set_precondition_conflicts(&self, enabled: bool) {
self.state.lock().unwrap().precondition_conflicts = enabled;
}
pub fn requests(&self) -> Vec<RecordedRequest> {
self.state.lock().unwrap().requests.clone()
}
pub fn container_exists(&self) -> bool {
self.state.lock().unwrap().container_exists
}
fn respond(status: StatusCode, error_code: Option<&str>, body: Vec<u8>) -> AsyncRawResponse {
let mut headers = Headers::new();
if let Some(code) = error_code {
headers.insert(ERROR_CODE, code.to_string());
}
AsyncRawResponse::from_bytes(status, headers, body)
}
fn list_body(state: &FakeState) -> Vec<u8> {
let mut xml = String::from(
r#"<?xml version="1.0" encoding="utf-8"?><EnumerationResults ServiceEndpoint="https://example.blob.core.windows.net/" ContainerName="litellm-cache"><Blobs>"#,
);
for name in state.blobs.keys() {
xml.push_str(&format!(
"<Blob><Name>{name}</Name><Properties><BlobType>BlockBlob</BlobType></Properties></Blob>"
));
}
xml.push_str("</Blobs><NextMarker /></EnumerationResults>");
xml.into_bytes()
}
}
#[async_trait::async_trait]
impl HttpClient for FakeBlobService {
async fn execute_request(&self, request: &Request) -> azure_core::Result<AsyncRawResponse> {
let mut state = self.state.lock().unwrap();
let path = request.url().path().to_string();
let query = request.url().query().unwrap_or_default().to_string();
let if_none_match = request
.headers()
.get_optional_str(&IF_NONE_MATCH)
.map(str::to_owned);
state.requests.push(RecordedRequest {
method: request.method(),
path: path.clone(),
query: query.clone(),
if_none_match: if_none_match.clone(),
});
if state.failing {
return Ok(Self::respond(
StatusCode::Forbidden,
Some("AuthorizationFailure"),
Vec::new(),
));
}
let container_path = format!("/{CONTAINER}");
let blob_name = path
.strip_prefix(&format!("{container_path}/"))
.map(str::to_owned);
let is_container = path == container_path && query.contains("restype=container");
let response = match (request.method(), is_container, blob_name) {
(Method::Put, true, None) if state.container_exists => Self::respond(
StatusCode::Conflict,
Some("ContainerAlreadyExists"),
Vec::new(),
),
(Method::Put, true, None) => {
state.container_exists = true;
Self::respond(StatusCode::Created, None, Vec::new())
}
(Method::Get, true, None) if query.contains("comp=list") => {
Self::respond(StatusCode::Ok, None, Self::list_body(&state))
}
(Method::Get, true, None) if state.container_exists => {
Self::respond(StatusCode::Ok, None, Vec::new())
}
(Method::Get, true, None) => {
Self::respond(StatusCode::NotFound, Some("ContainerNotFound"), Vec::new())
}
(Method::Put, false, Some(name)) => {
if if_none_match.as_deref() == Some("*") && state.blobs.contains_key(&name) {
if state.precondition_conflicts {
Self::respond(
StatusCode::PreconditionFailed,
Some("ConditionNotMet"),
Vec::new(),
)
} else {
Self::respond(StatusCode::Conflict, Some("BlobAlreadyExists"), Vec::new())
}
} else {
let bytes = match request.body() {
Body::Bytes(bytes) => bytes.to_vec(),
Body::SeekableStream(_) => panic!("unexpected streaming upload"),
};
state.blobs.insert(name, bytes);
Self::respond(StatusCode::Created, None, Vec::new())
}
}
(Method::Get, false, Some(name)) => match state.blobs.get(&name) {
Some(bytes) => Self::respond(StatusCode::Ok, None, bytes.clone()),
None => Self::respond(StatusCode::NotFound, Some("BlobNotFound"), Vec::new()),
},
(Method::Delete, false, Some(name)) => match state.blobs.remove(&name) {
Some(_) => Self::respond(StatusCode::Accepted, None, Vec::new()),
None => Self::respond(StatusCode::NotFound, Some("BlobNotFound"), Vec::new()),
},
(method, _, _) => panic!("unexpected request {method:?} {path}?{query}"),
};
Ok(response)
}
}
pub async fn connect<C: CacheCodec>(
service: &FakeBlobService,
account_url: &str,
codec: C,
handle: Handle,
) -> Result<AzureBlobCache<C>, Error> {
AzureBlobCache::connect_with_options(
account_url,
CONTAINER,
None,
ClientOptions {
transport: Some(Transport::new(Arc::new(service.clone()))),
..ClientOptions::default()
},
codec,
handle,
)
.await
}
/// A cache on its own fake service and runtime, so sync methods run outside any runtime.
pub struct Fixture<C> {
pub runtime: Runtime,
pub service: FakeBlobService,
pub cache: Arc<AzureBlobCache<C>>,
}
impl<C: CacheCodec> Fixture<C> {
pub fn new(service: FakeBlobService, codec: C) -> Self {
let runtime = Runtime::new().unwrap();
let cache = runtime
.block_on(connect(
&service,
ACCOUNT_URL,
codec,
runtime.handle().clone(),
))
.unwrap();
Self {
runtime,
service,
cache: Arc::new(cache),
}
}
pub fn stored_json(&self, key: &str) -> serde_json::Value {
serde_json::from_slice(&self.service.blob(key).expect("blob should exist")).unwrap()
}
}

View file

@ -0,0 +1,90 @@
use std::sync::Arc;
use azure_core::http::{ClientOptions, Transport};
use litellm_cache::{BaseCache, ExactCacheContext, JsonCodec};
use litellm_cache_azure_blob::{AzureBlobCache, ReqwestTransport};
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use tokio::runtime::Handle;
use wiremock::{
Mock, MockServer, ResponseTemplate,
matchers::{body_json, header, method, path, query_param},
};
#[fixture]
async fn server() -> MockServer {
let server = MockServer::start().await;
Mock::given(method("PUT"))
.and(path("/litellm-cache"))
.and(query_param("restype", "container"))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
server
}
async fn connect(server: &MockServer) -> AzureBlobCache<JsonCodec<Value>> {
AzureBlobCache::connect_with_options(
&server.uri(),
"litellm-cache",
None,
ClientOptions {
transport: Some(Transport::new(Arc::new(ReqwestTransport(
reqwest::Client::new(),
)))),
..ClientOptions::default()
},
JsonCodec::new(),
Handle::current(),
)
.await
.unwrap()
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn uploads_go_through_the_host_client(#[future(awt)] server: MockServer) {
Mock::given(method("PUT"))
.and(path("/litellm-cache/key"))
.and(header("if-none-match", "*"))
.and(body_json(json!({"answer": 1})))
.respond_with(ResponseTemplate::new(201))
.expect(1)
.mount(&server)
.await;
connect(&server)
.await
.set_cache("key", json!({"answer": 1}), &ExactCacheContext::default())
.unwrap();
}
#[rstest]
#[case::hit(
ResponseTemplate::new(200).set_body_json(json!({"answer": 2})),
Some(json!({"answer": 2}))
)]
#[case::blob_not_found(
ResponseTemplate::new(404).insert_header("x-ms-error-code", "BlobNotFound"),
None
)]
#[tokio::test(flavor = "multi_thread")]
async fn downloads_map_the_host_client_response(
#[future(awt)] server: MockServer,
#[case] response: ResponseTemplate,
#[case] expected: Option<Value>,
) {
Mock::given(method("GET"))
.and(path("/litellm-cache/key"))
.respond_with(response)
.mount(&server)
.await;
assert_eq!(
connect(&server)
.await
.async_get_cache("key", &ExactCacheContext::default())
.await
.unwrap(),
expected
);
}

View file

@ -0,0 +1,20 @@
[package]
name = "litellm-cache-disk"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-cache.workspace = true
py_literal = "0.4.0"
rand.workspace = true
rusqlite = { version = "0.40", features = ["bundled"] }
serde-pickle = "1.2"
serde_json.workspace = true
tokio.workspace = true
[dev-dependencies]
litellm-cache-testing.workspace = true
rstest.workspace = true
tempfile = "3.27.0"

View file

@ -0,0 +1,10 @@
use litellm_cache::Error;
use crate::StoredValue;
pub trait ValueAdapter: Send + Sync + 'static {
fn read(&self, value: StoredValue) -> Result<Option<Vec<u8>>, Error>;
fn write(&self, payload: Vec<u8>) -> StoredValue;
fn counter_seed(&self, value: Option<StoredValue>) -> Result<f64, Error>;
fn counter_value(&self, value: f64) -> StoredValue;
}

View file

@ -0,0 +1,283 @@
use std::{
path::Path,
sync::Arc,
time::{Duration, SystemTime, UNIX_EPOCH},
};
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheCodec, CounterCache, DeleteCache, DisconnectCache,
Error, ExactCacheContext, FlushCache,
};
use crate::{DiskStore, DiskcacheSqliteStore, PythonDiskCacheAdapter, StoredValue, ValueAdapter};
pub struct DiskCache<S, D = DiskcacheSqliteStore, A = PythonDiskCacheAdapter> {
store: Arc<D>,
adapter: Arc<A>,
codec: S,
}
impl<S: CacheCodec> DiskCache<S> {
pub fn open(directory: impl AsRef<Path>, codec: S) -> Result<Self, Error> {
Ok(Self {
store: Arc::new(DiskcacheSqliteStore::open(directory)?),
adapter: Arc::new(PythonDiskCacheAdapter),
codec,
})
}
}
impl<S: CacheCodec, D: DiskStore> DiskCache<S, D, PythonDiskCacheAdapter> {
pub fn with_store(store: D, codec: S) -> Self {
Self {
store: Arc::new(store),
adapter: Arc::new(PythonDiskCacheAdapter),
codec,
}
}
}
impl<S: CacheCodec, D: DiskStore, A: ValueAdapter> DiskCache<S, D, A> {
pub fn with_adapter(store: D, adapter: A, codec: S) -> Self {
Self {
store: Arc::new(store),
adapter: Arc::new(adapter),
codec,
}
}
pub fn directory(&self) -> &Path {
self.store.directory()
}
fn decode_stored(&self, value: StoredValue) -> Result<Option<S::Value>, Error> {
let Some(bytes) = self.adapter.read(value)? else {
return Ok(None);
};
self.codec.decode(&bytes).map(Some)
}
async fn run_blocking<T, F>(store: Arc<D>, operation: F) -> Result<T, Error>
where
T: Send + 'static,
F: FnOnce(&D) -> Result<T, Error> + Send + 'static,
{
tokio::task::spawn_blocking(move || operation(&store))
.await
.map_err(|_| Error::Unavailable)?
}
}
impl<S: CacheCodec, D: DiskStore, A: ValueAdapter> BaseCache for DiskCache<S, D, A> {
type Value = S::Value;
type Context = ExactCacheContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl
}
fn set_cache(
&self,
key: &str,
value: Self::Value,
context: &Self::Context,
) -> Result<(), Error> {
let value = self.adapter.write(self.codec.encode(&value)?);
let expire_time = context.ttl.map(|ttl| unix_now() + ttl.as_secs_f64());
self.store.set(key, value, expire_time, unix_now())
}
fn get_cache(&self, key: &str, _: &Self::Context) -> Result<Option<Self::Value>, Error> {
self.store
.get(key, unix_now())?
.map(|value| self.decode_stored(value))
.transpose()
.map(|value| value.flatten())
}
async fn async_set_cache(
&self,
key: &str,
value: Self::Value,
context: ExactCacheContext,
) -> Result<(), Error> {
let value = self.adapter.write(self.codec.encode(&value)?);
let ttl = context.ttl;
let key = key.to_string();
Self::run_blocking(Arc::clone(&self.store), move |store| {
let expire_time = ttl.map(|ttl| unix_now() + ttl.as_secs_f64());
store.set(&key, value, expire_time, unix_now())
})
.await
}
async fn async_get_cache(
&self,
key: &str,
_: &ExactCacheContext,
) -> Result<Option<Self::Value>, Error> {
let key = key.to_string();
let value = Self::run_blocking(Arc::clone(&self.store), move |store| {
store.get(&key, unix_now())
})
.await?;
value
.map(|value| self.decode_stored(value))
.transpose()
.map(|value| value.flatten())
}
async fn async_set_cache_pipeline(
&self,
entries: Vec<(String, Self::Value)>,
context: ExactCacheContext,
) -> Result<(), Error> {
let entries = entries
.into_iter()
.map(|(key, value)| {
self.codec
.encode(&value)
.map(|value| (key, self.adapter.write(value)))
})
.collect::<Result<Vec<_>, _>>()?;
let expire_after = context.ttl;
Self::run_blocking(Arc::clone(&self.store), move |store| {
for (key, value) in entries {
let expire_time = expire_after.map(|ttl| unix_now() + ttl.as_secs_f64());
store.set(&key, value, expire_time, unix_now())?;
}
Ok(())
})
.await
}
}
impl<S: CacheCodec, D: DiskStore, A: ValueAdapter> BatchCache for DiskCache<S, D, A> {
fn batch_get_cache(
&self,
keys: &[String],
context: &ExactCacheContext,
) -> Result<Vec<BatchEntry<Self::Value>>, Error> {
keys.iter()
.map(|key| match self.get_cache(key, context) {
Ok(Some(value)) => Ok(BatchEntry::Hit(value)),
Ok(None) => Ok(BatchEntry::Miss),
Err(Error::InvalidEntry) => Ok(BatchEntry::Invalid),
Err(error) => Err(error),
})
.collect()
}
async fn async_batch_get_cache(
&self,
keys: Vec<String>,
_: ExactCacheContext,
) -> Result<Vec<BatchEntry<Self::Value>>, Error> {
let values = Self::run_blocking(Arc::clone(&self.store), move |store| {
keys.into_iter()
.map(|key| store.get(&key, unix_now()).map(|value| (key, value)))
.collect::<Result<Vec<_>, _>>()
})
.await?;
values
.into_iter()
.map(|(_, value)| match value {
None => Ok(BatchEntry::Miss),
Some(value) => match self.decode_stored(value) {
Ok(Some(value)) => Ok(BatchEntry::Hit(value)),
Ok(None) => Ok(BatchEntry::Miss),
Err(Error::InvalidEntry) => Ok(BatchEntry::Invalid),
Err(error) => Err(error),
},
})
.collect()
}
}
impl<S: CacheCodec, D: DiskStore, A: ValueAdapter> DeleteCache for DiskCache<S, D, A> {
fn delete_cache(&self, key: &str) -> Result<(), Error> {
self.store.pop(key, unix_now()).map(|_| ())
}
async fn async_delete_cache(&self, key: &str) -> Result<(), Error> {
let key = key.to_string();
Self::run_blocking(Arc::clone(&self.store), move |store| {
store.pop(&key, unix_now()).map(|_| ())
})
.await
}
}
impl<S: CacheCodec, D: DiskStore, A: ValueAdapter> FlushCache for DiskCache<S, D, A> {
fn flush_cache(&self) -> Result<(), Error> {
self.store.clear()
}
async fn async_flush_cache(&self) -> Result<(), Error> {
Self::run_blocking(Arc::clone(&self.store), |store| store.clear()).await
}
}
impl<S: CacheCodec, D: DiskStore, A: ValueAdapter> DisconnectCache for DiskCache<S, D, A> {
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
}
impl<S: CacheCodec, D: DiskStore, A: ValueAdapter> CounterCache for DiskCache<S, D, A> {
fn increment_cache(
&self,
key: &str,
amount: f64,
context: ExactCacheContext,
) -> Result<f64, Error> {
increment(
self.adapter.as_ref(),
self.store.as_ref(),
key,
amount,
context.ttl,
)
}
async fn async_increment(
&self,
key: &str,
amount: f64,
context: ExactCacheContext,
_refresh_ttl: bool,
) -> Result<f64, Error> {
let key = key.to_string();
let adapter = Arc::clone(&self.adapter);
Self::run_blocking(Arc::clone(&self.store), move |store| {
increment(adapter.as_ref(), store, &key, amount, context.ttl)
})
.await
}
}
fn increment<A: ValueAdapter, D: DiskStore>(
adapter: &A,
store: &D,
key: &str,
amount: f64,
ttl: Option<Duration>,
) -> Result<f64, Error> {
let mut result = None;
let mut apply = |current: Option<StoredValue>| {
let initial = adapter.counter_seed(current)?;
let value = initial + amount;
let stored = adapter.counter_value(value);
result = Some(value);
Ok((stored, ttl.map(|ttl| unix_now() + ttl.as_secs_f64())))
};
store.update(key, unix_now(), &mut apply)?;
result.ok_or(Error::InvalidEntry)
}
fn unix_now() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs_f64()
}

View file

@ -0,0 +1,11 @@
mod adapter;
mod cache;
mod python;
mod sqlite;
mod store;
pub use adapter::ValueAdapter;
pub use cache::DiskCache;
pub use python::PythonDiskCacheAdapter;
pub use sqlite::DiskcacheSqliteStore;
pub use store::{DiskStore, StoredValue};

View file

@ -0,0 +1,77 @@
mod value;
use litellm_cache::Error;
use py_literal::Value;
use crate::{StoredValue, ValueAdapter};
#[derive(Clone, Copy, Debug, Default)]
pub struct PythonDiskCacheAdapter;
impl PythonDiskCacheAdapter {
fn python_get_cache(value: StoredValue) -> Result<Option<Value>, Error> {
let value = match value {
StoredValue::Bytes(value) => Value::Bytes(value),
StoredValue::Text(value) => Value::String(value),
StoredValue::Integer(value) => Value::Integer(value.into()),
StoredValue::Float(value) => Value::Float(value),
StoredValue::Pickle(value) => value::from_pickle(&value)?,
};
if !value::is_truthy(&value) {
return Ok(None);
}
match value {
Value::String(text) => Ok(Some(
value::from_json_text(&text).unwrap_or(Value::String(text)),
)),
Value::Bytes(bytes) => match std::str::from_utf8(&bytes) {
Ok(text) => Ok(Some(
value::from_json_text(text).unwrap_or(Value::Bytes(bytes)),
)),
Err(_) => Ok(Some(Value::Bytes(bytes))),
},
value => Ok(Some(value)),
}
}
}
impl ValueAdapter for PythonDiskCacheAdapter {
fn read(&self, value: StoredValue) -> Result<Option<Vec<u8>>, Error> {
match value {
StoredValue::Text(value) => Ok((!value.is_empty()).then(|| value.into_bytes())),
StoredValue::Bytes(value) => Ok((!value.is_empty()).then_some(value)),
value => {
let Some(value) = Self::python_get_cache(value)? else {
return Ok(None);
};
value::to_json(&value).map(Some)
}
}
}
fn write(&self, payload: Vec<u8>) -> StoredValue {
StoredValue::Bytes(payload)
}
fn counter_seed(&self, value: Option<StoredValue>) -> Result<f64, Error> {
let Some(value) = value else {
return Ok(0.0);
};
let Some(value) = Self::python_get_cache(value)? else {
return Ok(0.0);
};
Ok(if value::is_int(&value) {
value::to_f64(&value).unwrap_or(0.0)
} else {
0.0
})
}
fn counter_value(&self, value: f64) -> StoredValue {
if value.fract() == 0.0 && value >= i64::MIN as f64 && value <= i64::MAX as f64 {
StoredValue::Integer(value as i64)
} else {
StoredValue::Float(value)
}
}
}

View file

@ -0,0 +1,173 @@
use litellm_cache::Error;
use py_literal::Value;
use serde_json::{Map, Number};
pub(crate) fn from_pickle(bytes: &[u8]) -> Result<Value, Error> {
let value = serde_pickle::value_from_slice(bytes, Default::default())
.map_err(|_| Error::InvalidEntry)?;
from_pickle_value(value)
}
fn from_pickle_value(value: serde_pickle::Value) -> Result<Value, Error> {
match value {
serde_pickle::Value::None => Ok(Value::None),
serde_pickle::Value::Bool(value) => Ok(Value::Boolean(value)),
serde_pickle::Value::I64(value) => integer(value.to_string()),
serde_pickle::Value::Int(value) => integer(value.to_string()),
serde_pickle::Value::F64(value) => Ok(Value::Float(value)),
serde_pickle::Value::String(value) => Ok(Value::String(value)),
serde_pickle::Value::Bytes(value) => Ok(Value::Bytes(value)),
serde_pickle::Value::List(values) => values
.into_iter()
.map(from_pickle_value)
.collect::<Result<Vec<_>, _>>()
.map(Value::List),
serde_pickle::Value::Tuple(values) => values
.into_iter()
.map(from_pickle_value)
.collect::<Result<Vec<_>, _>>()
.map(Value::Tuple),
serde_pickle::Value::Set(values) => values
.into_iter()
.map(from_pickle_hashable)
.collect::<Result<Vec<_>, _>>()
.map(Value::Set),
serde_pickle::Value::FrozenSet(values) => values
.into_iter()
.map(from_pickle_hashable)
.collect::<Result<Vec<_>, _>>()
.map(Value::Set),
serde_pickle::Value::Dict(values) => values
.into_iter()
.map(|(key, value)| Ok((from_pickle_hashable(key)?, from_pickle_value(value)?)))
.collect::<Result<Vec<_>, Error>>()
.map(Value::Dict),
}
}
fn from_pickle_hashable(value: serde_pickle::HashableValue) -> Result<Value, Error> {
Ok(match value {
serde_pickle::HashableValue::None => Value::None,
serde_pickle::HashableValue::Bool(value) => Value::Boolean(value),
serde_pickle::HashableValue::I64(value) => integer(value.to_string())?,
serde_pickle::HashableValue::Int(value) => integer(value.to_string())?,
serde_pickle::HashableValue::F64(value) => Value::Float(value),
serde_pickle::HashableValue::Bytes(value) => Value::Bytes(value),
serde_pickle::HashableValue::String(value) => Value::String(value),
serde_pickle::HashableValue::Tuple(values) => Value::Tuple(
values
.into_iter()
.map(from_pickle_hashable)
.collect::<Result<Vec<_>, _>>()?,
),
serde_pickle::HashableValue::FrozenSet(values) => Value::Set(
values
.into_iter()
.map(from_pickle_hashable)
.collect::<Result<Vec<_>, _>>()?,
),
})
}
fn integer(value: String) -> Result<Value, Error> {
value.parse().map_err(|_| Error::InvalidEntry)
}
pub(crate) fn from_json(value: serde_json::Value) -> Value {
match value {
serde_json::Value::Null => Value::None,
serde_json::Value::Bool(value) => Value::Boolean(value),
serde_json::Value::Number(value) => {
if value.is_i64() || value.is_u64() {
integer(value.to_string())
.unwrap_or(Value::Float(value.as_f64().unwrap_or(f64::NAN)))
} else {
Value::Float(value.as_f64().unwrap_or(f64::NAN))
}
}
serde_json::Value::String(value) => Value::String(value),
serde_json::Value::Array(values) => {
Value::List(values.into_iter().map(from_json).collect())
}
serde_json::Value::Object(values) => Value::Dict(
values
.into_iter()
.map(|(key, value)| (Value::String(key), from_json(value)))
.collect(),
),
}
}
pub(crate) fn from_json_text(value: &str) -> Result<Value, Error> {
serde_json::from_str(value)
.map(from_json)
.map_err(|_| Error::InvalidEntry)
}
pub(crate) fn is_truthy(value: &Value) -> bool {
match value {
Value::None => false,
Value::Boolean(value) => *value,
Value::Integer(value) => value.to_string() != "0",
Value::Float(value) => *value != 0.0,
Value::Complex(value) => value.re != 0.0 || value.im != 0.0,
Value::String(value) => !value.is_empty(),
Value::Bytes(value) => !value.is_empty(),
Value::Tuple(value) | Value::List(value) | Value::Set(value) => !value.is_empty(),
Value::Dict(value) => !value.is_empty(),
}
}
pub(crate) fn is_int(value: &Value) -> bool {
matches!(value, Value::Integer(_) | Value::Boolean(_))
}
pub(crate) fn to_f64(value: &Value) -> Option<f64> {
match value {
Value::Integer(value) => value.to_string().parse().ok(),
Value::Boolean(value) => Some(if *value { 1.0 } else { 0.0 }),
_ => None,
}
}
pub(crate) fn to_json(value: &Value) -> Result<Vec<u8>, Error> {
serde_json::to_vec(&to_json_value(value)?).map_err(|_| Error::InvalidEntry)
}
fn to_json_value(value: &Value) -> Result<serde_json::Value, Error> {
Ok(match value {
Value::None => serde_json::Value::Null,
Value::Boolean(value) => serde_json::Value::Bool(*value),
Value::Integer(value) => serde_json::Value::Number(
value
.to_string()
.parse::<Number>()
.map_err(|_| Error::InvalidEntry)?,
),
Value::Float(value) => {
serde_json::Value::Number(Number::from_f64(*value).ok_or(Error::InvalidEntry)?)
}
Value::Complex(_) | Value::Bytes(_) => return Err(Error::InvalidEntry),
Value::String(value) => serde_json::Value::String(value.clone()),
Value::Tuple(values) | Value::List(values) | Value::Set(values) => {
serde_json::Value::Array(
values
.iter()
.map(to_json_value)
.collect::<Result<Vec<_>, _>>()?,
)
}
Value::Dict(values) => {
let values = values
.iter()
.map(|(key, value)| {
let Value::String(key) = key else {
return Err(Error::InvalidEntry);
};
Ok((key.clone(), to_json_value(value)?))
})
.collect::<Result<Map<String, serde_json::Value>, _>>()?;
serde_json::Value::Object(values)
}
})
}

View file

@ -0,0 +1,805 @@
use std::{
collections::HashMap,
fs::{self, OpenOptions},
io::Write,
path::{Path, PathBuf},
sync::Mutex,
};
use litellm_cache::Error;
use rand::RngCore;
use rusqlite::{Connection, OptionalExtension, params, types::Value};
use crate::{DiskStore, StoredValue};
const MODE_RAW: i64 = 1;
const MODE_BINARY: i64 = 2;
const MODE_TEXT: i64 = 3;
const MODE_PICKLE: i64 = 4;
const DEFAULT_DISK_MIN_FILE_SIZE: i64 = 2_i64.pow(15);
const DEFAULT_SIZE_LIMIT: i64 = 2_i64.pow(30);
const DEFAULT_CULL_LIMIT: i64 = 10;
pub struct DiskcacheSqliteStore {
directory: PathBuf,
connection: Mutex<Connection>,
min_file_size: usize,
eviction_policy: String,
size_limit: i64,
cull_limit: i64,
statistics: bool,
}
struct StoredColumns {
size: i64,
mode: i64,
filename: Option<String>,
value: Option<Value>,
}
struct Row {
rowid: i64,
mode: i64,
filename: Option<String>,
value: Value,
}
impl DiskcacheSqliteStore {
pub fn open(directory: impl AsRef<Path>) -> Result<Self, Error> {
let directory = directory.as_ref().to_path_buf();
fs::create_dir_all(&directory).map_err(|_| Error::Unavailable)?;
let directory = std::path::absolute(&directory).map_err(|_| Error::Unavailable)?;
let database = directory.join("cache.db");
let connection = Connection::open(database).map_err(|_| Error::Unavailable)?;
connection
.busy_timeout(std::time::Duration::from_secs(60))
.map_err(|_| Error::Unavailable)?;
let mut settings = read_settings(&connection)?;
for (key, value) in default_settings() {
settings.entry(key).or_insert(value);
}
for (key, value) in settings
.iter()
.filter(|(key, _)| key.starts_with("sqlite_"))
{
apply_pragma(&connection, key, value)?;
}
connection
.execute_batch(
"CREATE TABLE IF NOT EXISTS Settings (
key TEXT NOT NULL UNIQUE,
value
)",
)
.map_err(|_| Error::Unavailable)?;
for (key, value) in &settings {
if !matches!(key.as_str(), "count" | "size" | "hits" | "misses") {
connection
.execute(
"INSERT OR REPLACE INTO Settings VALUES (?, ?)",
params![key, value],
)
.map_err(|_| Error::Unavailable)?;
}
}
for (key, value) in [
("count", Value::Integer(0)),
("size", Value::Integer(0)),
("hits", Value::Integer(0)),
("misses", Value::Integer(0)),
] {
connection
.execute(
"INSERT OR IGNORE INTO Settings VALUES (?, ?)",
params![key, value],
)
.map_err(|_| Error::Unavailable)?;
}
connection
.execute_batch(
"CREATE TABLE IF NOT EXISTS Cache (
rowid INTEGER PRIMARY KEY,
key BLOB,
raw INTEGER,
store_time REAL,
expire_time REAL,
access_time REAL,
access_count INTEGER DEFAULT 0,
tag BLOB,
size INTEGER DEFAULT 0,
mode INTEGER DEFAULT 0,
filename TEXT,
value BLOB
);
CREATE UNIQUE INDEX IF NOT EXISTS Cache_key_raw ON Cache(key, raw);
CREATE INDEX IF NOT EXISTS Cache_expire_time ON Cache(expire_time);",
)
.map_err(|_| Error::Unavailable)?;
let eviction_policy = setting_string(&settings, "eviction_policy")
.unwrap_or_else(|| "least-recently-stored".to_string());
match eviction_policy.as_str() {
"none" => {}
"least-recently-stored" => {
connection
.execute_batch(
"CREATE INDEX IF NOT EXISTS Cache_store_time ON Cache(store_time)",
)
.map_err(|_| Error::Unavailable)?;
}
"least-recently-used" => {
connection
.execute_batch(
"CREATE INDEX IF NOT EXISTS Cache_access_time ON Cache(access_time)",
)
.map_err(|_| Error::Unavailable)?;
}
"least-frequently-used" => {
connection
.execute_batch(
"CREATE INDEX IF NOT EXISTS Cache_access_count ON Cache(access_count)",
)
.map_err(|_| Error::Unavailable)?;
}
_ => return Err(Error::Unavailable),
}
connection
.execute_batch(
"CREATE TRIGGER IF NOT EXISTS Settings_count_insert
AFTER INSERT ON Cache FOR EACH ROW BEGIN
UPDATE Settings SET value = value + 1
WHERE key = \"count\"; END;
CREATE TRIGGER IF NOT EXISTS Settings_count_delete
AFTER DELETE ON Cache FOR EACH ROW BEGIN
UPDATE Settings SET value = value - 1
WHERE key = \"count\"; END;
CREATE TRIGGER IF NOT EXISTS Settings_size_insert
AFTER INSERT ON Cache FOR EACH ROW BEGIN
UPDATE Settings SET value = value + NEW.size
WHERE key = \"size\"; END;
CREATE TRIGGER IF NOT EXISTS Settings_size_update
AFTER UPDATE ON Cache FOR EACH ROW BEGIN
UPDATE Settings
SET value = value + NEW.size - OLD.size
WHERE key = \"size\"; END;
CREATE TRIGGER IF NOT EXISTS Settings_size_delete
AFTER DELETE ON Cache FOR EACH ROW BEGIN
UPDATE Settings SET value = value - OLD.size
WHERE key = \"size\"; END;",
)
.map_err(|_| Error::Unavailable)?;
let min_file_size = setting_i64(&settings, "disk_min_file_size")
.unwrap_or(DEFAULT_DISK_MIN_FILE_SIZE)
.try_into()
.map_err(|_| Error::Unavailable)?;
let size_limit = setting_i64(&settings, "size_limit").unwrap_or(DEFAULT_SIZE_LIMIT);
let cull_limit = setting_i64(&settings, "cull_limit").unwrap_or(DEFAULT_CULL_LIMIT);
let statistics = setting_i64(&settings, "statistics").unwrap_or_default() != 0;
Ok(Self {
directory,
connection: Mutex::new(connection),
min_file_size,
eviction_policy,
size_limit,
cull_limit,
statistics,
})
}
fn set_locked(
&self,
connection: &Connection,
key: &str,
columns: StoredColumns,
expire_time: Option<f64>,
now: f64,
) -> Result<Vec<String>, Error> {
let mut cleanup = Vec::new();
if let Some(old_filename) = connection
.query_row(
"SELECT filename FROM Cache WHERE key = ? AND raw = 1",
params![key],
|row| row.get::<_, Option<String>>(0),
)
.optional()
.map_err(|_| Error::Unavailable)?
.flatten()
{
cleanup.push(old_filename);
}
let (size, mode, filename, value) =
(columns.size, columns.mode, columns.filename, columns.value);
let rowid = connection
.query_row(
"SELECT rowid FROM Cache WHERE key = ? AND raw = 1",
params![key],
|row| row.get::<_, i64>(0),
)
.optional()
.map_err(|_| Error::Unavailable)?;
if let Some(rowid) = rowid {
connection
.execute(
"UPDATE Cache SET store_time = ?, expire_time = ?, access_time = ?,
access_count = 0, tag = NULL, size = ?, mode = ?, filename = ?, value = ?
WHERE rowid = ?",
params![now, expire_time, now, size, mode, filename, value, rowid],
)
.map_err(|_| Error::Unavailable)?;
} else {
connection
.execute(
"INSERT INTO Cache(
key, raw, store_time, expire_time, access_time, access_count,
tag, size, mode, filename, value
) VALUES (?, 1, ?, ?, ?, 0, NULL, ?, ?, ?, ?)",
params![key, now, expire_time, now, size, mode, filename, value],
)
.map_err(|_| Error::Unavailable)?;
}
cleanup.extend(self.cull(connection, now)?);
Ok(cleanup)
}
fn cull(&self, connection: &Connection, now: f64) -> Result<Vec<String>, Error> {
if self.cull_limit <= 0 {
return Ok(Vec::new());
}
let mut cleanup = Vec::new();
let expired = connection
.prepare(
"SELECT rowid, filename FROM Cache
WHERE expire_time IS NOT NULL AND expire_time < ?
ORDER BY expire_time LIMIT ?",
)
.map_err(|_| Error::Unavailable)?
.query_map(params![now, self.cull_limit], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, Option<String>>(1)?))
})
.map_err(|_| Error::Unavailable)?
.collect::<Result<Vec<_>, _>>()
.map_err(|_| Error::Unavailable)?;
for (_, filename) in &expired {
if let Some(filename) = filename {
cleanup.push(filename.clone());
}
}
for (rowid, _) in &expired {
connection
.execute("DELETE FROM Cache WHERE rowid = ?", params![rowid])
.map_err(|_| Error::Unavailable)?;
}
let remaining = self.cull_limit - i64::try_from(expired.len()).unwrap_or(self.cull_limit);
if remaining <= 0 || self.volume(connection)? < self.size_limit {
return Ok(cleanup);
}
let order = match self.eviction_policy.as_str() {
"none" => return Ok(cleanup),
"least-recently-stored" => "store_time",
"least-recently-used" => "access_time",
"least-frequently-used" => "access_count",
_ => return Err(Error::Unavailable),
};
let rows = connection
.prepare(&format!(
"SELECT rowid, filename FROM Cache ORDER BY {order} LIMIT ?"
))
.map_err(|_| Error::Unavailable)?
.query_map(params![remaining], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, Option<String>>(1)?))
})
.map_err(|_| Error::Unavailable)?
.collect::<Result<Vec<_>, _>>()
.map_err(|_| Error::Unavailable)?;
for (_, filename) in &rows {
if let Some(filename) = filename {
cleanup.push(filename.clone());
}
}
for (rowid, _) in rows {
connection
.execute("DELETE FROM Cache WHERE rowid = ?", params![rowid])
.map_err(|_| Error::Unavailable)?;
}
Ok(cleanup)
}
fn volume(&self, connection: &Connection) -> Result<i64, Error> {
let page_count: i64 = connection
.query_row("PRAGMA page_count", [], |row| row.get(0))
.map_err(|_| Error::Unavailable)?;
let page_size: i64 = connection
.query_row("PRAGMA page_size", [], |row| row.get(0))
.map_err(|_| Error::Unavailable)?;
let size: i64 = connection
.query_row("SELECT value FROM Settings WHERE key = 'size'", [], |row| {
row.get(0)
})
.map_err(|_| Error::Unavailable)?;
Ok(page_count.saturating_mul(page_size).saturating_add(size))
}
}
impl DiskStore for DiskcacheSqliteStore {
fn directory(&self) -> &Path {
&self.directory
}
fn get(&self, key: &str, now: f64) -> Result<Option<StoredValue>, Error> {
let connection = self.connection.lock().map_err(|_| Error::Unavailable)?;
let select = "SELECT rowid, expire_time, mode, filename, value FROM Cache
WHERE key = ? AND raw = 1 AND (expire_time IS NULL OR expire_time > ?)";
let row = connection
.query_row(select, params![key, now], row_from_query)
.optional()
.map_err(|_| Error::Unavailable)?;
if !self.statistics && !has_get_update(&self.eviction_policy) {
return row
.map(|row| fetch_row(&self.directory, row))
.transpose()
.map(|value| value.flatten());
}
transactional(&connection, |connection| {
let row = connection
.query_row(select, params![key, now], row_from_query)
.optional()
.map_err(|_| Error::Unavailable)?;
let Some(row) = row else {
if self.statistics {
connection
.execute(
"UPDATE Settings SET value = value + 1 WHERE key = 'misses'",
[],
)
.map_err(|_| Error::Unavailable)?;
}
return Ok(None);
};
let rowid = row.rowid;
let value = fetch_row(&self.directory, row);
let hit = value.as_ref().is_ok_and(Option::is_some);
if hit && self.statistics {
connection
.execute(
"UPDATE Settings SET value = value + 1 WHERE key = 'hits'",
[],
)
.map_err(|_| Error::Unavailable)?;
} else if !hit && self.statistics {
connection
.execute(
"UPDATE Settings SET value = value + 1 WHERE key = 'misses'",
[],
)
.map_err(|_| Error::Unavailable)?;
}
if has_get_update(&self.eviction_policy) && hit {
let update = match self.eviction_policy.as_str() {
"least-recently-used" => "UPDATE Cache SET access_time = ? WHERE rowid = ?",
"least-frequently-used" => {
"UPDATE Cache SET access_count = access_count + 1 WHERE rowid = ?"
}
_ => return Err(Error::Unavailable),
};
if self.eviction_policy == "least-recently-used" {
connection
.execute(update, params![now, rowid])
.map_err(|_| Error::Unavailable)?;
} else {
connection
.execute(update, params![rowid])
.map_err(|_| Error::Unavailable)?;
}
}
value
})
}
fn set(
&self,
key: &str,
value: StoredValue,
expire_time: Option<f64>,
now: f64,
) -> Result<(), Error> {
let columns = store_value(&self.directory, self.min_file_size, value)?;
let new_filename = columns.filename.clone();
let connection = self.connection.lock().map_err(|_| Error::Unavailable)?;
let result = transactional(&connection, |connection| {
self.set_locked(connection, key, columns, expire_time, now)
});
match result {
Ok(cleanup) => {
cleanup_files(&self.directory, cleanup);
Ok(())
}
Err(error) => {
if let Some(filename) = new_filename {
remove_file(&self.directory, &filename);
}
Err(error)
}
}
}
fn pop(&self, key: &str, now: f64) -> Result<Option<StoredValue>, Error> {
let connection = self.connection.lock().map_err(|_| Error::Unavailable)?;
let selected = transactional(&connection, |connection| {
let row = connection
.query_row(
"SELECT rowid, expire_time, mode, filename, value FROM Cache
WHERE key = ? AND raw = 1
AND (expire_time IS NULL OR expire_time > ?)",
params![key, now],
row_from_query,
)
.optional()
.map_err(|_| Error::Unavailable)?;
let Some(row) = row else {
return Ok(None);
};
connection
.execute("DELETE FROM Cache WHERE rowid = ?", params![row.rowid])
.map_err(|_| Error::Unavailable)?;
Ok(Some(row))
})?;
let Some(row) = selected else {
return Ok(None);
};
let filename = row.filename.clone();
let result = fetch_row(&self.directory, row)?;
if let Some(filename) = filename {
remove_file(&self.directory, &filename);
}
Ok(result)
}
fn clear(&self) -> Result<(), Error> {
let connection = self.connection.lock().map_err(|_| Error::Unavailable)?;
let mut last_rowid = 0_i64;
loop {
let batch = transactional(&connection, |connection| {
let rows = connection
.prepare(
"SELECT rowid, filename FROM Cache
WHERE rowid > ? ORDER BY rowid LIMIT 100",
)
.map_err(|_| Error::Unavailable)?
.query_map(params![last_rowid], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, Option<String>>(1)?))
})
.map_err(|_| Error::Unavailable)?
.collect::<Result<Vec<_>, _>>()
.map_err(|_| Error::Unavailable)?;
if rows.is_empty() {
return Ok(rows);
}
let ids = rows
.iter()
.map(|(rowid, _)| rowid.to_string())
.collect::<Vec<_>>()
.join(",");
connection
.execute(&format!("DELETE FROM Cache WHERE rowid IN ({ids})"), [])
.map_err(|_| Error::Unavailable)?;
Ok(rows)
})?;
if batch.is_empty() {
return Ok(());
}
last_rowid = batch.last().map(|(rowid, _)| *rowid).unwrap_or(last_rowid);
cleanup_files(
&self.directory,
batch
.into_iter()
.filter_map(|(_, filename)| filename)
.collect(),
);
}
}
fn update(
&self,
key: &str,
now: f64,
apply: &mut dyn FnMut(Option<StoredValue>) -> Result<(StoredValue, Option<f64>), Error>,
) -> Result<(), Error> {
let connection = self.connection.lock().map_err(|_| Error::Unavailable)?;
let mut created_filename = None;
let result = transactional(&connection, |connection| {
let current = connection
.query_row(
"SELECT rowid, expire_time, mode, filename, value FROM Cache
WHERE key = ? AND raw = 1
AND (expire_time IS NULL OR expire_time > ?)",
params![key, now],
row_from_query,
)
.optional()
.map_err(|_| Error::Unavailable)?
.map(|row| fetch_row(&self.directory, row))
.transpose()?
.flatten();
let (value, expire_time) = apply(current)?;
let columns = store_value(&self.directory, self.min_file_size, value)?;
created_filename = columns.filename.clone();
let cleanup = self.set_locked(connection, key, columns, expire_time, now)?;
Ok(cleanup)
});
match result {
Ok(cleanup) => {
cleanup_files(&self.directory, cleanup);
Ok(())
}
Err(error) => {
if let Some(filename) = created_filename {
remove_file(&self.directory, &filename);
}
Err(error)
}
}
}
}
fn default_settings() -> HashMap<String, Value> {
HashMap::from([
("statistics".to_string(), Value::Integer(0)),
("tag_index".to_string(), Value::Integer(0)),
(
"eviction_policy".to_string(),
Value::Text("least-recently-stored".to_string()),
),
("size_limit".to_string(), Value::Integer(DEFAULT_SIZE_LIMIT)),
("cull_limit".to_string(), Value::Integer(DEFAULT_CULL_LIMIT)),
("sqlite_auto_vacuum".to_string(), Value::Integer(1)),
("sqlite_cache_size".to_string(), Value::Integer(8192)),
(
"sqlite_journal_mode".to_string(),
Value::Text("wal".to_string()),
),
(
"sqlite_mmap_size".to_string(),
Value::Integer(2_i64.pow(26)),
),
("sqlite_synchronous".to_string(), Value::Integer(1)),
(
"disk_min_file_size".to_string(),
Value::Integer(DEFAULT_DISK_MIN_FILE_SIZE),
),
("disk_pickle_protocol".to_string(), Value::Integer(5)),
])
}
fn read_settings(connection: &Connection) -> Result<HashMap<String, Value>, Error> {
let mut statement = match connection.prepare("SELECT key, value FROM Settings") {
Ok(statement) => statement,
Err(_) => return Ok(HashMap::new()),
};
statement
.query_map([], |row| Ok((row.get(0)?, row.get(1)?)))
.map_err(|_| Error::Unavailable)?
.collect::<Result<HashMap<_, _>, _>>()
.map_err(|_| Error::Unavailable)
}
fn apply_pragma(connection: &Connection, key: &str, value: &Value) -> Result<(), Error> {
let pragma = key.strip_prefix("sqlite_").ok_or(Error::Unavailable)?;
match value {
Value::Integer(value) => connection
.pragma_update(None, pragma, value)
.map_err(|_| Error::Unavailable),
Value::Text(value) => connection
.pragma_update(None, pragma, value)
.map_err(|_| Error::Unavailable),
_ => Err(Error::Unavailable),
}
}
fn setting_i64(settings: &HashMap<String, Value>, key: &str) -> Option<i64> {
match settings.get(key) {
Some(Value::Integer(value)) => Some(*value),
_ => None,
}
}
fn setting_string(settings: &HashMap<String, Value>, key: &str) -> Option<String> {
match settings.get(key) {
Some(Value::Text(value)) => Some(value.clone()),
_ => None,
}
}
fn has_get_update(policy: &str) -> bool {
matches!(policy, "least-recently-used" | "least-frequently-used")
}
fn row_from_query(row: &rusqlite::Row<'_>) -> rusqlite::Result<Row> {
Ok(Row {
rowid: row.get(0)?,
mode: row.get(2)?,
filename: row.get(3)?,
value: row.get(4)?,
})
}
fn fetch_row(directory: &Path, row: Row) -> Result<Option<StoredValue>, Error> {
match row.mode {
MODE_RAW => match row.value {
Value::Blob(value) => Ok(Some(StoredValue::Bytes(value))),
Value::Text(value) => Ok(Some(StoredValue::Text(value))),
Value::Integer(value) => Ok(Some(StoredValue::Integer(value))),
Value::Real(value) => Ok(Some(StoredValue::Float(value))),
Value::Null => Err(Error::InvalidEntry),
},
MODE_BINARY | MODE_PICKLE => {
let bytes = match row.value {
Value::Blob(value) => value,
Value::Null => {
let Some(value) = read_file(directory, row.filename.as_deref())? else {
return Ok(None);
};
value
}
_ => return Err(Error::InvalidEntry),
};
Ok(Some(if row.mode == MODE_BINARY {
StoredValue::Bytes(bytes)
} else {
StoredValue::Pickle(bytes)
}))
}
MODE_TEXT => {
let bytes = match row.value {
Value::Null => {
let Some(value) = read_file(directory, row.filename.as_deref())? else {
return Ok(None);
};
value
}
Value::Blob(value) => value,
Value::Text(value) => value.into_bytes(),
_ => return Err(Error::InvalidEntry),
};
Ok(Some(StoredValue::Text(
String::from_utf8(bytes).map_err(|_| Error::InvalidEntry)?,
)))
}
_ => Err(Error::InvalidEntry),
}
}
fn read_file(directory: &Path, filename: Option<&str>) -> Result<Option<Vec<u8>>, Error> {
let Some(filename) = filename else {
return Err(Error::InvalidEntry);
};
match fs::read(directory.join(filename)) {
Ok(value) => Ok(Some(value)),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(None),
Err(_) => Err(Error::Unavailable),
}
}
fn store_value(
directory: &Path,
min_file_size: usize,
value: StoredValue,
) -> Result<StoredColumns, Error> {
match value {
StoredValue::Integer(value) => Ok(StoredColumns {
size: 0,
mode: MODE_RAW,
filename: None,
value: Some(Value::Integer(value)),
}),
StoredValue::Float(value) => Ok(StoredColumns {
size: 0,
mode: MODE_RAW,
filename: None,
value: Some(Value::Real(value)),
}),
StoredValue::Text(value) if value.chars().count() < min_file_size => Ok(StoredColumns {
size: 0,
mode: MODE_RAW,
filename: None,
value: Some(Value::Text(value)),
}),
StoredValue::Text(value) => {
let bytes = value.into_bytes();
let filename = write_file(directory, &bytes)?;
Ok(StoredColumns {
size: i64::try_from(bytes.len()).map_err(|_| Error::Unavailable)?,
mode: MODE_TEXT,
filename: Some(filename),
value: None,
})
}
StoredValue::Bytes(value) if value.len() < min_file_size => Ok(StoredColumns {
size: 0,
mode: MODE_RAW,
filename: None,
value: Some(Value::Blob(value)),
}),
StoredValue::Bytes(value) => {
let filename = write_file(directory, &value)?;
Ok(StoredColumns {
size: i64::try_from(value.len()).map_err(|_| Error::Unavailable)?,
mode: MODE_BINARY,
filename: Some(filename),
value: None,
})
}
StoredValue::Pickle(value) if value.len() < min_file_size => Ok(StoredColumns {
size: 0,
mode: MODE_PICKLE,
filename: None,
value: Some(Value::Blob(value)),
}),
StoredValue::Pickle(value) => {
let filename = write_file(directory, &value)?;
Ok(StoredColumns {
size: i64::try_from(value.len()).map_err(|_| Error::Unavailable)?,
mode: MODE_PICKLE,
filename: Some(filename),
value: None,
})
}
}
}
fn write_file(directory: &Path, bytes: &[u8]) -> Result<String, Error> {
let mut random = [0_u8; 16];
rand::rngs::OsRng.fill_bytes(&mut random);
let hex = random
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>();
let filename = format!("{}/{}/{}.val", &hex[..2], &hex[2..4], &hex[4..]);
let path = directory.join(&filename);
if let Some(parent) = path.parent() {
fs::create_dir_all(parent).map_err(|_| Error::Unavailable)?;
}
let mut file = OpenOptions::new()
.write(true)
.create_new(true)
.open(path)
.map_err(|_| Error::Unavailable)?;
file.write_all(bytes).map_err(|_| Error::Unavailable)?;
Ok(filename)
}
fn cleanup_files(directory: &Path, filenames: Vec<String>) {
for filename in filenames {
remove_file(directory, &filename);
}
}
fn remove_file(directory: &Path, filename: &str) {
let path = directory.join(filename);
let _ = fs::remove_file(&path);
}
fn transactional<T>(
connection: &Connection,
operation: impl FnOnce(&Connection) -> Result<T, Error>,
) -> Result<T, Error> {
connection
.execute_batch("BEGIN IMMEDIATE")
.map_err(|_| Error::Unavailable)?;
match operation(connection) {
Ok(value) => {
connection
.execute_batch("COMMIT")
.map_err(|_| Error::Unavailable)?;
Ok(value)
}
Err(error) => {
let _ = connection.execute_batch("ROLLBACK");
Err(error)
}
}
}

View file

@ -0,0 +1,32 @@
use std::path::Path;
use litellm_cache::Error;
#[derive(Clone, Debug, PartialEq)]
pub enum StoredValue {
Bytes(Vec<u8>),
Text(String),
Integer(i64),
Float(f64),
Pickle(Vec<u8>),
}
pub trait DiskStore: Send + Sync + 'static {
fn directory(&self) -> &Path;
fn get(&self, key: &str, now: f64) -> Result<Option<StoredValue>, Error>;
fn set(
&self,
key: &str,
value: StoredValue,
expire_time: Option<f64>,
now: f64,
) -> Result<(), Error>;
fn pop(&self, key: &str, now: f64) -> Result<Option<StoredValue>, Error>;
fn clear(&self) -> Result<(), Error>;
fn update(
&self,
key: &str,
now: f64,
apply: &mut dyn FnMut(Option<StoredValue>) -> Result<(StoredValue, Option<f64>), Error>,
) -> Result<(), Error>;
}

View file

@ -0,0 +1,524 @@
use std::{
fs,
path::{Path, PathBuf},
sync::Arc,
thread,
time::Duration,
};
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheCodec, CounterCache, DeleteCache, DisconnectCache,
ExactCacheContext, FlushCache, JsonCodec,
};
use litellm_cache_disk::{DiskCache, DiskStore, DiskcacheSqliteStore, StoredValue, ValueAdapter};
use rstest::{fixture, rstest};
use rusqlite::Connection;
use serde_json::{Value, json};
use tempfile::TempDir;
struct Sandbox {
directory: TempDir,
}
#[fixture]
fn sandbox() -> Sandbox {
Sandbox {
directory: tempfile::tempdir().unwrap(),
}
}
impl Sandbox {
fn store(&self) -> DiskcacheSqliteStore {
DiskcacheSqliteStore::open(self.directory.path()).unwrap()
}
fn cache<V>(&self) -> DiskCache<JsonCodec<V>>
where
JsonCodec<V>: CacheCodec,
{
DiskCache::open(self.directory.path(), JsonCodec::new()).unwrap()
}
fn db(&self) -> Connection {
Connection::open(self.directory.path().join("cache.db")).unwrap()
}
fn value_files(&self) -> Vec<PathBuf> {
fn visit(directory: &Path, files: &mut Vec<PathBuf>) {
for entry in fs::read_dir(directory).unwrap() {
let path = entry.unwrap().path();
if path.is_dir() {
visit(&path, files);
} else if path.extension().is_some_and(|extension| extension == "val") {
files.push(path);
}
}
}
let mut files = Vec::new();
visit(self.directory.path(), &mut files);
files
}
}
#[rstest]
fn relative_store_directory_is_absolutized(sandbox: Sandbox) {
let relative = PathBuf::from(format!(
".litellm-cache-disk-{}",
sandbox
.directory
.path()
.file_name()
.unwrap()
.to_string_lossy()
));
let store = DiskcacheSqliteStore::open(&relative).unwrap();
assert!(store.directory().is_absolute());
assert!(store.directory().ends_with(&relative));
let directory = store.directory().to_path_buf();
drop(store);
fs::remove_dir_all(directory).unwrap();
}
#[derive(Clone, Copy, Debug, Default)]
struct TextAdapter;
impl ValueAdapter for TextAdapter {
fn read(&self, value: StoredValue) -> Result<Option<Vec<u8>>, litellm_cache::Error> {
match value {
StoredValue::Text(value) => Ok(Some(value.into_bytes())),
_ => Ok(None),
}
}
fn write(&self, payload: Vec<u8>) -> StoredValue {
StoredValue::Text(String::from_utf8(payload).unwrap())
}
fn counter_seed(&self, _: Option<StoredValue>) -> Result<f64, litellm_cache::Error> {
Ok(0.0)
}
fn counter_value(&self, value: f64) -> StoredValue {
if value.fract() == 0.0 {
StoredValue::Integer(value as i64)
} else {
StoredValue::Float(value)
}
}
}
#[rstest]
fn roundtrip_persists_and_reopens(sandbox: Sandbox) {
let context = ExactCacheContext::default();
let opened = sandbox.cache::<Value>();
opened
.set_cache("key", json!({"answer": 42}), &context)
.unwrap();
assert_eq!(
opened.get_cache("key", &context).unwrap(),
Some(json!({"answer": 42}))
);
drop(opened);
let reopened = sandbox.cache::<Value>();
assert_eq!(
reopened.get_cache("key", &context).unwrap(),
Some(json!({"answer": 42}))
);
}
#[rstest]
fn ttl_and_expired_culling_match_cache_contract(sandbox: Sandbox) {
let store = sandbox.store();
store
.set(
"expired",
StoredValue::Bytes(b"old".to_vec()),
Some(10.0),
0.0,
)
.unwrap();
assert_eq!(store.get("expired", 10.0).unwrap(), None);
store
.set("new", StoredValue::Bytes(b"new".to_vec()), None, 11.0)
.unwrap();
assert_eq!(
sandbox
.db()
.query_row("SELECT COUNT(*) FROM Cache", [], |row| row.get::<_, i64>(0))
.unwrap(),
1
);
assert_eq!(
sandbox
.db()
.query_row(
"SELECT value FROM Settings WHERE key = 'count'",
[],
|row| row.get::<_, i64>(0)
)
.unwrap(),
1
);
}
#[rstest]
fn batch_preserves_order_and_classifies_misses_and_invalid_values(sandbox: Sandbox) {
let store = sandbox.store();
store
.set(
"hit",
StoredValue::Bytes(br#"{"ok":true}"#.to_vec()),
None,
0.0,
)
.unwrap();
store
.set(
"invalid",
StoredValue::Pickle(vec![0x80, 0x05, 0x2e]),
None,
0.0,
)
.unwrap();
let entries = sandbox
.cache::<Value>()
.batch_get_cache(
&["hit".into(), "missing".into(), "invalid".into()],
&ExactCacheContext::default(),
)
.unwrap();
assert_eq!(
entries,
vec![
BatchEntry::Hit(json!({"ok": true})),
BatchEntry::Miss,
BatchEntry::Invalid
]
);
}
#[rstest]
#[case(StoredValue::Bytes(Vec::new()))]
#[case(StoredValue::Text(String::new()))]
#[case(StoredValue::Integer(0))]
#[case(StoredValue::Float(0.0))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x4e, 0x2e]))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x89, 0x2e]))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x4b, 0x00, 0x2e]))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x47, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x2e]))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x04, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x8c, 0x00, 0x94, 0x2e]))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x5d, 0x94, 0x2e]))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x7d, 0x94, 0x2e]))]
#[case(StoredValue::Pickle(vec![0x80, 0x05, 0x29, 0x2e]))]
fn falsy_values_are_misses(sandbox: Sandbox, #[case] value: StoredValue) {
sandbox.store().set("key", value, None, 0.0).unwrap();
assert_eq!(
sandbox
.cache::<Value>()
.get_cache("key", &ExactCacheContext::default())
.unwrap(),
None
);
}
#[rstest]
#[case(Some(StoredValue::Integer(2)), 1.5, 3.5, "real")]
#[case(Some(StoredValue::Integer(2)), 1.0, 3.0, "integer")]
#[case(Some(StoredValue::Float(3.5)), 1.0, 1.0, "integer")]
#[case(Some(StoredValue::Text("not a number".into())), 2.0, 2.0, "integer")]
#[case(Some(StoredValue::Text("5".into())), 2.0, 7.0, "integer")]
#[case(Some(StoredValue::Text("3.5".into())), 2.0, 2.0, "integer")]
#[case(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x88, 0x2e])), 1.0, 2.0, "integer")]
#[case(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x4b, 0x02, 0x2e])), 1.0, 3.0, "integer")]
#[case(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x7d, 0x94, 0x8c, 0x01, 0x61, 0x94, 0x4b, 0x01, 0x73, 0x2e])), 1.0, 1.0, "integer")]
#[case(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x4e, 0x2e])), 4.0, 4.0, "integer")]
fn counters_follow_python_initialization(
sandbox: Sandbox,
#[case] initial: Option<StoredValue>,
#[case] amount: f64,
#[case] expected: f64,
#[case] sqlite_type: &str,
) {
if let Some(initial) = initial {
sandbox.store().set("counter", initial, None, 0.0).unwrap();
}
let cache = sandbox.cache::<f64>();
assert_eq!(
cache
.increment_cache("counter", amount, ExactCacheContext::default())
.unwrap(),
expected
);
assert_eq!(
sandbox
.db()
.query_row(
"SELECT typeof(value) FROM Cache WHERE key = 'counter'",
[],
|row| row.get::<_, String>(0)
)
.unwrap(),
sqlite_type
);
}
#[rstest]
fn counters_are_atomic_across_concurrent_callers(sandbox: Sandbox) {
let cache = Arc::new(sandbox.cache::<f64>());
let workers = (0..8)
.map(|_| {
let cache = Arc::clone(&cache);
thread::spawn(move || {
for _ in 0..25 {
cache
.increment_cache("counter", 1.0, ExactCacheContext::default())
.unwrap();
}
})
})
.collect::<Vec<_>>();
for worker in workers {
worker.join().unwrap();
}
assert_eq!(
cache
.increment_cache("counter", 0.0, ExactCacheContext::default())
.unwrap(),
200.0
);
}
#[rstest]
fn fractional_then_integer_increment_follows_python_behavior(sandbox: Sandbox) {
let cache = sandbox.cache::<f64>();
assert_eq!(
cache
.increment_cache("counter", 3.5, ExactCacheContext::default())
.unwrap(),
3.5
);
assert_eq!(
cache
.increment_cache("counter", 1.0, ExactCacheContext::default())
.unwrap(),
1.0
);
}
#[rstest]
fn increment_ttl_replacement_clears_expiry_without_ttl(sandbox: Sandbox) {
let cache = sandbox.cache::<f64>();
cache
.increment_cache(
"counter",
1.0,
ExactCacheContext {
ttl: Some(Duration::from_secs(60)),
},
)
.unwrap();
assert!(
sandbox
.db()
.query_row(
"SELECT expire_time IS NOT NULL FROM Cache WHERE key = 'counter'",
[],
|row| row.get::<_, bool>(0)
)
.unwrap()
);
cache
.increment_cache("counter", 1.0, ExactCacheContext::default())
.unwrap();
assert!(
!sandbox
.db()
.query_row(
"SELECT expire_time IS NOT NULL FROM Cache WHERE key = 'counter'",
[],
|row| row.get::<_, bool>(0)
)
.unwrap()
);
}
#[rstest]
fn custom_adapter_controls_storage_and_reads(sandbox: Sandbox) {
let cache = DiskCache::with_adapter(sandbox.store(), TextAdapter, JsonCodec::<Value>::new());
cache
.set_cache("key", json!({"answer": 42}), &ExactCacheContext::default())
.unwrap();
assert!(matches!(
sandbox.store().get("key", 0.0).unwrap(),
Some(StoredValue::Text(_))
));
assert_eq!(
cache
.get_cache("key", &ExactCacheContext::default())
.unwrap(),
Some(json!({"answer": 42}))
);
}
#[rstest]
fn delete_flush_and_spilled_file_replacement_clean_up_storage(sandbox: Sandbox) {
let large = vec![b'x'; 32 * 1024];
sandbox
.store()
.set("large", StoredValue::Bytes(large.clone()), None, 0.0)
.unwrap();
assert_eq!(sandbox.value_files().len(), 1);
sandbox
.store()
.set(
"large",
StoredValue::Bytes(vec![b'y'; 32 * 1024]),
None,
0.0,
)
.unwrap();
assert_eq!(sandbox.value_files().len(), 1);
sandbox.store().pop("large", 0.0).unwrap();
assert!(sandbox.value_files().is_empty());
sandbox
.store()
.set("a", StoredValue::Bytes(large.clone()), None, 0.0)
.unwrap();
sandbox
.store()
.set("b", StoredValue::Bytes(large), None, 0.0)
.unwrap();
sandbox.store().clear().unwrap();
assert!(sandbox.value_files().is_empty());
}
#[rstest]
#[tokio::test]
async fn async_operations_disconnect_and_delete_match_sync_operations(sandbox: Sandbox) {
let cache = sandbox.cache::<Value>();
let context = ExactCacheContext {
ttl: Some(Duration::from_secs(60)),
};
cache
.async_set_cache("a", json!(1), context.clone())
.await
.unwrap();
cache
.async_set_cache_pipeline(
vec![("b".into(), json!(2)), ("c".into(), json!(3))],
context.clone(),
)
.await
.unwrap();
assert_eq!(
cache.async_get_cache("a", &context).await.unwrap(),
Some(json!(1))
);
assert_eq!(
cache
.async_batch_get_cache(vec!["c".into(), "missing".into()], context.clone())
.await
.unwrap(),
vec![BatchEntry::Hit(json!(3)), BatchEntry::Miss]
);
cache.async_delete_cache("a").await.unwrap();
cache.async_flush_cache().await.unwrap();
cache.disconnect().await.unwrap();
}
#[derive(Clone, Copy, Debug)]
enum Increment {
Sync,
Async { refresh_ttl: bool },
}
impl Increment {
async fn apply(
self,
cache: &DiskCache<JsonCodec<Value>>,
key: &str,
amount: f64,
context: ExactCacheContext,
) -> f64 {
match self {
Self::Sync => cache.increment_cache(key, amount, context).unwrap(),
Self::Async { refresh_ttl } => cache
.async_increment(key, amount, context, refresh_ttl)
.await
.unwrap(),
}
}
}
#[rstest]
#[case::sync_missing(Increment::Sync, None, 3.0, 3.0)]
#[case::sync_existing_int(Increment::Sync, Some(json!(7)), 5.0, 12.0)]
#[case::sync_non_int(Increment::Sync, Some(json!("not-a-number")), 4.0, 4.0)]
#[case::async_missing(Increment::Async { refresh_ttl: false }, None, 2.0, 2.0)]
#[case::async_existing_int(Increment::Async { refresh_ttl: false }, Some(json!(10)), 5.0, 15.0)]
#[case::async_non_int(Increment::Async { refresh_ttl: false }, Some(json!("corrupt")), 9.0, 9.0)]
#[case::async_refresh_ttl_is_ignored(Increment::Async { refresh_ttl: true }, Some(json!(1)), 1.0, 2.0)]
#[tokio::test]
async fn increments_read_back_through_get_cache(
sandbox: Sandbox,
#[case] increment: Increment,
#[case] initial: Option<Value>,
#[case] amount: f64,
#[case] expected: f64,
) {
let cache = sandbox.cache::<Value>();
let context = ExactCacheContext::default();
if let Some(initial) = initial {
cache
.async_set_cache("counter", initial, context.clone())
.await
.unwrap();
}
assert_eq!(
increment
.apply(&cache, "counter", amount, context.clone())
.await,
expected
);
assert_eq!(
cache.get_cache("counter", &context).unwrap(),
Some(json!(expected as i64))
);
}
#[rstest]
#[case::without_refresh(false)]
#[case::with_refresh(true)]
#[tokio::test]
async fn async_increment_rewrites_ttl_on_every_write(sandbox: Sandbox, #[case] refresh_ttl: bool) {
let cache = sandbox.cache::<Value>();
let expiry = || {
sandbox
.db()
.query_row(
"SELECT expire_time IS NOT NULL FROM Cache WHERE key = 'counter'",
[],
|row| row.get::<_, bool>(0),
)
.unwrap()
};
let ttl = ExactCacheContext {
ttl: Some(Duration::from_secs(60)),
};
cache
.async_increment("counter", 1.0, ttl.clone(), refresh_ttl)
.await
.unwrap();
assert!(expiry());
cache
.async_increment("counter", 1.0, ExactCacheContext::default(), refresh_ttl)
.await
.unwrap();
assert!(!expiry());
cache
.async_increment("counter", 1.0, ttl, refresh_ttl)
.await
.unwrap();
assert!(expiry());
}

View file

@ -0,0 +1,82 @@
use litellm_cache::{ExactCacheContext, JsonCodec};
use litellm_cache_disk::DiskCache;
use litellm_cache_testing as contract;
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use tempfile::TempDir;
struct Disk {
cache: DiskCache<JsonCodec<Value>>,
_directory: TempDir,
}
#[fixture]
fn disk() -> Disk {
let directory = tempfile::tempdir().unwrap();
Disk {
cache: DiskCache::open(directory.path(), JsonCodec::new()).unwrap(),
_directory: directory,
}
}
#[fixture]
fn context() -> ExactCacheContext {
ExactCacheContext::default()
}
const PREFIX: &str = "contract:";
#[rstest]
#[tokio::test]
async fn hit_and_miss(disk: Disk, context: ExactCacheContext) {
contract::hit_and_miss(&disk.cache, context, PREFIX, json!({"answer": 42})).await;
}
#[rstest]
#[tokio::test]
async fn sync_async_equivalence(disk: Disk, context: ExactCacheContext) {
contract::sync_async_equivalence(&disk.cache, context, PREFIX, json!("first"), json!([2]))
.await;
}
#[rstest]
#[tokio::test]
async fn overwrite_replaces(disk: Disk, context: ExactCacheContext) {
contract::overwrite_replaces(&disk.cache, context, PREFIX, json!(1), json!({"b": 2})).await;
}
#[rstest]
#[tokio::test]
async fn pipeline_writes_every_entry(disk: Disk, context: ExactCacheContext) {
contract::pipeline_writes_every_entry(
&disk.cache,
context,
PREFIX,
vec![json!("a"), json!(2), json!({"c": true})],
)
.await;
}
#[rstest]
#[tokio::test]
async fn batch_preserves_order(disk: Disk, context: ExactCacheContext) {
contract::batch_preserves_order(&disk.cache, context, PREFIX, json!("first"), json!(2)).await;
}
#[rstest]
#[tokio::test]
async fn delete_removes_key(disk: Disk, context: ExactCacheContext) {
contract::delete_removes_key(&disk.cache, context, PREFIX, json!("value")).await;
}
#[rstest]
#[tokio::test]
async fn flush_clears(disk: Disk, context: ExactCacheContext) {
contract::flush_clears(&disk.cache, context, PREFIX, json!("value")).await;
}
#[rstest]
#[tokio::test]
async fn counter_accumulates(disk: Disk, context: ExactCacheContext) {
contract::counter_accumulates(&disk.cache, context, PREFIX).await;
}

View file

@ -0,0 +1,113 @@
use litellm_cache::Error;
use litellm_cache_disk::{PythonDiskCacheAdapter, StoredValue, ValueAdapter};
use rstest::rstest;
enum ReadExpectation {
Bytes(&'static [u8]),
Miss,
Invalid,
}
#[rstest]
#[case::pickled_dictionary_with_string_keys(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x7d, 0x94, 0x8c, 0x01, 0x61, 0x94, 0x4b, 0x01, 0x73, 0x2e]),
ReadExpectation::Bytes(br#"{"a":1}"#)
)]
#[case::pickled_list_of_integers(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x09, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x5d, 0x94, 0x28, 0x4b, 0x01, 0x4b, 0x02, 0x65, 0x2e]),
ReadExpectation::Bytes(br#"[1,2]"#)
)]
#[case::pickled_tuple_of_integers(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x07, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x4b, 0x01, 0x4b, 0x02, 0x86, 0x94, 0x2e]),
ReadExpectation::Bytes(br#"[1,2]"#)
)]
#[case::pickled_set_of_integers(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x09, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x8f, 0x94, 0x28, 0x4b, 0x01, 0x4b, 0x02, 0x90, 0x2e]),
ReadExpectation::Bytes(br#"[1,2]"#)
)]
#[case::pickled_response_envelope(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x30, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x7d, 0x94, 0x28, 0x8c, 0x09, 0x74, 0x69, 0x6d, 0x65, 0x73, 0x74, 0x61, 0x6d, 0x70, 0x94, 0x47, 0x3f, 0xf8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x8c, 0x08, 0x72, 0x65, 0x73, 0x70, 0x6f, 0x6e, 0x73, 0x65, 0x94, 0x8c, 0x08, 0x7b, 0x22, 0x61, 0x22, 0x3a, 0x20, 0x31, 0x7d, 0x94, 0x75, 0x2e]),
ReadExpectation::Bytes(br#"{"response":"{\"a\": 1}","timestamp":1.5}"#)
)]
#[case::non_json_text(
StoredValue::Text("not json".into()),
ReadExpectation::Bytes(b"not json")
)]
#[case::json_text(
StoredValue::Text("{\"a\": 1}".into()),
ReadExpectation::Bytes(br#"{"a": 1}"#)
)]
#[case::non_utf8_bytes(
StoredValue::Bytes(vec![0xff, 0xfe]),
ReadExpectation::Bytes(&[0xff, 0xfe])
)]
#[case::integer_seven(StoredValue::Integer(7), ReadExpectation::Bytes(b"7"))]
#[case::float_one_point_five(StoredValue::Float(1.5), ReadExpectation::Bytes(b"1.5"))]
#[case::pickled_true(
StoredValue::Pickle(vec![0x80, 0x05, 0x88, 0x2e]),
ReadExpectation::Bytes(b"true")
)]
#[case::pickled_negative_integer(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x06, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x4a, 0xfd, 0xff, 0xff, 0xff, 0x2e]),
ReadExpectation::Bytes(b"-3")
)]
#[case::pickled_bytes(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x09, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x43, 0x05, 0x62, 0x79, 0x74, 0x65, 0x73, 0x94, 0x2e]),
ReadExpectation::Invalid
)]
#[case::pickled_dictionary_with_integer_key(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x7d, 0x94, 0x4b, 0x01, 0x8c, 0x01, 0x61, 0x94, 0x73, 0x2e]),
ReadExpectation::Invalid
)]
#[case::pickled_complex(
StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x2e, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x8c, 0x08, 0x62, 0x75, 0x69, 0x6c, 0x74, 0x69, 0x6e, 0x73, 0x94, 0x8c, 0x07, 0x63, 0x6f, 0x6d, 0x70, 0x6c, 0x65, 0x78, 0x94, 0x93, 0x94, 0x47, 0x3f, 0xf0, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x47, 0x40, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x86, 0x94, 0x52, 0x94, 0x2e]),
ReadExpectation::Invalid
)]
#[case::truncated_pickle(
StoredValue::Pickle(vec![0x80, 0x05, 0x2e]),
ReadExpectation::Invalid
)]
#[case::empty_bytes(StoredValue::Bytes(Vec::new()), ReadExpectation::Miss)]
#[case::empty_text(StoredValue::Text(String::new()), ReadExpectation::Miss)]
#[case::zero_integer(StoredValue::Integer(0), ReadExpectation::Miss)]
#[case::zero_float(StoredValue::Float(0.0), ReadExpectation::Miss)]
#[case::pickled_none(StoredValue::Pickle(vec![0x80, 0x05, 0x4e, 0x2e]), ReadExpectation::Miss)]
#[case::pickled_false(StoredValue::Pickle(vec![0x80, 0x05, 0x89, 0x2e]), ReadExpectation::Miss)]
#[case::pickled_zero(StoredValue::Pickle(vec![0x80, 0x05, 0x4b, 0x00, 0x2e]), ReadExpectation::Miss)]
#[case::pickled_zero_float(StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x47, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x2e]), ReadExpectation::Miss)]
#[case::pickled_empty_string(StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x04, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x8c, 0x00, 0x94, 0x2e]), ReadExpectation::Miss)]
#[case::pickled_empty_list(StoredValue::Pickle(vec![0x80, 0x05, 0x5d, 0x94, 0x2e]), ReadExpectation::Miss)]
#[case::pickled_empty_dictionary(StoredValue::Pickle(vec![0x80, 0x05, 0x7d, 0x94, 0x2e]), ReadExpectation::Miss)]
#[case::pickled_empty_tuple(StoredValue::Pickle(vec![0x80, 0x05, 0x29, 0x2e]), ReadExpectation::Miss)]
fn python_read_cases(#[case] row: StoredValue, #[case] expected: ReadExpectation) {
let result = PythonDiskCacheAdapter.read(row);
match expected {
ReadExpectation::Bytes(expected) => assert_eq!(result.unwrap().unwrap(), expected),
ReadExpectation::Miss => assert_eq!(result.unwrap(), None),
ReadExpectation::Invalid => assert!(matches!(result, Err(Error::InvalidEntry))),
}
}
#[rstest]
#[case::integer_two(Some(StoredValue::Integer(2)), 2.0)]
#[case::float_three_point_five(Some(StoredValue::Float(3.5)), 0.0)]
#[case::text_not_a_number(Some(StoredValue::Text("not a number".into())), 0.0)]
#[case::text_five(Some(StoredValue::Text("5".into())), 5.0)]
#[case::text_three_point_five(Some(StoredValue::Text("3.5".into())), 0.0)]
#[case::pickled_true(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x88, 0x2e])), 1.0)]
#[case::pickled_two(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x4b, 0x02, 0x2e])), 2.0)]
#[case::pickled_dictionary(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x95, 0x0a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x7d, 0x94, 0x8c, 0x01, 0x61, 0x94, 0x4b, 0x01, 0x73, 0x2e])), 0.0)]
#[case::missing(None, 0.0)]
#[case::pickled_none(Some(StoredValue::Pickle(vec![0x80, 0x05, 0x4e, 0x2e])), 0.0)]
fn python_counter_seed_cases(#[case] row: Option<StoredValue>, #[case] expected: f64) {
assert_eq!(PythonDiskCacheAdapter.counter_seed(row).unwrap(), expected);
}
#[rstest]
#[case::integer_three(3.0, StoredValue::Integer(3))]
#[case::fractional_three_point_five(3.5, StoredValue::Float(3.5))]
#[case::negative_zero(-0.0, StoredValue::Integer(0))]
#[case::large_float(1e300, StoredValue::Float(1e300))]
fn python_counter_value_cases(#[case] value: f64, #[case] expected: StoredValue) {
assert_eq!(PythonDiskCacheAdapter.counter_value(value), expected);
}

View file

@ -0,0 +1,22 @@
[package]
name = "litellm-cache-gcs"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
futures-util.workspace = true
litellm-auth-gcp.workspace = true
litellm-auth-types.workspace = true
litellm-cache.workspace = true
percent-encoding.workspace = true
reqwest.workspace = true
tokio.workspace = true
[dev-dependencies]
litellm-cache-testing.workspace = true
rstest.workspace = true
serde_json.workspace = true
tokio.workspace = true
wiremock = "0.6.5"

View file

@ -0,0 +1,258 @@
use std::{future::Future, sync::Arc, time::Duration};
use futures_util::future::try_join_all;
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheCodec, DisconnectCache, Error, ExactCacheContext,
FlushCache,
};
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode};
use reqwest::Client;
use crate::{GcpTokenSource, TokenSource};
pub const DEFAULT_ENDPOINT: &str = "https://storage.googleapis.com";
const OBJECT_NAME_ENCODE_SET: &AsciiSet = &NON_ALPHANUMERIC
.remove(b'-')
.remove(b'_')
.remove(b'.')
.remove(b'~');
pub fn key_prefix(gcs_path: Option<&str>) -> String {
match gcs_path {
Some(path) if !path.is_empty() => format!("{}/", path.trim_end_matches('/')),
_ => String::new(),
}
}
#[derive(Clone, Debug)]
pub struct GcsConfig {
pub bucket_name: String,
pub gcs_path: Option<String>,
pub path_service_account: Option<String>,
pub endpoint: String,
}
impl GcsConfig {
pub fn new(bucket_name: impl Into<String>) -> Self {
Self {
bucket_name: bucket_name.into(),
gcs_path: None,
path_service_account: None,
endpoint: DEFAULT_ENDPOINT.to_string(),
}
}
}
pub struct GcsCache<S: CacheCodec> {
config: GcsConfig,
key_prefix: String,
client: Client,
token: Arc<dyn TokenSource>,
codec: S,
}
impl<S: CacheCodec> GcsCache<S> {
pub fn new(config: GcsConfig, client: Client, codec: S) -> Self {
let token = Arc::new(GcpTokenSource::new(config.path_service_account.clone()));
Self::with_token_source(config, client, codec, token)
}
pub fn with_token_source(
config: GcsConfig,
client: Client,
codec: S,
token: Arc<dyn TokenSource>,
) -> Self {
let key_prefix = key_prefix(config.gcs_path.as_deref());
Self {
config,
key_prefix,
client,
token,
codec,
}
}
pub fn bucket_name(&self) -> &str {
&self.config.bucket_name
}
pub fn key_prefix(&self) -> &str {
&self.key_prefix
}
pub fn path_service_account(&self) -> Option<&str> {
self.config.path_service_account.as_deref()
}
pub fn object_name(&self, key: &str) -> String {
format!("{}{}", self.key_prefix, key)
}
fn encoded_object_name(&self, key: &str) -> String {
percent_encode(self.object_name(key).as_bytes(), OBJECT_NAME_ENCODE_SET).to_string()
}
fn endpoint(&self, path: &str) -> String {
format!("{}{}", self.config.endpoint.trim_end_matches('/'), path)
}
async fn async_set(&self, key: &str, value: S::Value) -> Result<(), Error> {
let token = self.token.bearer_token().await?;
let payload = self.codec.encode(&value)?;
let url = self.endpoint(&format!(
"/upload/storage/v1/b/{}/o?uploadType=media&name={}",
self.config.bucket_name,
self.encoded_object_name(key)
));
let response = self
.client
.post(url)
.bearer_auth(token)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.body(payload)
.send()
.await
.map_err(|_| Error::Unavailable)?;
if !response.status().is_success() {
return Err(Error::Unavailable);
}
Ok(())
}
async fn async_get(&self, key: &str) -> Result<Option<S::Value>, Error> {
let token = self.token.bearer_token().await?;
let url = self.endpoint(&format!(
"/storage/v1/b/{}/o/{}?alt=media",
self.config.bucket_name,
self.encoded_object_name(key)
));
let response = self
.client
.get(url)
.bearer_auth(token)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.send()
.await
.map_err(|_| Error::Unavailable)?;
if response.status() == reqwest::StatusCode::NOT_FOUND {
return Ok(None);
}
if !response.status().is_success() {
return Err(Error::Unavailable);
}
let body = response.bytes().await.map_err(|_| Error::Unavailable)?;
self.codec
.decode(&body)
.map(Some)
.map_err(|_| Error::InvalidEntry)
}
fn run_sync<T, F>(future: F) -> Result<T, Error>
where
F: Future<Output = Result<T, Error>> + Send,
T: Send,
{
let run = |future: F| {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|_| Error::Unavailable)
.and_then(|runtime| runtime.block_on(future))
};
match tokio::runtime::Handle::try_current() {
Ok(handle) if handle.runtime_flavor() == tokio::runtime::RuntimeFlavor::MultiThread => {
tokio::task::block_in_place(|| handle.block_on(future))
}
Ok(_) => std::thread::scope(|scope| {
scope
.spawn(|| run(future))
.join()
.map_err(|_| Error::Unavailable)
.and_then(|result| result)
}),
Err(_) => run(future),
}
}
}
impl<S: CacheCodec> BaseCache for GcsCache<S> {
type Value = S::Value;
type Context = ExactCacheContext;
fn get_ttl(&self, _: &Self::Context) -> Option<Duration> {
None
}
fn set_cache(&self, key: &str, value: Self::Value, _: &Self::Context) -> Result<(), Error> {
Self::run_sync(self.async_set(key, value))
}
fn get_cache(&self, key: &str, _: &Self::Context) -> Result<Option<Self::Value>, Error> {
Self::run_sync(self.async_get(key))
}
async fn async_set_cache(
&self,
key: &str,
value: Self::Value,
_: Self::Context,
) -> Result<(), Error> {
self.async_set(key, value).await
}
async fn async_get_cache(
&self,
key: &str,
_: &Self::Context,
) -> Result<Option<Self::Value>, Error> {
self.async_get(key).await
}
async fn async_set_cache_pipeline(
&self,
entries: Vec<(String, Self::Value)>,
context: Self::Context,
) -> Result<(), Error> {
try_join_all(entries.into_iter().map(|(key, value)| {
let context = context.clone();
async move { self.async_set_cache(&key, value, context).await }
}))
.await
.map(|_| ())
}
}
impl<S: CacheCodec> DisconnectCache for GcsCache<S> {
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
}
impl<S: CacheCodec> BatchCache for GcsCache<S> {
async fn async_batch_get_cache(
&self,
keys: Vec<String>,
context: Self::Context,
) -> Result<Vec<BatchEntry<Self::Value>>, Error> {
try_join_all(keys.into_iter().map(|key| {
let context = context.clone();
async move {
match self.async_get_cache(&key, &context).await {
Ok(Some(value)) => Ok(BatchEntry::Hit(value)),
Ok(None) => Ok(BatchEntry::Miss),
Err(Error::InvalidEntry) => Ok(BatchEntry::Invalid),
Err(error) => Err(error),
}
}
}))
.await
}
}
impl<S: CacheCodec> FlushCache for GcsCache<S> {
fn flush_cache(&self) -> Result<(), Error> {
Ok(())
}
}

View file

@ -0,0 +1,5 @@
mod cache;
mod token;
pub use cache::{DEFAULT_ENDPOINT, GcsCache, GcsConfig, key_prefix};
pub use token::{GcpTokenSource, StaticTokenSource, TokenSource};

View file

@ -0,0 +1,44 @@
use std::{future::Future, pin::Pin};
use litellm_auth_gcp::{VertexAuth, VertexConfig};
use litellm_auth_types::{InputSource, SecretValue, Sourced};
use litellm_cache::Error;
pub trait TokenSource: Send + Sync + 'static {
fn bearer_token(&self) -> Pin<Box<dyn Future<Output = Result<String, Error>> + Send + '_>>;
}
pub struct GcpTokenSource {
auth: VertexAuth,
config: VertexConfig,
}
impl GcpTokenSource {
pub fn new(path_service_account: Option<String>) -> Self {
let credentials = path_service_account
.map(|path| Sourced::new(SecretValue::new(path), InputSource::Deployment));
Self {
auth: VertexAuth::default(),
config: VertexConfig::new(credentials, None, None),
}
}
}
impl TokenSource for GcpTokenSource {
fn bearer_token(&self) -> Pin<Box<dyn Future<Output = Result<String, Error>> + Send + '_>> {
Box::pin(async move {
self.auth
.access_token(&self.config, &|name| std::env::var(name).ok())
.await
.map_err(|_| Error::Unavailable)
})
}
}
pub struct StaticTokenSource(pub String);
impl TokenSource for StaticTokenSource {
fn bearer_token(&self) -> Pin<Box<dyn Future<Output = Result<String, Error>> + Send + '_>> {
Box::pin(async move { Ok(self.0.clone()) })
}
}

View file

@ -0,0 +1,279 @@
mod support;
use std::{future::Future, pin::Pin, sync::Arc, time::Duration};
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheContext, DisconnectCache, Error, ExactCacheContext,
FlushCache,
};
use litellm_cache_gcs::{GcsCache, GcsConfig, TokenSource, key_prefix};
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use support::FakeBucket;
use wiremock::{
Mock, MockServer, ResponseTemplate,
matchers::{body_bytes, header, method, path, query_param},
};
#[fixture]
async fn server() -> MockServer {
MockServer::start().await
}
fn context() -> ExactCacheContext {
ExactCacheContext::default()
}
#[rstest]
#[tokio::test]
async fn set_writes_encoded_object_and_headers(#[future(awt)] server: MockServer) {
Mock::given(method("POST"))
.and(path("/upload/storage/v1/b/bucket/o"))
.and(query_param("uploadType", "media"))
.and(header("authorization", "Bearer tok"))
.and(header("content-type", "application/json"))
.and(body_bytes(br#"{"value":"entry"}"#))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
support::cache(&server, Some("cache/"))
.set_cache("team:a b/c", json!({"value": "entry"}), &context())
.unwrap();
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].url.query(),
Some("uploadType=media&name=cache%2Fteam%3Aa%20b%2Fc")
);
}
#[rstest]
#[case::hit(
"hit",
ResponseTemplate::new(200).set_body_json(json!({"value": "entry"})),
Ok(Some(json!({"value": "entry"})))
)]
#[case::missing("missing", ResponseTemplate::new(404), Ok(None))]
#[case::server_error("server-error", ResponseTemplate::new(500), Err(Error::Unavailable))]
#[case::invalid(
"invalid",
ResponseTemplate::new(200).set_body_string("not json"),
Err(Error::InvalidEntry)
)]
#[tokio::test]
async fn get_maps_statuses_and_decode_failures(
#[future(awt)] server: MockServer,
#[case] key: &str,
#[case] response: ResponseTemplate,
#[case] expected: Result<Option<Value>, Error>,
) {
Mock::given(method("GET"))
.and(path(format!("/storage/v1/b/bucket/o/{key}")))
.and(query_param("alt", "media"))
.respond_with(response)
.mount(&server)
.await;
let cache = support::cache(&server, None);
assert_eq!(cache.get_cache(key, &context()), expected);
assert_eq!(cache.async_get_cache(key, &context()).await, expected);
}
#[rstest]
#[case::none(None, "")]
#[case::trailing_slash(Some("a/b/"), "a/b/")]
#[case::no_trailing_slash(Some("a/b"), "a/b/")]
#[case::empty(Some(""), "")]
fn key_prefix_normalizes_paths(#[case] gcs_path: Option<&str>, #[case] expected: &str) {
assert_eq!(key_prefix(gcs_path), expected);
}
#[rstest]
#[tokio::test]
async fn cache_exposes_its_configuration(#[future(awt)] server: MockServer) {
let cache = GcsCache::new(
GcsConfig {
path_service_account: Some("/secrets/sa.json".into()),
..support::config(&server, Some("folder"))
},
reqwest::Client::new(),
litellm_cache::JsonCodec::<Value>::new(),
);
assert_eq!(cache.bucket_name(), "bucket");
assert_eq!(cache.key_prefix(), "folder/");
assert_eq!(cache.path_service_account(), Some("/secrets/sa.json"));
assert_eq!(cache.object_name("k"), "folder/k");
}
#[rstest]
#[case::punctuation("a~b-c_d.e/f g%h", "uploadType=media&name=p%2Fa~b-c_d.e%2Ff%20g%25h")]
#[case::utf8("ключ", "uploadType=media&name=p%2F%D0%BA%D0%BB%D1%8E%D1%87")]
#[tokio::test]
async fn object_names_use_python_quote_encoding(
#[future(awt)] server: MockServer,
#[case] key: &str,
#[case] query: &str,
) {
Mock::given(method("POST"))
.and(path("/upload/storage/v1/b/bucket/o"))
.and(query_param("uploadType", "media"))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
support::cache(&server, Some("p/"))
.async_set_cache(key, json!({"value": key}), context())
.await
.unwrap();
let requests = server.received_requests().await.unwrap();
assert_eq!(requests[0].url.query(), Some(query));
}
#[rstest]
#[tokio::test]
async fn object_names_are_encoded_in_the_download_path(#[future(awt)] server: MockServer) {
Mock::given(method("GET"))
.and(path("/storage/v1/b/bucket/o/p%2Fa%3Ab%20c"))
.and(query_param("alt", "media"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!(1)))
.expect(2)
.mount(&server)
.await;
let cache = support::cache(&server, Some("p"));
assert_eq!(cache.get_cache("a:b c", &context()), Ok(Some(json!(1))));
assert_eq!(
cache.async_get_cache("a:b c", &context()).await,
Ok(Some(json!(1)))
);
}
#[rstest]
#[tokio::test]
async fn ignores_ttl_and_writes_pipeline_concurrently(#[future(awt)] server: MockServer) {
for key in ["one", "two", "three"] {
Mock::given(method("POST"))
.and(path("/upload/storage/v1/b/bucket/o"))
.and(query_param("uploadType", "media"))
.and(query_param("name", key))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&server)
.await;
}
let cache = support::cache(&server, None);
let with_ttl = context().with_ttl(Some(Duration::from_secs(5)));
assert_eq!(cache.get_ttl(&context()), None);
assert_eq!(cache.get_ttl(&with_ttl), None);
cache
.async_set_cache_pipeline(
vec![
("one".into(), json!({"key": "one"})),
("two".into(), json!({"key": "two"})),
("three".into(), json!({"key": "three"})),
],
with_ttl,
)
.await
.unwrap();
}
#[rstest]
#[tokio::test]
async fn batch_get_preserves_hits_misses_and_invalid_entries(#[future(awt)] server: MockServer) {
Mock::given(method("GET"))
.and(path("/storage/v1/b/bucket/o/hit"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": "entry"})))
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/storage/v1/b/bucket/o/missing"))
.respond_with(ResponseTemplate::new(404))
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/storage/v1/b/bucket/o/invalid"))
.respond_with(ResponseTemplate::new(200).set_body_string("not json"))
.mount(&server)
.await;
let cache = support::cache(&server, None);
let keys = vec!["hit".to_string(), "missing".into(), "invalid".into()];
let expected = vec![
BatchEntry::Hit(json!({"value": "entry"})),
BatchEntry::Miss,
BatchEntry::Invalid,
];
assert_eq!(cache.batch_get_cache(&keys, &context()).unwrap(), expected);
assert_eq!(
cache.async_batch_get_cache(keys, context()).await.unwrap(),
expected
);
}
#[rstest]
#[tokio::test]
async fn flush_and_disconnect_are_noops_like_python(#[future(awt)] server: MockServer) {
let cache = support::cache(&server, None);
assert_eq!(cache.flush_cache(), Ok(()));
assert_eq!(cache.async_flush_cache().await, Ok(()));
assert_eq!(cache.disconnect().await, Ok(()));
assert!(server.received_requests().await.unwrap().is_empty());
}
fn round_trip(cache: &support::JsonGcsCache) -> Result<Option<Value>, Error> {
cache.set_cache("key", json!({"value": "entry"}), &context())?;
cache.get_cache("key", &context())
}
#[rstest]
fn sync_operations_work_without_an_active_runtime() {
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.unwrap();
let server = runtime.block_on(FakeBucket::serve());
let cache = support::cache(&server, None);
assert_eq!(round_trip(&cache), Ok(Some(json!({"value": "entry"}))));
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn sync_operations_work_inside_a_multi_thread_runtime() {
let server = FakeBucket::serve().await;
let cache = support::cache(&server, None);
assert_eq!(round_trip(&cache), Ok(Some(json!({"value": "entry"}))));
}
#[rstest]
#[tokio::test]
async fn sync_operations_work_inside_a_current_thread_runtime() {
let server = FakeBucket::serve().await;
let cache = support::cache(&server, None);
assert_eq!(round_trip(&cache), Ok(Some(json!({"value": "entry"}))));
}
struct FailingTokenSource;
impl TokenSource for FailingTokenSource {
fn bearer_token(&self) -> Pin<Box<dyn Future<Output = Result<String, Error>> + Send + '_>> {
Box::pin(async { Err(Error::Unavailable) })
}
}
#[rstest]
#[tokio::test]
async fn token_source_failure_skips_http(#[future(awt)] server: MockServer) {
let cache = support::cache_with_token(&server, None, Arc::new(FailingTokenSource));
assert_eq!(
cache.get_cache("key", &context()).unwrap_err(),
Error::Unavailable
);
assert_eq!(
cache
.async_set_cache("key", json!(1), context())
.await
.unwrap_err(),
Error::Unavailable
);
assert_eq!(server.received_requests().await.unwrap().len(), 0);
}

View file

@ -0,0 +1,65 @@
mod support;
use litellm_cache::ExactCacheContext;
use litellm_cache_testing as contract;
use rstest::{fixture, rstest};
use serde_json::json;
use support::{FakeBucket, JsonGcsCache};
use wiremock::MockServer;
struct Gcs {
cache: JsonGcsCache,
_server: MockServer,
}
#[fixture]
async fn gcs() -> Gcs {
let server = FakeBucket::serve().await;
Gcs {
cache: support::cache(&server, Some("contract")),
_server: server,
}
}
#[fixture]
fn context() -> ExactCacheContext {
ExactCacheContext::default()
}
const PREFIX: &str = "contract:";
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn hit_and_miss(#[future(awt)] gcs: Gcs, context: ExactCacheContext) {
contract::hit_and_miss(&gcs.cache, context, PREFIX, json!({"answer": 42})).await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn sync_async_equivalence(#[future(awt)] gcs: Gcs, context: ExactCacheContext) {
contract::sync_async_equivalence(&gcs.cache, context, PREFIX, json!("first"), json!([2])).await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn overwrite_replaces(#[future(awt)] gcs: Gcs, context: ExactCacheContext) {
contract::overwrite_replaces(&gcs.cache, context, PREFIX, json!(1), json!({"b": 2})).await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn pipeline_writes_every_entry(#[future(awt)] gcs: Gcs, context: ExactCacheContext) {
contract::pipeline_writes_every_entry(
&gcs.cache,
context,
PREFIX,
vec![json!("a"), json!(2), json!({"c": true})],
)
.await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn batch_preserves_order(#[future(awt)] gcs: Gcs, context: ExactCacheContext) {
contract::batch_preserves_order(&gcs.cache, context, PREFIX, json!("first"), json!(2)).await;
}

View file

@ -0,0 +1,87 @@
#![allow(dead_code)]
use std::{
collections::HashMap,
sync::{Arc, Mutex},
};
use litellm_cache::JsonCodec;
use litellm_cache_gcs::{GcsCache, GcsConfig, StaticTokenSource, TokenSource};
use percent_encoding::percent_decode_str;
use serde_json::Value;
use wiremock::{Mock, MockServer, Request, Respond, ResponseTemplate, http::Method, matchers::any};
pub type JsonGcsCache = GcsCache<JsonCodec<Value>>;
pub fn config(server: &MockServer, gcs_path: Option<&str>) -> GcsConfig {
GcsConfig {
bucket_name: "bucket".into(),
gcs_path: gcs_path.map(str::to_string),
path_service_account: None,
endpoint: server.uri(),
}
}
pub fn cache_with_token(
server: &MockServer,
gcs_path: Option<&str>,
token: Arc<dyn TokenSource>,
) -> JsonGcsCache {
GcsCache::with_token_source(
config(server, gcs_path),
reqwest::Client::new(),
JsonCodec::new(),
token,
)
}
pub fn cache(server: &MockServer, gcs_path: Option<&str>) -> JsonGcsCache {
cache_with_token(server, gcs_path, Arc::new(StaticTokenSource("tok".into())))
}
/// An in-memory bucket speaking the JSON API's media upload and `alt=media` download.
#[derive(Clone, Default)]
pub struct FakeBucket {
objects: Arc<Mutex<HashMap<String, Vec<u8>>>>,
}
impl FakeBucket {
pub async fn serve() -> MockServer {
let server = MockServer::start().await;
Mock::given(any())
.respond_with(Self::default())
.mount(&server)
.await;
server
}
}
impl Respond for FakeBucket {
fn respond(&self, request: &Request) -> ResponseTemplate {
let mut objects = self.objects.lock().unwrap();
match request.method {
Method::POST => {
let name = request
.url
.query_pairs()
.find_map(|(key, value)| (key == "name").then(|| value.into_owned()))
.expect("uploads carry the object name");
objects.insert(name, request.body.clone());
ResponseTemplate::new(200)
}
Method::GET => {
let encoded = request
.url
.path()
.strip_prefix("/storage/v1/b/bucket/o/")
.expect("downloads address an object");
let name = percent_decode_str(encoded).decode_utf8().unwrap();
match objects.get(name.as_ref()) {
Some(body) => ResponseTemplate::new(200).set_body_bytes(body.clone()),
None => ResponseTemplate::new(404),
}
}
_ => ResponseTemplate::new(405),
}
}
}

View file

@ -7,8 +7,8 @@ repository.workspace = true
[dependencies]
litellm-cache.workspace = true
serde_json.workspace = true
[dev-dependencies]
litellm-cache-testing.workspace = true
rstest.workspace = true
tokio.workspace = true

View file

@ -1,18 +1,20 @@
use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap};
use std::sync::{Arc, Mutex};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use std::{
cmp::Reverse,
collections::{BinaryHeap, HashMap, HashSet},
hash::Hash,
sync::{Arc, Mutex},
time::{Duration, SystemTime, UNIX_EPOCH},
};
use litellm_cache::{
BaseCache, CacheConnectionResult, CacheConnectionStatus, CacheEntry, CacheFuture, CacheKwargs,
Error,
BaseCache, BatchCache, ClaimCache, CounterCache, DeleteCache, DisconnectCache, Error,
ExactCacheContext, FlushCache, SetCache, TtlCache,
};
const DEFAULT_MAX_SIZE_IN_MEMORY: usize = 200;
const DEFAULT_TTL: Duration = Duration::from_secs(600);
type ValueMeasure<V> = Arc<dyn Fn(&V) -> Result<usize, Error> + Send + Sync>;
type ValueValidator<V> = Arc<dyn Fn(&V) -> Result<(), Error> + Send + Sync>;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CacheWrite {
@ -33,7 +35,6 @@ pub struct InMemoryCache<V: Clone> {
default_ttl: Duration,
max_entry_bytes: Option<usize>,
measure_value: Option<ValueMeasure<V>>,
validate_value: Option<ValueValidator<V>>,
now: Arc<dyn Fn() -> Duration + Send + Sync>,
}
@ -74,10 +75,11 @@ impl<V: Clone> InMemoryCache<V> {
expiration_heap: BinaryHeap::new(),
}),
max_size_in_memory: max_size_in_memory.unwrap_or(DEFAULT_MAX_SIZE_IN_MEMORY),
default_ttl: default_ttl.unwrap_or(DEFAULT_TTL),
default_ttl: default_ttl
.filter(|ttl| !ttl.is_zero())
.unwrap_or(DEFAULT_TTL),
max_entry_bytes,
measure_value,
validate_value: None,
now: Arc::new(now),
}
}
@ -91,26 +93,9 @@ impl<V: Clone> InMemoryCache<V> {
if self.max_size_in_memory == 0 {
return Ok(CacheWrite::Disabled);
}
if let Some(validate) = &self.validate_value {
validate(&value)?;
}
if let (Some(limit), Some(measure)) = (self.max_entry_bytes, &self.measure_value)
&& measure(&value)? > limit
{
return Ok(CacheWrite::TooLarge);
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::evict(&mut state, self.max_size_in_memory, now);
let key = key.into();
state.values.insert(key.clone(), value);
let expiration = state.expirations.get(&key).copied();
if expiration.is_none_or(|expiration| expiration < now) {
let expiration = now + ttl.unwrap_or(self.default_ttl);
state.expirations.insert(key.clone(), expiration);
state.expiration_heap.push(Reverse((expiration, key)));
}
Ok(CacheWrite::Stored)
self.store(&mut state, key.into(), value, ttl, now)
}
pub fn get_cache(&self, key: &str) -> Result<Option<V>, Error> {
@ -126,6 +111,78 @@ impl<V: Clone> InMemoryCache<V> {
Ok(state.values.get(key).cloned())
}
/// `check_value_size`: whether `value` fits `max_entry_bytes`. Always `true` without a
/// limit and a measure, since typed values have no generic size.
pub fn check_value_size(&self, value: &V) -> Result<bool, Error> {
match (self.max_entry_bytes, &self.measure_value) {
(Some(limit), Some(measure)) => Ok(measure(value)? <= limit),
_ => Ok(true),
}
}
/// `evict_cache`: drops expired entries, then the earliest-expiring ones until a new key
/// fits.
pub fn evict_cache(&self) -> Result<(), Error> {
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::evict(&mut state, self.max_size_in_memory, now, None);
Ok(())
}
/// `evict_element_if_expired`: `true` when `key` had expired and was removed.
pub fn evict_element_if_expired(&self, key: &str) -> Result<bool, Error> {
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
let expired = state
.expirations
.get(key)
.is_some_and(|expiration| *expiration < now);
if expired {
Self::remove(&mut state, key);
}
Ok(expired)
}
/// `allow_ttl_override`: a write may set the TTL when the key has none or it has passed.
pub fn allow_ttl_override(&self, key: &str) -> Result<bool, Error> {
let now = (self.now)();
Ok(self
.expires_at(key)?
.is_none_or(|expiration| expiration < now))
}
/// The number of stored entries, expired ones included until they are evicted.
pub fn len(&self) -> Result<usize, Error> {
Ok(self
.state
.lock()
.map_err(|_| Error::Unavailable)?
.values
.len())
}
pub fn is_empty(&self) -> Result<bool, Error> {
Ok(self.len()? == 0)
}
/// Entries in the expiration heap, stale ones included; bounded by eviction.
pub fn expiration_heap_len(&self) -> Result<usize, Error> {
Ok(self
.state
.lock()
.map_err(|_| Error::Unavailable)?
.expiration_heap
.len())
}
pub fn max_size_in_memory(&self) -> usize {
self.max_size_in_memory
}
pub fn max_entry_bytes(&self) -> Option<usize> {
self.max_entry_bytes
}
pub fn expires_at(&self, key: &str) -> Result<Option<Duration>, Error> {
Ok(self
.state
@ -136,6 +193,25 @@ impl<V: Clone> InMemoryCache<V> {
.copied())
}
pub async fn async_get_ttl(&self, key: &str) -> Result<Option<Duration>, Error> {
self.expires_at(key)
}
pub async fn async_get_oldest_n_keys(&self, count: usize) -> Result<Vec<String>, Error> {
let state = self.state.lock().map_err(|_| Error::Unavailable)?;
let mut expirations = state
.expirations
.iter()
.map(|(key, expiration)| (key.clone(), *expiration))
.collect::<Vec<_>>();
expirations.sort_unstable_by_key(|(_, expiration)| *expiration);
Ok(expirations
.into_iter()
.take(count)
.map(|(key, _)| key)
.collect())
}
pub fn delete_cache(&self, key: &str) -> Result<(), Error> {
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::remove(&mut state, key);
@ -150,7 +226,9 @@ impl<V: Clone> InMemoryCache<V> {
Ok(())
}
fn evict(state: &mut CacheState<V>, capacity: usize, now: Duration) {
/// Writing an existing `key` never evicts another entry, unlike Python, which pops the
/// earliest-expiring entry whenever the cache is full.
fn evict(state: &mut CacheState<V>, capacity: usize, now: Duration, key: Option<&str>) {
while let Some(Reverse((expiration, key))) = state.expiration_heap.peek().cloned() {
if state.expirations.get(&key).copied() != Some(expiration) {
state.expiration_heap.pop();
@ -161,6 +239,9 @@ impl<V: Clone> InMemoryCache<V> {
break;
}
}
if key.is_some_and(|key| state.values.contains_key(key)) {
return;
}
while state.values.len() >= capacity {
let Some(Reverse((expiration, key))) = state.expiration_heap.pop() else {
break;
@ -171,84 +252,182 @@ impl<V: Clone> InMemoryCache<V> {
}
}
fn set_expiration(state: &mut CacheState<V>, key: &str, expiration: Duration) {
if state.expirations.get(key).copied() != Some(expiration) {
state.expirations.insert(key.into(), expiration);
state
.expiration_heap
.push(Reverse((expiration, key.into())));
}
}
fn remove(state: &mut CacheState<V>, key: &str) {
state.values.remove(key);
state.expirations.remove(key);
}
}
impl InMemoryCache<CacheEntry> {
pub fn response_cache(capacity: usize, ttl: Duration, max_entry_bytes: usize) -> Self {
Self::response_cache_with_clock(capacity, ttl, max_entry_bytes, || {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
})
/// `get_cache` under the held lock: an expired entry is removed and reads as missing.
fn live(state: &mut CacheState<V>, key: &str, now: Duration) -> Option<V> {
if state
.expirations
.get(key)
.is_some_and(|expiration| *expiration < now)
{
Self::remove(state, key);
}
state.values.get(key).cloned()
}
pub fn response_cache_with_clock(
capacity: usize,
ttl: Duration,
max_entry_bytes: usize,
now: impl Fn() -> Duration + Send + Sync + 'static,
) -> Self {
let mut cache = Self::with_clock_and_size_measurement(
Some(capacity),
Some(ttl),
Some(max_entry_bytes),
Some(Arc::new(|entry: &CacheEntry| {
serde_json::to_vec(entry)
.map(|bytes| bytes.len())
.map_err(|_| Error::InvalidEntry)
})),
now,
/// Python `set_cache` under the held lock: evict first (even when `key` already exists),
/// then skip oversized values, then write, keeping a live key's expiry.
fn store(
&self,
state: &mut CacheState<V>,
key: String,
value: V,
ttl: Option<Duration>,
now: Duration,
) -> Result<CacheWrite, Error> {
Self::evict(state, self.max_size_in_memory, now, None);
if !self.check_value_size(&value)? {
return Ok(CacheWrite::TooLarge);
}
let expiration = state.expirations.get(&key).copied();
if expiration.is_none_or(|expiration| expiration < now) {
Self::set_expiration(state, &key, now + ttl.unwrap_or(self.default_ttl));
}
state.values.insert(key, value);
Ok(CacheWrite::Stored)
}
}
impl<V> ClaimCache for InMemoryCache<V>
where
V: Clone + PartialEq + Send + Sync + 'static,
{
fn claim_cache(
&self,
key: &str,
candidate: V,
eligible: &[V],
context: ExactCacheContext,
) -> Result<V, Error> {
if self.max_size_in_memory == 0 {
return Ok(candidate);
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
Self::evict(&mut state, self.max_size_in_memory, now, Some(key));
let existing = state
.values
.get(key)
.filter(|existing| eligible.is_empty() || eligible.contains(existing))
.cloned();
if let Some(existing) = &existing
&& eligible.is_empty()
&& *existing != candidate
{
return Ok(existing.clone());
}
let winner = existing.unwrap_or(candidate);
Self::set_expiration(
&mut state,
key,
now + self.get_ttl(&context).unwrap_or(self.default_ttl),
);
cache.validate_value = Some(Arc::new(|entry: &CacheEntry| {
entry
.timestamp
.is_finite()
.then_some(())
.ok_or(Error::InvalidEntry)
}));
cache
state.values.insert(key.into(), winner.clone());
Ok(winner)
}
}
impl BaseCache for InMemoryCache<CacheEntry> {
type Value = CacheEntry;
impl CounterCache for InMemoryCache<f64> {
fn increment_cache(
&self,
key: &str,
amount: f64,
context: ExactCacheContext,
) -> Result<f64, Error> {
if self.max_size_in_memory == 0 {
return Ok(amount);
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
let value = Self::live(&mut state, key, now).unwrap_or_default() + amount;
self.store(&mut state, key.into(), value, self.get_ttl(&context), now)?;
Ok(value)
}
}
fn default_ttl(&self) -> Duration {
self.default_ttl
impl<V: Clone + Send + Sync + 'static> BaseCache for InMemoryCache<V> {
type Value = V;
type Context = ExactCacheContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl.or(Some(self.default_ttl))
}
fn set_cache(&self, key: &str, value: Self::Value, kwargs: CacheKwargs) -> Result<(), Error> {
let ttl = self.get_ttl(&kwargs);
fn set_cache(
&self,
key: &str,
value: Self::Value,
context: &ExactCacheContext,
) -> Result<(), Error> {
let ttl = self.get_ttl(context).unwrap_or(self.default_ttl);
self.set_cache(key, value, Some(ttl)).map(|_| ())
}
fn get_cache(&self, key: &str, _: &CacheKwargs) -> Result<Option<Self::Value>, Error> {
fn get_cache(&self, key: &str, _: &ExactCacheContext) -> Result<Option<Self::Value>, Error> {
self.get_cache(key)
}
}
fn delete_cache(&self, key: &str) -> Result<(), Error> {
self.delete_cache(key)
}
fn flush_cache(&self) -> Result<(), Error> {
self.flush_cache()
}
fn disconnect(&self) -> CacheFuture<'_, ()> {
Box::pin(async { Ok(()) })
}
fn test_connection(&self) -> CacheFuture<'_, CacheConnectionResult> {
Box::pin(async {
Ok(CacheConnectionResult {
status: CacheConnectionStatus::Success,
message: "In-memory cache connection test successful".into(),
error: None,
})
})
impl<V: Clone + Send + Sync + 'static> DisconnectCache for InMemoryCache<V> {
async fn disconnect(&self) -> Result<(), Error> {
Ok(())
}
}
impl<V: Clone + Send + Sync + 'static> BatchCache for InMemoryCache<V> {}
impl<V: Clone + Send + Sync + 'static> DeleteCache for InMemoryCache<V> {
fn delete_cache(&self, key: &str) -> Result<(), Error> {
InMemoryCache::delete_cache(self, key)
}
}
impl<V: Clone + Send + Sync + 'static> FlushCache for InMemoryCache<V> {
fn flush_cache(&self) -> Result<(), Error> {
InMemoryCache::flush_cache(self)
}
}
impl<V: Clone + Send + Sync + 'static> TtlCache for InMemoryCache<V> {
async fn async_get_ttl(&self, key: &str) -> Result<Option<Duration>, Error> {
InMemoryCache::async_get_ttl(self, key).await
}
}
impl<T> SetCache for InMemoryCache<HashSet<T>>
where
T: Clone + Eq + Hash + Send + Sync + 'static,
{
type SetValue = T;
type SetResult = Vec<T>;
async fn async_set_cache_sadd(
&self,
key: &str,
values: Vec<Self::SetValue>,
ttl: Option<Duration>,
) -> Result<Self::SetResult, Error> {
if self.max_size_in_memory == 0 {
return Ok(values);
}
let now = (self.now)();
let mut state = self.state.lock().map_err(|_| Error::Unavailable)?;
let mut stored = Self::live(&mut state, key, now).unwrap_or_default();
stored.extend(values.iter().cloned());
self.store(&mut state, key.into(), stored, ttl, now)?;
Ok(values)
}
}

View file

@ -1,158 +1,731 @@
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use std::{
collections::HashSet,
sync::{
Arc,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
use litellm_cache::{BaseCache, CacheConnectionStatus, CacheEntry, Error};
use litellm_cache::{
BaseCache, BatchCache, BatchEntry, CacheBackend, ClaimCache, CounterCache, DeleteCache,
DisconnectCache, Error, ExactCacheContext, FlushCache, IncrementOperation, SetCache, TtlCache,
get_cache, set_cache,
};
use litellm_cache_memory::{CacheWrite, InMemoryCache};
use rstest::{fixture, rstest};
type Clock = Arc<AtomicU64>;
#[fixture]
fn clock() -> Arc<AtomicU64> {
fn clock() -> Clock {
Arc::new(AtomicU64::new(100))
}
fn cache(clock: Arc<AtomicU64>, capacity: usize) -> InMemoryCache<String> {
fn cache_with<V: Clone>(clock: &Clock, capacity: usize) -> InMemoryCache<V> {
let clock = clock.clone();
InMemoryCache::with_clock(Some(capacity), Some(Duration::from_secs(60)), move || {
Duration::from_secs(clock.load(Ordering::SeqCst))
Duration::from_millis(clock.load(Ordering::SeqCst) * 1000)
})
}
fn cache(clock: &Clock, capacity: usize) -> InMemoryCache<String> {
cache_with(clock, capacity)
}
fn at(clock: &Clock, seconds: u64) {
clock.store(seconds, Ordering::SeqCst);
}
fn secs(seconds: u64) -> Option<Duration> {
Some(Duration::from_secs(seconds))
}
fn ttl(seconds: u64) -> ExactCacheContext {
ExactCacheContext { ttl: secs(seconds) }
}
fn measured(capacity: usize) -> InMemoryCache<String> {
InMemoryCache::with_clock_and_size_measurement(
Some(capacity),
secs(60),
Some(4),
Some(Arc::new(|value: &String| {
if value.is_empty() {
return Err(Error::InvalidEntry);
}
Ok(value.len())
})),
|| Duration::from_secs(100),
)
}
#[rstest]
fn default_explicit_and_override_ttls_follow_python_rules(clock: Arc<AtomicU64>) {
let cache = cache(clock.clone(), 4);
fn default_explicit_and_override_ttls_follow_python_rules(clock: Clock) {
let cache = cache(&clock, 4);
cache.set_cache("key", "first".into(), None).unwrap();
assert_eq!(
cache.expires_at("key").unwrap(),
Some(Duration::from_secs(160))
);
cache
.set_cache("key", "second".into(), Some(Duration::from_secs(10)))
.unwrap();
assert_eq!(
cache.expires_at("key").unwrap(),
Some(Duration::from_secs(160))
);
clock.store(160, Ordering::SeqCst);
assert_eq!(cache.expires_at("key").unwrap(), secs(160));
cache.set_cache("key", "second".into(), secs(10)).unwrap();
assert_eq!(cache.expires_at("key").unwrap(), secs(160));
at(&clock, 160);
assert_eq!(cache.get_cache("key").unwrap(), Some("second".into()));
clock.store(161, Ordering::SeqCst);
at(&clock, 161);
assert_eq!(cache.get_cache("key").unwrap(), None);
cache
.set_cache("key", "third".into(), Some(Duration::from_secs(10)))
.unwrap();
assert_eq!(
cache.expires_at("key").unwrap(),
Some(Duration::from_secs(171))
);
assert_eq!(cache.expires_at("key").unwrap(), None);
cache.set_cache("key", "third".into(), secs(10)).unwrap();
assert_eq!(cache.expires_at("key").unwrap(), secs(171));
}
#[rstest]
fn write_at_expiry_boundary_refreshes_ttl(clock: Arc<AtomicU64>) {
let cache = cache(clock.clone(), 4);
cache
.set_cache("key", "first".into(), Some(Duration::from_secs(10)))
.unwrap();
clock.store(110, Ordering::SeqCst);
cache
.set_cache("key", "second".into(), Some(Duration::from_secs(10)))
.unwrap();
assert_eq!(
cache.expires_at("key").unwrap(),
Some(Duration::from_secs(120))
);
clock.store(115, Ordering::SeqCst);
#[case::unset(None, secs(600))]
#[case::zero_falls_back_like_python_or(Some(Duration::ZERO), secs(600))]
#[case::explicit(secs(5), secs(5))]
fn default_ttl_falls_back_to_ten_minutes(
#[case] default_ttl: Option<Duration>,
#[case] expected: Option<Duration>,
) {
let cache = InMemoryCache::<String>::with_clock(None, default_ttl, || Duration::ZERO);
assert_eq!(cache.get_ttl(&ExactCacheContext::default()), expected);
cache.set_cache("key", "value".into(), None).unwrap();
assert_eq!(cache.expires_at("key").unwrap(), expected);
assert_eq!(cache.max_size_in_memory(), 200);
}
#[rstest]
fn write_at_expiry_boundary_refreshes_ttl(clock: Clock) {
let cache = cache(&clock, 4);
cache.set_cache("key", "first".into(), secs(10)).unwrap();
at(&clock, 110);
cache.set_cache("key", "second".into(), secs(10)).unwrap();
assert_eq!(cache.expires_at("key").unwrap(), secs(120));
at(&clock, 115);
assert_eq!(cache.get_cache("key").unwrap(), Some("second".into()));
}
#[rstest]
fn capacity_evicts_earliest_and_ignores_stale_heap_entries(clock: Arc<AtomicU64>) {
let cache = cache(clock, 2);
cache
.set_cache("early", "a".into(), Some(Duration::from_secs(10)))
.unwrap();
cache
.set_cache("late", "b".into(), Some(Duration::from_secs(20)))
.unwrap();
fn expired_key_without_a_read_allows_a_ttl_override(clock: Clock) {
let cache = cache(&clock, 4);
cache.set_cache("key", "first".into(), secs(1)).unwrap();
assert_eq!(cache.allow_ttl_override("key"), Ok(false));
at(&clock, 102);
assert_eq!(cache.allow_ttl_override("key"), Ok(true));
cache.set_cache("key", "second".into(), secs(1)).unwrap();
assert_eq!(cache.expires_at("key").unwrap(), secs(103));
assert_eq!(cache.allow_ttl_override("missing"), Ok(true));
}
#[rstest]
fn capacity_evicts_earliest_and_ignores_stale_heap_entries(clock: Clock) {
let cache = cache(&clock, 2);
cache.set_cache("early", "a".into(), secs(10)).unwrap();
cache.set_cache("late", "b".into(), secs(20)).unwrap();
cache.delete_cache("early").unwrap();
cache
.set_cache("new", "c".into(), Some(Duration::from_secs(30)))
.unwrap();
cache.set_cache("new", "c".into(), secs(30)).unwrap();
assert_eq!(cache.get_cache("late").unwrap(), Some("b".into()));
cache
.set_cache("last", "d".into(), Some(Duration::from_secs(40)))
.unwrap();
cache.set_cache("last", "d".into(), secs(40)).unwrap();
assert_eq!(cache.get_cache("late").unwrap(), None);
}
#[test]
fn disabled_size_limited_and_synchronized_response_writes_are_observable() {
let disabled = InMemoryCache::<CacheEntry>::response_cache(0, Duration::from_secs(60), 80);
assert_eq!(
disabled
.set_cache(
"a",
CacheEntry {
timestamp: 1.0,
response: serde_json::json!("x")
},
None
)
.unwrap(),
CacheWrite::Disabled
);
let cache = InMemoryCache::<CacheEntry>::response_cache(2, Duration::from_secs(60), 80);
assert_eq!(
#[rstest]
fn max_size_is_respected_when_every_item_has_a_long_ttl(clock: Clock) {
let cache = cache(&clock, 3);
for index in 0..3 {
at(&clock, 100 + index);
cache
.set_cache(
"large",
CacheEntry {
timestamp: 1.0,
response: serde_json::json!("x".repeat(100))
},
None
format!("key_{index}"),
format!("value_{index}"),
secs(86_400),
)
.unwrap(),
CacheWrite::TooLarge
);
.unwrap();
}
assert_eq!(cache.len(), Ok(3));
cache
.set_cache(
"small",
CacheEntry {
timestamp: 1.0,
response: serde_json::json!("ok"),
},
None,
)
.set_cache("key_3", "value_3".into(), secs(86_400))
.unwrap();
assert!(cache.get_cache("small").unwrap().is_some());
assert_eq!(
cache
.set_cache(
"invalid",
CacheEntry {
timestamp: f64::NAN,
response: serde_json::json!("bad"),
},
None,
)
.unwrap_err(),
Error::InvalidEntry
);
cache.delete_cache("small").unwrap();
cache.flush_cache().unwrap();
assert_eq!(cache.len(), Ok(3));
assert_eq!(cache.get_cache("key_0").unwrap(), None);
assert_eq!(cache.expires_at("key_0").unwrap(), None);
for key in ["key_1", "key_2", "key_3"] {
assert!(cache.get_cache(key).unwrap().is_some(), "{key}");
}
}
#[tokio::test]
async fn connection_test_matches_python_result_contract() {
let cache = InMemoryCache::<CacheEntry>::default();
let result = BaseCache::test_connection(&cache).await.unwrap();
assert_eq!(result.status, CacheConnectionStatus::Success);
assert_eq!(result.message, "In-memory cache connection test successful");
assert_eq!(result.error, None);
#[rstest]
fn expired_items_are_evicted_before_live_ones(clock: Clock) {
let cache = cache(&clock, 3);
cache.set_cache("expired_1", "1".into(), secs(1)).unwrap();
cache.set_cache("expired_2", "2".into(), secs(1)).unwrap();
cache
.set_cache("long_lived", "3".into(), secs(86_400))
.unwrap();
assert_eq!(cache.len(), Ok(3));
at(&clock, 102);
cache
.set_cache("new_item", "4".into(), secs(86_400))
.unwrap();
assert_eq!(cache.len(), Ok(2));
assert_eq!(cache.get_cache("long_lived").unwrap(), Some("3".into()));
assert_eq!(cache.get_cache("new_item").unwrap(), Some("4".into()));
for key in ["expired_1", "expired_2"] {
assert_eq!(cache.expires_at(key).unwrap(), None, "{key}");
}
}
#[rstest]
fn injected_clock_controls_expiry_and_eviction(clock: Clock) {
let cache = cache(&clock, 2);
at(&clock, 0);
cache
.set_cache("first", "original".into(), secs(10))
.unwrap();
at(&clock, 9);
cache.set_cache("second", "survivor".into(), None).unwrap();
assert_eq!(cache.get_cache("first").unwrap(), Some("original".into()));
at(&clock, 11);
assert_eq!(cache.get_cache("first").unwrap(), None);
cache
.set_cache("third", "replacement".into(), None)
.unwrap();
assert_eq!(cache.get_cache("second").unwrap(), Some("survivor".into()));
at(&clock, 70);
cache.set_cache("fourth", "new".into(), None).unwrap();
assert_eq!(cache.get_cache("second").unwrap(), None);
assert_eq!(
serde_json::to_value(result).unwrap(),
serde_json::json!({
"status": "success",
"message": "In-memory cache connection test successful"
})
cache.get_cache("third").unwrap(),
Some("replacement".into())
);
assert_eq!(cache.get_cache("fourth").unwrap(), Some("new".into()));
}
#[rstest]
fn rewriting_one_key_keeps_one_heap_entry(clock: Clock) {
let cache = cache(&clock, 10);
for index in 0..1_000 {
cache
.set_cache("hot_key", format!("value_{index}"), secs(60))
.unwrap();
}
assert_eq!(cache.expiration_heap_len(), Ok(1));
}
#[rstest]
fn repeated_increments_keep_one_heap_entry_per_expiration() {
let cache = InMemoryCache::<f64>::new(Some(4), None);
for _ in 0..100 {
cache
.increment_cache("counter", 1.0, ExactCacheContext::default())
.unwrap();
}
assert_eq!(cache.expiration_heap_len(), Ok(1));
}
#[rstest]
fn reinserting_expired_keys_below_capacity_prunes_the_heap(clock: Clock) {
let cache = cache(&clock, 200);
for cycle in 0..3 {
for index in 0..5 {
cache
.set_cache(format!("key_{index}"), format!("value_{cycle}"), secs(1))
.unwrap();
}
at(&clock, 100 + 2 * (cycle + 1));
}
for index in 0..5 {
cache
.set_cache(format!("key_{index}"), "final".into(), secs(1))
.unwrap();
}
assert_eq!(cache.len(), Ok(5));
assert_eq!(cache.expiration_heap_len(), Ok(5));
}
#[rstest]
fn evict_cache_drops_expired_entries_then_makes_room(clock: Clock) {
let cache = cache(&clock, 2);
assert_eq!(cache.is_empty(), Ok(true));
cache.set_cache("short", "a".into(), secs(1)).unwrap();
cache.set_cache("long", "b".into(), secs(50)).unwrap();
at(&clock, 102);
cache.evict_cache().unwrap();
assert_eq!(cache.len(), Ok(1));
assert_eq!(cache.expires_at("short").unwrap(), None);
cache.set_cache("longer", "c".into(), secs(90)).unwrap();
cache.evict_cache().unwrap();
assert_eq!(cache.len(), Ok(1));
assert_eq!(cache.get_cache("long").unwrap(), None);
assert_eq!(cache.get_cache("longer").unwrap(), Some("c".into()));
}
#[rstest]
fn evict_element_if_expired_reports_removal(clock: Clock) {
let cache = cache(&clock, 4);
cache.set_cache("key", "value".into(), secs(10)).unwrap();
assert_eq!(cache.evict_element_if_expired("key"), Ok(false));
assert_eq!(cache.evict_element_if_expired("missing"), Ok(false));
at(&clock, 110);
assert_eq!(cache.evict_element_if_expired("key"), Ok(false));
at(&clock, 111);
assert_eq!(cache.evict_element_if_expired("key"), Ok(true));
assert_eq!(cache.len(), Ok(0));
assert_eq!(cache.expires_at("key").unwrap(), None);
}
#[rstest]
#[case::fits("ok", Ok(true))]
#[case::at_limit("four", Ok(true))]
#[case::too_large("oversized", Ok(false))]
#[case::measure_error("", Err(Error::InvalidEntry))]
fn check_value_size_applies_the_entry_limit(
#[case] value: &str,
#[case] expected: Result<bool, Error>,
) {
assert_eq!(measured(2).check_value_size(&value.to_string()), expected);
}
#[rstest]
fn values_are_unbounded_without_a_measure() {
let cache = InMemoryCache::<String>::default();
assert_eq!(cache.max_entry_bytes(), None);
assert_eq!(cache.check_value_size(&"x".repeat(1 << 20)), Ok(true));
}
#[rstest]
fn disabled_size_limited_and_validated_writes_are_observable() {
assert_eq!(
measured(0).set_cache("a", "x".into(), None).unwrap(),
CacheWrite::Disabled
);
let cache = measured(2);
assert_eq!(cache.max_entry_bytes(), Some(4));
assert_eq!(
cache.set_cache("large", "oversized".into(), None).unwrap(),
CacheWrite::TooLarge
);
assert_eq!(cache.get_cache("large").unwrap(), None);
assert_eq!(
cache.set_cache("small", "ok".into(), None).unwrap(),
CacheWrite::Stored
);
assert_eq!(cache.get_cache("small").unwrap(), Some("ok".into()));
assert_eq!(
cache.set_cache("invalid", String::new(), None),
Err(Error::InvalidEntry)
);
assert_eq!(cache.get_cache("invalid").unwrap(), None);
cache.delete_cache("small").unwrap();
assert_eq!(cache.get_cache("small").unwrap(), None);
}
#[rstest]
#[tokio::test]
async fn disconnect_is_a_no_op_that_keeps_entries() {
let cache = InMemoryCache::<String>::default();
cache.set_cache("key", "value".into(), None).unwrap();
cache.disconnect().await.unwrap();
assert_eq!(cache.get_cache("key").unwrap(), Some("value".into()));
}
#[rstest]
#[tokio::test]
async fn generic_consumers_share_typed_values_and_honor_expiration(clock: Clock) {
let cache: CacheBackend<InMemoryCache<String>> = Arc::new(self::cache(&clock, 4));
let reader = Arc::clone(&cache);
let context = ttl(5);
set_cache(cache.as_ref(), "sync", "first".into(), &context).unwrap();
assert_eq!(
get_cache(reader.as_ref(), "sync", &context).unwrap(),
Some("first".into())
);
cache
.batch_cache_write("async", "second".into(), context.clone())
.await
.unwrap();
cache
.async_set_cache_pipeline(vec![("batch".into(), "third".into())], context.clone())
.await
.unwrap();
drop(cache);
for (key, value) in [("sync", "first"), ("async", "second"), ("batch", "third")] {
assert_eq!(
reader.async_get_cache(key, &context).await.unwrap(),
Some(value.into())
);
}
reader.async_delete_cache("async").await.unwrap();
assert_eq!(
reader.async_get_cache("async", &context).await.unwrap(),
None
);
at(&clock, 106);
assert_eq!(get_cache(reader.as_ref(), "sync", &context).unwrap(), None);
assert_eq!(
reader.async_get_cache("batch", &context).await.unwrap(),
None
);
}
#[rstest]
#[case::context_ttl(ttl(5), secs(105))]
#[case::default_ttl(ExactCacheContext::default(), secs(160))]
#[tokio::test]
async fn pipeline_writes_use_the_context_ttl_or_the_default(
clock: Clock,
#[case] context: ExactCacheContext,
#[case] expected: Option<Duration>,
) {
let cache = cache(&clock, 4);
cache
.async_set_cache_pipeline(
vec![("a".into(), "1".into()), ("b".into(), "2".into())],
context,
)
.await
.unwrap();
assert_eq!(cache.expires_at("a").unwrap(), expected);
assert_eq!(cache.expires_at("b").unwrap(), expected);
}
#[rstest]
#[tokio::test]
async fn batch_reads_return_one_entry_per_key_and_drop_expired_ones(clock: Clock) {
let cache = cache(&clock, 4);
cache.set_cache("short", "a".into(), secs(1)).unwrap();
cache.set_cache("long", "b".into(), secs(50)).unwrap();
let keys = vec!["short".to_string(), "missing".into(), "long".into()];
assert_eq!(
cache
.batch_get_cache(&keys, &ExactCacheContext::default())
.unwrap(),
[
BatchEntry::Hit("a".to_string()),
BatchEntry::Miss,
BatchEntry::Hit("b".into()),
]
);
at(&clock, 102);
assert_eq!(
cache
.async_batch_get_cache(keys, ExactCacheContext::default())
.await
.unwrap(),
[
BatchEntry::Miss,
BatchEntry::Miss,
BatchEntry::Hit("b".into())
]
);
}
#[rstest]
#[tokio::test]
async fn flush_clears_values_and_expirations(clock: Clock) {
let cache = cache(&clock, 4);
cache.set_cache("a", "1".into(), None).unwrap();
cache.set_cache("b", "2".into(), None).unwrap();
cache.flush_cache().unwrap();
assert_eq!(cache.len(), Ok(0));
assert_eq!(cache.expiration_heap_len(), Ok(0));
cache.set_cache("c", "3".into(), None).unwrap();
FlushCache::async_flush_cache(&cache).await.unwrap();
assert_eq!(cache.is_empty(), Ok(true));
assert_eq!(
cache.async_get_oldest_n_keys(5).await.unwrap(),
Vec::<String>::new()
);
}
#[rstest]
fn claims_are_atomic_and_refresh_eligible_winners(clock: Clock) {
let cache = cache(&clock, 4);
let context = ttl(10);
assert_eq!(
cache
.claim_cache("affinity", "first".to_string(), &[], context.clone())
.unwrap(),
"first"
);
at(&clock, 103);
assert_eq!(
cache
.claim_cache("affinity", "second".to_string(), &[], context.clone())
.unwrap(),
"first"
);
assert_eq!(cache.expires_at("affinity").unwrap(), secs(110));
at(&clock, 105);
assert_eq!(
cache
.claim_cache(
"affinity",
"second".to_string(),
&["first".to_string(), "second".to_string()],
context,
)
.unwrap(),
"first"
);
assert_eq!(cache.expires_at("affinity").unwrap(), secs(115));
}
#[rstest]
fn counters_increment_under_one_lock() {
let cache = InMemoryCache::<f64>::default();
assert_eq!(
CounterCache::increment_cache(&cache, "counter", 1.5, ExactCacheContext::default())
.unwrap(),
1.5
);
assert_eq!(
CounterCache::increment_cache(&cache, "counter", 2.0, ExactCacheContext::default())
.unwrap(),
3.5
);
}
#[rstest]
fn concurrent_increments_are_atomic() {
let cache = Arc::new(InMemoryCache::<f64>::default());
cache.set_cache("counter", 1000.0, None).unwrap();
let threads = (0..8)
.map(|_| {
let cache = cache.clone();
std::thread::spawn(move || {
cache
.increment_cache("counter", 1.0, ExactCacheContext::default())
.unwrap()
})
})
.collect::<Vec<_>>();
for thread in threads {
thread.join().unwrap();
}
assert_eq!(cache.get_cache("counter").unwrap(), Some(1008.0));
}
#[rstest]
#[case::window_semantics(false)]
#[case::refresh_ttl_is_ignored(true)]
#[tokio::test]
async fn async_increment_delegates_to_the_locked_sync_path(
clock: Clock,
#[case] refresh_ttl: bool,
) {
let cache = cache_with::<f64>(&clock, 4);
assert_eq!(
cache
.async_increment("counter", 2.0, ttl(10), refresh_ttl)
.await,
Ok(2.0)
);
at(&clock, 105);
assert_eq!(
cache
.async_increment("counter", 3.0, ttl(10), refresh_ttl)
.await,
Ok(5.0)
);
assert_eq!(cache.get_cache("counter").unwrap(), Some(5.0));
assert_eq!(cache.expires_at("counter").unwrap(), secs(110));
}
#[rstest]
fn expired_counters_restart_from_zero_with_a_new_ttl(clock: Clock) {
let cache = cache_with::<f64>(&clock, 4);
cache.increment_cache("counter", 2.0, ttl(10)).unwrap();
at(&clock, 111);
assert_eq!(cache.increment_cache("counter", 1.0, ttl(10)), Ok(1.0));
assert_eq!(cache.expires_at("counter").unwrap(), secs(121));
}
/// Python `InMemoryCache.set_cache` runs `evict_cache()` before every insert, and step 2 evicts
/// the earliest expiry while `len(cache_dict) >= max_size_in_memory`, even when the key being
/// written already exists.
#[rstest]
fn overwriting_an_existing_key_at_capacity_evicts_the_earliest_expiry_like_python(clock: Clock) {
let cache = cache(&clock, 2);
cache.set_cache("hot", "1".into(), secs(10)).unwrap();
cache.set_cache("cold", "2".into(), secs(20)).unwrap();
cache.set_cache("cold", "3".into(), None).unwrap();
assert_eq!(cache.get_cache("hot").unwrap(), None);
assert_eq!(cache.get_cache("cold").unwrap(), Some("3".into()));
}
/// `claim_cache` has no Python counterpart; it never evicts another entry for a key it holds.
#[rstest]
fn claiming_an_existing_key_at_capacity_keeps_other_entries(clock: Clock) {
let cache = cache(&clock, 2);
cache.set_cache("hot", "1".into(), secs(10)).unwrap();
cache.set_cache("cold", "2".into(), secs(20)).unwrap();
cache
.claim_cache("cold", "4".into(), &[], ExactCacheContext::default())
.unwrap();
assert_eq!(cache.get_cache("hot").unwrap(), Some("1".into()));
assert_eq!(cache.get_cache("cold").unwrap(), Some("2".into()));
}
/// Python `increment_cache` is `get_cache` then `set_cache`, so at capacity the write evicts
/// the earliest expiry first: equal expiries tie-break on the key, and the value read before
/// eviction is the one written back.
#[rstest]
fn incrementing_at_capacity_evicts_the_earliest_expiry_like_python(clock: Clock) {
let cache = cache_with::<f64>(&clock, 2);
for key in ["a", "b", "a", "b"] {
cache
.increment_cache(key, 1.0, ExactCacheContext::default())
.unwrap();
}
assert_eq!(cache.get_cache("a").unwrap(), None);
assert_eq!(cache.get_cache("b").unwrap(), Some(2.0));
}
#[rstest]
#[tokio::test]
async fn disabled_cache_does_not_retain_claims_counters_or_sets() {
let claims = InMemoryCache::<String>::new(Some(0), None);
assert_eq!(
claims
.claim_cache("key", "first".into(), &[], ExactCacheContext::default())
.unwrap(),
"first"
);
assert_eq!(claims.get_cache("key").unwrap(), None);
let counters = InMemoryCache::<f64>::new(Some(0), None);
assert_eq!(
counters
.increment_cache("key", 2.0, ExactCacheContext::default())
.unwrap(),
2.0
);
assert_eq!(counters.get_cache("key").unwrap(), None);
let sets = InMemoryCache::<HashSet<String>>::new(Some(0), None);
assert_eq!(
sets.async_set_cache_sadd("key", vec!["a".into()], None)
.await
.unwrap(),
["a"]
);
assert_eq!(sets.get_cache("key").unwrap(), None);
}
#[rstest]
#[tokio::test]
async fn ttl_and_oldest_key_operations_use_the_stored_expirations(clock: Clock) {
let cache = cache(&clock, 3);
cache.set_cache("later", "2".into(), secs(20)).unwrap();
cache.set_cache("first", "1".into(), secs(10)).unwrap();
cache.set_cache("latest", "3".into(), secs(30)).unwrap();
assert_eq!(cache.async_get_ttl("first").await.unwrap(), secs(110));
assert_eq!(
TtlCache::async_get_ttl(&cache, "later").await.unwrap(),
secs(120)
);
assert_eq!(cache.async_get_oldest_n_keys(1).await.unwrap(), ["first"]);
assert_eq!(
cache.async_get_oldest_n_keys(10).await.unwrap(),
["first", "later", "latest"]
);
assert_eq!(
cache.async_get_oldest_n_keys(0).await.unwrap(),
Vec::<String>::new()
);
assert_eq!(cache.async_get_ttl("missing").await.unwrap(), None);
}
#[rstest]
#[tokio::test]
async fn increment_pipeline_preserves_operation_order(clock: Clock) {
let cache = cache_with::<f64>(&clock, 3);
let operation = |key: &str, amount, ttl| IncrementOperation {
key: key.into(),
amount,
ttl: secs(ttl),
};
assert_eq!(
cache
.async_increment_pipeline(vec![
operation("a", 1.0, 10),
operation("b", 5.0, 30),
operation("a", 2.0, 20),
])
.await
.unwrap(),
[1.0, 5.0, 3.0]
);
assert_eq!(cache.get_cache("a").unwrap(), Some(3.0));
assert_eq!(cache.expires_at("a").unwrap(), secs(110));
assert_eq!(cache.expires_at("b").unwrap(), secs(130));
assert_eq!(
cache.async_increment_pipeline(Vec::new()).await.unwrap(),
Vec::<f64>::new()
);
}
#[rstest]
#[tokio::test]
async fn set_capability_preserves_python_result_and_deduplicates_storage(clock: Clock) {
let cache = cache_with::<HashSet<String>>(&clock, 4);
let inserted = vec!["a".to_string(), "a".into(), "b".into()];
assert_eq!(
cache
.async_set_cache_sadd("members", inserted.clone(), secs(10))
.await
.unwrap(),
inserted
);
assert_eq!(
cache
.async_set_cache_sadd("members", vec!["c".into()], secs(99))
.await
.unwrap(),
["c"]
);
assert_eq!(
cache.get_cache("members").unwrap(),
Some(HashSet::from(["a".into(), "b".into(), "c".into()]))
);
assert_eq!(cache.expires_at("members").unwrap(), secs(110));
at(&clock, 111);
cache
.async_set_cache_sadd("members", vec!["d".into()], None)
.await
.unwrap();
assert_eq!(
cache.get_cache("members").unwrap(),
Some(HashSet::from(["d".into()]))
);
assert_eq!(cache.expires_at("members").unwrap(), secs(171));
}
#[rstest]
#[tokio::test]
async fn oversized_set_additions_are_not_stored() {
let cache = InMemoryCache::<HashSet<String>>::with_clock_and_size_measurement(
Some(4),
None,
Some(2),
Some(Arc::new(|value: &HashSet<String>| Ok(value.len()))),
|| Duration::ZERO,
);
cache
.async_set_cache_sadd("members", vec!["a".into(), "b".into()], None)
.await
.unwrap();
assert_eq!(
cache
.async_set_cache_sadd("members", vec!["c".into()], None)
.await
.unwrap(),
["c"]
);
assert_eq!(
cache.get_cache("members").unwrap(),
Some(HashSet::from(["a".into(), "b".into()]))
);
}

View file

@ -0,0 +1,98 @@
use std::time::Duration;
use litellm_cache::ExactCacheContext;
use litellm_cache_memory::InMemoryCache;
use litellm_cache_testing as contract;
use rstest::{fixture, rstest};
#[fixture]
fn strings() -> InMemoryCache<String> {
InMemoryCache::new(Some(16), None)
}
#[fixture]
fn counters() -> InMemoryCache<f64> {
InMemoryCache::new(Some(16), None)
}
#[fixture]
fn context() -> ExactCacheContext {
ExactCacheContext {
ttl: Some(Duration::from_secs(60)),
}
}
#[rstest]
#[tokio::test]
async fn hit_and_miss(strings: InMemoryCache<String>, context: ExactCacheContext) {
contract::hit_and_miss(&strings, context, "memory:", "value".into()).await;
}
#[rstest]
#[tokio::test]
async fn sync_async_equivalence(strings: InMemoryCache<String>, context: ExactCacheContext) {
contract::sync_async_equivalence(
&strings,
context,
"memory:",
"first".into(),
"second".into(),
)
.await;
}
#[rstest]
#[tokio::test]
async fn overwrite_replaces(strings: InMemoryCache<String>, context: ExactCacheContext) {
contract::overwrite_replaces(
&strings,
context,
"memory:",
"first".into(),
"second".into(),
)
.await;
}
#[rstest]
#[tokio::test]
async fn pipeline_writes_every_entry(strings: InMemoryCache<String>, context: ExactCacheContext) {
contract::pipeline_writes_every_entry(
&strings,
context,
"memory:",
vec!["a".into(), "b".into(), "c".into()],
)
.await;
}
#[rstest]
#[tokio::test]
async fn batch_preserves_order(strings: InMemoryCache<String>, context: ExactCacheContext) {
contract::batch_preserves_order(
&strings,
context,
"memory:",
"first".into(),
"second".into(),
)
.await;
}
#[rstest]
#[tokio::test]
async fn delete_removes_key(strings: InMemoryCache<String>, context: ExactCacheContext) {
contract::delete_removes_key(&strings, context, "memory:", "value".into()).await;
}
#[rstest]
#[tokio::test]
async fn flush_clears(strings: InMemoryCache<String>, context: ExactCacheContext) {
contract::flush_clears(&strings, context, "memory:", "value".into()).await;
}
#[rstest]
#[tokio::test]
async fn counter_accumulates(counters: InMemoryCache<f64>, context: ExactCacheContext) {
contract::counter_accumulates(&counters, context, "memory:").await;
}

View file

@ -0,0 +1,25 @@
[package]
name = "litellm-cache-qdrant-semantic"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
futures-util.workspace = true
litellm-cache.workspace = true
qdrant-client = { workspace = true, features = ["serde"] }
reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tokio.workspace = true
uuid.workspace = true
[dev-dependencies]
futures-executor = "0.3"
litellm-cache-testing.workspace = true
rstest.workspace = true
tonic = "0.14"
tonic-prost = "0.14"
tokio-stream = "0.1"

View file

@ -0,0 +1,277 @@
use futures_util::future::try_join_all;
use litellm_cache::{
BaseCache, CacheCodec, Error, SemanticCacheContext,
semantic::{Embedder, SemanticCache, SemanticLookup, prompt_from_messages},
};
use qdrant_client::{
Payload, Qdrant,
qdrant::{
BinaryQuantizationBuilder, CompressionRatio, Condition, CreateCollectionBuilder,
CreateFieldIndexCollectionBuilder, Distance, FieldType, Filter, PointStruct,
ProductQuantizationBuilder, QuantizationSearchParamsBuilder, ScalarQuantizationBuilder,
SearchParamsBuilder, SearchPointsBuilder, UpsertPointsBuilder, VectorParamsBuilder,
},
};
use serde_json::{Map, Value, json};
use uuid::Uuid;
use crate::{QdrantSemanticConfig, Quantization};
pub struct QdrantSemanticCache<E: Embedder, C: CacheCodec> {
client: Qdrant,
embedder: E,
codec: C,
config: QdrantSemanticConfig,
runtime: tokio::runtime::Handle,
}
impl<E: Embedder, C: CacheCodec> QdrantSemanticCache<E, C> {
pub async fn connect(
client: Qdrant,
embedder: E,
codec: C,
config: QdrantSemanticConfig,
runtime: tokio::runtime::Handle,
) -> Result<Self, Error> {
let exists = client
.collection_exists(config.collection_name.clone())
.await
.map_err(|_| Error::Unavailable)?;
if !exists {
client
.create_collection(
CreateCollectionBuilder::new(config.collection_name.clone())
.vectors_config(VectorParamsBuilder::new(
config.vector_size,
Distance::Cosine,
))
.quantization_config(quantization(&config.quantization)),
)
.await
.map_err(|_| Error::Unavailable)?;
}
let _ = client
.create_field_index(CreateFieldIndexCollectionBuilder::new(
config.collection_name.clone(),
"litellm_cache_key".to_owned(),
FieldType::Keyword,
))
.await;
Ok(Self {
client,
embedder,
codec,
config,
runtime,
})
}
pub fn collection_name(&self) -> &str {
&self.config.collection_name
}
pub fn similarity_threshold(&self) -> f64 {
self.config.similarity_threshold
}
pub fn vector_size(&self) -> u64 {
self.config.vector_size
}
pub fn embedder(&self) -> &E {
&self.embedder
}
/// Python reads `kwargs["messages"]` unguarded, so a request without messages fails.
fn prompt(context: &SemanticCacheContext) -> Result<String, Error> {
prompt_from_messages(context).ok_or(Error::MissingPrompt)
}
async fn set(
&self,
key: &str,
value: C::Value,
context: &SemanticCacheContext,
) -> Result<(), Error> {
let prompt = Self::prompt(context)?;
let vector = self
.embedder
.async_embed(&prompt, context.metadata.as_ref())
.await?;
let response =
String::from_utf8(self.codec.encode(&value)?).map_err(|_| Error::InvalidEntry)?;
let payload = Payload::try_from(json!({
"litellm_cache_key": key,
"text": prompt,
"response": response,
}))
.map_err(|_| Error::InvalidEntry)?;
self.client
.upsert_points(
UpsertPointsBuilder::new(
self.collection_name(),
vec![PointStruct::new(
Uuid::new_v4().to_string(),
vector,
payload,
)],
)
.wait(true),
)
.await
.map_err(|_| Error::Unavailable)?;
Ok(())
}
async fn get(
&self,
key: &str,
context: &SemanticCacheContext,
) -> Result<SemanticLookup<C::Value>, Error> {
let prompt = Self::prompt(context)?;
let vector = self
.embedder
.async_embed(&prompt, context.metadata.as_ref())
.await?;
let result = self
.client
.search_points(
SearchPointsBuilder::new(self.collection_name(), vector, 1)
.with_payload(true)
.filter(Filter::must([Condition::matches(
"litellm_cache_key",
key.to_owned(),
)]))
.params(
SearchParamsBuilder::default().quantization(
QuantizationSearchParamsBuilder::default()
.ignore(false)
.rescore(true)
.oversampling(3.0),
),
),
)
.await
.map_err(|_| Error::Unavailable)?;
let Some(point) = result.result.into_iter().next() else {
return Ok(SemanticLookup::miss(Some(0.0)));
};
let payload: Map<String, Value> = Payload::from(point.payload).into();
if !payload
.get("litellm_cache_key")
.is_some_and(|cached| python_str(cached).as_deref() == Some(key))
{
return Ok(SemanticLookup::miss(Some(0.0)));
}
let similarity = f64::from(point.score);
if similarity < self.config.similarity_threshold {
return Ok(SemanticLookup::miss(Some(similarity)));
}
let response = payload
.get("response")
.and_then(Value::as_str)
.ok_or(Error::InvalidEntry)?;
Ok(SemanticLookup {
value: Some(self.codec.decode(response.as_bytes())?),
similarity: Some(similarity),
})
}
}
fn quantization(value: &Quantization) -> qdrant_client::qdrant::quantization_config::Quantization {
match value {
Quantization::Binary => BinaryQuantizationBuilder::new(false).into(),
Quantization::Scalar => ScalarQuantizationBuilder::default()
.quantile(0.99)
.always_ram(false)
.into(),
Quantization::Product => ProductQuantizationBuilder::new(CompressionRatio::X16.into())
.always_ram(false)
.into(),
}
}
impl<E: Embedder, C: CacheCodec> BaseCache for QdrantSemanticCache<E, C> {
type Value = C::Value;
type Context = SemanticCacheContext;
fn get_ttl(&self, _: &Self::Context) -> Option<std::time::Duration> {
None
}
fn set_cache(
&self,
key: &str,
value: Self::Value,
context: &Self::Context,
) -> Result<(), Error> {
self.runtime.block_on(self.set(key, value, context))
}
fn get_cache(&self, key: &str, context: &Self::Context) -> Result<Option<Self::Value>, Error> {
self.get_cache_with_similarity(key, context)
.map(|lookup| lookup.value)
}
async fn async_set_cache(
&self,
key: &str,
value: Self::Value,
context: Self::Context,
) -> Result<(), Error> {
self.set(key, value, &context).await
}
async fn async_get_cache(
&self,
key: &str,
context: &Self::Context,
) -> Result<Option<Self::Value>, Error> {
self.get(key, context).await.map(|lookup| lookup.value)
}
async fn async_set_cache_pipeline(
&self,
entries: Vec<(String, Self::Value)>,
context: Self::Context,
) -> Result<(), Error> {
try_join_all(entries.into_iter().map(|(key, value)| {
let context = context.clone();
async move { self.async_set_cache(&key, value, context).await }
}))
.await
.map(|_| ())
}
}
/// Python stamps the top point's score, even below the threshold, and `0.0` when there is no
/// point or it belongs to another key. A request without messages fails before any search.
impl<E: Embedder, C: CacheCodec> SemanticCache for QdrantSemanticCache<E, C> {
fn get_cache_with_similarity(
&self,
key: &str,
context: &Self::Context,
) -> Result<SemanticLookup<Self::Value>, Error> {
self.runtime.block_on(self.get(key, context))
}
async fn async_get_cache_with_similarity(
&self,
key: &str,
context: &Self::Context,
) -> Result<SemanticLookup<Self::Value>, Error> {
self.get(key, context).await
}
}
/// `str(value)` for the scalar payload values `_payload_matches_cache_key` compares; `None` for
/// null (a pre-isolation point without a key) and for containers, which never equal a key.
fn python_str(value: &Value) -> Option<String> {
match value {
Value::String(text) => Some(text.clone()),
Value::Number(number) => Some(number.to_string()),
Value::Bool(true) => Some("True".into()),
Value::Bool(false) => Some("False".into()),
Value::Null | Value::Array(_) | Value::Object(_) => None,
}
}

View file

@ -0,0 +1,13 @@
#[derive(Clone, Debug, PartialEq)]
pub enum Quantization {
Binary,
Scalar,
Product,
}
pub struct QdrantSemanticConfig {
pub collection_name: String,
pub similarity_threshold: f64,
pub vector_size: u64,
pub quantization: Quantization,
}

View file

@ -0,0 +1,75 @@
use std::time::Duration;
use litellm_cache::{Error, semantic::Embedder};
use reqwest::Client;
use serde_json::Value;
pub struct OpenAiEmbedder {
client: Client,
api_base: String,
api_key: String,
model: String,
timeout: Option<Duration>,
}
pub struct OpenAiEmbedderConfig {
pub api_base: String,
pub api_key: String,
pub model: String,
pub timeout: Option<Duration>,
}
impl OpenAiEmbedder {
pub fn new(client: Client, config: OpenAiEmbedderConfig) -> Self {
Self {
client,
api_base: config.api_base.trim_end_matches('/').to_owned(),
api_key: config.api_key,
model: config.model,
timeout: config.timeout,
}
}
pub fn model(&self) -> &str {
&self.model
}
}
/// An OpenAI-compatible `/embeddings` call. It has no router to route on, so `metadata` is
/// unused, and it only embeds asynchronously: sync cache calls block on the cache's runtime.
impl Embedder for OpenAiEmbedder {
async fn async_embed(&self, input: &str, _metadata: Option<&Value>) -> Result<Vec<f32>, Error> {
let request = self
.client
.post(format!("{}/embeddings", self.api_base))
.bearer_auth(&self.api_key)
.json(&serde_json::json!({
"model": self.model,
"input": input,
"encoding_format": "float",
}));
let response = if let Some(timeout) = self.timeout {
request.timeout(timeout)
} else {
request
}
.send()
.await
.map_err(|_| Error::Unavailable)?
.error_for_status()
.map_err(|_| Error::Unavailable)?;
let body: Value = response.json().await.map_err(|_| Error::Unavailable)?;
body.get("data")
.and_then(Value::as_array)
.and_then(|data| data.first())
.and_then(|item| item.get("embedding"))
.and_then(Value::as_array)
.and_then(|embedding| {
embedding
.iter()
.map(|value| value.as_f64().map(|value| value as f32))
.collect::<Option<Vec<_>>>()
})
.ok_or(Error::Unavailable)
}
}

View file

@ -0,0 +1,7 @@
mod cache;
mod config;
mod embedder;
pub use cache::QdrantSemanticCache;
pub use config::{QdrantSemanticConfig, Quantization};
pub use embedder::{OpenAiEmbedder, OpenAiEmbedderConfig};

View file

@ -0,0 +1,91 @@
//! `overwrite_replaces` does not apply: like Python, every write upserts a new `uuid4` point,
//! so a second write with the same prompt adds a tie instead of replacing the first.
mod support;
use std::future::Future;
use litellm_cache::{JsonCodec, SemanticCacheContext, semantic::PreparedEmbedding};
use litellm_cache_qdrant_semantic::{QdrantSemanticCache, QdrantSemanticConfig, Quantization};
use litellm_cache_testing as contract;
use qdrant_client::Qdrant;
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use support::{FakeQdrant, FakeState};
type Cache = QdrantSemanticCache<PreparedEmbedding, JsonCodec<Value>>;
const PREFIX: &str = "contract:";
#[fixture]
fn context() -> SemanticCacheContext {
SemanticCacheContext {
messages: Some(json!([{"role": "user", "content": "contract prompt"}])),
..Default::default()
}
}
/// Runs a contract against a fresh fake Qdrant. The sync cache methods block on the runtime, so
/// the contract is polled on a blocking thread outside the runtime's own executor.
async fn run<F, Fut>(check: F)
where
F: FnOnce(Cache) -> Fut + Send + 'static,
Fut: Future<Output = ()>,
{
let server = FakeQdrant::start(FakeState::default()).await;
let runtime = tokio::runtime::Handle::current();
let cache = QdrantSemanticCache::connect(
Qdrant::from_url(&server.url()).build().unwrap(),
PreparedEmbedding(vec![0.6, 0.8]),
JsonCodec::new(),
QdrantSemanticConfig {
collection_name: "contract".to_owned(),
similarity_threshold: 0.9,
vector_size: 2,
quantization: Quantization::Binary,
},
runtime.clone(),
)
.await
.unwrap();
tokio::task::spawn_blocking(move || {
let _guard = runtime.enter();
futures_executor::block_on(check(cache));
})
.await
.unwrap();
server.stop();
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn hit_and_miss(context: SemanticCacheContext) {
run(|cache| async move {
contract::hit_and_miss(&cache, context, PREFIX, json!({"answer": 42})).await;
})
.await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn sync_async_equivalence(context: SemanticCacheContext) {
run(|cache| async move {
contract::sync_async_equivalence(&cache, context, PREFIX, json!("first"), json!([2])).await;
})
.await;
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn pipeline_writes_every_entry(context: SemanticCacheContext) {
run(|cache| async move {
contract::pipeline_writes_every_entry(
&cache,
context,
PREFIX,
vec![json!("a"), json!(2), json!({"c": true})],
)
.await;
})
.await;
}

View file

@ -0,0 +1,191 @@
use std::{
sync::{Arc, Mutex},
time::Duration,
};
use litellm_cache::{Error, semantic::Embedder};
use litellm_cache_qdrant_semantic::{OpenAiEmbedder, OpenAiEmbedderConfig};
use rstest::rstest;
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
struct TestHttpServer {
address: std::net::SocketAddr,
request: Arc<Mutex<Option<Vec<u8>>>>,
task: tokio::task::JoinHandle<()>,
}
impl TestHttpServer {
async fn response(status: &str, body: &str) -> Self {
Self::response_after(status, body, Duration::ZERO).await
}
async fn response_after(status: &str, body: &str, delay: Duration) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let request = Arc::new(Mutex::new(None));
let captured = request.clone();
let status = status.to_owned();
let body = body.to_owned();
let task = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let request_bytes = read_request(&mut stream).await;
*captured.lock().unwrap() = Some(request_bytes);
tokio::time::sleep(delay).await;
let response = format!(
"HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
stream.write_all(response.as_bytes()).await.unwrap();
});
Self {
address,
request,
task,
}
}
fn base_url(&self) -> String {
format!("http://{}", self.address)
}
}
impl Drop for TestHttpServer {
fn drop(&mut self) {
self.task.abort();
}
}
async fn read_request(stream: &mut tokio::net::TcpStream) -> Vec<u8> {
let mut bytes = Vec::new();
let header_end = loop {
let mut chunk = [0_u8; 1024];
let count = stream.read(&mut chunk).await.unwrap();
assert_ne!(count, 0);
bytes.extend_from_slice(&chunk[..count]);
if let Some(end) = bytes.windows(4).position(|window| window == b"\r\n\r\n") {
break end + 4;
}
};
let headers = String::from_utf8_lossy(&bytes[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
line.split_once(':')
.filter(|(name, _)| name.eq_ignore_ascii_case("content-length"))
.map(|(_, value)| value.trim())
})
.unwrap()
.parse::<usize>()
.unwrap();
while bytes.len() < header_end + content_length {
let mut chunk = [0_u8; 1024];
let count = stream.read(&mut chunk).await.unwrap();
assert_ne!(count, 0);
bytes.extend_from_slice(&chunk[..count]);
}
bytes
}
fn config(base: String, timeout: Option<Duration>) -> OpenAiEmbedderConfig {
OpenAiEmbedderConfig {
api_base: base,
api_key: "test-key".to_owned(),
model: "test-model".to_owned(),
timeout,
}
}
#[rstest]
#[tokio::test]
async fn posts_embeddings_request_and_parses_vector() {
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
let embedder = OpenAiEmbedder::new(
reqwest::Client::new(),
config(
format!("{}/", server.base_url()),
Some(Duration::from_secs(1)),
),
);
assert_eq!(embedder.model(), "test-model");
assert_eq!(
embedder
.async_embed("hello", Some(&json!({"ignored": true})))
.await
.unwrap(),
vec![0.1, 0.2]
);
let request = server.request.lock().unwrap().clone().unwrap();
let request_text = String::from_utf8(request).unwrap();
assert!(request_text.starts_with("POST /embeddings HTTP/1.1\r\n"));
assert!(request_text.contains("\r\nauthorization: Bearer test-key\r\n"));
let body = request_text.split("\r\n\r\n").nth(1).unwrap();
let body: Value = serde_json::from_str(body).unwrap();
assert_eq!(body["model"], "test-model");
assert_eq!(body["input"], "hello");
assert_eq!(body["encoding_format"], "float");
}
#[rstest]
#[case::error_status("500 Internal Server Error", "{}", 0, None, Err(Error::Unavailable))]
#[case::timed_out(
"200 OK",
r#"{"data":[{"embedding":[0.1,0.2]}]}"#,
500,
Some(Duration::from_millis(200)),
Err(Error::Unavailable)
)]
#[case::within_timeout(
"200 OK",
r#"{"data":[{"embedding":[0.1,0.2]}]}"#,
100,
Some(Duration::from_secs(1)),
Ok(vec![0.1, 0.2])
)]
#[case::missing_embedding("200 OK", r#"{"data":[]}"#, 0, None, Err(Error::Unavailable))]
#[tokio::test]
async fn status_timeout_and_body_errors_are_unavailable(
#[case] status: &str,
#[case] body: &str,
#[case] delay_ms: u64,
#[case] timeout: Option<Duration>,
#[case] expected: Result<Vec<f32>, Error>,
) {
let server =
TestHttpServer::response_after(status, body, Duration::from_millis(delay_ms)).await;
let embedder = OpenAiEmbedder::new(reqwest::Client::new(), config(server.base_url(), timeout));
assert_eq!(embedder.async_embed("hello", None).await, expected);
}
#[rstest]
fn sync_embedding_is_unsupported() {
let embedder = OpenAiEmbedder::new(
reqwest::Client::new(),
config("http://127.0.0.1:9".to_owned(), None),
);
assert_eq!(
embedder.embed("hello", None),
Err(Error::UnsupportedOperation)
);
}
#[rstest]
#[tokio::test]
async fn uses_the_injected_client() {
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
let client = reqwest::Client::builder()
.user_agent("litellm-embedder-test")
.build()
.unwrap();
let embedder = OpenAiEmbedder::new(client, config(server.base_url(), None));
assert_eq!(
embedder.async_embed("hello", None).await.unwrap(),
vec![0.1, 0.2]
);
let request = server.request.lock().unwrap().clone().unwrap();
let request_text = String::from_utf8(request).unwrap();
assert!(request_text.contains("\r\nuser-agent: litellm-embedder-test\r\n"));
}

View file

@ -0,0 +1,525 @@
mod support;
use std::{
collections::HashMap,
sync::{Arc, Mutex},
time::Duration,
};
use litellm_cache::{
BaseCache, CacheContext, Error, JsonCodec, SemanticCacheContext,
semantic::{Embedder, SemanticCache, SemanticLookup},
};
use litellm_cache_qdrant_semantic::{QdrantSemanticCache, QdrantSemanticConfig, Quantization};
use qdrant_client::{
Payload, Qdrant,
qdrant::{self, CompressionRatio, Distance, PointId, QuantizationType, Value, VectorParams},
};
use rstest::{fixture, rstest};
use serde_json::{Value as JsonValue, json};
use support::{FakeQdrant, FakeState, StoredPoint};
type Calls = Arc<Mutex<Vec<(String, Option<JsonValue>)>>>;
type Cache = QdrantSemanticCache<FixedEmbedder, JsonCodec<JsonValue>>;
/// Embeds known prompts, fails on anything else, and records every call.
#[derive(Clone)]
struct FixedEmbedder {
vectors: Arc<HashMap<String, Vec<f32>>>,
calls: Calls,
}
impl FixedEmbedder {
fn new(vectors: impl IntoIterator<Item = (&'static str, Vec<f32>)>) -> Self {
Self {
vectors: Arc::new(
vectors
.into_iter()
.map(|(prompt, vector)| (prompt.to_owned(), vector))
.collect(),
),
calls: Calls::default(),
}
}
}
impl Embedder for FixedEmbedder {
async fn async_embed(
&self,
input: &str,
metadata: Option<&JsonValue>,
) -> Result<Vec<f32>, Error> {
self.calls
.lock()
.unwrap()
.push((input.to_owned(), metadata.cloned()));
self.vectors.get(input).cloned().ok_or(Error::Unavailable)
}
}
fn config(quantization: Quantization) -> QdrantSemanticConfig {
QdrantSemanticConfig {
collection_name: "semantic".to_owned(),
similarity_threshold: 0.9,
vector_size: 2,
quantization,
}
}
fn context(prompt: &str) -> SemanticCacheContext {
SemanticCacheContext {
messages: Some(json!([{"role": "user", "content": prompt}])),
..Default::default()
}
}
#[fixture]
fn entry() -> JsonValue {
json!({"timestamp": 1.0, "response": {"answer": 42}})
}
async fn connect(
server: &FakeQdrant,
vectors: impl IntoIterator<Item = (&'static str, Vec<f32>)>,
) -> Cache {
let client = Qdrant::from_url(&server.url()).build().unwrap();
QdrantSemanticCache::connect(
client,
FixedEmbedder::new(vectors),
JsonCodec::new(),
config(Quantization::Binary),
tokio::runtime::Handle::current(),
)
.await
.unwrap()
}
#[rstest]
#[case::binary(Quantization::Binary)]
#[case::scalar(Quantization::Scalar)]
#[case::product(Quantization::Product)]
#[tokio::test(flavor = "multi_thread")]
async fn connect_sets_collection_quantization_and_index(#[case] quantization: Quantization) {
let server = FakeQdrant::start(FakeState::default()).await;
let client = Qdrant::from_url(&server.url()).build().unwrap();
QdrantSemanticCache::connect(
client,
FixedEmbedder::new([]),
JsonCodec::<JsonValue>::new(),
config(quantization.clone()),
tokio::runtime::Handle::current(),
)
.await
.unwrap();
let state = server.state.lock().unwrap();
let request = &state.created_collections[0];
let Some(qdrant::vectors_config::Config::Params(VectorParams { size, distance, .. })) = request
.vectors_config
.as_ref()
.and_then(|config| config.config.clone())
else {
panic!("missing vector params");
};
assert_eq!(size, 2);
assert_eq!(distance, Distance::Cosine as i32);
let quantization_config = request
.quantization_config
.as_ref()
.unwrap()
.quantization
.unwrap();
#[expect(
deprecated,
reason = "the test verifies Qdrant's legacy always_ram quantization contract"
)]
match (quantization, quantization_config) {
(Quantization::Binary, qdrant::quantization_config::Quantization::Binary(binary)) => {
assert_eq!(binary.always_ram, Some(false));
}
(Quantization::Scalar, qdrant::quantization_config::Quantization::Scalar(scalar)) => {
assert_eq!(scalar.r#type, QuantizationType::Int8 as i32);
assert_eq!(scalar.quantile, Some(0.99));
assert_eq!(scalar.always_ram, Some(false));
}
(Quantization::Product, qdrant::quantization_config::Quantization::Product(product)) => {
assert_eq!(product.compression, CompressionRatio::X16 as i32);
assert_eq!(product.always_ram, Some(false));
}
_ => panic!("unexpected quantization"),
}
assert!(state.index_creations >= 1);
assert_eq!(state.field_indexes[0].collection_name, "semantic");
assert_eq!(state.field_indexes[0].field_name, "litellm_cache_key");
assert_eq!(
state.field_indexes[0].field_type,
Some(qdrant::FieldType::Keyword as i32)
);
server.stop();
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn existing_collection_skips_create_and_index_failure_is_non_fatal() {
let server = FakeQdrant::start(FakeState {
collections: ["semantic".to_owned()].into_iter().collect(),
fail_field_index: true,
..Default::default()
})
.await;
let cache = connect(&server, [("hello", vec![1.0, 0.0])]).await;
assert_eq!(cache.collection_name(), "semantic");
assert_eq!(cache.similarity_threshold(), 0.9);
assert_eq!(cache.vector_size(), 2);
let state = server.state.lock().unwrap();
assert!(state.created_collections.is_empty());
assert!(state.index_creations >= 1);
server.stop();
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn async_and_sync_set_get_store_exact_payload(entry: JsonValue) {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = Arc::new(connect(&server, [("hello", vec![1.0, 0.0])]).await);
let ctx = SemanticCacheContext {
metadata: Some(json!({"tenant": "team"})),
..context("hello")
};
cache
.async_set_cache("key", entry.clone(), ctx.clone())
.await
.unwrap();
assert_eq!(
cache.async_get_cache("key", &ctx).await.unwrap().as_ref(),
Some(&entry)
);
{
let state = server.state.lock().unwrap();
let payload = &state.points[0].payload;
let mut payload_keys = payload.keys().cloned().collect::<Vec<_>>();
payload_keys.sort();
assert_eq!(payload_keys, ["litellm_cache_key", "response", "text"]);
assert_eq!(payload["litellm_cache_key"], Value::from("key"));
assert_eq!(payload["text"], Value::from("hello"));
assert_eq!(payload["response"], Value::from(entry.to_string()));
}
let sync_entry = entry.clone();
let sync_cache = cache.clone();
let sync_ctx = ctx.clone();
tokio::task::spawn_blocking(move || {
sync_cache
.set_cache("sync", sync_entry.clone(), &sync_ctx)
.unwrap();
assert_eq!(
sync_cache.get_cache("sync", &sync_ctx).unwrap(),
Some(sync_entry)
);
})
.await
.unwrap();
assert_eq!(
*cache.embedder().calls.lock().unwrap(),
vec![("hello".to_owned(), ctx.metadata.clone()); 4]
);
server.stop();
}
#[rstest]
#[case::content_parts_skip_images(
json!([
{"role": "user", "content": "hello"},
{
"role": "user",
"content": [
{"type": "text", "text": "world"},
{"type": "image_url", "image_url": {"url": "ignored"}},
{"type": "text", "text": "!"},
],
},
]),
"helloworld!"
)]
#[case::search_results_and_compact_citations(
json!([{
"role": "tool",
"content": null,
"search_results": [{
"source": "source",
"title": "title",
"content": [{"text": "body"}],
"citations": {"page": 1, "section": "intro"},
}],
}]),
r#"sourcetitlebody{"page":1,"section":"intro"}"#
)]
#[tokio::test(flavor = "multi_thread")]
async fn prompt_matches_python_message_rules(
#[case] messages: JsonValue,
#[case] prompt: &'static str,
entry: JsonValue,
) {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, [(prompt, vec![1.0, 0.0])]).await;
let context = SemanticCacheContext {
messages: Some(messages),
..Default::default()
};
cache.async_set_cache("key", entry, context).await.unwrap();
assert_eq!(cache.embedder().calls.lock().unwrap()[0].0, prompt);
assert_eq!(
server.state.lock().unwrap().points[0].payload["text"],
Value::from(prompt)
);
server.stop();
}
#[rstest]
#[case::no_messages(SemanticCacheContext::default())]
#[case::empty_messages(SemanticCacheContext { messages: Some(json!([])), ..Default::default() })]
#[case::responses_input_is_not_read(SemanticCacheContext { input: Some(json!("hello")), ..Default::default() })]
#[tokio::test(flavor = "multi_thread")]
async fn requests_without_messages_are_missing_a_prompt(
#[case] context: SemanticCacheContext,
entry: JsonValue,
) {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, [("hello", vec![1.0, 0.0])]).await;
assert_eq!(
cache.async_set_cache("key", entry, context.clone()).await,
Err(Error::MissingPrompt)
);
assert_eq!(
cache.async_get_cache("key", &context).await,
Err(Error::MissingPrompt)
);
assert!(cache.embedder().calls.lock().unwrap().is_empty());
server.stop();
}
#[rstest]
#[case::other_key("other", "hello", None)]
#[case::below_similarity_threshold("key", "near", None)]
#[tokio::test(flavor = "multi_thread")]
async fn misses_and_payload_validation_are_safe(
#[case] key: &str,
#[case] prompt: &str,
#[case] numeric_key_point: Option<u64>,
entry: JsonValue,
) {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(
&server,
[("hello", vec![1.0, 0.0]), ("near", vec![0.7, 0.71414286])],
)
.await;
cache
.async_set_cache("key", entry, context("hello"))
.await
.unwrap();
if let Some(id) = numeric_key_point {
server.insert_point(StoredPoint {
id: Some(PointId::from(id)),
vector: vec![1.0, 0.0],
payload: Payload::try_from(json!({
"litellm_cache_key": id,
"response": "{}",
}))
.unwrap()
.into(),
});
}
assert_eq!(
cache.async_get_cache(key, &context(prompt)).await.unwrap(),
None
);
server.stop();
}
#[rstest]
#[case::hit("key", context("hello"), Ok((true, Some(1.0))))]
#[case::below_similarity_threshold("key", context("near"), Ok((false, Some(0.7))))]
#[case::no_results("other", context("hello"), Ok((false, Some(0.0))))]
#[case::no_prompt("key", SemanticCacheContext::default(), Err(Error::MissingPrompt))]
#[tokio::test(flavor = "multi_thread")]
async fn lookup_reports_python_semantic_similarity(
#[case] key: &'static str,
#[case] context: SemanticCacheContext,
#[case] expected: Result<(bool, Option<f64>), Error>,
#[values(false, true)] use_async: bool,
entry: JsonValue,
) {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = Arc::new(
connect(
&server,
[("hello", vec![1.0, 0.0]), ("near", vec![0.7, 0.71414286])],
)
.await,
);
cache
.async_set_cache("key", entry.clone(), self::context("hello"))
.await
.unwrap();
server.insert_point(StoredPoint {
id: Some(PointId::from(99_u64)),
vector: vec![1.0, 0.0],
payload: Payload::try_from(json!({"litellm_cache_key": 99, "response": "{}"}))
.unwrap()
.into(),
});
let lookup = if use_async {
cache.async_get_cache_with_similarity(key, &context).await
} else {
let cache = Arc::clone(&cache);
tokio::task::spawn_blocking(move || cache.get_cache_with_similarity(key, &context))
.await
.unwrap()
};
match (lookup, expected) {
(Ok(SemanticLookup { value, similarity }), Ok((hit, expected))) => {
assert_eq!(value, hit.then_some(entry));
assert_eq!(similarity.is_some(), expected.is_some());
if let (Some(similarity), Some(expected)) = (similarity, expected) {
assert!((similarity - expected).abs() < 1e-6, "{similarity}");
}
}
(lookup, expected) => assert_eq!(lookup.map(|_| ()), expected.map(|_| ())),
}
server.stop();
}
#[rstest]
#[case::codec_decodes_the_payload(Some(json!("{\"a\":1}")), Ok(Some(json!({"a": 1}))))]
#[case::undecodable_response(Some(json!("not json")), Err(Error::InvalidEntry))]
#[case::non_string_response(Some(json!(1)), Err(Error::InvalidEntry))]
#[case::missing_response(None, Err(Error::InvalidEntry))]
#[tokio::test(flavor = "multi_thread")]
async fn stored_responses_go_through_the_codec(
#[case] response: Option<JsonValue>,
#[case] expected: Result<Option<JsonValue>, Error>,
) {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, [("hello", vec![1.0, 0.0])]).await;
let mut payload = serde_json::Map::new();
payload.insert("litellm_cache_key".to_owned(), json!("key"));
if let Some(response) = response {
payload.insert("response".to_owned(), response);
}
server.insert_point(StoredPoint {
id: Some(PointId::from(1_u64)),
vector: vec![1.0, 0.0],
payload: Payload::try_from(JsonValue::Object(payload))
.unwrap()
.into(),
});
assert_eq!(
cache.async_get_cache("key", &context("hello")).await,
expected
);
server.stop();
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn embedding_failures_propagate() {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, []).await;
assert_eq!(
cache.async_get_cache("key", &context("unknown")).await,
Err(Error::Unavailable)
);
server.stop();
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn ttl_is_ignored_and_entries_do_not_expire(entry: JsonValue) {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, [("one", vec![1.0, 0.0])]).await;
let ctx = context("one").with_ttl(Some(Duration::from_secs(1)));
assert_eq!(cache.get_ttl(&ctx), None);
cache
.async_set_cache("ttl", entry, ctx.clone())
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(1_100)).await;
assert!(cache.async_get_cache("ttl", &ctx).await.unwrap().is_some());
server.stop();
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn pipeline_upserts_each_entry_and_waits_for_indexing() {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, [("one", vec![1.0, 0.0])]).await;
cache
.async_set_cache_pipeline(
vec![
("one".to_owned(), json!({"n": 1})),
("two".to_owned(), json!({"n": 2})),
],
context("one"),
)
.await
.unwrap();
for (key, value) in [("one", json!({"n": 1})), ("two", json!({"n": 2}))] {
assert_eq!(
cache.async_get_cache(key, &context("one")).await.unwrap(),
Some(value)
);
}
assert_eq!(
server.state.lock().unwrap().upsert_waits,
vec![Some(true), Some(true)]
);
server.stop();
}
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn stopped_qdrant_server_maps_to_unavailable() {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, [("hello", vec![1.0, 0.0])]).await;
server.stop();
tokio::time::sleep(Duration::from_millis(50)).await;
assert_eq!(
cache.async_get_cache("key", &context("hello")).await,
Err(Error::Unavailable)
);
}
/// `_payload_matches_cache_key` compares `str(cached_key) == str(key)`, so a point whose stored
/// key is the number 99 answers a lookup for `"99"`.
#[rstest]
#[tokio::test(flavor = "multi_thread")]
async fn numeric_stored_cache_keys_match_like_python_str() {
let server = FakeQdrant::start(FakeState::default()).await;
let cache = connect(&server, [("hello", vec![1.0, 0.0])]).await;
server.insert_point(StoredPoint {
id: Some(PointId::from(99_u64)),
vector: vec![1.0, 0.0],
payload: Payload::try_from(json!({"litellm_cache_key": 99, "response": "{}"}))
.unwrap()
.into(),
});
let lookup = cache
.async_get_cache_with_similarity("99", &context("hello"))
.await
.unwrap();
assert_eq!(lookup.value, Some(json!({})));
assert!((lookup.similarity.unwrap() - 1.0).abs() < 1e-6);
server.stop();
}

View file

@ -0,0 +1,343 @@
#![allow(dead_code)]
use std::{
collections::{HashMap, HashSet},
net::SocketAddr,
sync::{Arc, Mutex},
};
use qdrant_client::qdrant::{
self, CollectionExists, CollectionExistsRequest, CollectionExistsResponse,
CollectionOperationResponse, CreateCollection, CreateFieldIndexCollection, Filter, PointId,
PointsOperationResponse, ScoredPoint, SearchPoints, SearchResponse, Value, Vector, Vectors,
collections_server::{Collections, CollectionsServer},
points_server::{Points, PointsServer},
};
use tokio::sync::oneshot;
use tokio_stream::wrappers::TcpListenerStream;
use tonic::{Request, Response, Status, transport::Server};
#[derive(Clone, Debug)]
pub struct StoredPoint {
pub id: Option<PointId>,
pub vector: Vec<f32>,
pub payload: HashMap<String, Value>,
}
#[derive(Default)]
pub struct FakeState {
pub collections: HashSet<String>,
pub created_collections: Vec<CreateCollection>,
pub field_indexes: Vec<CreateFieldIndexCollection>,
pub points: Vec<StoredPoint>,
pub upsert_waits: Vec<Option<bool>>,
pub index_creations: usize,
pub fail_field_index: bool,
}
#[derive(Clone)]
pub struct FakeQdrant {
pub state: Arc<Mutex<FakeState>>,
pub address: SocketAddr,
shutdown: Arc<Mutex<Option<oneshot::Sender<()>>>>,
}
impl FakeQdrant {
pub async fn start(state: FakeState) -> Self {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let state = Arc::new(Mutex::new(state));
let service = FakeService {
state: state.clone(),
};
let (shutdown_tx, shutdown_rx) = oneshot::channel();
tokio::spawn(async move {
Server::builder()
.add_service(CollectionsServer::new(service.clone()))
.add_service(PointsServer::new(service))
.serve_with_incoming_shutdown(TcpListenerStream::new(listener), async {
let _ = shutdown_rx.await;
})
.await
.unwrap();
});
Self {
state,
address,
shutdown: Arc::new(Mutex::new(Some(shutdown_tx))),
}
}
pub fn url(&self) -> String {
format!("http://{}", self.address)
}
pub fn stop(&self) {
self.shutdown
.lock()
.unwrap()
.take()
.unwrap()
.send(())
.unwrap();
}
pub fn insert_point(&self, point: StoredPoint) {
self.state.lock().unwrap().points.push(point);
}
}
#[derive(Clone)]
struct FakeService {
state: Arc<Mutex<FakeState>>,
}
macro_rules! unimplemented_collections {
($($name:ident, $request:ty, $response:ty);* $(;)?) => {
$(
fn $name<'life0, 'async_trait>(
&'life0 self,
_: Request<$request>,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<Response<$response>, Status>,
> + Send
+ 'async_trait,
>,
>
where
'life0: 'async_trait,
Self: 'async_trait,
{
Box::pin(async { Err(Status::unimplemented(stringify!($name))) })
}
)*
};
}
macro_rules! unimplemented_points {
($($name:ident, $request:ty, $response:ty);* $(;)?) => {
$(
fn $name<'life0, 'async_trait>(
&'life0 self,
_: Request<$request>,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<Response<$response>, Status>,
> + Send
+ 'async_trait,
>,
>
where
'life0: 'async_trait,
Self: 'async_trait,
{
Box::pin(async { Err(Status::unimplemented(stringify!($name))) })
}
)*
};
}
#[tonic::async_trait]
impl Collections for FakeService {
async fn create(
&self,
request: Request<CreateCollection>,
) -> Result<Response<CollectionOperationResponse>, Status> {
let request = request.into_inner();
let mut state = self.state.lock().unwrap();
state.collections.insert(request.collection_name.clone());
state.created_collections.push(request);
Ok(Response::new(CollectionOperationResponse {
result: true,
..Default::default()
}))
}
async fn collection_exists(
&self,
request: Request<CollectionExistsRequest>,
) -> Result<Response<CollectionExistsResponse>, Status> {
let exists = self
.state
.lock()
.unwrap()
.collections
.contains(&request.into_inner().collection_name);
Ok(Response::new(CollectionExistsResponse {
result: Some(CollectionExists { exists }),
..Default::default()
}))
}
unimplemented_collections!(
get, qdrant::GetCollectionInfoRequest, qdrant::GetCollectionInfoResponse;
list, qdrant::ListCollectionsRequest, qdrant::ListCollectionsResponse;
update, qdrant::UpdateCollection, qdrant::CollectionOperationResponse;
delete, qdrant::DeleteCollection, qdrant::CollectionOperationResponse;
update_aliases, qdrant::ChangeAliases, qdrant::CollectionOperationResponse;
list_collection_aliases, qdrant::ListCollectionAliasesRequest, qdrant::ListAliasesResponse;
list_aliases, qdrant::ListAliasesRequest, qdrant::ListAliasesResponse;
collection_cluster_info, qdrant::CollectionClusterInfoRequest, qdrant::CollectionClusterInfoResponse;
update_collection_cluster_setup, qdrant::UpdateCollectionClusterSetupRequest, qdrant::UpdateCollectionClusterSetupResponse;
create_shard_key, qdrant::CreateShardKeyRequest, qdrant::CreateShardKeyResponse;
delete_shard_key, qdrant::DeleteShardKeyRequest, qdrant::DeleteShardKeyResponse;
list_shard_keys, qdrant::ListShardKeysRequest, qdrant::ListShardKeysResponse;
);
}
#[tonic::async_trait]
impl Points for FakeService {
async fn create_field_index(
&self,
request: Request<CreateFieldIndexCollection>,
) -> Result<Response<PointsOperationResponse>, Status> {
let mut state = self.state.lock().unwrap();
state.index_creations += 1;
state.field_indexes.push(request.into_inner());
if state.fail_field_index {
return Err(Status::internal("field index failure"));
}
Ok(Response::new(PointsOperationResponse::default()))
}
async fn upsert(
&self,
request: Request<qdrant::UpsertPoints>,
) -> Result<Response<PointsOperationResponse>, Status> {
let request = request.into_inner();
let mut state = self.state.lock().unwrap();
state.upsert_waits.push(request.wait);
for point in request.points {
let stored = StoredPoint {
id: point.id.clone(),
vector: dense_vector(point.vectors)?,
payload: point.payload,
};
if let Some(existing) = state
.points
.iter_mut()
.find(|existing| existing.id == stored.id)
{
*existing = stored;
} else {
state.points.push(stored);
}
}
Ok(Response::new(PointsOperationResponse::default()))
}
async fn search(
&self,
request: Request<SearchPoints>,
) -> Result<Response<SearchResponse>, Status> {
let request = request.into_inner();
let key_filter = keyword_filter(request.filter.as_ref());
let state = self.state.lock().unwrap();
let mut results = state
.points
.iter()
.filter(|point| {
key_filter.as_ref().is_none_or(|(field, expected)| {
point
.payload
.get(field)
.and_then(|value| {
let value: serde_json::Value = value.clone().into();
value
.as_str()
.map(str::to_owned)
.or_else(|| value.as_i64().map(|value| value.to_string()))
})
.is_some_and(|value| value == *expected)
})
})
.map(|point| ScoredPoint {
id: point.id.clone(),
payload: point.payload.clone(),
score: cosine(&request.vector, &point.vector),
..Default::default()
})
.collect::<Vec<_>>();
results.sort_by(|left, right| right.score.total_cmp(&left.score));
results.truncate(request.limit as usize);
Ok(Response::new(SearchResponse {
result: results,
..Default::default()
}))
}
unimplemented_points!(
delete, qdrant::DeletePoints, qdrant::PointsOperationResponse;
get, qdrant::GetPoints, qdrant::GetResponse;
update_vectors, qdrant::UpdatePointVectors, qdrant::PointsOperationResponse;
delete_vectors, qdrant::DeletePointVectors, qdrant::PointsOperationResponse;
set_payload, qdrant::SetPayloadPoints, qdrant::PointsOperationResponse;
overwrite_payload, qdrant::SetPayloadPoints, qdrant::PointsOperationResponse;
delete_payload, qdrant::DeletePayloadPoints, qdrant::PointsOperationResponse;
clear_payload, qdrant::ClearPayloadPoints, qdrant::PointsOperationResponse;
delete_field_index, qdrant::DeleteFieldIndexCollection, qdrant::PointsOperationResponse;
create_vector_name, qdrant::CreateVectorNameRequest, qdrant::PointsOperationResponse;
delete_vector_name, qdrant::DeleteVectorNameRequest, qdrant::PointsOperationResponse;
search_batch, qdrant::SearchBatchPoints, qdrant::SearchBatchResponse;
search_groups, qdrant::SearchPointGroups, qdrant::SearchGroupsResponse;
scroll, qdrant::ScrollPoints, qdrant::ScrollResponse;
recommend, qdrant::RecommendPoints, qdrant::RecommendResponse;
recommend_batch, qdrant::RecommendBatchPoints, qdrant::RecommendBatchResponse;
recommend_groups, qdrant::RecommendPointGroups, qdrant::RecommendGroupsResponse;
discover, qdrant::DiscoverPoints, qdrant::DiscoverResponse;
discover_batch, qdrant::DiscoverBatchPoints, qdrant::DiscoverBatchResponse;
count, qdrant::CountPoints, qdrant::CountResponse;
update_batch, qdrant::UpdateBatchPoints, qdrant::UpdateBatchResponse;
query, qdrant::QueryPoints, qdrant::QueryResponse;
query_batch, qdrant::QueryBatchPoints, qdrant::QueryBatchResponse;
query_groups, qdrant::QueryPointGroups, qdrant::QueryGroupsResponse;
facet, qdrant::FacetCounts, qdrant::FacetResponse;
search_matrix_pairs, qdrant::SearchMatrixPoints, qdrant::SearchMatrixPairsResponse;
search_matrix_offsets, qdrant::SearchMatrixPoints, qdrant::SearchMatrixOffsetsResponse;
);
}
fn dense_vector(vectors: Option<Vectors>) -> Result<Vec<f32>, Status> {
let Some(Vectors {
vectors_options:
Some(qdrant::vectors::VectorsOptions::Vector(Vector {
vector: Some(qdrant::vector::Vector::Dense(qdrant::DenseVector { data })),
..
})),
}) = vectors
else {
return Err(Status::invalid_argument("expected dense vector"));
};
Ok(data)
}
fn keyword_filter(filter: Option<&Filter>) -> Option<(String, String)> {
filter?
.must
.iter()
.find_map(|condition| match condition.condition_one_of.as_ref()? {
qdrant::condition::ConditionOneOf::Field(field) => {
let qdrant::r#match::MatchValue::Keyword(value) =
field.r#match.as_ref()?.match_value.as_ref()?
else {
return None;
};
Some((field.key.clone(), value.clone()))
}
_ => None,
})
}
fn cosine(left: &[f32], right: &[f32]) -> f32 {
let dot = left
.iter()
.zip(right)
.map(|(left, right)| left * right)
.sum::<f32>();
let left_norm = left.iter().map(|value| value * value).sum::<f32>().sqrt();
let right_norm = right.iter().map(|value| value * value).sum::<f32>().sqrt();
dot / (left_norm * right_norm)
}

View file

@ -0,0 +1,19 @@
[package]
name = "litellm-cache-redis-semantic"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-cache.workspace = true
litellm-cache-redis.workspace = true
redis = { version = "1.7.0", features = ["tls-rustls"] }
sha2.workspace = true
[dev-dependencies]
litellm-cache-testing.workspace = true
redis-test = "1.0.4"
rstest.workspace = true
serde_json.workspace = true
tokio.workspace = true

View file

@ -0,0 +1,404 @@
use std::{
sync::Arc,
time::{Duration, SystemTime, UNIX_EPOCH},
};
use litellm_cache::{
BaseCache, CacheCodec, Error, SemanticCacheContext,
semantic::{Embedder, SemanticCache, SemanticLookup, prompt_from_context},
};
use litellm_cache_redis::{
RedisTopology,
connection::{ConnectionRef, Connections},
};
use sha2::{Digest, Sha256};
use crate::{
RedisSemanticConfig,
index::{CACHE_KEY_FIELD, Index, VECTOR_FIELD},
reply::{bytes_field, first_document, number_field, string_field},
};
struct Inner {
index: Index,
distance_threshold: f64,
clock: fn() -> f64,
}
impl Inner {
fn new(config: RedisSemanticConfig, clock: fn() -> f64) -> Self {
Self {
index: Index::new(config.index_name),
distance_threshold: 1.0 - f64::from(config.similarity_threshold),
clock,
}
}
fn store(
&self,
connection: &mut ConnectionRef<'_>,
tag: &str,
response: Vec<u8>,
prompt: &str,
vector: &[f32],
ttl: Option<Duration>,
) -> Result<(), Error> {
let index = self.index.ensure(connection, vector.len())?;
let entry_id = entry_id(prompt, tag);
let hash_key = format!("{index}:{entry_id}");
redis::cmd("HSET")
.arg(&hash_key)
.arg("entry_id")
.arg(&entry_id)
.arg("prompt")
.arg(prompt)
.arg("response")
.arg(response)
.arg(VECTOR_FIELD)
.arg(vector_buffer(vector))
.arg("inserted_at")
.arg(format!("{}", (self.clock)()))
.arg("updated_at")
.arg(format!("{}", (self.clock)()))
.arg(CACHE_KEY_FIELD)
.arg(tag)
.query::<()>(connection)
.map_err(|_| Error::Unavailable)?;
if let Some(ttl) = ttl {
redis::cmd("EXPIRE")
.arg(&hash_key)
.arg(ttl_seconds(ttl))
.query::<()>(connection)
.map_err(|_| Error::Unavailable)?;
}
Ok(())
}
fn lookup(
&self,
connection: &mut ConnectionRef<'_>,
tag: &str,
vector: &[f32],
) -> Result<SemanticLookup<Vec<u8>>, Error> {
let index = self.index.ensure(connection, vector.len())?;
let query = format!(
"(@{CACHE_KEY_FIELD}:{{{}}})=>[KNN 1 @{VECTOR_FIELD} $vector AS vector_distance]",
escape_tag(tag)
);
let result = redis::cmd("FT.SEARCH")
.arg(&index)
.arg(query)
.arg("RETURN")
.arg(8)
.arg("entry_id")
.arg("prompt")
.arg("response")
.arg("inserted_at")
.arg("updated_at")
.arg("metadata")
.arg(CACHE_KEY_FIELD)
.arg("vector_distance")
.arg("SORTBY")
.arg("vector_distance")
.arg("ASC")
.arg("DIALECT")
.arg(2)
.arg("LIMIT")
.arg(0)
.arg(1)
.arg("PARAMS")
.arg(2)
.arg("vector")
.arg(vector_buffer(vector))
.query::<redis::Value>(connection)
.map_err(|_| Error::Unavailable)?;
let Some(fields) = first_document(&result) else {
return Ok(SemanticLookup::miss(Some(0.0)));
};
if string_field(fields, CACHE_KEY_FIELD).as_deref() != Some(tag) {
return Ok(SemanticLookup::miss(Some(0.0)));
}
// redisvl's range query only returns entries within the distance threshold, so a
// farther hit reads as no result.
let Some(distance) = number_field(fields, "vector_distance")
.filter(|distance| *distance <= self.distance_threshold)
else {
return Ok(SemanticLookup::miss(Some(0.0)));
};
let Some(response) = bytes_field(fields, "response") else {
return Ok(SemanticLookup::miss(Some(0.0)));
};
Ok(SemanticLookup {
value: Some(response),
similarity: Some(1.0 - distance),
})
}
}
/// `RedisSemanticCache`: a redisvl-compatible semantic index on Redis Stack. Values go through
/// the injected codec, so the response layer decides what a cached entry is.
pub struct RedisSemanticCache<E, S, C = redis::Connection> {
connections: Arc<Connections<C>>,
embedder: E,
codec: S,
inner: Arc<Inner>,
}
impl<E: Embedder, S: CacheCodec> RedisSemanticCache<E, S> {
pub fn new(
url: &str,
embedder: E,
codec: S,
config: RedisSemanticConfig,
) -> Result<Self, Error> {
Ok(Self {
connections: Arc::new(Connections::open(url, &RedisTopology::Standalone)?),
embedder,
codec,
inner: Arc::new(Inner::new(config, timestamp)),
})
}
}
impl<E, S, C> RedisSemanticCache<E, S, C>
where
E: Embedder,
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
pub fn with_connection(
connection: C,
embedder: E,
codec: S,
config: RedisSemanticConfig,
) -> Self {
Self {
connections: Arc::new(Connections::fixed(connection)),
embedder,
codec,
inner: Arc::new(Inner::new(config, timestamp)),
}
}
pub fn with_clock(self, clock: fn() -> f64) -> Self {
let config = RedisSemanticConfig {
index_name: self.index_name().to_owned(),
similarity_threshold: self.similarity_threshold(),
};
Self {
inner: Arc::new(Inner::new(config, clock)),
..self
}
}
pub fn embedder(&self) -> &E {
&self.embedder
}
pub fn index_name(&self) -> &str {
self.inner.index.name()
}
pub fn similarity_threshold(&self) -> f32 {
(1.0 - self.inner.distance_threshold) as f32
}
fn tag<'a>(key: &'a str, context: &'a SemanticCacheContext) -> &'a str {
context.scope.as_deref().unwrap_or(key)
}
fn decode(&self, lookup: SemanticLookup<Vec<u8>>) -> Result<SemanticLookup<S::Value>, Error> {
Ok(SemanticLookup {
value: lookup
.value
.map(|bytes| self.codec.decode(&bytes))
.transpose()?,
similarity: lookup.similarity,
})
}
}
impl<E, S, C> BaseCache for RedisSemanticCache<E, S, C>
where
E: Embedder,
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
type Value = S::Value;
type Context = SemanticCacheContext;
fn get_ttl(&self, context: &Self::Context) -> Option<Duration> {
context.ttl
}
fn set_cache(
&self,
key: &str,
value: Self::Value,
context: &Self::Context,
) -> Result<(), Error> {
let Some(prompt) = prompt_from_context(context) else {
return Ok(());
};
let response = self.codec.encode(&value)?;
let vector = self.embedder.embed(&prompt, context.metadata.as_ref())?;
let tag = Self::tag(key, context);
self.connections.execute(|connection| {
self.inner
.store(connection, tag, response, &prompt, &vector, context.ttl)
})
}
fn get_cache(&self, key: &str, context: &Self::Context) -> Result<Option<Self::Value>, Error> {
self.get_cache_with_similarity(key, context)
.map(|lookup| lookup.value)
}
async fn async_set_cache(
&self,
key: &str,
value: Self::Value,
context: Self::Context,
) -> Result<(), Error> {
let Some(prompt) = prompt_from_context(&context) else {
return Ok(());
};
let response = self.codec.encode(&value)?;
let vector = self
.embedder
.async_embed(&prompt, context.metadata.as_ref())
.await?;
let tag = Self::tag(key, &context).to_owned();
let inner = Arc::clone(&self.inner);
Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
inner.store(connection, &tag, response, &prompt, &vector, context.ttl)
})
.await
}
async fn async_get_cache(
&self,
key: &str,
context: &Self::Context,
) -> Result<Option<Self::Value>, Error> {
self.async_get_cache_with_similarity(key, context)
.await
.map(|lookup| lookup.value)
}
}
/// Python stamps a similarity of `0.0` when there is no prompt or no hit in the key's scope.
impl<E, S, C> SemanticCache for RedisSemanticCache<E, S, C>
where
E: Embedder,
S: CacheCodec,
C: redis::ConnectionLike + Send + 'static,
{
fn get_cache_with_similarity(
&self,
key: &str,
context: &Self::Context,
) -> Result<SemanticLookup<Self::Value>, Error> {
let Some(prompt) = prompt_from_context(context) else {
return Ok(SemanticLookup::miss(Some(0.0)));
};
let vector = self.embedder.embed(&prompt, context.metadata.as_ref())?;
let tag = Self::tag(key, context);
let lookup = self
.connections
.execute(|connection| self.inner.lookup(connection, tag, &vector))?;
self.decode(lookup)
}
async fn async_get_cache_with_similarity(
&self,
key: &str,
context: &Self::Context,
) -> Result<SemanticLookup<Self::Value>, Error> {
let Some(prompt) = prompt_from_context(context) else {
return Ok(SemanticLookup::miss(Some(0.0)));
};
let vector = self
.embedder
.async_embed(&prompt, context.metadata.as_ref())
.await?;
let tag = Self::tag(key, context).to_owned();
let inner = Arc::clone(&self.inner);
let lookup = Connections::run_blocking(Arc::clone(&self.connections), move |connection| {
inner.lookup(connection, &tag, &vector)
})
.await?;
self.decode(lookup)
}
}
fn timestamp() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs_f64())
.unwrap_or_default()
}
fn entry_id(prompt: &str, tag: &str) -> String {
let mut digest = Sha256::new();
digest.update(prompt.as_bytes());
digest.update(CACHE_KEY_FIELD.as_bytes());
digest.update(tag.as_bytes());
format!("{:x}", digest.finalize())
}
fn vector_buffer(vector: &[f32]) -> Vec<u8> {
vector
.iter()
.flat_map(|component| component.to_le_bytes())
.collect()
}
fn escape_tag(value: &str) -> String {
let mut escaped = String::with_capacity(value.len());
for ch in value.chars() {
if matches!(
ch,
',' | '.'
| '<'
| '>'
| '{'
| '}'
| '['
| ']'
| '\\'
| '"'
| '\''
| ':'
| ';'
| '!'
| '@'
| '#'
| '$'
| '%'
| '^'
| '&'
| '*'
| '('
| ')'
| '-'
| '+'
| '='
| '~'
| '|'
| '/'
| ' '
| '?'
) {
escaped.push('\\');
}
escaped.push(ch);
}
escaped
}
fn ttl_seconds(ttl: Duration) -> u64 {
ttl.as_secs()
.saturating_add(u64::from(ttl.subsec_nanos() > 0))
.max(1)
}

View file

@ -0,0 +1,8 @@
/// `RedisSemanticCache.DEFAULT_REDIS_INDEX_NAME`.
pub const DEFAULT_INDEX_NAME: &str = "litellm_semantic_cache_index";
#[derive(Clone, Debug)]
pub struct RedisSemanticConfig {
pub index_name: String,
pub similarity_threshold: f32,
}

View file

@ -0,0 +1,205 @@
use std::sync::OnceLock;
use litellm_cache::Error;
use litellm_cache_redis::connection::ConnectionRef;
use crate::reply::{number_value, string_value};
pub(crate) const CACHE_KEY_FIELD: &str = "litellm_cache_key";
pub(crate) const VECTOR_FIELD: &str = "prompt_vector";
/// The redisvl `SemanticCache` index, resolved once per cache: the configured name when its
/// schema fits, else `<name>_isolated`, recreated when that one is stale too.
pub(crate) struct Index {
name: String,
resolved: OnceLock<String>,
}
impl Index {
pub(crate) fn new(name: String) -> Self {
Self {
name,
resolved: OnceLock::new(),
}
}
pub(crate) fn name(&self) -> &str {
&self.name
}
pub(crate) fn ensure(
&self,
connection: &mut ConnectionRef<'_>,
dims: usize,
) -> Result<String, Error> {
if let Some(name) = self.resolved.get() {
return Ok(name.clone());
}
let name = match index_compatible(connection, &self.name, dims)? {
Some(true) => self.name.clone(),
Some(false) => self.isolated(connection, dims)?,
None => match create_index(connection, &self.name, dims) {
Ok(()) => self.name.clone(),
Err(_) => match index_compatible(connection, &self.name, dims)? {
Some(true) => self.name.clone(),
Some(false) => self.isolated(connection, dims)?,
None => return Err(Error::Unavailable),
},
},
};
let _ = self.resolved.set(name.clone());
Ok(name)
}
fn isolated(&self, connection: &mut ConnectionRef<'_>, dims: usize) -> Result<String, Error> {
let name = format!("{}_isolated", self.name);
match index_compatible(connection, &name, dims)? {
Some(true) => Ok(name),
Some(false) => {
redis::cmd("FT.DROPINDEX")
.arg(&name)
.query::<()>(connection)
.map_err(|_| Error::Unavailable)?;
create_index(connection, &name, dims)?;
Ok(name)
}
None => {
create_index(connection, &name, dims)?;
Ok(name)
}
}
}
}
fn create_index(connection: &mut ConnectionRef<'_>, name: &str, dims: usize) -> Result<(), Error> {
redis::cmd("FT.CREATE")
.arg(name)
.arg("ON")
.arg("HASH")
.arg("PREFIX")
.arg(1)
.arg(name)
.arg("SCORE")
.arg(1.0)
.arg("SCHEMA")
.arg("prompt")
.arg("TEXT")
.arg("WEIGHT")
.arg(1)
.arg("response")
.arg("TEXT")
.arg("WEIGHT")
.arg(1)
.arg("inserted_at")
.arg("NUMERIC")
.arg("updated_at")
.arg("NUMERIC")
.arg(VECTOR_FIELD)
.arg("VECTOR")
.arg("FLAT")
.arg(6)
.arg("TYPE")
.arg("FLOAT32")
.arg("DIM")
.arg(dims)
.arg("DISTANCE_METRIC")
.arg("COSINE")
.arg(CACHE_KEY_FIELD)
.arg("TAG")
.arg("SEPARATOR")
.arg(",")
.query::<()>(connection)
.map_err(|_| Error::Unavailable)
}
fn index_compatible(
connection: &mut ConnectionRef<'_>,
name: &str,
dims: usize,
) -> Result<Option<bool>, Error> {
let info = match redis::cmd("FT.INFO")
.arg(name)
.query::<redis::Value>(connection)
{
Ok(info) => info,
Err(error) if unknown_index(&error) => return Ok(None),
Err(_) => return Err(Error::Unavailable),
};
Ok(Some(schema_compatible(&info, dims)))
}
fn unknown_index(error: &redis::RedisError) -> bool {
let message = error.to_string().to_lowercase();
message.contains("unknown") && message.contains("index")
}
struct Attribute {
name: Option<String>,
field_type: Option<String>,
dim: Option<f64>,
data_type: Option<String>,
distance_metric: Option<String>,
}
fn attribute(value: &redis::Value) -> Option<Attribute> {
let redis::Value::Array(pairs) = value else {
return None;
};
let mut attribute = Attribute {
name: None,
field_type: None,
dim: None,
data_type: None,
distance_metric: None,
};
for pair in pairs.as_chunks::<2>().0 {
match string_value(&pair[0]).as_deref() {
Some("identifier") => attribute.name = string_value(&pair[1]),
Some("type") => attribute.field_type = string_value(&pair[1]),
Some("dim") => attribute.dim = number_value(&pair[1]),
Some("data_type") => attribute.data_type = string_value(&pair[1]),
Some("distance_metric") => attribute.distance_metric = string_value(&pair[1]),
_ => {}
}
}
Some(attribute)
}
fn schema_compatible(info: &redis::Value, dims: usize) -> bool {
let redis::Value::Array(entries) = info else {
return false;
};
let attributes = entries
.as_chunks::<2>()
.0
.iter()
.find(|pair| string_value(&pair[0]).as_deref() == Some("attributes"))
.map(|pair| &pair[1]);
let Some(redis::Value::Array(attributes)) = attributes else {
return false;
};
let fields = attributes.iter().filter_map(attribute).collect::<Vec<_>>();
let has_field = |name: &str, field_type: &str| {
fields.iter().any(|field| {
field.name.as_deref() == Some(name) && field.field_type.as_deref() == Some(field_type)
})
};
has_field("prompt", "TEXT")
&& has_field("response", "TEXT")
&& has_field("inserted_at", "NUMERIC")
&& has_field("updated_at", "NUMERIC")
&& has_field(CACHE_KEY_FIELD, "TAG")
&& fields.iter().any(|field| {
field.name.as_deref() == Some(VECTOR_FIELD)
&& field.field_type.as_deref() == Some("VECTOR")
&& field.dim == Some(dims as f64)
&& field
.data_type
.as_deref()
.is_some_and(|data| data.eq_ignore_ascii_case("float32"))
&& field
.distance_metric
.as_deref()
.is_some_and(|metric| metric.eq_ignore_ascii_case("cosine"))
})
}

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