mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
merge: sync with upstream main
This commit is contained in:
commit
f7d6bb579d
807 changed files with 61755 additions and 14485 deletions
|
|
@ -17,3 +17,6 @@ rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
|||
|
||||
[target.aarch64-apple-darwin]
|
||||
rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
||||
|
||||
[env]
|
||||
SQLX_OFFLINE = "true"
|
||||
|
|
|
|||
|
|
@ -3419,7 +3419,7 @@ workflows:
|
|||
name: integration-<< matrix.suite >>
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, mcp, sdk, cost, browser]
|
||||
suite: [management, accounting, database, providers, mcp, sdk, cost, security, browser]
|
||||
- integration_contracts:
|
||||
name: integration-extensions
|
||||
suite: extensions
|
||||
|
|
|
|||
|
|
@ -89,6 +89,7 @@ legacy_paths() {
|
|||
proxy-db-auth-checks)
|
||||
echo tests/unit/proxy/auth/test_auth_checks.py
|
||||
echo tests/unit/proxy/auth/test_user_api_key_auth.py
|
||||
echo tests/unit/proxy/test_credential_slot_registry.py
|
||||
echo tests/unit/proxy/test_deprecated_key_grace_period.py ;;
|
||||
proxy-db-budgets)
|
||||
echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ case "$subject" in
|
|||
;;
|
||||
esac
|
||||
|
||||
ALLOWED_TYPES="feat|fix|docs|style|refactor|perf|test|build|ci|chore|revert"
|
||||
ALLOWED_TYPES="feat|fix|docs|style|refactor|perf|test|build|ci|chore|revert|security"
|
||||
# Description must not start with an uppercase letter — kept in sync with the
|
||||
# subjectPattern in .github/workflows/conventional-commits.yml so the local
|
||||
# hook is the strictly tighter of the two gates. (Without this guard, a commit
|
||||
|
|
@ -61,7 +61,7 @@ cat >&2 <<EOF
|
|||
Expected: <type>(<scope>)!: <description>
|
||||
(description must start with a lowercase letter)
|
||||
|
||||
Allowed types: feat, fix, docs, style, refactor, perf, test, build, ci, chore, revert
|
||||
Allowed types: feat, fix, docs, style, refactor, perf, test, build, ci, chore, revert, security
|
||||
Examples:
|
||||
feat(router): add weighted round-robin strategy
|
||||
fix(bedrock): decouple STS region from aws_region_name
|
||||
|
|
|
|||
8
.github/CODEOWNERS
vendored
8
.github/CODEOWNERS
vendored
|
|
@ -1,10 +1,2 @@
|
|||
/ui/ @yuneng-berri @ryan-crabbe-berri
|
||||
/litellm/proxy/_experimental/out/ @yuneng-berri @ryan-crabbe-berri
|
||||
/ui/Dockerfile
|
||||
/ui/nginx.conf
|
||||
/ui/litellm-dashboard/src/lib/http/schema.d.ts
|
||||
/ui/litellm-dashboard/tsconfig.tsbuildinfo
|
||||
/model_prices_and_context_window.json @mateo-berri @ryan-crabbe-berri @kerry-berri
|
||||
/litellm/model_prices_and_context_window_backup.json @mateo-berri @ryan-crabbe-berri @kerry-berri
|
||||
/litellm-proxy-extras/litellm_proxy_extras/migrations/ @yuneng-berri @ryan-crabbe-berri
|
||||
/.github/CODEOWNERS @yuneng-berri
|
||||
|
|
|
|||
7
.github/ci-coverage-allowlist.yml
vendored
7
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -111,3 +111,10 @@ dockerfiles:
|
|||
An example image under cookbook/ that is documentation rather than a shipped artifact
|
||||
paths:
|
||||
- cookbook/litellm-ollama-docker-image/Dockerfile
|
||||
- reason: >-
|
||||
The Rust gateway image compiles the whole workspace in release mode, which is too slow for
|
||||
a per-pull-request job while the gateway binary is still being assembled; the Rust lint,
|
||||
clippy, and compile jobs already cover the code it packages. Revisit when the gateway is
|
||||
published
|
||||
paths:
|
||||
- litellm-rust/crates/gateway/Dockerfile
|
||||
|
|
|
|||
1
.github/workflows/conventional-commits.yml
vendored
1
.github/workflows/conventional-commits.yml
vendored
|
|
@ -41,6 +41,7 @@ jobs:
|
|||
ci
|
||||
chore
|
||||
revert
|
||||
security
|
||||
requireScope: false
|
||||
subjectPattern: ^(?![A-Z]).+$
|
||||
subjectPatternError: |
|
||||
|
|
|
|||
13
.github/workflows/create-rc-branch.yml
vendored
13
.github/workflows/create-rc-branch.yml
vendored
|
|
@ -15,6 +15,8 @@ jobs:
|
|||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
outputs:
|
||||
version: ${{ steps.version.outputs.version }}
|
||||
steps:
|
||||
- name: Require main
|
||||
env:
|
||||
|
|
@ -64,3 +66,14 @@ jobs:
|
|||
sha: context.sha,
|
||||
});
|
||||
core.info(`Created branch ${branchName} at ${context.sha}`);
|
||||
|
||||
linear-release:
|
||||
name: Move the Linear release to rc
|
||||
needs: create-rc-branch
|
||||
permissions:
|
||||
contents: read
|
||||
uses: ./.github/workflows/linear-release.yml
|
||||
with:
|
||||
rc_version: ${{ needs.create-rc-branch.outputs.version }}
|
||||
secrets:
|
||||
LINEAR_API_KEY: ${{ secrets.LINEAR_API_KEY }}
|
||||
|
|
|
|||
131
.github/workflows/linear-release.yml
vendored
Normal file
131
.github/workflows/linear-release.yml
vendored
Normal file
|
|
@ -0,0 +1,131 @@
|
|||
name: Linear Release
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- "rc/**"
|
||||
release:
|
||||
types: [published]
|
||||
workflow_call:
|
||||
inputs:
|
||||
rc_version:
|
||||
description: "X.Y.0 release whose rc branch was just cut"
|
||||
required: true
|
||||
type: string
|
||||
secrets:
|
||||
LINEAR_API_KEY:
|
||||
required: true
|
||||
|
||||
permissions: {}
|
||||
|
||||
jobs:
|
||||
linear-release:
|
||||
name: Linear Release
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Plan
|
||||
id: plan
|
||||
env:
|
||||
EVENT: ${{ github.event_name }}
|
||||
REF_NAME: ${{ github.ref_name }}
|
||||
BEFORE: ${{ github.event.before }}
|
||||
CREATED: ${{ github.event.created }}
|
||||
RC_VERSION: ${{ inputs.rc_version }}
|
||||
RELEASE_TAG: ${{ github.event.release.tag_name }}
|
||||
PRERELEASE: ${{ github.event.release.prerelease }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
sync_base="${BEFORE}"
|
||||
if [ "${CREATED}" = "true" ]; then
|
||||
sync_base=""
|
||||
fi
|
||||
if [ -n "${RC_VERSION}" ]; then
|
||||
echo "version=${RC_VERSION}" >> "$GITHUB_OUTPUT"
|
||||
echo "stage=rc" >> "$GITHUB_OUTPUT"
|
||||
elif [ "${EVENT}" = "release" ]; then
|
||||
if [ "${PRERELEASE}" = "true" ] || ! echo "${RELEASE_TAG}" | grep -qE '^v[0-9]+\.[0-9]+\.0$'; then
|
||||
echo "::notice::${RELEASE_TAG} is not an X.Y.0 stable release; nothing to complete"
|
||||
exit 0
|
||||
fi
|
||||
echo "version=${RELEASE_TAG#v}" >> "$GITHUB_OUTPUT"
|
||||
echo "complete=true" >> "$GITHUB_OUTPUT"
|
||||
elif [ "${REF_NAME}" = "main" ]; then
|
||||
version="$(python3 .github/scripts/read_rc_version.py | cut -d= -f2)"
|
||||
status=0
|
||||
git ls-remote --exit-code --heads origin "rc/${version}" > /dev/null || status=$?
|
||||
case "${status}" in
|
||||
0)
|
||||
IFS=. read -r major minor _ <<< "${version}"
|
||||
version="${major}.$((minor + 1)).0"
|
||||
;;
|
||||
2) ;;
|
||||
*)
|
||||
echo "::error::could not check whether rc/${version} exists (git ls-remote exit ${status})"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
echo "version=${version}" >> "$GITHUB_OUTPUT"
|
||||
echo "sync_base=${sync_base}" >> "$GITHUB_OUTPUT"
|
||||
echo "main=true" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "version=${REF_NAME#rc/}" >> "$GITHUB_OUTPUT"
|
||||
echo "sync_base=${sync_base}" >> "$GITHUB_OUTPUT"
|
||||
echo "stage=rc" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Sync commits into the release
|
||||
if: steps.plan.outputs.sync_base != ''
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: sync
|
||||
name: LiteLLM ${{ steps.plan.outputs.version }}
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
base_ref: ${{ steps.plan.outputs.sync_base }}
|
||||
cli_version: v0.18.0
|
||||
|
||||
- name: Keep the main stage unless the rc branch was cut during this run
|
||||
id: main_stage
|
||||
if: steps.plan.outputs.main == 'true'
|
||||
env:
|
||||
VERSION: ${{ steps.plan.outputs.version }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
status=0
|
||||
git ls-remote --exit-code --heads origin "rc/${VERSION}" > /dev/null || status=$?
|
||||
case "${status}" in
|
||||
0) echo "::notice::rc/${VERSION} was cut during this run; leaving the release in its rc stage" ;;
|
||||
2) echo "stage=main" >> "$GITHUB_OUTPUT" ;;
|
||||
*)
|
||||
echo "::error::could not check whether rc/${VERSION} exists (git ls-remote exit ${status})"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
- name: Move the release to its stage
|
||||
if: steps.plan.outputs.stage != '' || steps.main_stage.outputs.stage != ''
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: update
|
||||
stage: ${{ steps.plan.outputs.stage || steps.main_stage.outputs.stage }}
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
cli_version: v0.18.0
|
||||
|
||||
- name: Complete the release
|
||||
if: steps.plan.outputs.complete == 'true'
|
||||
uses: linear/linear-release-action@d4af10092984f9bc6d5efa075b242bdf01333463 # v0.18.0
|
||||
with:
|
||||
access_key: ${{ secrets.LINEAR_API_KEY }}
|
||||
command: complete
|
||||
version: ${{ steps.plan.outputs.version }}
|
||||
cli_version: v0.18.0
|
||||
6
Makefile
6
Makefile
|
|
@ -4,7 +4,7 @@
|
|||
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
|
||||
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
|
||||
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
|
||||
test-rust-extension \
|
||||
test-rust-extension rust-sqlx-prepare \
|
||||
info lint lint-inner lint-dev lint-checks format \
|
||||
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
|
||||
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
|
||||
|
|
@ -56,6 +56,7 @@ help:
|
|||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
@echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container"
|
||||
@echo ""
|
||||
@echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide"
|
||||
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
|
||||
|
|
@ -306,6 +307,9 @@ test-rust-extension:
|
|||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
||||
rust-sqlx-prepare:
|
||||
cd litellm-rust && cargo run -p litellm-db-testing --bin sqlx-prepare
|
||||
|
||||
test: install-test-deps
|
||||
$(UV_RUN) pytest tests/
|
||||
|
||||
|
|
|
|||
|
|
@ -1,37 +0,0 @@
|
|||
# Publish MCP servers in the AI Hub
|
||||
|
||||
Set `litellm_settings.public_mcp_servers` to the concrete IDs of the servers you want listed in the public AI Hub. Pin `server_id` in each configuration entry so the publication list stays stable across deployments
|
||||
|
||||
```yaml
|
||||
mcp_servers:
|
||||
documentation:
|
||||
server_id: documentation-mcp
|
||||
url: https://mcp.example.com/mcp
|
||||
transport: http
|
||||
available_on_public_internet: true
|
||||
|
||||
litellm_settings:
|
||||
public_mcp_hub_strict_whitelist: true
|
||||
public_mcp_servers:
|
||||
- documentation-mcp
|
||||
```
|
||||
|
||||
Use `documentation-mcp`, the `server_id`, in the publication list. The configuration key `documentation`, display names, and aliases are not publication IDs. Database-created servers use the ID returned by `/v1/mcp/server`
|
||||
|
||||
The dashboard's **AI Hub > MCP Hub > Manage MCP Hub Visibility** dialog edits this same list. Its YAML example includes the selected server IDs. With database-backed configuration (`store_model_in_db: true`), a value declared in YAML is owned by that file: edit the file and reload, or remove that key from YAML to let the dashboard manage it in the database. File-backed deployments can save the list directly to their configuration file
|
||||
|
||||
To remove all explicit entries, save an empty selection in the dialog or configure:
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
public_mcp_hub_strict_whitelist: true
|
||||
public_mcp_servers: []
|
||||
```
|
||||
|
||||
## Hub listing and network access
|
||||
|
||||
The **Hub listing** column in AI Hub identifies servers that appear in `/public/mcp_hub`. The dashboard derives this status from the current registry and publication settings. Setting `mcp_info.is_public` on a server does not publish it; that response field is derived metadata. `mcp_info.is_public_explicit` identifies registered servers included in the explicit publication list
|
||||
|
||||
Gateway cards and server details show **All Networks** when `available_on_public_internet` is enabled or the server is explicitly published in `public_mcp_servers`. They show **Internal Only** when both are false. The per-server flag defaults to `true`; explicit publication overrides a disabled flag for compatibility. Older proxies that omit the metadata needed to determine access show **Unknown**. These labels describe allowed client IPs; authentication and tool permissions still apply
|
||||
|
||||
The default `public_mcp_hub_strict_whitelist: true` lists only registered servers in `public_mcp_servers`. Legacy mode (`false`) additionally lists registered servers with `available_on_public_internet: true`. In legacy mode, clearing the explicit publication list leaves these automatically listed servers visible. Enable strict mode when the publication list should fully determine hub visibility
|
||||
|
|
@ -24,7 +24,7 @@ model_list:
|
|||
- model_name: sagemaker-completion-model
|
||||
litellm_params:
|
||||
model: sagemaker/berri-benchmarking-Llama-2-70b-chat-hf-4
|
||||
input_cost_per_second: 0.000420
|
||||
cost_per_second: 0.000420
|
||||
- model_name: text-embedding-ada-002
|
||||
litellm_params:
|
||||
model: azure/azure-embedding-model
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.71"
|
||||
version = "0.1.72"
|
||||
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.71"
|
||||
version = "0.1.72"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,97 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "enabled" BOOLEAN NOT NULL DEFAULT true,
|
||||
ADD COLUMN IF NOT EXISTS "execution_mode" TEXT NOT NULL DEFAULT 'autonomous',
|
||||
ADD COLUMN IF NOT EXISTS "identity_managed" BOOLEAN NOT NULL DEFAULT false;
|
||||
|
||||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN IF NOT EXISTS "billing_agent_id" TEXT;
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_AgentIdentity" (
|
||||
"agent_id" TEXT NOT NULL,
|
||||
"active" BOOLEAN NOT NULL DEFAULT true,
|
||||
"provider" TEXT NOT NULL,
|
||||
"issuer" TEXT NOT NULL,
|
||||
"tenant_id" TEXT NOT NULL,
|
||||
"client_id" TEXT NOT NULL,
|
||||
"service_principal_id" TEXT,
|
||||
"required_roles" TEXT[] DEFAULT ARRAY[]::TEXT[],
|
||||
"required_scopes" TEXT[] DEFAULT ARRAY['user_impersonation']::TEXT[],
|
||||
"revision" TEXT NOT NULL,
|
||||
"last_authenticated_at" TIMESTAMP(3),
|
||||
|
||||
CONSTRAINT "LiteLLM_AgentIdentity_pkey" PRIMARY KEY ("agent_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgentIdentity" (
|
||||
"binding_id" TEXT NOT NULL,
|
||||
"agent_id" TEXT,
|
||||
"provider" TEXT NOT NULL,
|
||||
"issuer" TEXT NOT NULL,
|
||||
"tenant_id" TEXT NOT NULL,
|
||||
"client_id" TEXT NOT NULL,
|
||||
|
||||
CONSTRAINT "LiteLLM_RetiredAgentIdentity_pkey" PRIMARY KEY ("binding_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_RetiredAgent" (
|
||||
"original_agent_id" TEXT NOT NULL,
|
||||
"retired_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_RetiredAgent_pkey" PRIMARY KEY ("original_agent_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE IF NOT EXISTS "LiteLLM_VerifiedSubject" (
|
||||
"subject_id" TEXT NOT NULL,
|
||||
"issuer" TEXT NOT NULL,
|
||||
"tenant_id" TEXT NOT NULL,
|
||||
"oid" TEXT NOT NULL,
|
||||
"kind" TEXT NOT NULL DEFAULT 'human',
|
||||
"user_id" TEXT,
|
||||
"verified_via" TEXT NOT NULL DEFAULT 'sso_interactive',
|
||||
"verified_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_VerifiedSubject_pkey" PRIMARY KEY ("subject_id")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_AgentIdentity"("provider", "tenant_id", "client_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentIdentity_issuer_service_principal_id_key" ON "LiteLLM_AgentIdentity"("issuer", "service_principal_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_RetiredAgentIdentity_provider_tenant_id_client_id_key" ON "LiteLLM_RetiredAgentIdentity"("provider", "tenant_id", "client_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_user_id_idx" ON "LiteLLM_VerifiedSubject"("user_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_VerifiedSubject_issuer_tenant_id_oid_key" ON "LiteLLM_VerifiedSubject"("issuer", "tenant_id", "oid");
|
||||
|
||||
-- AddForeignKey
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_AgentIdentity_agent_id_fkey') THEN
|
||||
ALTER TABLE "LiteLLM_AgentIdentity" ADD CONSTRAINT "LiteLLM_AgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE CASCADE ON UPDATE CASCADE;
|
||||
END IF;
|
||||
END $$;
|
||||
|
||||
-- AddForeignKey
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_RetiredAgentIdentity_agent_id_fkey') THEN
|
||||
ALTER TABLE "LiteLLM_RetiredAgentIdentity" ADD CONSTRAINT "LiteLLM_RetiredAgentIdentity_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_AgentsTable"("agent_id") ON DELETE SET NULL ON UPDATE CASCADE;
|
||||
END IF;
|
||||
END $$;
|
||||
|
||||
-- AddForeignKey
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'LiteLLM_VerifiedSubject_user_id_fkey') THEN
|
||||
ALTER TABLE "LiteLLM_VerifiedSubject" ADD CONSTRAINT "LiteLLM_VerifiedSubject_user_id_fkey" FOREIGN KEY ("user_id") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE CASCADE ON UPDATE CASCADE;
|
||||
END IF;
|
||||
END $$;
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "pinned_tools" JSONB DEFAULT '{}';
|
||||
|
|
@ -0,0 +1,19 @@
|
|||
CREATE TABLE IF NOT EXISTS "LiteLLM_DailyModelUsage" (
|
||||
"date" TEXT NOT NULL,
|
||||
"model_group" TEXT NOT NULL,
|
||||
"model" TEXT NOT NULL,
|
||||
"custom_llm_provider" TEXT NOT NULL,
|
||||
"task_type" TEXT NOT NULL,
|
||||
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
|
||||
"prompt_tokens" BIGINT NOT NULL DEFAULT 0,
|
||||
"completion_tokens" BIGINT NOT NULL DEFAULT 0,
|
||||
"request_count" BIGINT NOT NULL DEFAULT 0,
|
||||
"successful_requests" BIGINT NOT NULL DEFAULT 0,
|
||||
"failed_requests" BIGINT NOT NULL DEFAULT 0,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL,
|
||||
CONSTRAINT "LiteLLM_DailyModelUsage_pkey" PRIMARY KEY ("date", "model_group", "model", "custom_llm_provider", "task_type")
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_DailyModelUsage_date_idx" ON "LiteLLM_DailyModelUsage"("date");
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_DailyModelUsage_model_group_idx" ON "LiteLLM_DailyModelUsage"("model_group");
|
||||
|
|
@ -78,6 +78,11 @@ model LiteLLM_AgentsTable {
|
|||
object_permission_id String?
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
spend Float @default(0.0)
|
||||
identity_managed Boolean @default(false)
|
||||
enabled Boolean @default(true)
|
||||
execution_mode String @default("autonomous")
|
||||
identity LiteLLM_AgentIdentity?
|
||||
retired_identities LiteLLM_RetiredAgentIdentity[]
|
||||
tpm_limit Int?
|
||||
rpm_limit Int?
|
||||
session_tpm_limit Int?
|
||||
|
|
@ -88,6 +93,56 @@ model LiteLLM_AgentsTable {
|
|||
updated_by String
|
||||
}
|
||||
|
||||
model LiteLLM_AgentIdentity {
|
||||
agent_id String @id
|
||||
active Boolean @default(true)
|
||||
agent LiteLLM_AgentsTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
|
||||
provider String
|
||||
issuer String
|
||||
tenant_id String
|
||||
client_id String
|
||||
service_principal_id String?
|
||||
required_roles String[] @default([])
|
||||
required_scopes String[] @default(["user_impersonation"])
|
||||
revision String @default(uuid())
|
||||
last_authenticated_at DateTime?
|
||||
@@unique([provider, tenant_id, client_id])
|
||||
@@unique([issuer, service_principal_id])
|
||||
}
|
||||
|
||||
model LiteLLM_RetiredAgentIdentity {
|
||||
binding_id String @id @default(uuid())
|
||||
agent_id String?
|
||||
agent LiteLLM_AgentsTable? @relation(fields: [agent_id], references: [agent_id], onDelete: SetNull)
|
||||
provider String
|
||||
issuer String
|
||||
tenant_id String
|
||||
client_id String
|
||||
@@unique([provider, tenant_id, client_id])
|
||||
}
|
||||
|
||||
model LiteLLM_RetiredAgent {
|
||||
original_agent_id String @id
|
||||
retired_at DateTime @default(now())
|
||||
}
|
||||
|
||||
model LiteLLM_VerifiedSubject {
|
||||
subject_id String @id @default(uuid())
|
||||
issuer String
|
||||
tenant_id String
|
||||
oid String
|
||||
kind String @default("human")
|
||||
user_id String?
|
||||
user LiteLLM_UserTable? @relation(fields: [user_id], references: [user_id], onDelete: Cascade)
|
||||
verified_via String @default("sso_interactive")
|
||||
verified_at DateTime @default(now())
|
||||
@@unique([issuer, tenant_id, oid])
|
||||
@@index([user_id])
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
model LiteLLM_OrganizationTable {
|
||||
organization_id String @id @default(uuid())
|
||||
organization_alias String
|
||||
|
|
@ -241,6 +296,7 @@ model LiteLLM_DeletedTeamTable {
|
|||
|
||||
// Track spend, rate limit, budget Users
|
||||
model LiteLLM_UserTable {
|
||||
verified_subjects LiteLLM_VerifiedSubject[]
|
||||
user_id String @id
|
||||
user_alias String?
|
||||
team_id String?
|
||||
|
|
@ -322,6 +378,7 @@ model LiteLLM_MCPServerTable {
|
|||
allowed_tools String[] @default([])
|
||||
tool_name_to_display_name Json? @default("{}")
|
||||
tool_name_to_description Json? @default("{}")
|
||||
pinned_tools Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
static_headers Json? @default("{}")
|
||||
// Admin-configured environment variables interpolated into static_headers
|
||||
|
|
@ -674,6 +731,7 @@ model LiteLLM_SpendLogs {
|
|||
session_id String?
|
||||
status String?
|
||||
mcp_namespaced_tool_name String?
|
||||
billing_agent_id String?
|
||||
agent_id String?
|
||||
proxy_server_request Json? @default("{}")
|
||||
litellm_call_id String?
|
||||
|
|
@ -1259,6 +1317,26 @@ model LiteLLM_DailyToolSpend {
|
|||
@@id([date, tool_name])
|
||||
}
|
||||
|
||||
model LiteLLM_DailyModelUsage {
|
||||
date String
|
||||
model_group String
|
||||
model String
|
||||
custom_llm_provider String
|
||||
task_type String
|
||||
spend Float @default(0.0)
|
||||
prompt_tokens BigInt @default(0)
|
||||
completion_tokens BigInt @default(0)
|
||||
request_count BigInt @default(0)
|
||||
successful_requests BigInt @default(0)
|
||||
failed_requests BigInt @default(0)
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
|
||||
@@id([date, model_group, model, custom_llm_provider, task_type])
|
||||
@@index([date])
|
||||
@@index([model_group])
|
||||
}
|
||||
|
||||
// Gateway request counts recorded at the ASGI edge by
|
||||
// BillableRequestMetricsMiddleware. This is the source of truth for SGR
|
||||
// (successful gateway requests): it counts what the proxy actually answered,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.102"
|
||||
version = "0.4.103"
|
||||
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.102"
|
||||
version = "0.4.103"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
---
|
||||
name: rust-tracing
|
||||
description: Add or change Rust diagnostic tracing in litellm-rust, including route spans, subscriber layers, and Python logger delivery
|
||||
---
|
||||
|
||||
# Rust tracing
|
||||
|
||||
Use upstream `tracing` throughout Rust, including `#[tracing::instrument]`, events, and span propagation. Centralize collection and delivery infrastructure in `crates/tracing`. Direct upstream imports still reach our configured subscriber; re-exporting macros does not control delivery. Do not introduce Rust `log` or `pyo3-log` for this path
|
||||
|
||||
`litellm-tracing` owns shared subscriber layers, span field collection, and diagnostic processing. Keep adapters composable as `tracing_subscriber::Layer`s, with `Logger` providing host setup. Runtime-specific delivery belongs in the host bridge. The Python bridge delivers directly to the existing Python SDK logger, preserving its handlers, filtering, redaction, and request correlation. Keep Python dependencies out of `crates/tracing`
|
||||
|
||||
Hosts configure subscribers. Keep Python execution scoped to its captured dispatch rather than installing a process-wide subscriber. Propagate both span context and dispatch across spawned work and returned streams
|
||||
|
||||
In core, instrument execution shared by native calls and hosted machines. Use consistent route, model, provider, streaming, and outcome fields. Put status recording at shared provider boundaries instead of scattering basic logging through handlers. Keep upstream HTTP status separate from route success
|
||||
|
||||
Use `skip_all` and explicitly selected fields. Basic tracing excludes bodies, credentials, headers, and raw error strings. Avoid automatic `ret` or `err` capture of sensitive values. Keep payload diagnostics separate and subject to existing redaction
|
||||
|
||||
A returned stream retains its route span until exhaustion, error, or drop, with exactly one terminal outcome. Builder construction does not start a trace. Never hold a span entry guard across an await. Diagnostic tracing remains separate from lifecycle callbacks and `CustomLogger` dispatch
|
||||
|
||||
Use `litellm_tracing::sink_layer` to compose a sink with other subscriber layers. It inherits span fields into events and emits span-close summaries with elapsed time. Test observable records, concurrent isolation, dynamic filtering, sensitive-field exclusion, and stream cancellation when changing this behavior
|
||||
|
||||
Consult the [tracing API](https://docs.rs/tracing/latest/tracing/) and [subscriber layers](https://docs.rs/tracing-subscriber/latest/tracing_subscriber/layer/index.html) for implementation details
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
# Rust workspace rules
|
||||
|
||||
For diagnostic tracing changes, follow [.agents/skills/rust-tracing/SKILL.md](.agents/skills/rust-tracing/SKILL.md)
|
||||
|
||||
## Test placement
|
||||
|
||||
- Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;`
|
||||
|
|
|
|||
1455
litellm-rust/Cargo.lock
generated
1455
litellm-rust/Cargo.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -13,11 +13,16 @@ litellm-config = { path = "crates/config" }
|
|||
litellm-router = { path = "crates/router" }
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
litellm-gateway-inference = { path = "crates/gateway-inference" }
|
||||
litellm-gateway-auth = { path = "crates/gateway-auth" }
|
||||
litellm-gateway-management = { path = "crates/gateway-management" }
|
||||
litellm-gateway-ui = { path = "crates/gateway-ui" }
|
||||
litellm-coroutine = { path = "crates/coroutine" }
|
||||
litellm-host = { path = "crates/host" }
|
||||
litellm-host-http = { path = "crates/host-http" }
|
||||
litellm-host-native = { path = "crates/host-native" }
|
||||
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
|
||||
litellm-framing = { path = "crates/framer" }
|
||||
litellm-auth = { path = "crates/auth" }
|
||||
|
|
@ -34,8 +39,10 @@ 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-llms-types = { path = "crates/llms-types" }
|
||||
litellm-core-utils = { path = "crates/core-utils" }
|
||||
litellm-db = { path = "crates/db" }
|
||||
litellm-db-testing = { path = "crates/db-testing" }
|
||||
litellm-cache = { path = "crates/cache" }
|
||||
litellm-cache-azure-blob = { path = "crates/cache-azure-blob" }
|
||||
litellm-cache-memory = { path = "crates/cache-memory" }
|
||||
|
|
@ -56,6 +63,8 @@ litellm-python-compat = { path = "crates/python-compat" }
|
|||
|
||||
tracing = "0.1"
|
||||
axum = { version = "0.8.9", default-features = false, features = ["http1", "tokio", "multipart"] }
|
||||
axum-login = "0.18.0"
|
||||
tower-sessions = { version = "0.14.0", features = ["memory-store"] }
|
||||
bytes = "1"
|
||||
http = "1"
|
||||
google-cloud-auth = { version = "1.16.0", default-features = false }
|
||||
|
|
@ -65,6 +74,7 @@ proptest = "1.7.0"
|
|||
pyo3 = "0.29.2"
|
||||
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
|
||||
rand = "0.8"
|
||||
macro_rules_attribute = "0.2.3"
|
||||
schemars = "1"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2", "stream"] }
|
||||
qdrant-client = { version = "1.19.0", default-features = false }
|
||||
|
|
@ -80,6 +90,7 @@ serde = { version = "1.0", features = ["derive"] }
|
|||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
serde_with = { version = "=3.16.1", default-features = false, features = ["std", "macros"] }
|
||||
sha2 = "0.10"
|
||||
sqlx = { version = "0.9.0", default-features = false, features = ["json", "macros", "postgres", "runtime-tokio", "chrono", "tls-rustls-ring-native-roots"] }
|
||||
subtle = "2"
|
||||
thiserror = "2.0"
|
||||
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
# The Tokio runtime is reached only through `host-python/src/execution.rs`, whose fork gate
|
||||
# The Tokio runtime is reached only through `host-python/src/runtime.rs`, whose fork gate
|
||||
# must see every entry. Going around it makes a fork-after-use hang instead of raising.
|
||||
disallowed-methods = [
|
||||
{ path = "pyo3_async_runtimes::tokio::get_runtime", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
|
|
@ -12,6 +12,13 @@ disallowed-methods = [
|
|||
{ path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" },
|
||||
{ path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" },
|
||||
{ path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" },
|
||||
{ path = "sqlx::query", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_with", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as_with", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar_with", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::raw_sql", reason = "raw_sql is unchecked; use the checked query macros" },
|
||||
]
|
||||
|
||||
# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS,
|
||||
|
|
|
|||
|
|
@ -39,7 +39,33 @@ impl TokenProviderHandle {
|
|||
Self(caller)
|
||||
}
|
||||
|
||||
pub fn from_callback<F, Fut>(acquire: F) -> Self
|
||||
where
|
||||
F: Fn() -> Fut + Send + Sync + 'static,
|
||||
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
|
||||
{
|
||||
Self::new(Arc::new(CallbackTokenProvider(acquire)))
|
||||
}
|
||||
|
||||
pub async fn acquire(&self) -> Result<ResolvedCredential, Error> {
|
||||
self.0.acquire().await
|
||||
}
|
||||
}
|
||||
|
||||
struct CallbackTokenProvider<F>(F);
|
||||
|
||||
impl<F> std::fmt::Debug for CallbackTokenProvider<F> {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("CallbackTokenProvider")
|
||||
}
|
||||
}
|
||||
|
||||
impl<F, Fut> TokenProvider for CallbackTokenProvider<F>
|
||||
where
|
||||
F: Fn() -> Fut + Send + Sync,
|
||||
Fut: Future<Output = Result<ResolvedCredential, Error>> + Send + 'static,
|
||||
{
|
||||
fn acquire(&self) -> TokenFuture<'_> {
|
||||
Box::pin((self.0)())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
104
litellm-rust/crates/auth-types/tests/token.rs
Normal file
104
litellm-rust/crates/auth-types/tests/token.rs
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
use std::{
|
||||
error::Error as StdError,
|
||||
future::{Future, poll_fn},
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicBool, AtomicUsize, Ordering},
|
||||
},
|
||||
task::Poll,
|
||||
time::{Duration, SystemTime},
|
||||
};
|
||||
|
||||
use litellm_auth_types::{
|
||||
Error, ErrorDetail, ResolvedCredential, SecretValue, TokenProviderHandle,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
fn credential(index: usize, access_token: bool) -> ResolvedCredential {
|
||||
let token = SecretValue::new(format!("credential-{index}"));
|
||||
if access_token {
|
||||
return ResolvedCredential::AccessToken {
|
||||
token,
|
||||
expires_on: Some(SystemTime::UNIX_EPOCH + Duration::from_secs(index as u64)),
|
||||
};
|
||||
}
|
||||
ResolvedCredential::Static(token)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::static_secret(false)]
|
||||
#[case::access_token(true)]
|
||||
#[tokio::test]
|
||||
async fn callbacks_acquire_fresh_credentials_on_demand(#[case] access_token: bool) {
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
let callback_calls = calls.clone();
|
||||
let provider = TokenProviderHandle::from_callback(move || {
|
||||
let index = callback_calls.fetch_add(1, Ordering::SeqCst);
|
||||
async move {
|
||||
tokio::task::yield_now().await;
|
||||
Ok(credential(index, access_token))
|
||||
}
|
||||
});
|
||||
let cloned = provider.clone();
|
||||
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 0);
|
||||
assert_eq!(
|
||||
provider.acquire().await.unwrap(),
|
||||
credential(0, access_token)
|
||||
);
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(cloned.acquire().await.unwrap(), credential(1, access_token));
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 2);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn callback_errors_preserve_the_original_source() {
|
||||
let provider = TokenProviderHandle::from_callback(|| async {
|
||||
Err(Error::CredentialAcquisition(ErrorDetail::failed(
|
||||
"caller credential",
|
||||
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
|
||||
)))
|
||||
});
|
||||
|
||||
let error = provider.acquire().await.unwrap_err();
|
||||
assert!(matches!(error, Error::CredentialAcquisition(_)));
|
||||
let source = std::iter::successors(Some(&error as &(dyn StdError + 'static)), |error| {
|
||||
(*error).source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<std::io::Error>())
|
||||
.unwrap();
|
||||
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
|
||||
}
|
||||
|
||||
struct Release(Arc<AtomicBool>);
|
||||
|
||||
impl Drop for Release {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn cancelling_acquisition_drops_the_callback_future() {
|
||||
let released = Arc::new(AtomicBool::new(false));
|
||||
let callback_released = released.clone();
|
||||
let provider = TokenProviderHandle::from_callback(move || {
|
||||
let released = callback_released.clone();
|
||||
async move {
|
||||
let _release = Release(released);
|
||||
std::future::pending().await
|
||||
}
|
||||
});
|
||||
|
||||
let mut acquisition = Box::pin(provider.acquire());
|
||||
poll_fn(|context| {
|
||||
assert!(acquisition.as_mut().poll(context).is_pending());
|
||||
assert!(!released.load(Ordering::SeqCst));
|
||||
Poll::Ready(())
|
||||
})
|
||||
.await;
|
||||
drop(acquisition);
|
||||
assert!(released.load(Ordering::SeqCst));
|
||||
}
|
||||
|
|
@ -9,7 +9,7 @@ repository.workspace = true
|
|||
litellm-cache.workspace = true
|
||||
py_literal = "0.4.0"
|
||||
rand.workspace = true
|
||||
rusqlite = { version = "0.40", features = ["bundled"] }
|
||||
rusqlite = { version = "0.39", features = ["bundled"] }
|
||||
serde-pickle = "1.2"
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ async fn set_writes_encoded_object_and_headers(#[future(awt)] server: MockServer
|
|||
)]
|
||||
#[case::missing("missing", ResponseTemplate::new(404), Ok(None))]
|
||||
#[case::server_error("server-error", ResponseTemplate::new(500), Err(Error::Unavailable))]
|
||||
#[case::unauthorized("unauthorized", ResponseTemplate::new(401), Err(Error::Unavailable))]
|
||||
#[case::invalid(
|
||||
"invalid",
|
||||
ResponseTemplate::new(200).set_body_string("not json"),
|
||||
|
|
|
|||
29
litellm-rust/crates/cache-response/AGENTS.md
Normal file
29
litellm-rust/crates/cache-response/AGENTS.md
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
# Response caching
|
||||
|
||||
Design this crate for shared Rust execution used by the Python SDK and the Rust gateway. The Python SDK will remain, with more core execution moving to Rust and Python callbacks staying in Python. The Rust gateway is still evolving and is intended to replace the Python proxy. Keep response-cache policy independent of Python, HTTP serving, and either proxy's configuration format
|
||||
|
||||
Separate what is cached, how a hit is matched, and where entries are stored. Chat Completions, Messages, Responses, and embeddings are API workloads. Exact and semantic matching are lookup behaviors. Memory, Redis, disk, and object stores are storage choices. Embeddings are inference too, so do not use an inference-cache name to imply a category that excludes embeddings. Consult the existing Python cache and caching handler for behavior and compatibility contracts without copying their class structure
|
||||
|
||||
Storage traits, codecs, and backend capabilities belong in `litellm-cache` and the storage crates. Keep storage reusable for value types beyond LLM responses. This crate owns response entries, matching and freshness semantics, the Python-compatible response codec, and deferred-write policy. Core owns route-specific request identity, response encoding and reconstruction, embedding partial-hit orchestration, and stream capture and replay. Boundaries own configuration translation, resource construction, and caller identity
|
||||
|
||||
Construct and inject the response-cache service at the Python bridge or gateway boundary, as with the HTTP client. Reuse it across calls. Core and provider code must not discover cache configuration through Python globals, process configuration, or backend-specific factories
|
||||
|
||||
Keep `ResponseCache<B>` generic over its storage backend. Preserve typed backend contexts and capability bounds internally. Inject an object-safe service into core for runtime backend selection, so storage types do not spread through route and host types. Keep API request and response types statically typed. Add a generic parameter only where it preserves a useful type relationship or capability
|
||||
|
||||
Keep the core service contract narrow. Lookup and store must not require connection testing, ping, flush, deletion, counters, queues, or scripts. Require batch operations where a consumer needs partial hits, and keep management capabilities on their own interfaces. An exact-only adapter must remain explicit about its matching restriction. Supporting semantic matching requires a defined lookup-context and embedding execution contract, not just a renamed trait
|
||||
|
||||
Separate reusable resources from per-call policy. Backend configuration, namespace, default expiry, and entry limits belong to the configured service or backend. Read/write controls, expiry and freshness overrides, and authenticated caller scope belong to the call. Passing call options must not replace or mutate the route's configured service
|
||||
|
||||
Keep cache misses and storage failures distinguishable in return values. Core owns the decision to continue with provider execution after a cache failure. A read can reject an entry for freshness while the backend still retains it. Preserve the timestamp at which a response was produced when writing it later
|
||||
|
||||
Define lookup placement explicitly relative to authorization, deployment and credential resolution, and request-transforming callbacks. Cache identity must account for every input that affects reuse, including API surface and caller scope, while preserving intentional Python caching groups. Preserve existing keys and response formats unless changing them is an explicit migration decision
|
||||
|
||||
Cache normalized provider results before caller-specific response transformations. Hits must still run the applicable response processing, success callbacks, and cache-hit accounting. Keep callback execution in the host. Python cache implementations and semantic embedders that require the caller's task must use the existing host-operation mechanism rather than Python calls from a Rust worker. Preserve legacy fallback until that contract is supported
|
||||
|
||||
Keep unary caching independent of stream-only methods. Store streams only after successful exhaustion and protocol completion. Errors, incomplete streams, cancellation, and oversized entries must not populate the cache. Embedding batches need ordered partial results and reconstruction around the uncached inputs
|
||||
|
||||
Test each contract in its owner: storage capabilities in backend tests, envelopes and freshness here, reuse and replay in core, Python callback and fallback behavior at the bridge, and HTTP behavior at the gateway. Run backend contract checks and Python response-codec fixtures before exposing a new backend
|
||||
|
||||
`ScopedCache` requires an explicit shared or isolated scope at construction. Per-call `CachePolicy` controls reads, writes, expiry, and freshness without replacing the attached scope or service. `CacheOptions` binds that policy to an explicit scope for storage requests and has no default sharing policy. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec
|
||||
|
||||
Response storage is not the source of budget or rate-limit coordination dependencies. Keep counters, reservations, and atomic admission operations out of `ResponseCacheService`, including when both services happen to use Redis
|
||||
|
|
@ -13,9 +13,12 @@ serde_json.workspace = true
|
|||
sha2.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-cache-gcs.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-memory.workspace = true
|
||||
litellm-cache-redis.workspace = true
|
||||
redis = "1.7.0"
|
||||
redis-test = "1.0.4"
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -1,51 +0,0 @@
|
|||
# Response cache
|
||||
|
||||
`ResponseCache<B>` adds request keys, independent read/write controls, response envelopes, and freshness checks to any `B: BaseCache<Value = CacheEntry>`
|
||||
|
||||
## Ownership
|
||||
|
||||
`litellm-cache` defines typed storage, codec, and capability traits. `BaseCache` is only get, set, TTL, and pipeline writes. Everything else is an optional capability a backend implements only where its Python class defines the method: `DisconnectCache`, `ConnectionCache` (`test_connection`), `PingCache`, `BatchCache`, `DeleteCache`, `FlushCache`, counters, queues, TTL, scan, and scripts. Memory, Redis, disk, S3, GCS, and Azure Blob implement those traits without depending on response policy, so other consumers can store their own value types in the same backends
|
||||
|
||||
Semantic backends (Redis, Valkey, Qdrant) are generic over their embedder and codec, and share one prompt and embedding contract from `litellm_cache::semantic`. They take a `SemanticCacheContext`, so `ResponseCache` drives them the same way it drives exact backends
|
||||
|
||||
`litellm-cache-response` owns response keys, controls, entries, the Python-compatible response codec, and `WriteBuffer`, the backend-neutral deferred-write policy. It has no runtime dependency on a specific cache backend or Python
|
||||
|
||||
`ExactResponseCache` is the object-safe view of a `ResponseCache` over an exact backend. `ConnectionProbe` is the object-safe `test_connection`, implemented only when the backend implements `ConnectionCache`, so a host holds one next to its `ExactResponseCache` and reports the operation as unsupported otherwise, as Python's `BaseCache` does. Lookup, store, batch, and flush never require it
|
||||
|
||||
## Native Rust use
|
||||
|
||||
```rust
|
||||
use std::{sync::Arc, time::Duration};
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{CacheKeyInput, ResponseCache, ResponseCacheRequest};
|
||||
use serde_json::json;
|
||||
|
||||
let cache = ResponseCache::new(Arc::new(InMemoryCache::default()));
|
||||
let request = ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some("example:key".into()),
|
||||
..Default::default()
|
||||
});
|
||||
let now = Duration::from_secs(100);
|
||||
cache.store(&request, json!({"answer": 7}), now)?;
|
||||
assert_eq!(cache.async_lookup(&request, now).await?, Some(json!({"answer": 7})));
|
||||
```
|
||||
|
||||
For Redis, inject `RedisCache::new(url, ttl, ResponseCacheCodec)` instead. Namespaces are optional and existing namespace prefixes are preserved
|
||||
|
||||
Callers supply Unix time for response freshness. Backend TTL uses its own clock. A read can reject an entry through `max_age` even while the backend still retains it
|
||||
|
||||
## Python integration
|
||||
|
||||
The bridge activates backends through the Rust catalog in `litellm/rust_bridge/catalog.py`. Every cache rule ships as `PYTHON_ONLY`, so SDK, Router, and proxy calls stay on Python and construct no native cache resources until a rule is changed
|
||||
|
||||
When a rule selects a backend, the Python `Cache` facade builds the native runtime from its own configuration and routes its storage calls (sync and async lookup and store, and pipelined batch store) to it. Stream replay, embedding partial-hit merging, response reconstruction, and callbacks stay in Python on top of that native store. The Python backend object remains for its direct API
|
||||
|
||||
Object responses are written as they are, and every other response shape is written as a serialized string, which is the pair of shapes Python reads. A string on the wire is therefore always a serialized response, so string-valued responses round trip. Typed backends such as memory never pass through the codec
|
||||
|
||||
Native cache handles must be recreated after fork. Native errors propagate to the host, which owns the existing fail-open and logging policy
|
||||
|
||||
## Adding another backend
|
||||
|
||||
Implement `BaseCache` for the backend with its associated value type and the capability traits its Python class supports, and accept a `CacheCodec` when wire serialization is needed. `ResponseCache<B>` then works without another response implementation
|
||||
|
||||
Run the `litellm-cache-testing` contract checks the backend's capabilities allow, and run response fixtures with `ResponseCacheCodec`, including both Python envelope encodings, before adding a catalog rule
|
||||
|
|
@ -4,6 +4,7 @@ mod codec;
|
|||
mod embedding;
|
||||
mod exact;
|
||||
mod response;
|
||||
mod service;
|
||||
|
||||
pub use buffer::WriteBuffer;
|
||||
pub use caching::{
|
||||
|
|
@ -14,3 +15,8 @@ pub use codec::ResponseCacheCodec;
|
|||
pub use embedding::PartialHits;
|
||||
pub use exact::{ConnectionProbe, ExactResponseCache};
|
||||
pub use response::{ResponseCache, ResponseCacheRequest};
|
||||
|
||||
pub use service::{
|
||||
CacheOptions, CachePolicy, CacheScope, ResponseCacheConfig, ResponseCacheService,
|
||||
ResponseEnvelope, ScopedCache,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -7,7 +7,9 @@ use litellm_cache::{
|
|||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{CacheControls, CacheEntry, CacheKeyInput, PartialHits, cache_key};
|
||||
use crate::{
|
||||
CacheControls, CacheEntry, CacheKeyInput, PartialHits, ResponseCacheConfig, cache_key,
|
||||
};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponseCacheRequest<C: CacheContext = litellm_cache::ExactCacheContext> {
|
||||
|
|
@ -50,6 +52,7 @@ where
|
|||
B::Context: Default + PartialEq,
|
||||
{
|
||||
backend: Arc<B>,
|
||||
config: ResponseCacheConfig,
|
||||
}
|
||||
|
||||
impl<B> ResponseCache<B>
|
||||
|
|
@ -58,7 +61,18 @@ where
|
|||
B::Context: Default + PartialEq,
|
||||
{
|
||||
pub fn new(backend: Arc<B>) -> Self {
|
||||
Self { backend }
|
||||
Self {
|
||||
backend,
|
||||
config: ResponseCacheConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_config(self, config: ResponseCacheConfig) -> Self {
|
||||
Self { config, ..self }
|
||||
}
|
||||
|
||||
pub fn config(&self) -> &ResponseCacheConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
pub fn backend(&self) -> &B {
|
||||
|
|
@ -221,7 +235,7 @@ where
|
|||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
if !request.controls.writes() {
|
||||
if !request.controls.writes() || !self.fits(&response) {
|
||||
return Ok(());
|
||||
}
|
||||
self.backend.set_cache(
|
||||
|
|
@ -240,7 +254,7 @@ where
|
|||
response: Value,
|
||||
now: Duration,
|
||||
) -> Result<(), Error> {
|
||||
if !request.controls.writes() {
|
||||
if !request.controls.writes() || !self.fits(&response) {
|
||||
return Ok(());
|
||||
}
|
||||
self.backend
|
||||
|
|
@ -277,7 +291,7 @@ where
|
|||
) -> Result<(), Error> {
|
||||
let writable = entries
|
||||
.into_iter()
|
||||
.filter(|(request, _, _)| request.controls.writes())
|
||||
.filter(|(request, response, _)| request.controls.writes() && self.fits(response))
|
||||
.map(|(request, response, now)| {
|
||||
(
|
||||
cache_key(&request.key),
|
||||
|
|
@ -312,6 +326,11 @@ where
|
|||
Ok(())
|
||||
}
|
||||
|
||||
fn fits(&self, response: &Value) -> bool {
|
||||
self.config.max_entry_bytes == usize::MAX
|
||||
|| response.to_string().len() <= self.config.max_entry_bytes
|
||||
}
|
||||
|
||||
fn partial_hits(
|
||||
requests: &[ResponseCacheRequest<B::Context>],
|
||||
readable: Vec<(usize, &ResponseCacheRequest<B::Context>)>,
|
||||
|
|
|
|||
185
litellm-rust/crates/cache-response/src/service.rs
Normal file
185
litellm-rust/crates/cache-response/src/service.rs
Normal file
|
|
@ -0,0 +1,185 @@
|
|||
use std::{future::Future, pin::Pin, time::Duration};
|
||||
|
||||
use litellm_cache::{BaseCache, Error, ExactCacheContext};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
CacheControls, CacheEntry, CacheKeyField, CacheKeyInput, ResponseCache, ResponseCacheRequest,
|
||||
};
|
||||
|
||||
type CacheFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponseCacheConfig {
|
||||
pub namespace: String,
|
||||
pub max_entry_bytes: usize,
|
||||
}
|
||||
|
||||
impl Default for ResponseCacheConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
namespace: String::new(),
|
||||
max_entry_bytes: usize::MAX,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait ResponseCacheService: Send + Sync {
|
||||
fn config(&self) -> &ResponseCacheConfig;
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, Option<Value>>;
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, ()>;
|
||||
}
|
||||
|
||||
impl<B> ResponseCacheService for ResponseCache<B>
|
||||
where
|
||||
B: BaseCache<Value = CacheEntry, Context = ExactCacheContext>,
|
||||
{
|
||||
fn config(&self) -> &ResponseCacheConfig {
|
||||
self.config()
|
||||
}
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, Option<Value>> {
|
||||
Box::pin(self.async_lookup(request, now))
|
||||
}
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
request: &'a ResponseCacheRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> CacheFuture<'a, ()> {
|
||||
Box::pin(self.async_store(request, response, now))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum CacheScope {
|
||||
Shared,
|
||||
Isolated(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Default)]
|
||||
pub struct CachePolicy {
|
||||
pub caching: Option<bool>,
|
||||
pub no_cache: bool,
|
||||
pub no_store: bool,
|
||||
pub ttl: Option<Duration>,
|
||||
pub max_age: Option<Duration>,
|
||||
}
|
||||
|
||||
impl CachePolicy {
|
||||
pub fn enabled(&self) -> bool {
|
||||
self.caching != Some(false) && !(self.no_cache && self.no_store)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CacheOptions {
|
||||
pub policy: CachePolicy,
|
||||
pub scope: CacheScope,
|
||||
}
|
||||
|
||||
impl CacheOptions {
|
||||
pub fn new(scope: CacheScope) -> Self {
|
||||
Self {
|
||||
policy: CachePolicy::default(),
|
||||
scope,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn request(self, namespace: &str, surface: &str, mut input: Value) -> ResponseCacheRequest {
|
||||
input.sort_all_objects();
|
||||
let scope = match self.scope {
|
||||
CacheScope::Shared => String::new(),
|
||||
CacheScope::Isolated(scope) => serde_json::json!(["isolated", scope]).to_string(),
|
||||
};
|
||||
ResponseCacheRequest {
|
||||
key: CacheKeyInput {
|
||||
namespace: Some(format!("{namespace}:inference-v2")),
|
||||
fields: [
|
||||
("surface", surface.to_owned()),
|
||||
("scope", scope),
|
||||
("request", input.to_string()),
|
||||
]
|
||||
.into_iter()
|
||||
.map(|(name, value)| CacheKeyField {
|
||||
name: name.into(),
|
||||
value: Some(value),
|
||||
api_parameter: true,
|
||||
internal_parameter: false,
|
||||
})
|
||||
.collect(),
|
||||
..Default::default()
|
||||
},
|
||||
controls: CacheControls {
|
||||
configured: true,
|
||||
supported_call_type: true,
|
||||
native_backend: true,
|
||||
default_on: true,
|
||||
caching: self.policy.caching,
|
||||
no_cache: self.policy.no_cache,
|
||||
no_store: self.policy.no_store,
|
||||
..Default::default()
|
||||
},
|
||||
context: ExactCacheContext {
|
||||
ttl: self.policy.ttl,
|
||||
},
|
||||
max_age: self.policy.max_age,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize, serde::Deserialize)]
|
||||
pub struct ResponseEnvelope<T> {
|
||||
version: u32,
|
||||
surface: String,
|
||||
output: T,
|
||||
}
|
||||
|
||||
impl<T> ResponseEnvelope<T> {
|
||||
pub fn new(surface: &str, output: T) -> Self {
|
||||
Self {
|
||||
version: 1,
|
||||
surface: surface.into(),
|
||||
output,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decode(self, surface: &str) -> Option<T> {
|
||||
(self.version == 1 && self.surface == surface).then_some(self.output)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ScopedCache {
|
||||
pub service: std::sync::Arc<dyn ResponseCacheService>,
|
||||
pub scope: CacheScope,
|
||||
}
|
||||
|
||||
impl ScopedCache {
|
||||
pub fn new(service: std::sync::Arc<dyn ResponseCacheService>, scope: CacheScope) -> Self {
|
||||
Self { service, scope }
|
||||
}
|
||||
|
||||
pub fn options(&self, policy: Option<CachePolicy>) -> CacheOptions {
|
||||
CacheOptions {
|
||||
policy: policy.unwrap_or_default(),
|
||||
scope: self.scope.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -18,7 +18,7 @@ use litellm_cache_response::{
|
|||
WriteBuffer, cache_key,
|
||||
};
|
||||
use redis_test::MockCmd;
|
||||
use rstest::rstest;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
use support::{keyed, memory, redis, request};
|
||||
|
||||
|
|
@ -648,3 +648,129 @@ async fn write_buffer_clear_drops_pending_entries(memory: Memory, request: Respo
|
|||
assert_eq!(memory.lookup(&request, now).unwrap(), None);
|
||||
assert_eq!(memory.lookup(&other, now).unwrap(), None);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::python_sync("{'timestamp': 100.0, 'response': '{\"answer\": 7}'}")]
|
||||
#[case::python_async(r#"{"timestamp":100.0,"response":{"answer":7}}"#)]
|
||||
#[case::bare_response(r#"{"answer":7}"#)]
|
||||
#[tokio::test]
|
||||
async fn gcs_reads_python_entries_and_writes_python_compatible_envelopes(
|
||||
#[case] encoded: &str,
|
||||
#[values(false, true)] asynchronous: bool,
|
||||
#[future(awt)] gcs: (wiremock::MockServer, Gcs),
|
||||
) {
|
||||
use wiremock::{
|
||||
Mock, ResponseTemplate,
|
||||
matchers::{body_json, header, method, path, query_param},
|
||||
};
|
||||
|
||||
let (server, cache) = gcs;
|
||||
let response = json!({"answer": 7});
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fpython"))
|
||||
.and(query_param("alt", "media"))
|
||||
.and(header("authorization", "Bearer token"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string(encoded))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/upload/storage/v1/b/bucket/o"))
|
||||
.and(query_param("uploadType", "media"))
|
||||
.and(query_param("name", "cache/native"))
|
||||
.and(header("authorization", "Bearer token"))
|
||||
.and(header("content-type", "application/json"))
|
||||
.and(body_json(json!({"timestamp": 102.0, "response": response})))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let lookup = if asynchronous {
|
||||
cache
|
||||
.async_lookup(&keyed("python"), Duration::from_secs(102))
|
||||
.await
|
||||
} else {
|
||||
cache.lookup(&keyed("python"), Duration::from_secs(102))
|
||||
};
|
||||
assert_eq!(lookup.unwrap(), Some(response.clone()));
|
||||
let request = ResponseCacheRequest {
|
||||
context: litellm_cache::ExactCacheContext {
|
||||
ttl: Some(Duration::from_secs(12)),
|
||||
},
|
||||
..keyed("native")
|
||||
};
|
||||
let stored = if asynchronous {
|
||||
cache
|
||||
.async_store(&request, response, Duration::from_secs(102))
|
||||
.await
|
||||
} else {
|
||||
cache.store(&request, response, Duration::from_secs(102))
|
||||
};
|
||||
assert_eq!(stored, Ok(()));
|
||||
let requests = server.received_requests().await.unwrap();
|
||||
let upload = requests
|
||||
.iter()
|
||||
.find(|request| request.method.as_str() == "POST")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
upload.url.query(),
|
||||
Some("uploadType=media&name=cache%2Fnative")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn gcs_batch_reads_preserve_order_and_treat_invalid_entries_as_misses(
|
||||
#[future(awt)] gcs: (wiremock::MockServer, Gcs),
|
||||
) {
|
||||
use wiremock::{
|
||||
Mock, ResponseTemplate,
|
||||
matchers::{method, path},
|
||||
};
|
||||
|
||||
let (server, cache) = gcs;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fhit"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(json!({"timestamp": 100.0, "response": {"answer":7}})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Finvalid"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string("not an entry"))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/storage/v1/b/bucket/o/cache%2Fmissing"))
|
||||
.respond_with(ResponseTemplate::new(404))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let requests = [keyed("hit"), keyed("missing"), keyed("invalid")];
|
||||
let partial = cache
|
||||
.async_lookup_batch(&requests, Duration::from_secs(102))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(partial.values, vec![Some(json!({"answer":7})), None, None]);
|
||||
assert_eq!(partial.missing_indices, vec![1, 2]);
|
||||
}
|
||||
|
||||
type Gcs = ResponseCache<litellm_cache_gcs::GcsCache<litellm_cache_response::ResponseCacheCodec>>;
|
||||
|
||||
#[fixture]
|
||||
async fn gcs() -> (wiremock::MockServer, Gcs) {
|
||||
let server = wiremock::MockServer::start().await;
|
||||
let cache = ResponseCache::new(Arc::new(litellm_cache_gcs::GcsCache::with_token_source(
|
||||
litellm_cache_gcs::GcsConfig {
|
||||
bucket_name: "bucket".into(),
|
||||
gcs_path: Some("cache".into()),
|
||||
path_service_account: None,
|
||||
endpoint: server.uri(),
|
||||
},
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_cache_response::ResponseCacheCodec,
|
||||
Arc::new(litellm_cache_gcs::StaticTokenSource("token".into())),
|
||||
)));
|
||||
(server, cache)
|
||||
}
|
||||
|
|
|
|||
185
litellm-rust/crates/cache-response/tests/service.rs
Normal file
185
litellm-rust/crates/cache-response/tests/service.rs
Normal file
|
|
@ -0,0 +1,185 @@
|
|||
use std::{
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicU64, Ordering},
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use litellm_cache::ExactCacheContext;
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{
|
||||
CacheEntry, CacheKeyInput, ResponseCache, ResponseCacheConfig, ResponseCacheRequest,
|
||||
ResponseCacheService,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn service_honors_per_call_expiry_and_freshness() {
|
||||
let clock = Arc::new(AtomicU64::new(0));
|
||||
let cache_clock = clock.clone();
|
||||
let cache: Arc<dyn ResponseCacheService> = Arc::new(ResponseCache::new(Arc::new(
|
||||
InMemoryCache::with_clock(Some(100), Some(Duration::from_secs(60)), move || {
|
||||
Duration::from_secs(cache_clock.load(Ordering::SeqCst))
|
||||
}),
|
||||
)));
|
||||
let request = ResponseCacheRequest {
|
||||
context: ExactCacheContext {
|
||||
ttl: Some(Duration::from_secs(5)),
|
||||
},
|
||||
..ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some("entry".into()),
|
||||
..Default::default()
|
||||
})
|
||||
};
|
||||
cache
|
||||
.store(&request, json!({"answer":7}), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
cache.lookup(&request, Duration::ZERO).await.unwrap(),
|
||||
Some(json!({"answer":7}))
|
||||
);
|
||||
let stale_request = ResponseCacheRequest {
|
||||
max_age: Some(Duration::from_secs(1)),
|
||||
..request.clone()
|
||||
};
|
||||
clock.store(2, Ordering::SeqCst);
|
||||
assert_eq!(
|
||||
cache
|
||||
.lookup(&stale_request, Duration::from_secs(2))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
assert!(
|
||||
cache
|
||||
.lookup(&request, Duration::from_secs(2))
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some()
|
||||
);
|
||||
clock.store(6, Ordering::SeqCst);
|
||||
assert_eq!(
|
||||
cache
|
||||
.lookup(&request, Duration::from_secs(6))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn entry_limit_applies_to_sync_async_and_batch_writes() {
|
||||
let storage = Arc::new(InMemoryCache::<CacheEntry>::default());
|
||||
let cache = ResponseCache::new(storage.clone()).with_config(ResponseCacheConfig {
|
||||
namespace: "service-test".into(),
|
||||
max_entry_bytes: json!({"answer":7}).to_string().len(),
|
||||
});
|
||||
let small = json!({"answer":7});
|
||||
let large = json!({"answer":"too large"});
|
||||
let request = |key: &str| {
|
||||
ResponseCacheRequest::new(CacheKeyInput {
|
||||
preset: Some(key.into()),
|
||||
..Default::default()
|
||||
})
|
||||
};
|
||||
cache
|
||||
.store(&request("sync"), large.clone(), Duration::ZERO)
|
||||
.unwrap();
|
||||
cache
|
||||
.async_store(&request("async"), large.clone(), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
cache
|
||||
.async_store_batch(
|
||||
vec![
|
||||
(request("batch-large"), large),
|
||||
(request("batch-small"), small.clone()),
|
||||
],
|
||||
Duration::ZERO,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let service: Arc<dyn ResponseCacheService> = Arc::new(cache);
|
||||
service
|
||||
.store(&request("service"), small.clone(), Duration::ZERO)
|
||||
.await
|
||||
.unwrap();
|
||||
for key in ["sync", "async", "batch-large"] {
|
||||
assert!(storage.get_cache(key).unwrap().is_none());
|
||||
}
|
||||
for key in ["batch-small", "service"] {
|
||||
assert_eq!(
|
||||
service.lookup(&request(key), Duration::ZERO).await.unwrap(),
|
||||
Some(small.clone())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::same_scope("tenant-a", "tenant-a", true)]
|
||||
#[case::different_scope("tenant-a", "tenant-b", false)]
|
||||
#[case::empty_isolated_scope("", "", true)]
|
||||
#[tokio::test]
|
||||
async fn isolated_policy_controls_actual_entry_reuse(
|
||||
#[case] first: &str,
|
||||
#[case] second: &str,
|
||||
#[case] hit: bool,
|
||||
#[values(false, true)] override_policy: bool,
|
||||
) {
|
||||
use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache};
|
||||
let service = Arc::new(ResponseCache::new(Arc::new(
|
||||
InMemoryCache::<CacheEntry>::default(),
|
||||
)));
|
||||
let request = |scope| {
|
||||
ScopedCache::new(service.clone(), scope)
|
||||
.options(override_policy.then_some(CachePolicy {
|
||||
ttl: Some(Duration::from_secs(30)),
|
||||
..CachePolicy::default()
|
||||
}))
|
||||
.request("test", "messages", json!({"prompt":"hello"}))
|
||||
};
|
||||
service
|
||||
.async_store(
|
||||
&request(CacheScope::Isolated(first.into())),
|
||||
json!({"answer":7}),
|
||||
Duration::ZERO,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
service
|
||||
.async_lookup(
|
||||
&request(CacheScope::Isolated(second.into())),
|
||||
Duration::ZERO
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
hit.then(|| json!({"answer":7}))
|
||||
);
|
||||
assert_eq!(
|
||||
service
|
||||
.async_lookup(&request(CacheScope::Shared), Duration::ZERO)
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::valid(1, "messages", Some(7))]
|
||||
#[case::unknown_version(2, "messages", None)]
|
||||
#[case::another_surface(1, "responses", None)]
|
||||
fn envelopes_require_a_matching_surface_and_version(
|
||||
#[case] version: u32,
|
||||
#[case] surface: &str,
|
||||
#[case] expected: Option<u32>,
|
||||
) {
|
||||
let envelope: litellm_cache_response::ResponseEnvelope<u32> =
|
||||
serde_json::from_value(json!({"version":version,"surface":surface,"output":7})).unwrap();
|
||||
assert_eq!(envelope.decode("messages"), expected);
|
||||
}
|
||||
|
|
@ -1,12 +1,13 @@
|
|||
- Target invariants, not completion claims
|
||||
- This crate is the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on)
|
||||
- This crate owns compatibility for all existing Python callbacks and loggers, including `CustomLogger`. `mapping.rs` owns the executable call bindings and the inventory of Python-owned hooks. A Python-owned entry records an existing path, never permission to invoke it a second time. The native call adapter preserves the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on)
|
||||
- Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here
|
||||
- SDK request policy (credential inheritance, the budget and retry-count limits) is the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call
|
||||
- SDK request policy (credential inheritance, the budget and retry-count limits) is a separate hook supplied by `python-bridge`; compose it after this adapter so logging adopts the final keyword view before policy mutates or rejects it
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonCallHooks` using the shared `CallEvent`; they never learn which Python objects consume a call
|
||||
- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json`
|
||||
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
|
||||
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; routes hand it over through `run_legacy_call` and keep no copy
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; shared bridge composition hands it to `LegacyLogging`; routes use the neutral call boundary
|
||||
- `LoggingOperation` selects legacy logging entrypoints and response handling. It belongs here rather than in shared inference data contracts
|
||||
- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation
|
||||
- Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -3,16 +3,13 @@
|
|||
//! lifetime. No other callback host has that obligation, which is why nothing outside
|
||||
//! this crate holds them.
|
||||
|
||||
use litellm_host::{machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call};
|
||||
use litellm_host_python::lookup;
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
};
|
||||
|
||||
use crate::{LegacyLogging, LegacySurface};
|
||||
|
||||
pub struct PublicCall {
|
||||
args: Py<PyTuple>,
|
||||
kwargs: Py<PyDict>,
|
||||
|
|
@ -34,6 +31,10 @@ impl PublicCall {
|
|||
})
|
||||
}
|
||||
|
||||
pub fn arguments(&self, py: Python<'_>) -> Py<PyDict> {
|
||||
self.kwargs.clone_ref(py)
|
||||
}
|
||||
|
||||
pub(crate) fn args(&self) -> &Py<PyTuple> {
|
||||
&self.args
|
||||
}
|
||||
|
|
@ -64,34 +65,6 @@ impl PublicCall {
|
|||
}
|
||||
}
|
||||
|
||||
/// Runs one native call under the legacy `Logging` contract: the protocol host projects from
|
||||
/// the keyword view the contract prepares and `preflight` rewrites, and the contract
|
||||
/// observes the call.
|
||||
pub fn run_legacy_call<H, M>(
|
||||
py: Python<'_>,
|
||||
surface: LegacySurface,
|
||||
call: PublicCall,
|
||||
machine: M,
|
||||
host: H,
|
||||
preflight: Preflight,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
H: ProtocolHost + 'static,
|
||||
M: Machine<Protocol = H::Protocol, Complete = <H::Protocol as Protocol>::Response> + 'static,
|
||||
{
|
||||
let arguments = call.kwargs.clone_ref(py);
|
||||
run_call(
|
||||
py,
|
||||
machine,
|
||||
host,
|
||||
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
|
||||
preflight,
|
||||
arguments,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
@ -110,7 +83,7 @@ mod tests {
|
|||
(call, locals)
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn capture_copies_the_keyword_dict_without_copying_its_values() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
//! the deferred and worker-submitted success paths, and the sync-callbacks-for-async-calls
|
||||
//! duplication. All of it expires with the legacy callback contract.
|
||||
|
||||
use litellm_host::event::{RequestContext, WireRequest};
|
||||
use litellm_host::interceptors::{RequestContext, WireRequest};
|
||||
use litellm_host_python::to_py;
|
||||
use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict};
|
||||
|
||||
|
|
|
|||
|
|
@ -1,26 +1,32 @@
|
|||
//! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the
|
||||
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
|
||||
//! proxy release. All of it sits behind one
|
||||
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
|
||||
//! core never learn which Python object is on the other end. The SDK's own request policy
|
||||
//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this
|
||||
//! crate's.
|
||||
//! [`PythonCallHooks`](litellm_host_python::PythonCallHooks), so the driver, the routes and
|
||||
//! core never learn which Python object is on the other end.
|
||||
//!
|
||||
//! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`]
|
||||
//! is where those objects live, and [`run_legacy_call`] is how a route hands them over
|
||||
//! without keeping a copy.
|
||||
//! is where those objects live.
|
||||
|
||||
mod adapter;
|
||||
mod call;
|
||||
mod callbacks;
|
||||
mod deferred;
|
||||
mod logger;
|
||||
mod mapping;
|
||||
mod python;
|
||||
pub(crate) use adapter::LegacyLogging;
|
||||
pub use adapter::{LegacySurface, PassThroughStream};
|
||||
pub use call::{PublicCall, run_legacy_call};
|
||||
pub use adapter::LegacyLogging;
|
||||
pub use call::PublicCall;
|
||||
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
|
||||
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
|
||||
pub use mapping::{CallBoundary, CallbackMapping, Dispatch, callback_mappings};
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum LoggingOperation {
|
||||
Completion,
|
||||
Responses,
|
||||
Messages,
|
||||
Ocr,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod test_support;
|
||||
|
|
|
|||
288
litellm-rust/crates/callbacks-legacy-python/src/mapping.rs
Normal file
288
litellm-rust/crates/callbacks-legacy-python/src/mapping.rs
Normal file
|
|
@ -0,0 +1,288 @@
|
|||
use litellm_host::{
|
||||
hooks::CallHooks,
|
||||
interceptors::{RawResponse, RequestContext, WireRequest},
|
||||
lifecycle::{ExecutionEvent, FailureOrigin, Timing},
|
||||
};
|
||||
use litellm_host_python::{HookStep, PythonCallEvent, PythonRuntime};
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
|
||||
use crate::LegacyLogging;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum CallBoundary {
|
||||
PrepareArguments,
|
||||
BeforeProviderRequest,
|
||||
AfterProviderResponse,
|
||||
TransformResponse,
|
||||
Succeeded,
|
||||
Failed,
|
||||
StreamOpened,
|
||||
StreamChunk,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Dispatch {
|
||||
Call(CallBoundary),
|
||||
Python(&'static str),
|
||||
DeclarationOnly,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct CallbackMapping {
|
||||
pub callback: &'static str,
|
||||
pub dispatch: Dispatch,
|
||||
}
|
||||
|
||||
struct Binding<H> {
|
||||
boundary: CallBoundary,
|
||||
invoke: H,
|
||||
callbacks: &'static [&'static str],
|
||||
}
|
||||
|
||||
impl<H> Binding<H> {
|
||||
fn mappings(&self) -> impl Iterator<Item = CallbackMapping> {
|
||||
self.callbacks.iter().map(|callback| CallbackMapping {
|
||||
callback,
|
||||
dispatch: Dispatch::Call(self.boundary),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type Step<T> = PyResult<HookStep<LegacyLogging, T>>;
|
||||
type Prepare = fn(&mut LegacyLogging, Python<'_>, Py<PyDict>, f64) -> Step<Py<PyDict>>;
|
||||
type Before =
|
||||
fn(&mut LegacyLogging, Python<'_>, Box<WireRequest>, &RequestContext) -> Step<Box<WireRequest>>;
|
||||
type After = fn(&mut LegacyLogging, Python<'_>, &RawResponse) -> Step<()>;
|
||||
type Transform = fn(&mut LegacyLogging, Python<'_>, Py<PyAny>, Timing) -> Step<Py<PyAny>>;
|
||||
type Success = fn(&mut LegacyLogging, Python<'_>, Timing, &Py<PyAny>) -> Step<()>;
|
||||
type Failure = fn(&mut LegacyLogging, Python<'_>, Timing, FailureOrigin, &PyErr) -> Step<()>;
|
||||
type Open = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
|
||||
type Chunk = fn(&mut LegacyLogging, Python<'_>, &Py<PyAny>) -> PyResult<()>;
|
||||
|
||||
const PREPARE: Binding<Prepare> = Binding {
|
||||
boundary: CallBoundary::PrepareArguments,
|
||||
invoke: LegacyLogging::prepare_call,
|
||||
callbacks: &["async_pre_call_deployment_hook"],
|
||||
};
|
||||
|
||||
const BEFORE: Binding<Before> = Binding {
|
||||
boundary: CallBoundary::BeforeProviderRequest,
|
||||
invoke: LegacyLogging::pre_call,
|
||||
callbacks: &["log_pre_api_call", "log_input_event"],
|
||||
};
|
||||
|
||||
const AFTER: Binding<After> = Binding {
|
||||
boundary: CallBoundary::AfterProviderResponse,
|
||||
invoke: LegacyLogging::post_call,
|
||||
callbacks: &["log_post_api_call"],
|
||||
};
|
||||
|
||||
const TRANSFORM: Binding<Transform> = Binding {
|
||||
boundary: CallBoundary::TransformResponse,
|
||||
invoke: LegacyLogging::transform_public_response,
|
||||
callbacks: &["async_post_call_success_deployment_hook"],
|
||||
};
|
||||
|
||||
const SUCCESS: Binding<Success> = Binding {
|
||||
boundary: CallBoundary::Succeeded,
|
||||
invoke: LegacyLogging::succeeded,
|
||||
callbacks: &[
|
||||
"log_success_event",
|
||||
"async_log_success_event",
|
||||
"logging_hook",
|
||||
"async_logging_hook",
|
||||
"redact_standard_logging_payload_from_model_call_details",
|
||||
"log_event",
|
||||
"async_log_event",
|
||||
],
|
||||
};
|
||||
|
||||
const FAILURE: Binding<Failure> = Binding {
|
||||
boundary: CallBoundary::Failed,
|
||||
invoke: LegacyLogging::failed,
|
||||
callbacks: &[
|
||||
"async_post_call_failure_deployment_hook",
|
||||
"log_failure_event",
|
||||
"async_log_failure_event",
|
||||
"log_model_group_rate_limit_error",
|
||||
"log_event",
|
||||
"async_log_event",
|
||||
],
|
||||
};
|
||||
|
||||
const OPEN: Binding<Open> = Binding {
|
||||
boundary: CallBoundary::StreamOpened,
|
||||
invoke: LegacyLogging::stream_opened,
|
||||
callbacks: &[],
|
||||
};
|
||||
|
||||
const CHUNK: Binding<Chunk> = Binding {
|
||||
boundary: CallBoundary::StreamChunk,
|
||||
invoke: LegacyLogging::stream_chunk,
|
||||
callbacks: &[],
|
||||
};
|
||||
|
||||
pub fn callback_mappings() -> impl Iterator<Item = CallbackMapping> {
|
||||
PREPARE
|
||||
.mappings()
|
||||
.chain(BEFORE.mappings())
|
||||
.chain(AFTER.mappings())
|
||||
.chain(TRANSFORM.mappings())
|
||||
.chain(SUCCESS.mappings())
|
||||
.chain(FAILURE.mappings())
|
||||
.chain(OPEN.mappings())
|
||||
.chain(CHUNK.mappings())
|
||||
.chain(PYTHON_CALLBACKS.iter().copied())
|
||||
}
|
||||
|
||||
impl CallHooks<PythonRuntime> for LegacyLogging {
|
||||
fn prepare_arguments(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
arguments: Py<PyDict>,
|
||||
started_at: f64,
|
||||
) -> Step<Py<PyDict>> {
|
||||
(PREPARE.invoke)(self, py, arguments, started_at)
|
||||
}
|
||||
|
||||
fn arguments_prepared(&mut self, py: Python<'_>, arguments: &Py<PyDict>) -> PyResult<()> {
|
||||
self.adopt_arguments(py, arguments);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn before_provider_request(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
wire: Box<WireRequest>,
|
||||
context: &RequestContext,
|
||||
) -> Step<Box<WireRequest>> {
|
||||
(BEFORE.invoke)(self, py, wire, context)
|
||||
}
|
||||
|
||||
fn transform_response(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
response: Py<PyAny>,
|
||||
timing: Timing,
|
||||
) -> Step<Py<PyAny>> {
|
||||
(TRANSFORM.invoke)(self, py, response, timing)
|
||||
}
|
||||
|
||||
fn on_event(&mut self, py: Python<'_>, event: PythonCallEvent<'_>) -> Step<()> {
|
||||
match event {
|
||||
PythonCallEvent::Started { .. } | PythonCallEvent::Cancelled { .. } => {
|
||||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
PythonCallEvent::Execution(ExecutionEvent::ResultReady { facts }) => {
|
||||
self.result_ready(py, &facts)
|
||||
}
|
||||
PythonCallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
(AFTER.invoke)(self, py, raw)
|
||||
}
|
||||
PythonCallEvent::Succeeded { timing, response } => {
|
||||
(SUCCESS.invoke)(self, py, timing, response)
|
||||
}
|
||||
PythonCallEvent::Failed {
|
||||
timing,
|
||||
origin,
|
||||
error,
|
||||
} => (FAILURE.invoke)(self, py, timing, origin, error),
|
||||
}
|
||||
}
|
||||
|
||||
fn on_stream_open(&mut self, py: Python<'_>, head: &Py<PyAny>) -> PyResult<()> {
|
||||
(OPEN.invoke)(self, py, head)
|
||||
}
|
||||
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
(CHUNK.invoke)(self, py, chunk)
|
||||
}
|
||||
}
|
||||
|
||||
macro_rules! python_callbacks {
|
||||
($($dispatch:expr => [$($callback:literal),* $(,)?]),* $(,)?) => {
|
||||
const PYTHON_CALLBACKS: &[CallbackMapping] = &[
|
||||
$($(CallbackMapping { callback: $callback, dispatch: $dispatch },)*)*
|
||||
];
|
||||
};
|
||||
}
|
||||
|
||||
python_callbacks! {
|
||||
Dispatch::Python("litellm.router") => [
|
||||
"async_pre_routing_hook",
|
||||
"async_filter_deployments",
|
||||
"pre_call_check",
|
||||
"async_pre_call_check",
|
||||
],
|
||||
Dispatch::Python("litellm.router_utils.fallback_event_handlers") => [
|
||||
"log_success_fallback_event",
|
||||
"log_failure_fallback_event",
|
||||
],
|
||||
Dispatch::Python("litellm.proxy.utils") => [
|
||||
"async_pre_call_hook",
|
||||
"async_post_call_response_headers_hook",
|
||||
"async_post_call_failure_hook",
|
||||
"async_post_call_success_hook",
|
||||
"async_moderation_hook",
|
||||
"async_post_call_streaming_hook",
|
||||
"async_post_call_streaming_iterator_hook",
|
||||
"async_filter_listed_models",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.litellm_logging") => [
|
||||
"async_get_chat_completion_prompt",
|
||||
"get_chat_completion_prompt",
|
||||
"log_stream_event",
|
||||
"async_log_stream_event",
|
||||
"async_post_mcp_tool_call_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.llms.anthropic.pass_through.messages.handler") => [
|
||||
"async_pre_request_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.streaming_handler") => [
|
||||
"async_post_call_streaming_deployment_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.responses.streaming_iterator") => [
|
||||
"async_post_call_streaming_deployment_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.main") => [
|
||||
"translate_completion_input_params",
|
||||
"translate_completion_output_params",
|
||||
"translate_completion_output_params_streaming",
|
||||
],
|
||||
Dispatch::Python("litellm.integrations.argilla") => ["async_dataset_hook"],
|
||||
Dispatch::Python("litellm.proxy.management_helpers.audit_logs") => ["async_log_audit_log_event"],
|
||||
Dispatch::Python("litellm.llms.custom_httpx.llm_http_handler") => [
|
||||
"async_should_run_agentic_loop",
|
||||
"async_run_agentic_loop",
|
||||
"async_build_agentic_loop_plan",
|
||||
"async_post_agentic_loop_response_hook",
|
||||
"async_agentic_loop_cleanup_hook",
|
||||
"async_should_run_chat_completion_agentic_loop",
|
||||
"async_run_chat_completion_agentic_loop",
|
||||
"async_build_chat_completion_agentic_loop_plan",
|
||||
],
|
||||
Dispatch::Python("litellm.litellm_core_utils.chat_completion_agentic_loop") => [
|
||||
"async_should_run_agentic_loop",
|
||||
"async_run_agentic_loop",
|
||||
"async_build_agentic_loop_plan",
|
||||
"async_post_agentic_loop_response_hook",
|
||||
"async_agentic_loop_cleanup_hook",
|
||||
],
|
||||
Dispatch::Python("litellm.llms.openai.openai") => [
|
||||
"async_should_run_chat_completion_agentic_loop",
|
||||
"async_run_chat_completion_agentic_loop",
|
||||
],
|
||||
Dispatch::Python("litellm.proxy.spend_tracking.cold_storage_handler") => [
|
||||
"get_proxy_server_request_from_cold_storage_with_object_key",
|
||||
],
|
||||
Dispatch::Python("litellm.integrations.custom_logger") => [
|
||||
"truncate_standard_logging_payload_content",
|
||||
"redacts_messages_itself",
|
||||
"handle_callback_failure",
|
||||
"get_callback_env_vars",
|
||||
],
|
||||
Dispatch::DeclarationOnly => [
|
||||
"async_log_pre_api_call",
|
||||
"async_log_input_event",
|
||||
],
|
||||
}
|
||||
|
|
@ -3,7 +3,7 @@ use std::ffi::CStr;
|
|||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use crate::{LegacyLogging, LegacySurface, PublicCall};
|
||||
use crate::{LegacyLogging, PublicCall};
|
||||
|
||||
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
|
||||
/// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
|
||||
|
|
@ -45,11 +45,14 @@ def contracted(name, fake):
|
|||
if not hasattr(legacy, 'is_internal'):
|
||||
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
|
||||
|
||||
def setup(call_type, args, kwargs, start, asynchronous):
|
||||
logger = kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger']
|
||||
logger.setup_call_type = call_type
|
||||
return types.SimpleNamespace(logger=logger, kwargs=kwargs)
|
||||
|
||||
|
||||
FAKES = {
|
||||
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
|
||||
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
|
||||
kwargs=kwargs,
|
||||
),
|
||||
'setup': setup,
|
||||
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
|
||||
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -82,10 +85,10 @@ FAKES = {
|
|||
),
|
||||
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
|
||||
'stream_opened': lambda logger: logger.record('stream_opened', None),
|
||||
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
|
||||
'stream_success': lambda logger, url_route, endpoint_type, request_body, chunks, start, end, first_chunk: logger.record(
|
||||
'stream_success', list(chunks)
|
||||
),
|
||||
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
|
||||
'stream_failure': lambda logger, endpoint_type, request_body, chunks, error: logger.record('stream_failure', error),
|
||||
}
|
||||
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
|
||||
for name, fake in FAKES.items():
|
||||
|
|
@ -186,14 +189,5 @@ pub(crate) fn legacy_call(
|
|||
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
|
||||
.unwrap_or_else(|| PyDict::new(py));
|
||||
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
|
||||
LegacyLogging::new(
|
||||
py,
|
||||
LegacySurface {
|
||||
call_type: "test",
|
||||
input_description: "test input",
|
||||
stream: None,
|
||||
},
|
||||
call,
|
||||
asynchronous,
|
||||
)
|
||||
LegacyLogging::new(py, crate::LoggingOperation::Ocr, call, asynchronous)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
mod error;
|
||||
mod includes;
|
||||
mod mcp;
|
||||
mod model;
|
||||
mod settings;
|
||||
mod value;
|
||||
|
|
@ -9,6 +10,7 @@ use std::{fmt, path::Path};
|
|||
use serde::Deserialize;
|
||||
|
||||
pub use error::Error;
|
||||
pub use mcp::{McpAuth, McpServer, McpTransport};
|
||||
pub use model::{LiteLlmParams, Model};
|
||||
pub use settings::{GeneralSettings, LiteLlmSettings, RouterSettings};
|
||||
pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
|
||||
|
|
@ -24,6 +26,7 @@ pub struct Config {
|
|||
pub callback_settings: Object,
|
||||
pub assistant_settings: Object,
|
||||
pub default_vertex_config: Object,
|
||||
pub mcp_servers: std::collections::BTreeMap<String, McpServer>,
|
||||
pub credential_list: Box<[Object]>,
|
||||
pub guardrails: Box<[Object]>,
|
||||
pub prompts: Box<[Object]>,
|
||||
|
|
@ -55,6 +58,7 @@ impl fmt::Debug for Config {
|
|||
.field("callback_settings", &self.callback_settings)
|
||||
.field("assistant_settings", &self.assistant_settings)
|
||||
.field("default_vertex_config", &self.default_vertex_config)
|
||||
.field("mcp_servers", &self.mcp_servers)
|
||||
.field("credential_list", &self.credential_list)
|
||||
.field("guardrails", &self.guardrails)
|
||||
.field("prompts", &self.prompts)
|
||||
|
|
|
|||
97
litellm-rust/crates/config/src/mcp.rs
Normal file
97
litellm-rust/crates/config/src/mcp.rs
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
use std::{collections::BTreeMap, fmt};
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::Object;
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct McpServer {
|
||||
pub server_id: Option<String>,
|
||||
pub alias: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub mcp_info: Object,
|
||||
pub transport: McpTransport,
|
||||
pub url: Option<SecretValue>,
|
||||
pub command: Option<String>,
|
||||
pub args: Box<[String]>,
|
||||
pub env: BTreeMap<String, SecretValue>,
|
||||
pub auth_type: Option<McpAuth>,
|
||||
#[serde(alias = "auth_value")]
|
||||
pub authentication_token: Option<SecretValue>,
|
||||
pub static_headers: BTreeMap<String, SecretValue>,
|
||||
pub upstream_token_header: Option<String>,
|
||||
pub allowed_tools: Option<Box<[String]>>,
|
||||
pub timeout: Option<f64>,
|
||||
pub max_concurrent_requests: Option<usize>,
|
||||
#[serde(flatten)]
|
||||
pub unsupported: Object,
|
||||
}
|
||||
|
||||
impl fmt::Debug for McpServer {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("McpServer")
|
||||
.field("transport", &self.transport)
|
||||
.field("auth_type", &self.auth_type)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("max_concurrent_requests", &self.max_concurrent_requests)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpTransport {
|
||||
#[default]
|
||||
Http,
|
||||
Sse,
|
||||
Stdio,
|
||||
}
|
||||
|
||||
impl McpTransport {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Http => "http",
|
||||
Self::Sse => "sse",
|
||||
Self::Stdio => "stdio",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpAuth {
|
||||
None,
|
||||
ApiKey,
|
||||
BearerToken,
|
||||
Basic,
|
||||
Authorization,
|
||||
Token,
|
||||
Oauth2,
|
||||
AwsSigv4,
|
||||
Oauth2TokenExchange,
|
||||
Oauth2IdJag,
|
||||
TruePassthrough,
|
||||
OauthDelegate,
|
||||
}
|
||||
|
||||
impl McpAuth {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::None => "none",
|
||||
Self::ApiKey => "api_key",
|
||||
Self::BearerToken => "bearer_token",
|
||||
Self::Basic => "basic",
|
||||
Self::Authorization => "authorization",
|
||||
Self::Token => "token",
|
||||
Self::Oauth2 => "oauth2",
|
||||
Self::AwsSigv4 => "aws_sigv4",
|
||||
Self::Oauth2TokenExchange => "oauth2_token_exchange",
|
||||
Self::Oauth2IdJag => "oauth2_id_jag",
|
||||
Self::TruePassthrough => "true_passthrough",
|
||||
Self::OauthDelegate => "oauth_delegate",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -35,6 +35,8 @@ pub struct GeneralSettings {
|
|||
pub dangerously_permit_weak_or_unset_master_key: Option<bool>,
|
||||
pub plugins: Option<Box<[Object]>>,
|
||||
pub coordination_redis: Option<Object>,
|
||||
pub mcp_allowed_hosts: Option<Box<[String]>>,
|
||||
pub mcp_allowed_origins: Box<[String]>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
|
@ -69,6 +71,8 @@ impl Default for GeneralSettings {
|
|||
dangerously_permit_weak_or_unset_master_key: None,
|
||||
plugins: None,
|
||||
coordination_redis: None,
|
||||
mcp_allowed_hosts: None,
|
||||
mcp_allowed_origins: Box::default(),
|
||||
additional_fields: AdditionalFields::new(),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -233,6 +233,9 @@ finetune_settings:
|
|||
- custom_llm_provider: openai
|
||||
mcp_tools:
|
||||
- name: lookup
|
||||
mcp_servers:
|
||||
docs:
|
||||
url: https://example.test/mcp
|
||||
vector_store_registry:
|
||||
- vector_store_name: docs
|
||||
worker_registry:
|
||||
|
|
@ -262,6 +265,7 @@ include:
|
|||
assert_eq!(config.files_settings.len(), 1);
|
||||
assert_eq!(config.finetune_settings.len(), 1);
|
||||
assert_eq!(config.mcp_tools.len(), 1);
|
||||
assert_eq!(config.mcp_servers.len(), 1);
|
||||
assert_eq!(config.vector_store_registry.len(), 1);
|
||||
assert_eq!(config.worker_registry.len(), 1);
|
||||
assert_eq!(config.agents.len(), 1);
|
||||
|
|
@ -341,3 +345,31 @@ fn resolves_nested_includes_once_in_breadth_first_order() {
|
|||
);
|
||||
assert!(config.include.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn mcp_config_redacts_nested_credentials_and_preserves_policy_for_validation() {
|
||||
let config = Config::from_yaml("mcp_servers:\n docs:\n url: https://example.test/private-secret/mcp\n authentication_token: upstream-secret\n static_headers: {x-token: header-secret}\n env: {TOKEN: env-secret}\n args: [argument-secret]\n client_secret: oauth-secret\n allowed_tools: [search]\n").unwrap();
|
||||
let server = &config.mcp_servers["docs"];
|
||||
assert_eq!(server.allowed_tools.as_deref().unwrap(), ["search"]);
|
||||
assert!(server.unsupported.contains_key("client_secret"));
|
||||
let debug = format!("{config:?}");
|
||||
for secret in [
|
||||
"private-secret",
|
||||
"upstream-secret",
|
||||
"header-secret",
|
||||
"env-secret",
|
||||
"oauth-secret",
|
||||
"argument-secret",
|
||||
] {
|
||||
assert!(!debug.contains(secret));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::transport("transport: invalid")]
|
||||
#[case::auth("auth_type: invalid")]
|
||||
#[case::concurrency("max_concurrent_requests: -1")]
|
||||
#[case::headers("static_headers: {x-token: [not, a, string]}")]
|
||||
fn rejects_invalid_typed_mcp_settings(#[case] setting: &str) {
|
||||
assert!(Config::from_yaml(&format!("mcp_servers:\n docs:\n {setting}\n")).is_err());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,11 +8,10 @@ repository.workspace = true
|
|||
[dependencies]
|
||||
fancy-regex.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
litellm-types.workspace = true
|
||||
litellm-llms-types.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
serde_path_to_error = "0.1"
|
||||
serde_with.workspace = true
|
||||
strum.workspace = true
|
||||
thiserror.workspace = true
|
||||
url.workspace = true
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use litellm_types::utils::{ChatCompletionsUsage, PromptTokensDetails};
|
||||
use litellm_llms_types::formats::chat_completions::{ChatCompletionsUsage, PromptTokensDetails};
|
||||
|
||||
/// OpenAI finish reasons, mirroring Python's `_FINISH_REASON_MAP` for the
|
||||
/// reasons the providers on this route can emit. Python warns and falls back to
|
||||
|
|
|
|||
|
|
@ -1,9 +1,26 @@
|
|||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct CustomLlmProvider<'a> {
|
||||
pub model: &'a str,
|
||||
pub custom_llm_provider: &'a str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub enum LlmProviders {
|
||||
Anthropic,
|
||||
AwsTextract,
|
||||
AzureAi,
|
||||
Bedrock,
|
||||
Cohere,
|
||||
Mistral,
|
||||
Openai,
|
||||
OpenaiLike,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
pub fn get_custom_llm_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use litellm_llms_types::headers::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub fn get_provider_specific_headers(
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@
|
|||
//! `_bedrock_converse_messages_pt` for the text-only surface this route
|
||||
//! accepts; anything richer is declined upstream by the capability gate.
|
||||
|
||||
use litellm_types::llms::openai::{ChatMessage, ChatMessageContent};
|
||||
use litellm_llms_types::formats::chat_completions::{ChatMessage, ChatMessageContent};
|
||||
use strum::IntoStaticStr;
|
||||
|
||||
pub const EMPTY_TEXT_PLACEHOLDER: &str =
|
||||
|
|
|
|||
|
|
@ -1,12 +1,3 @@
|
|||
use serde::{
|
||||
Deserializer,
|
||||
de::{Error, Visitor},
|
||||
};
|
||||
use serde_with::DeserializeAs;
|
||||
|
||||
pub struct LaxI64;
|
||||
pub struct FiniteF64;
|
||||
|
||||
pub fn parse_str_bool(value: &str) -> Option<bool> {
|
||||
let token = value.trim_matches(|character: char| {
|
||||
character.is_whitespace() || matches!(character, '\u{1c}'..='\u{1f}')
|
||||
|
|
@ -22,129 +13,12 @@ pub fn parse_redis_bool(value: &str) -> bool {
|
|||
value == "1" || value.eq_ignore_ascii_case("true") || value.eq_ignore_ascii_case("yes")
|
||||
}
|
||||
|
||||
impl<'de> DeserializeAs<'de, i64> for LaxI64 {
|
||||
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<i64, D::Error> {
|
||||
deserializer.deserialize_any(Self)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Visitor<'de> for LaxI64 {
|
||||
type Value = i64;
|
||||
|
||||
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("an integer in the i64 range")
|
||||
}
|
||||
|
||||
fn visit_i64<E: Error>(self, value: i64) -> Result<i64, E> {
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn visit_u64<E: Error>(self, value: u64) -> Result<i64, E> {
|
||||
i64::try_from(value).map_err(E::custom)
|
||||
}
|
||||
|
||||
fn visit_f64<E: Error>(self, value: f64) -> Result<i64, E> {
|
||||
integral_float(value).ok_or_else(|| E::custom("expected an integer in the i64 range"))
|
||||
}
|
||||
|
||||
fn visit_str<E: Error>(self, value: &str) -> Result<i64, E> {
|
||||
integer_string(value.trim())
|
||||
.ok_or_else(|| E::custom("expected an integer in the i64 range"))
|
||||
}
|
||||
|
||||
fn visit_bool<E: Error>(self, value: bool) -> Result<i64, E> {
|
||||
Ok(i64::from(value))
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> DeserializeAs<'de, f64> for FiniteF64 {
|
||||
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
|
||||
deserializer.deserialize_any(Self)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Visitor<'de> for FiniteF64 {
|
||||
type Value = f64;
|
||||
|
||||
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str("a finite number")
|
||||
}
|
||||
|
||||
fn visit_i64<E: Error>(self, value: i64) -> Result<f64, E> {
|
||||
Ok(value as f64)
|
||||
}
|
||||
|
||||
fn visit_u64<E: Error>(self, value: u64) -> Result<f64, E> {
|
||||
Ok(value as f64)
|
||||
}
|
||||
|
||||
fn visit_f64<E: Error>(self, value: f64) -> Result<f64, E> {
|
||||
value
|
||||
.is_finite()
|
||||
.then_some(value)
|
||||
.ok_or_else(|| E::custom("expected a finite number"))
|
||||
}
|
||||
|
||||
fn visit_str<E: Error>(self, value: &str) -> Result<f64, E> {
|
||||
self.visit_f64(value.trim().parse::<f64>().map_err(E::custom)?)
|
||||
}
|
||||
|
||||
fn visit_bool<E: Error>(self, value: bool) -> Result<f64, E> {
|
||||
Ok(f64::from(value))
|
||||
}
|
||||
}
|
||||
|
||||
fn integer_string(value: &str) -> Option<i64> {
|
||||
let integer = match value.split_once('.') {
|
||||
Some((integer, fraction)) => {
|
||||
if fraction.is_empty() || !fraction.bytes().all(|byte| byte == b'0') {
|
||||
return None;
|
||||
}
|
||||
integer
|
||||
}
|
||||
None => value,
|
||||
};
|
||||
if integer.starts_with('_') || integer.ends_with('_') || integer.contains("__") {
|
||||
return None;
|
||||
}
|
||||
let digits = integer.strip_prefix(['+', '-']).unwrap_or(integer);
|
||||
if digits.is_empty()
|
||||
|| digits.starts_with('_')
|
||||
|| !digits
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_digit() || byte == b'_')
|
||||
{
|
||||
return None;
|
||||
}
|
||||
integer.replace('_', "").parse().ok()
|
||||
}
|
||||
|
||||
fn integral_float(value: f64) -> Option<i64> {
|
||||
(value.is_finite()
|
||||
&& value.fract() == 0.0
|
||||
&& value >= i64::MIN as f64
|
||||
&& value < -(i64::MIN as f64))
|
||||
.then_some(value as i64)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::json;
|
||||
use serde_with::serde_as;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[serde_as]
|
||||
#[derive(Debug, Deserialize, Serialize, PartialEq)]
|
||||
struct Numbers {
|
||||
#[serde_as(deserialize_as = "Option<Vec<LaxI64>>")]
|
||||
integers: Option<Vec<i64>>,
|
||||
#[serde_as(deserialize_as = "Option<FiniteF64>")]
|
||||
float: Option<f64>,
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::trimmed_true(" True ", Some(true))]
|
||||
#[case::control_whitespace_true("\u{1c}TRUE\u{1f}", Some(true))]
|
||||
|
|
@ -160,73 +34,4 @@ mod tests {
|
|||
) {
|
||||
assert_eq!(parse_str_bool(input), expected, "{input:?}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn adapters_compose_and_serialize_as_numbers() {
|
||||
let numbers: Numbers = serde_json::from_value(json!({
|
||||
"integers": ["9007199254740993.0", "1_000", " +2.000 ", 3.0, true],
|
||||
"float": " 1.5 "
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(numbers).unwrap(),
|
||||
json!({
|
||||
"integers": [9_007_199_254_740_993_i64, 1000, 2, 3, 1], "float": 1.5
|
||||
})
|
||||
);
|
||||
for input in [json!({}), json!({"integers": null, "float": null})] {
|
||||
assert_eq!(
|
||||
serde_json::from_value::<Numbers>(input).unwrap(),
|
||||
Numbers {
|
||||
integers: None,
|
||||
float: None,
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn integer_bounds_and_invalid_values_are_checked() {
|
||||
for input in [
|
||||
json!(i64::MIN),
|
||||
json!(i64::MAX),
|
||||
json!(i64::MAX.to_string()),
|
||||
] {
|
||||
assert!(serde_json::from_value::<Numbers>(json!({"integers": [input]})).is_ok());
|
||||
}
|
||||
for input in [
|
||||
json!(u64::MAX),
|
||||
json!(9_223_372_036_854_775_808_u64),
|
||||
json!(9_223_372_036_854_775_808.0),
|
||||
json!("-9223372036854775809"),
|
||||
json!("1.0000000000000001"),
|
||||
json!("1e3"),
|
||||
json!("2."),
|
||||
json!(".0"),
|
||||
json!("_2"),
|
||||
json!("2__0"),
|
||||
json!(2.5),
|
||||
json!(null),
|
||||
json!({}),
|
||||
] {
|
||||
assert!(serde_json::from_value::<Numbers>(json!({"integers": [input]})).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn floats_reject_nonfinite_and_invalid_values() {
|
||||
for input in [
|
||||
json!("NaN"),
|
||||
json!("inf"),
|
||||
json!("-inf"),
|
||||
json!("1e999"),
|
||||
json!([]),
|
||||
] {
|
||||
assert!(serde_json::from_value::<Numbers>(json!({"float": input})).is_err());
|
||||
}
|
||||
for (input, expected) in [(json!(2), 2.0), (json!(2.5), 2.5), (json!(true), 1.0)] {
|
||||
let numbers: Numbers = serde_json::from_value(json!({"float": input})).unwrap();
|
||||
assert_eq!(numbers.float, Some(expected));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,20 +1,26 @@
|
|||
litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream<Item = Result<Bytes, Error>>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint
|
||||
litellm-core owns route orchestration. Messages and HTTP Responses return `litellm_host::call::CallOutput`, containing either a completed response or a stream head and chunks. OCR and currently non-streaming Chat Completions return their completed response directly
|
||||
|
||||
A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks<Error>` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver
|
||||
Hosts assemble route objects from shared `CoreResources`, HTTP settings, and secret sources. Each route owns its provider client and authentication dependencies. Gateway routes live for the gateway lifetime; Python assembles routes per call from its settings snapshot
|
||||
|
||||
Chat Completions, Messages, Responses, and OCR execute through their route objects. Calls pass `Interceptors` and an optional `ObservationSender` separately; use `&()` for no hooks and `None` for no observer. Construction does no work; preparation and lifecycle observation begin when the future is polled. Handlers accept `Interceptors`, never a concrete `ChannelInterceptors`. Native observers receive start and terminal events through the shared call runner; a stream retains its lifecycle until exhaustion, error, or drop. Hosted routes leave terminal observation to their driver
|
||||
|
||||
`route.rs` declares the concrete `Protocol` and implements a route method that accepts a typed request and constructs a `litellm_host::call::HostedMachine` with `hosted_call`. The shared call plumbing owns stream opening, delivery, backpressure, and detachment. Request decoding belongs to the boundary before the machine starts. Route closures only supply execution dependencies and route-specific host capabilities such as an OCR token provider. Use `run_hosted` for a native host so detachment is reported as cancellation. Python uses its own shared driver and preserves caller-task callback execution
|
||||
|
||||
Responses WebSocket sessions remain separate from the HTTP call driver because a connection can accept multiple requests while receiving events
|
||||
|
||||
## Crate layering
|
||||
|
||||
For Messages, Responses, Chat Completions, OCR, and other API formats, `core/src/<format>/` owns orchestration. Shared API data contracts belong in `litellm-types`, adapter contracts and shared transformation machinery in `llms/src/base_llm/<format>/`, and provider policy in `llms/src/<provider>/<format>/`. A repeated format directory name does not imply interchangeable responsibilities. Select concrete adapters here, then invoke their contracts instead of applying one provider's policy to every call. Route types describe call envelopes and execution state, not duplicate public payload schemas
|
||||
For Messages, Responses, Chat Completions, OCR, and other API formats, `core/src/<format>/` owns orchestration. Shared API data contracts belong in `litellm-llms-types`, adapter contracts and shared transformation machinery in `llms/src/base_llm/<format>/`, and provider policy in `llms/src/<provider>/<format>/`. A repeated format directory name does not imply interchangeable responsibilities. Select concrete adapters here, then invoke their contracts instead of applying one provider's policy to every call. Route types describe call envelopes and execution state, not duplicate public payload schemas
|
||||
|
||||
Each crate mirrors one top-level Python package, so a Rust path reads as its Python path with the crate name in place of the package directory. Dependencies only point down:
|
||||
Crates separate API data, transformations, transport, and orchestration. Python package names identify counterparts, not ownership. Dependencies only point down:
|
||||
|
||||
- `litellm-types` mirrors `litellm/types/`: pure serde data, no I/O
|
||||
- `litellm-llms-types` owns shared inference API contracts, grouped by format: pure serde data and shape validation, no I/O
|
||||
- `litellm-core-utils` mirrors `litellm/litellm_core_utils/`: pure helpers (provider resolution, prompt factory, call arguments, settings lookup and layer merge), no network I/O
|
||||
- `litellm-http` is Rust-only and route-neutral: settings resolution, the pooled `reqwest` clients, TLS, proxies, the SSRF-safe media fetcher, request and header helpers, and transport errors. Python's `litellm/llms/custom_httpx/` is split by responsibility instead of mirrored: its transport half lives here, its OCR handler in `litellm-llms`
|
||||
- `litellm-llms` mirrors `litellm/llms/`: `base_llm/<api>/transformation.rs`, `<provider>/<api>/transformation.rs`, and `base_llm/ocr/handler.rs` (the OCR request handler)
|
||||
- `litellm-core` mirrors the route packages (`litellm/ocr/`, `litellm/messages/`, ...): entrypoints, route request types, provider dispatch, the route machine, and hooks
|
||||
|
||||
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::hooks::RouteHooks`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
|
||||
A route module owns the call entrypoint, route request types (`*Request<'a>`), credential fallback, provider dispatch, and the handler glue that runs a provider config. Provider code never imports from core; when it needs the caller's hooks mid-call it goes through `litellm_llms::base_llm::ocr::handler::CallHooks`, the provider-level hooks OCR implements over its host until it folds into `litellm_host::interceptors::Interceptors`. Import every item from its canonical path. Never re-export another crate's items or give an item a second public path; the only re-export allowed is a private submodule surfacing its item at its module root (`mod error; pub use error::Error;`). Handlers belong in core or llms, never in a host crate
|
||||
|
||||
## Error placement
|
||||
|
||||
|
|
@ -27,3 +33,17 @@ Scope follows the concept, not the first caller. An error type under `litellm-ll
|
|||
`litellm_llms::Error` (`crates/llms/src/error.rs`) is the one transformation error for every provider and API. `base_llm/ocr/error.rs` is the recorded exception until OCR folds into it
|
||||
|
||||
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback execution of any kind. Core runs each route as a machine that yields host operations and call events; which integrations consume those events is the host's business.
|
||||
|
||||
## Response caching and accounting boundary
|
||||
|
||||
Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries a scope-free `CachePolicy` and observation; per-call policy never replaces the attached scope or service
|
||||
|
||||
Messages groups per-call dependencies in `CallContext` and explicitly sequences cache lookup, provider execution, result acceptance, and cache storage. Provider transport does not own cache orchestration. Stream capture remains in the shared cache implementation
|
||||
|
||||
Core owns request identity, typed response reconstruction and stream capture/replay. `cache-response` owns cache policy, namespacing, scope encoding, versioned envelopes and freshness. The SDK explicitly chooses shared scope. The gateway derives isolated scope from authenticated identity before attaching its service
|
||||
|
||||
Core delivers `ExecutionFacts` through the awaited `ResultReady` host operation for both provider and cached results, before public response processing or stream opening. Facts carry resolved model/provider and result source, including the hit key. Usage remains in the typed response or delivered stream, where completion and cancellation determine what was actually reported. Passive observation is not an accounting delivery mechanism
|
||||
|
||||
Core does not calculate prices, charge budgets, or update rate-limit counters. The legacy Python callback adapter translates execution facts into the existing Python logging contract; Python remains the accounting owner on that path. Native gateway accounting belongs to gateway dependencies, independently of `host-python`. Response-cache services expose no coordination counters or reservation APIs. A shared Redis deployment does not make response storage and accounting coordination the same dependency
|
||||
|
||||
Cache lookup follows provider preparation, credential resolution and the request interceptor. Keys describe the effective provider URL, authenticated headers and rewritten body. Signed requests bypass caching until the signing identity has a stable cache representation
|
||||
|
|
|
|||
|
|
@ -6,8 +6,12 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-cache.workspace = true
|
||||
litellm-cache-response.workspace = true
|
||||
litellm-framing.workspace = true
|
||||
tokio-util = { version = "0.7", features = ["codec"] }
|
||||
litellm-secrets.workspace = true
|
||||
litellm-types.workspace = true
|
||||
litellm-llms-types.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
litellm-host.workspace = true
|
||||
bytes.workspace = true
|
||||
|
|
@ -18,6 +22,7 @@ litellm-auth-aws.workspace = true
|
|||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
tracing.workspace = true
|
||||
moka.workspace = true
|
||||
mime_guess = "2.0.5"
|
||||
rand.workspace = true
|
||||
|
|
@ -35,8 +40,10 @@ url.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-cache-memory.workspace = true
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-host-native.workspace = true
|
||||
litellm-llms = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
rstest_reuse.workspace = true
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
) -> Result<Value, Error> {
|
||||
let env_lookup = |key: &str| request.secrets.get(key);
|
||||
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
|
||||
let response = crate::outbound::outbound_request(
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
authenticated,
|
||||
request.url.clone(),
|
||||
&request.body,
|
||||
|
|
@ -26,12 +26,12 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
.timeout
|
||||
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
|
||||
),
|
||||
)?
|
||||
.send(http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
let status = response.status();
|
||||
let text = response.text().await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
|
|
|
|||
|
|
@ -1,22 +1,54 @@
|
|||
use litellm_secrets::source::SecretSource;
|
||||
pub mod types;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub use handler::execute_audio_transcription_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
pub use prepare::prepare_audio_transcription_provider_call;
|
||||
use serde_json::Value;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::audio_transcription::types::AudioTranscriptionRequest;
|
||||
|
||||
pub async fn audio_transcription(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: &dyn SecretSource,
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
) -> Result<Value, Error> {
|
||||
let request = prepare_audio_transcription_provider_call(request, secrets).await?;
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
execute_audio_transcription_provider_call(&http, &resources.auth, request).await
|
||||
#[derive(Clone)]
|
||||
pub struct AudioTranscriptionRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl AudioTranscriptionRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "audio_transcription",
|
||||
model = request.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub async fn execute(&self, request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let request =
|
||||
prepare_audio_transcription_provider_call(request, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.model, &request.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<Value, Error>> = Box::pin(
|
||||
execute_audio_transcription_provider_call(&self.http, &self.auth, request),
|
||||
);
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_http::request::string_headers;
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::{
|
||||
|
|
@ -14,36 +13,36 @@ use super::Error;
|
|||
use crate::audio_transcription::types::{
|
||||
AudioTranscriptionRequest, ProviderAudioTranscriptionRequest,
|
||||
};
|
||||
use crate::provider::{LlmProviders, resolve_llm_provider};
|
||||
|
||||
fn provider_config(provider: &str) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
if provider == "bedrock" {
|
||||
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);
|
||||
fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
match provider {
|
||||
LlmProviders::Bedrock => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
let _ = provider;
|
||||
None
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub async fn prepare_audio_transcription_provider_call(
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderAudioTranscriptionRequest, Error> {
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
.or_else(|| {
|
||||
request
|
||||
.custom_llm_provider
|
||||
.map(|provider| CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for audio transcription request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let provider_info = resolve_llm_provider(
|
||||
request.model,
|
||||
request.custom_llm_provider,
|
||||
"audio transcription",
|
||||
)?;
|
||||
let model = provider_info.model.to_string();
|
||||
let config = provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
let config = provider_config(provider_info.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?;
|
||||
let snapshot = secrets.resolve(&config.secret_names()).await?;
|
||||
let env_lookup = |key: &str| snapshot.get(key);
|
||||
let forwarded = string_headers("audio transcription", request.extra_headers)?;
|
||||
|
|
@ -64,7 +63,7 @@ pub async fn prepare_audio_transcription_provider_call(
|
|||
config.transform_audio_transcription_request(&model, request.audio, filtered_params)?;
|
||||
Ok(ProviderAudioTranscriptionRequest {
|
||||
model,
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
|
||||
config,
|
||||
url,
|
||||
body: transformed.body,
|
||||
|
|
|
|||
376
litellm-rust/crates/core/src/caching.rs
Normal file
376
litellm-rust/crates/core/src/caching.rs
Normal file
|
|
@ -0,0 +1,376 @@
|
|||
use std::{
|
||||
future::Future,
|
||||
marker::PhantomData,
|
||||
sync::Arc,
|
||||
time::{Duration, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use litellm_cache_response::{
|
||||
CacheOptions, CachePolicy, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope,
|
||||
ScopedCache, cache_key,
|
||||
};
|
||||
use litellm_host::{
|
||||
call::{CallOutput, OutputOf},
|
||||
interceptors::{ExecutionFacts, Interceptors, ProviderIdentity, ResultSource, WireRequest},
|
||||
lifecycle::{CallEvent, ExecutionEvent},
|
||||
observation::ObservationSender,
|
||||
protocol::Protocol,
|
||||
};
|
||||
use serde::{Deserialize, Serialize, de::DeserializeOwned};
|
||||
use serde_json::Value;
|
||||
use tokio_util::codec::Decoder;
|
||||
|
||||
use crate::RouteError;
|
||||
|
||||
pub trait Cachable: Protocol<Error = RouteError> {
|
||||
const SURFACE: &'static str;
|
||||
|
||||
fn reusable(_response: &Self::Response) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
pub struct CacheRequest {
|
||||
pub identity: ProviderIdentity,
|
||||
pub input: Value,
|
||||
}
|
||||
|
||||
impl CacheRequest {
|
||||
pub fn from_wire(identity: ProviderIdentity, wire: Option<&WireRequest>) -> Self {
|
||||
Self {
|
||||
input: wire.map_or(Value::Null, |wire| {
|
||||
serde_json::json!({
|
||||
"provider": identity.provider,
|
||||
"model": identity.model,
|
||||
"url": wire.url,
|
||||
"headers": wire.headers,
|
||||
"body": wire.body,
|
||||
})
|
||||
}),
|
||||
identity,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait StreamCachable: Cachable {
|
||||
const TERMINAL_EVENT: &'static str;
|
||||
|
||||
fn replay(data: Bytes) -> Option<OutputOf<Self>>;
|
||||
fn bytes(chunk: &Self::Chunk) -> &[u8];
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", content = "value")]
|
||||
pub enum CachedOutput<R> {
|
||||
Response(R),
|
||||
Stream(String),
|
||||
}
|
||||
|
||||
struct CacheSession {
|
||||
service: Arc<dyn ResponseCacheService>,
|
||||
request: ResponseCacheRequest,
|
||||
}
|
||||
|
||||
impl CacheSession {
|
||||
fn prepare<P: Cachable>(
|
||||
service: Option<Arc<dyn ResponseCacheService>>,
|
||||
options: Option<CacheOptions>,
|
||||
request: &CacheRequest,
|
||||
) -> Option<Self> {
|
||||
let options = options.filter(|options| options.policy.enabled())?;
|
||||
let service = service?;
|
||||
let input = request.input.clone();
|
||||
let request = options.request(&service.config().namespace, P::SURFACE, input);
|
||||
Some(Self { service, request })
|
||||
}
|
||||
|
||||
async fn lookup<P: Cachable>(&self) -> Option<CachedOutput<P::Response>>
|
||||
where
|
||||
P::Response: DeserializeOwned,
|
||||
{
|
||||
if !self.request.controls.reads() {
|
||||
return None;
|
||||
}
|
||||
match self.service.lookup(&self.request, now()).await {
|
||||
Ok(Some(value)) => {
|
||||
serde_json::from_value::<ResponseEnvelope<CachedOutput<P::Response>>>(value)
|
||||
.ok()
|
||||
.and_then(|entry| entry.decode(P::SURFACE))
|
||||
}
|
||||
Ok(None) => None,
|
||||
Err(_) => {
|
||||
tracing::warn!("response cache lookup failed");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn store(&self, entry: Value) {
|
||||
if !self.request.controls.writes() {
|
||||
return;
|
||||
}
|
||||
if self
|
||||
.service
|
||||
.store(&self.request, entry, now())
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
tracing::warn!("response cache write failed");
|
||||
}
|
||||
}
|
||||
|
||||
async fn store_response<P: Cachable>(&self, response: &P::Response)
|
||||
where
|
||||
P::Response: Serialize,
|
||||
{
|
||||
if !self.request.controls.writes() || !P::reusable(response) {
|
||||
return;
|
||||
}
|
||||
if let Ok(value) = serde_json::to_value(response)
|
||||
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
|
||||
P::SURFACE,
|
||||
CachedOutput::Response(value),
|
||||
))
|
||||
{
|
||||
self.store(entry).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute_unary<P, F, Fut>(
|
||||
request: CacheRequest,
|
||||
cache: Option<Arc<dyn ResponseCacheService>>,
|
||||
options: Option<CacheOptions>,
|
||||
interceptors: &impl Interceptors<RouteError>,
|
||||
observers: Option<&ObservationSender>,
|
||||
provider: F,
|
||||
) -> Result<P::Response, RouteError>
|
||||
where
|
||||
P: Cachable,
|
||||
P::Response: Serialize + DeserializeOwned,
|
||||
F: FnOnce() -> Fut,
|
||||
Fut: Future<Output = Result<P::Response, RouteError>>,
|
||||
{
|
||||
let identity = request.identity.clone();
|
||||
crate::diagnostic::provider(&identity.model, &identity.provider);
|
||||
let session = CacheSession::prepare::<P>(cache, options, &request);
|
||||
let hit = match &session {
|
||||
Some(session) => session.lookup::<P>().await.and_then(|entry| match entry {
|
||||
CachedOutput::Response(response) => Some((response, cache_key(&session.request.key))),
|
||||
CachedOutput::Stream(_) => None,
|
||||
}),
|
||||
None => None,
|
||||
};
|
||||
let (response, source) = match hit {
|
||||
Some((response, key)) => (response, ResultSource::Cache { key }),
|
||||
None => (provider().await?, ResultSource::Provider),
|
||||
};
|
||||
let from_provider = source == ResultSource::Provider;
|
||||
publish(
|
||||
ExecutionFacts {
|
||||
provider: identity,
|
||||
source,
|
||||
},
|
||||
interceptors,
|
||||
observers,
|
||||
)
|
||||
.await?;
|
||||
if from_provider && let Some(session) = session {
|
||||
session.store_response::<P>(&response).await;
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
pub async fn execute_streaming<P, F, Fut>(
|
||||
request: CacheRequest,
|
||||
cache: Option<Arc<dyn ResponseCacheService>>,
|
||||
options: Option<CacheOptions>,
|
||||
interceptors: &impl Interceptors<RouteError>,
|
||||
observers: Option<&ObservationSender>,
|
||||
provider: F,
|
||||
) -> Result<OutputOf<P>, RouteError>
|
||||
where
|
||||
P: StreamCachable,
|
||||
P::Response: Serialize + DeserializeOwned,
|
||||
F: FnOnce() -> Fut,
|
||||
Fut: Future<Output = Result<OutputOf<P>, RouteError>>,
|
||||
{
|
||||
let identity = request.identity.clone();
|
||||
crate::diagnostic::provider(&identity.model, &identity.provider);
|
||||
let session = CacheSession::prepare::<P>(cache, options, &request);
|
||||
let cache = CallCache::<P> {
|
||||
session,
|
||||
protocol: PhantomData,
|
||||
};
|
||||
let hit = cache.lookup().await;
|
||||
let (output, source) = match hit {
|
||||
Some(hit) => hit,
|
||||
None => (provider().await?, ResultSource::Provider),
|
||||
};
|
||||
publish(
|
||||
ExecutionFacts {
|
||||
provider: identity,
|
||||
source: source.clone(),
|
||||
},
|
||||
interceptors,
|
||||
observers,
|
||||
)
|
||||
.await?;
|
||||
Ok(cache.finish(output, &source).await)
|
||||
}
|
||||
|
||||
pub(crate) struct CallCache<P> {
|
||||
session: Option<CacheSession>,
|
||||
protocol: PhantomData<P>,
|
||||
}
|
||||
|
||||
impl<P: StreamCachable> CallCache<P> {
|
||||
pub(crate) fn from_wire(
|
||||
cache: Option<&ScopedCache>,
|
||||
policy: CachePolicy,
|
||||
identity: &ProviderIdentity,
|
||||
wire: &WireRequest,
|
||||
) -> Self {
|
||||
let session = cache.and_then(|cache| {
|
||||
if !policy.enabled() {
|
||||
return None;
|
||||
}
|
||||
let options = cache.options(Some(policy));
|
||||
let request = CacheRequest::from_wire(identity.clone(), Some(wire));
|
||||
Some(CacheSession {
|
||||
request: options.request(
|
||||
&cache.service.config().namespace,
|
||||
P::SURFACE,
|
||||
request.input,
|
||||
),
|
||||
service: cache.service.clone(),
|
||||
})
|
||||
});
|
||||
Self {
|
||||
session,
|
||||
protocol: PhantomData,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn lookup(&self) -> Option<(OutputOf<P>, ResultSource)>
|
||||
where
|
||||
P::Response: DeserializeOwned,
|
||||
{
|
||||
let session = self.session.as_ref()?;
|
||||
let output = match session.lookup::<P>().await? {
|
||||
CachedOutput::Response(response) => CallOutput::Complete(response),
|
||||
CachedOutput::Stream(data) => P::replay(Bytes::from(data))?,
|
||||
};
|
||||
Some((
|
||||
output,
|
||||
ResultSource::Cache {
|
||||
key: cache_key(&session.request.key),
|
||||
},
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) async fn finish(self, output: OutputOf<P>, source: &ResultSource) -> OutputOf<P>
|
||||
where
|
||||
P::Response: Serialize,
|
||||
{
|
||||
let Some(session) = self.session.filter(|session| {
|
||||
*source == ResultSource::Provider && session.request.controls.writes()
|
||||
}) else {
|
||||
return output;
|
||||
};
|
||||
match output {
|
||||
CallOutput::Complete(response) => {
|
||||
session.store_response::<P>(&response).await;
|
||||
CallOutput::Complete(response)
|
||||
}
|
||||
CallOutput::Stream { head, chunks } => CallOutput::Stream {
|
||||
head,
|
||||
chunks: capture_stream::<P>(chunks, session),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn capture_stream<P: StreamCachable>(
|
||||
chunks: futures_util::stream::BoxStream<'static, Result<P::Chunk, RouteError>>,
|
||||
session: CacheSession,
|
||||
) -> futures_util::stream::BoxStream<'static, Result<P::Chunk, RouteError>> {
|
||||
stream::try_unfold(
|
||||
(chunks, Some(Vec::<u8>::new()), session),
|
||||
|(mut chunks, captured, session)| async move {
|
||||
match chunks.try_next().await? {
|
||||
Some(chunk) => {
|
||||
let captured = captured.and_then(|mut data| {
|
||||
let bytes = P::bytes(&chunk);
|
||||
if data.len().saturating_add(bytes.len())
|
||||
> session.service.config().max_entry_bytes
|
||||
{
|
||||
return None;
|
||||
}
|
||||
data.extend_from_slice(bytes);
|
||||
Some(data)
|
||||
});
|
||||
Ok(Some((chunk, (chunks, captured, session))))
|
||||
}
|
||||
None => {
|
||||
if let Some(data) = captured
|
||||
&& let Ok(text) = String::from_utf8(data)
|
||||
&& successful_stream(&text, P::TERMINAL_EVENT)
|
||||
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
|
||||
P::SURFACE,
|
||||
CachedOutput::<Value>::Stream(text),
|
||||
))
|
||||
{
|
||||
session.store(entry).await;
|
||||
}
|
||||
Ok::<_, RouteError>(None)
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
.boxed()
|
||||
}
|
||||
|
||||
fn now() -> Duration {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn successful_stream(text: &str, terminal: &str) -> bool {
|
||||
let mut pending = BytesMut::from(text.as_bytes());
|
||||
let mut codec = litellm_framing::sse::SseCodec::default();
|
||||
let mut complete = false;
|
||||
loop {
|
||||
let event = match codec.decode(&mut pending) {
|
||||
Ok(Some(event)) => event,
|
||||
Ok(None) => return complete && pending.is_empty(),
|
||||
Err(_) => return false,
|
||||
};
|
||||
let Ok(value) = serde_json::from_str::<Value>(&event.data) else {
|
||||
return false;
|
||||
};
|
||||
let Some(kind) = value.get("type").and_then(Value::as_str) else {
|
||||
return false;
|
||||
};
|
||||
if matches!(kind, "error" | "response.failed" | "response.incomplete") {
|
||||
return false;
|
||||
}
|
||||
complete |= kind == terminal;
|
||||
}
|
||||
}
|
||||
|
||||
async fn publish(
|
||||
facts: ExecutionFacts,
|
||||
interceptors: &impl Interceptors<RouteError>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<(), RouteError> {
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: facts.clone(),
|
||||
}));
|
||||
}
|
||||
interceptors.result_ready(facts).await
|
||||
}
|
||||
|
|
@ -8,15 +8,38 @@ use litellm_llms::{
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
const HEADER_CONTEXT: &str = "chat completions";
|
||||
|
||||
pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'static dyn BaseConfig> {
|
||||
pub(super) enum ChatProvider {
|
||||
Anthropic,
|
||||
Bedrock,
|
||||
OpenaiLike,
|
||||
}
|
||||
|
||||
impl ChatProvider {
|
||||
pub(super) fn config(self) -> &'static dyn BaseConfig {
|
||||
match self {
|
||||
Self::Anthropic => &ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
|
||||
Self::Bedrock => &BEDROCK_CHAT_COMPLETIONS_CONFIG,
|
||||
Self::OpenaiLike => &OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option<ChatProvider> {
|
||||
match provider {
|
||||
"anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG),
|
||||
"bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG),
|
||||
"openai_like" => Some(&OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG),
|
||||
_ => None,
|
||||
LlmProviders::Anthropic => Some(ChatProvider::Anthropic),
|
||||
LlmProviders::Bedrock => Some(ChatProvider::Bedrock),
|
||||
LlmProviders::OpenaiLike => Some(ChatProvider::OpenaiLike),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,16 +1,14 @@
|
|||
use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{Authenticated, resolve_auth},
|
||||
chat::transformation::ProviderChatResponseData,
|
||||
};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::Error;
|
||||
|
|
@ -23,7 +21,10 @@ pub(super) async fn execute(
|
|||
http: &Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderChatCompletionsRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
cache_options: Option<litellm_cache_response::CachePolicy>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let ProviderChatCompletionsRequest {
|
||||
model,
|
||||
|
|
@ -45,8 +46,12 @@ pub(super) async fn execute(
|
|||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
let identity = litellm_host::interceptors::ProviderIdentity {
|
||||
model: context.model.clone(),
|
||||
provider: context.custom_llm_provider.clone(),
|
||||
};
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
|
|
@ -55,54 +60,72 @@ pub(super) async fn execute(
|
|||
context,
|
||||
)
|
||||
.await?;
|
||||
let outbound = outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
let cache = cache.filter(|_| authenticated.signer.is_none());
|
||||
let cache_request =
|
||||
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
|
||||
crate::caching::execute_unary::<super::route::ChatCompletions, _, _>(
|
||||
cache_request,
|
||||
cache.as_ref().map(|cache| cache.service.clone()),
|
||||
cache.as_ref().map(|cache| cache.options(cache_options)),
|
||||
interceptors,
|
||||
observers,
|
||||
|| async move {
|
||||
let outbound = outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)?;
|
||||
|
||||
let response = crate::outbound::send(outbound, http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
if err.is_connect() || err.is_builder() {
|
||||
Error::Transport(litellm_http::transport::Error::Connect(err.to_string()))
|
||||
} else {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
}
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response.text().await.map_err(|err| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
})?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
let raw = RawResponse { body: text.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
.map_err(Error::from)
|
||||
.map_err(as_response_error)
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)?;
|
||||
|
||||
let response = outbound.send(http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
if err.is_connect() || err.is_builder() {
|
||||
Error::Transport(litellm_http::transport::Error::Connect(err.to_string()))
|
||||
} else {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
}
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response.text().await.map_err(|err| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
|
||||
})?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
.map_err(Error::from)
|
||||
.map_err(as_response_error)
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Re-tag an error raised while normalizing a response the provider already
|
||||
|
|
@ -167,8 +190,8 @@ mod tests {
|
|||
raw: Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RouteHooks<Error> for RecordingHooks {
|
||||
async fn before_send(
|
||||
impl Interceptors<Error> for RecordingHooks {
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
|
|
@ -187,8 +210,7 @@ mod tests {
|
|||
})
|
||||
}
|
||||
|
||||
async fn emit(&self, event: MachineEvent) -> Result<(), Error> {
|
||||
let MachineEvent::ResponseReceived { raw } = event;
|
||||
async fn after_provider_response(&self, raw: RawResponse) -> Result<(), Error> {
|
||||
self.raw.lock().unwrap().push(raw.body);
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -222,13 +244,16 @@ mod tests {
|
|||
)
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
let interceptors = RecordingHooks::default();
|
||||
|
||||
execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
None,
|
||||
None,
|
||||
&interceptors,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("chat completions call succeeds");
|
||||
|
|
@ -239,14 +264,17 @@ mod tests {
|
|||
assert_eq!(sent["system"], "added by the host");
|
||||
assert_eq!(request.headers["x-host"], "seen");
|
||||
assert_eq!(request.headers["x-api-key"], "sk-test");
|
||||
let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
let [context] =
|
||||
<[RequestContext; 1]>::try_from(interceptors.contexts.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| {
|
||||
panic!("before_provider_request runs once, saw {}", seen.len())
|
||||
});
|
||||
assert_eq!(
|
||||
(context.model.as_str(), context.custom_llm_provider.as_str()),
|
||||
("claude-sonnet-4-5", "anthropic")
|
||||
);
|
||||
assert_eq!(context.optional_params, json!({"max_tokens": 16}));
|
||||
assert_eq!(hooks.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
|
||||
assert_eq!(interceptors.raw.into_inner().unwrap(), [ANTHROPIC_MESSAGE]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -257,13 +285,16 @@ mod tests {
|
|||
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let hooks = RecordingHooks::default();
|
||||
let interceptors = RecordingHooks::default();
|
||||
|
||||
let error = execute(
|
||||
&Client::plain_for_test(),
|
||||
&AuthServices::default(),
|
||||
prepared(&upstream.uri()),
|
||||
&hooks,
|
||||
None,
|
||||
None,
|
||||
&interceptors,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect_err("the upstream failure fails the call");
|
||||
|
|
@ -272,10 +303,10 @@ mod tests {
|
|||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
|
||||
));
|
||||
assert!(hooks.raw.into_inner().unwrap().is_empty());
|
||||
assert!(interceptors.raw.into_inner().unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
|
||||
for original in [
|
||||
Error::MissingField("usage"),
|
||||
|
|
|
|||
|
|
@ -1,60 +1,93 @@
|
|||
//! The `/chat/completions` call, the Rust equivalent of Python's
|
||||
//! `litellm.completion()`.
|
||||
//!
|
||||
//! [`chat_completions`] is the top-level entrypoint: give it a model, the
|
||||
//! OpenAI-shaped message list, the provider-mapped optional params, and
|
||||
//! credentials, and it resolves the provider, translates the conversation,
|
||||
//! calls the provider, and returns a typed OpenAI-shaped response.
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
use litellm_host::observation::ObservationSender;
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod common_utils;
|
||||
pub(crate) mod handler;
|
||||
mod prepare;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request};
|
||||
use serde_json::{Map, Value};
|
||||
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
|
||||
use prepare::{prepare_provider_request, resolve_request};
|
||||
|
||||
use crate::chat_completions::types::ChatCompletionsRequest;
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub async fn chat_completions(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: &dyn SecretSource,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let resolved = resolve_request(request)?;
|
||||
let snapshot = secrets.resolve(&resolved.config.secret_names()).await?;
|
||||
let request = prepare_provider_request(resolved, snapshot)?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
#[derive(Clone)]
|
||||
pub struct ChatCompletionsRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
/// touching the network.
|
||||
///
|
||||
/// A host that keeps the Python implementation asks this first so it can emit
|
||||
/// its pre-call logging exactly once, on whichever path is about to run.
|
||||
/// Returns the decline reason, or `None` when the request is accepted.
|
||||
pub fn chat_completions_decline_reason(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
messages: Value,
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<&'static str> {
|
||||
let Ok(resolved) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
return Some("provider is not on the rust chat completions path");
|
||||
};
|
||||
let config = resolved.config;
|
||||
let Ok(messages) = parse_messages(messages) else {
|
||||
return Some("unreadable message list");
|
||||
};
|
||||
if messages.is_empty() {
|
||||
return Some("empty message list");
|
||||
impl ChatCompletionsRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
cache: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
litellm_host::lifecycle::observe_unary(
|
||||
observers.clone(),
|
||||
self.run_call(
|
||||
request.into(),
|
||||
cache_options,
|
||||
interceptors,
|
||||
observers.as_ref(),
|
||||
),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn run(
|
||||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
cache_options: Option<litellm_cache_response::CachePolicy>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let resolved = resolve_request(request)?;
|
||||
let snapshot = self
|
||||
.secrets
|
||||
.resolve(&resolved.config.secret_names())
|
||||
.await?;
|
||||
let prepared = prepare_provider_request(resolved, snapshot)?;
|
||||
crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ChatCompletionsResponse, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
prepared,
|
||||
self.cache.clone(),
|
||||
cache_options,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
}
|
||||
config
|
||||
.unsupported_reason(&messages, optional_params)
|
||||
.map(|reason| reason.0)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
|
||||
use litellm_llms_types::formats::chat_completions::ChatMessage;
|
||||
use litellm_secrets::source::Secrets;
|
||||
use litellm_types::llms::openai::ChatMessage;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{chat_completions_provider_config, string_headers},
|
||||
common_utils::{chat_completions_provider, string_headers},
|
||||
};
|
||||
use crate::chat_completions::types::{
|
||||
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
|
||||
};
|
||||
use crate::provider::resolve_llm_provider;
|
||||
|
||||
pub(super) struct ResolvedProvider {
|
||||
pub(super) model: String,
|
||||
|
|
@ -25,23 +25,13 @@ pub(super) fn resolve_provider_config<'a>(
|
|||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for chat completions request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
let provider_info = resolve_llm_provider(model, custom_llm_provider, "chat completions")?;
|
||||
let config = chat_completions_provider(provider_info.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?
|
||||
.config();
|
||||
Ok(ResolvedProvider {
|
||||
model: provider_info.model.to_string(),
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
|
@ -106,6 +96,7 @@ fn validate_environment(
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
secrets: Secrets,
|
||||
|
|
@ -194,13 +185,10 @@ mod tests {
|
|||
}
|
||||
}
|
||||
|
||||
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
|
||||
/// carry resolved credentials), so unwrap the failure case by hand.
|
||||
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
|
||||
match prepare_chat_completions_call(request) {
|
||||
Err(error) => error,
|
||||
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
|
||||
}
|
||||
fn preparation_error(request: ChatCompletionsRequest<'_>) -> Error {
|
||||
prepare_chat_completions_call(request)
|
||||
.err()
|
||||
.expect("request preparation should fail")
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -346,24 +334,19 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_an_unsupported_request_before_resolving_credentials() {
|
||||
let mut call = request(
|
||||
"claude-sonnet-4-5",
|
||||
Some("anthropic"),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
);
|
||||
call.api_key = None;
|
||||
// No api_key is set and no env is consulted: the gate must run first, so the
|
||||
// error is the decline rather than a missing-credential error.
|
||||
assert_eq!(decline(call), Error::Unsupported("streaming"));
|
||||
#[rstest::rstest]
|
||||
fn rejects_empty_messages_before_resolving_credentials() {
|
||||
let call = ChatCompletionsRequest {
|
||||
api_key: None,
|
||||
..request("claude-sonnet-4-5", Some("anthropic"), json!([]), json!({}))
|
||||
};
|
||||
assert!(matches!(preparation_error(call), Error::InvalidRequest(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_an_unknown_provider() {
|
||||
assert_eq!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"openai/gpt-4o",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -376,7 +359,7 @@ mod tests {
|
|||
#[test]
|
||||
fn rejects_a_model_with_no_resolvable_provider() {
|
||||
assert!(matches!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -389,7 +372,7 @@ mod tests {
|
|||
#[test]
|
||||
fn rejects_an_empty_or_malformed_message_list() {
|
||||
assert_eq!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([]),
|
||||
|
|
@ -402,7 +385,7 @@ mod tests {
|
|||
)
|
||||
);
|
||||
assert!(matches!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("not a list"),
|
||||
|
|
@ -422,7 +405,7 @@ mod tests {
|
|||
);
|
||||
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
|
||||
assert_eq!(
|
||||
decline(call),
|
||||
preparation_error(call),
|
||||
Error::Headers(litellm_http::request::HeaderError {
|
||||
context: "chat completions",
|
||||
name: "x-trace".to_string(),
|
||||
|
|
@ -520,18 +503,17 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::authorization("Authorization")]
|
||||
#[case::amz_date("x-amz-date")]
|
||||
#[case::security_token("x-amz-security-token")]
|
||||
#[case::date("Date")]
|
||||
#[tokio::test]
|
||||
async fn a_forwarded_header_the_signer_computes_declines_to_python() {
|
||||
// Reattaching the caller's copy next to the computed one puts the name on
|
||||
// the wire twice and Bedrock rejects the pair, so a request carrying one
|
||||
// has to go to Python instead of being signed here.
|
||||
for forwarded in [
|
||||
"Authorization",
|
||||
"x-amz-date",
|
||||
"x-amz-security-token",
|
||||
"Date",
|
||||
] {
|
||||
let mut call = request(
|
||||
async fn rejects_a_forwarded_header_the_signer_computes(#[case] forwarded: &str) {
|
||||
let call = ChatCompletionsRequest {
|
||||
api_key: None,
|
||||
extra_headers: Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])),
|
||||
..request(
|
||||
"bedrock/us-east-1/anthropic.claude-v2",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -540,29 +522,27 @@ mod tests {
|
|||
"aws_access_key_id": "AKIDEXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
|
||||
}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment,
|
||||
&|_| None,
|
||||
)
|
||||
.await
|
||||
.expect("resolves");
|
||||
let error = crate::chat_completions::handler::outbound_request(
|
||||
authenticated,
|
||||
prepared.url,
|
||||
&prepared.body,
|
||||
prepared.timeout,
|
||||
)
|
||||
.expect_err("{forwarded} should decline instead of being signed");
|
||||
assert!(
|
||||
matches!(error, Error::Unsupported(_)),
|
||||
"{forwarded} declined as {error:?}, which the host would not fall back on"
|
||||
);
|
||||
}
|
||||
};
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment,
|
||||
&|_| None,
|
||||
)
|
||||
.await
|
||||
.expect("resolves");
|
||||
let error = crate::chat_completions::handler::outbound_request(
|
||||
authenticated,
|
||||
prepared.url,
|
||||
&prepared.body,
|
||||
prepared.timeout,
|
||||
)
|
||||
.expect_err("conflicting signing headers must fail");
|
||||
assert!(
|
||||
matches!(error, Error::Unsupported(_)),
|
||||
"{forwarded} returned {error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -655,112 +635,4 @@ mod tests {
|
|||
"prepare did not carry the bearer token"
|
||||
);
|
||||
}
|
||||
|
||||
fn decline_reason(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
messages: Value,
|
||||
params: Value,
|
||||
) -> Option<&'static str> {
|
||||
let params = match params {
|
||||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
};
|
||||
crate::chat_completions::chat_completions_decline_reason(model, provider, messages, ¶ms)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_accepts_what_prepare_accepts() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
),
|
||||
Some("streaming")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"openai/gpt-4o",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("nope"),
|
||||
json!({})
|
||||
),
|
||||
Some("unreadable message list")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
|
||||
Some("empty message list")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
|
||||
// A gate that accepts what prepare then declines would make the host emit
|
||||
// its pre-call logging on a path that falls back, so pin the agreement.
|
||||
for (messages, params) in [
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 8}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
|
||||
json!({"temperature": 0.1}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
|
||||
json!({}),
|
||||
),
|
||||
] {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params.clone()
|
||||
),
|
||||
None,
|
||||
"gate declined {messages}"
|
||||
);
|
||||
prepare_chat_completions_call(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params,
|
||||
))
|
||||
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
82
litellm-rust/crates/core/src/chat_completions/route.rs
Normal file
82
litellm-rust/crates/core/src/chat_completions/route.rs
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
use litellm_host::observation::ObservationSender;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use litellm_host::{
|
||||
call::{CallOutput, HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
|
||||
|
||||
use super::{
|
||||
ChatCompletionsRoute, Error,
|
||||
types::{ChatCompletionsCall, ChatCompletionsRequest},
|
||||
};
|
||||
|
||||
pub struct ChatCompletions;
|
||||
|
||||
impl Protocol for ChatCompletions {
|
||||
type Response = ChatCompletionsResponse;
|
||||
type Error = Error;
|
||||
type Request = ChatCompletionsCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Infallible;
|
||||
type StreamHead = Infallible;
|
||||
}
|
||||
|
||||
impl ChatCompletionsRoute {
|
||||
pub fn machine(
|
||||
self,
|
||||
call: ChatCompletionsCall,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> HostedMachine<ChatCompletions> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
hosted_call(
|
||||
call,
|
||||
observers,
|
||||
move |call, _, interceptors, observers| async move {
|
||||
self.run_call(call, cache_options, &interceptors, observers.as_ref())
|
||||
.await
|
||||
.map(CallOutput::Complete)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "chat_completions",
|
||||
model = %call.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub(super) async fn run_call(
|
||||
&self,
|
||||
call: ChatCompletionsCall,
|
||||
cache_options: Option<litellm_cache_response::CachePolicy>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let request = ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers,
|
||||
timeout: call.timeout,
|
||||
};
|
||||
self.run(request, cache_options, interceptors, observers)
|
||||
.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::Cachable for ChatCompletions {
|
||||
const SURFACE: &'static str = "chat_completions";
|
||||
}
|
||||
|
|
@ -3,7 +3,7 @@ use std::time::Duration;
|
|||
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
|
||||
use litellm_types::llms::openai::ChatMessage;
|
||||
use litellm_llms_types::formats::chat_completions::ChatMessage;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
/// A `/chat/completions` call as it crosses into the core.
|
||||
|
|
@ -23,6 +23,32 @@ pub struct ChatCompletionsRequest<'a> {
|
|||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub struct ChatCompletionsCall {
|
||||
pub model: String,
|
||||
pub messages: Value,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
impl From<ChatCompletionsRequest<'_>> for ChatCompletionsCall {
|
||||
fn from(request: ChatCompletionsRequest<'_>) -> Self {
|
||||
Self {
|
||||
model: request.model.into(),
|
||||
messages: request.messages,
|
||||
optional_params: request.optional_params,
|
||||
api_key: request.api_key.map(str::to_owned),
|
||||
api_base: request.api_base.map(str::to_owned),
|
||||
custom_llm_provider: request.custom_llm_provider.map(str::to_owned),
|
||||
extra_headers: request.extra_headers,
|
||||
timeout: request.timeout,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ResolvedChatCompletionsRequest<'a> {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
|
|
|
|||
48
litellm-rust/crates/core/src/context.rs
Normal file
48
litellm-rust/crates/core/src/context.rs
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
use litellm_cache_response::CachePolicy;
|
||||
use litellm_host::{
|
||||
interceptors::{ExecutionFacts, Interceptors, RawResponse},
|
||||
lifecycle::{CallEvent, ExecutionEvent},
|
||||
observation::ObservationSender,
|
||||
};
|
||||
|
||||
use crate::{CallOptions, RouteError};
|
||||
|
||||
pub(crate) struct CallContext<'a, I> {
|
||||
pub interceptors: &'a I,
|
||||
pub observers: Option<ObservationSender>,
|
||||
pub cache: CachePolicy,
|
||||
}
|
||||
|
||||
impl<'a, I: Interceptors<RouteError>> CallContext<'a, I> {
|
||||
pub fn new(interceptors: &'a I, options: CallOptions) -> Self {
|
||||
Self {
|
||||
interceptors,
|
||||
observers: options.observers,
|
||||
cache: options.cache.unwrap_or_default(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> {
|
||||
if let Some(observers) = &self.observers {
|
||||
observers.emit(CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: facts.clone(),
|
||||
}));
|
||||
}
|
||||
self.interceptors.result_ready(facts).await
|
||||
}
|
||||
|
||||
pub async fn response_received(&self, body: &str) -> Result<(), RouteError> {
|
||||
let raw = RawResponse {
|
||||
body: body.to_owned(),
|
||||
};
|
||||
if let Some(observers) = &self.observers {
|
||||
observers.emit(CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
self.interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(RouteError::post_call)
|
||||
}
|
||||
}
|
||||
324
litellm-rust/crates/core/src/diagnostic.rs
Normal file
324
litellm-rust/crates/core/src/diagnostic.rs
Normal file
|
|
@ -0,0 +1,324 @@
|
|||
use std::{
|
||||
future::Future,
|
||||
pin::Pin,
|
||||
task::{Context, Poll},
|
||||
};
|
||||
|
||||
use futures_util::{Stream, stream::BoxStream};
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_tracing::Logger;
|
||||
use tracing::Span;
|
||||
|
||||
struct Completion {
|
||||
span: Span,
|
||||
outcome: &'static str,
|
||||
}
|
||||
|
||||
impl Completion {
|
||||
fn new(name: &str) -> Self {
|
||||
let current = Span::current();
|
||||
Self {
|
||||
span: if current
|
||||
.metadata()
|
||||
.is_some_and(|metadata| metadata.name() == name)
|
||||
{
|
||||
current
|
||||
} else {
|
||||
Span::none()
|
||||
},
|
||||
outcome: "cancelled",
|
||||
}
|
||||
}
|
||||
|
||||
fn finish(mut self, outcome: &'static str) {
|
||||
self.outcome = outcome;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Completion {
|
||||
fn drop(&mut self) {
|
||||
self.span.record("outcome", self.outcome);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn provider(model: &str, provider: &str) {
|
||||
let span = Span::current();
|
||||
span.record("resolved_model", model);
|
||||
span.record("provider", provider);
|
||||
}
|
||||
|
||||
pub(crate) async fn unary<R, E>(execute: impl Future<Output = Result<R, E>>) -> Result<R, E> {
|
||||
operation("litellm.route", execute).await
|
||||
}
|
||||
|
||||
pub(crate) async fn operation<R, E>(
|
||||
name: &str,
|
||||
execute: impl Future<Output = Result<R, E>>,
|
||||
) -> Result<R, E> {
|
||||
let completion = Completion::new(name);
|
||||
let result = execute.await;
|
||||
completion.finish(if result.is_ok() { "success" } else { "failure" });
|
||||
result
|
||||
}
|
||||
|
||||
pub(crate) async fn call<R, H, C, E>(
|
||||
execute: impl Future<Output = Result<CallOutput<R, H, C, E>, E>>,
|
||||
) -> Result<CallOutput<R, H, C, E>, E>
|
||||
where
|
||||
C: Send + 'static,
|
||||
E: Send + 'static,
|
||||
{
|
||||
let completion = Completion::new("litellm.route");
|
||||
match execute.await {
|
||||
Err(error) => {
|
||||
completion.finish("failure");
|
||||
Err(error)
|
||||
}
|
||||
Ok(CallOutput::Complete(response)) => {
|
||||
completion.span.record("stream", false);
|
||||
completion.finish("success");
|
||||
Ok(CallOutput::Complete(response))
|
||||
}
|
||||
Ok(CallOutput::Stream { head, chunks }) => {
|
||||
completion.span.record("stream", true);
|
||||
Ok(CallOutput::Stream {
|
||||
head,
|
||||
chunks: Box::pin(TracedStream {
|
||||
state: Some(StreamState {
|
||||
chunks,
|
||||
completion,
|
||||
logger: Logger::current(),
|
||||
}),
|
||||
}),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct StreamState<C, E> {
|
||||
chunks: BoxStream<'static, Result<C, E>>,
|
||||
completion: Completion,
|
||||
logger: Logger,
|
||||
}
|
||||
|
||||
impl<C, E> StreamState<C, E> {
|
||||
fn close(self, outcome: &'static str) {
|
||||
self.logger
|
||||
.scope(|| self.completion.span.in_scope(|| drop(self.chunks)));
|
||||
self.completion.finish(outcome);
|
||||
}
|
||||
}
|
||||
|
||||
struct TracedStream<C, E> {
|
||||
state: Option<StreamState<C, E>>,
|
||||
}
|
||||
|
||||
impl<C, E> Stream for TracedStream<C, E> {
|
||||
type Item = Result<C, E>;
|
||||
|
||||
fn poll_next(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
let Some(state) = self.state.as_mut() else {
|
||||
return Poll::Ready(None);
|
||||
};
|
||||
let next = state.logger.scope(|| {
|
||||
state
|
||||
.completion
|
||||
.span
|
||||
.in_scope(|| state.chunks.as_mut().poll_next(context))
|
||||
});
|
||||
let outcome = match &next {
|
||||
Poll::Ready(None) => "success",
|
||||
Poll::Ready(Some(Err(_))) => "failure",
|
||||
_ => return next,
|
||||
};
|
||||
if let Some(state) = self.state.take() {
|
||||
state.close(outcome);
|
||||
}
|
||||
next
|
||||
}
|
||||
}
|
||||
|
||||
impl<C, E> Drop for TracedStream<C, E> {
|
||||
fn drop(&mut self) {
|
||||
if let Some(state) = self.state.take() {
|
||||
state.close("cancelled");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{sync::mpsc, task::Context};
|
||||
|
||||
use futures_util::{StreamExt, task::noop_waker_ref};
|
||||
use litellm_tracing::{Metadata, Record, Sink};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
struct Capture(mpsc::Sender<Value>);
|
||||
|
||||
impl Sink for Capture {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
|
||||
*metadata.level() <= tracing::Level::INFO
|
||||
}
|
||||
fn emit(&self, record: &Record) {
|
||||
self.0.send(Value::Object(record.fields.clone())).unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn logger() -> (Logger, mpsc::Receiver<Value>) {
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
(Logger::new(Capture(sender)), receiver)
|
||||
}
|
||||
|
||||
struct Chunks(std::vec::IntoIter<Result<u8, &'static str>>);
|
||||
|
||||
impl Stream for Chunks {
|
||||
type Item = Result<u8, &'static str>;
|
||||
fn poll_next(mut self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
tracing::info!(event = "poll");
|
||||
Poll::Ready(self.0.next())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Chunks {
|
||||
fn drop(&mut self) {
|
||||
tracing::info!(event = "drop");
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.route",
|
||||
skip_all,
|
||||
fields(route = "fixture", stream, outcome)
|
||||
)]
|
||||
async fn streamed() -> Result<CallOutput<(), (), u8, &'static str>, &'static str> {
|
||||
call(async {
|
||||
Ok(CallOutput::Stream {
|
||||
head: (),
|
||||
chunks: Box::pin(Chunks(vec![Ok(1), Err("broken"), Ok(2)].into_iter())),
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn stream_errors_finish_once_and_poll_and_drop_use_the_captured_context(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let CallOutput::Stream { mut chunks, .. } = logger.instrument(streamed()).await.unwrap()
|
||||
else {
|
||||
panic!()
|
||||
};
|
||||
assert!(records.try_recv().is_err());
|
||||
tokio::spawn(async move {
|
||||
assert_eq!(chunks.next().await, Some(Ok(1)));
|
||||
assert_eq!(chunks.next().await, Some(Err("broken")));
|
||||
assert_eq!(chunks.next().await, None);
|
||||
let emitted = records.try_iter().collect::<Vec<_>>();
|
||||
assert_eq!(emitted.len(), 4);
|
||||
assert!(emitted.iter().all(|record| record["route"] == "fixture"));
|
||||
assert_eq!(emitted[0]["event"], "poll");
|
||||
assert_eq!(emitted[1]["event"], "poll");
|
||||
assert_eq!(emitted[2]["event"], "drop");
|
||||
assert_eq!(emitted[3]["outcome"], "failure");
|
||||
drop(chunks);
|
||||
assert!(records.try_recv().is_err());
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(route = "waiting", outcome))]
|
||||
async fn waiting(streaming: bool) {
|
||||
if streaming {
|
||||
let _: Result<CallOutput<(), (), u8, ()>, ()> = call(std::future::pending()).await;
|
||||
} else {
|
||||
let _: Result<(), ()> = unary(std::future::pending()).await;
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unary(false)]
|
||||
#[case::streaming(true)]
|
||||
fn cancellation_before_headers_closes_the_span(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
#[case] streaming: bool,
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let mut future = Box::pin(logger.instrument(waiting(streaming)));
|
||||
assert!(
|
||||
future
|
||||
.as_mut()
|
||||
.poll(&mut Context::from_waker(noop_waker_ref()))
|
||||
.is_pending()
|
||||
);
|
||||
assert!(records.try_recv().is_err());
|
||||
drop(future);
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["outcome"], "cancelled");
|
||||
assert_eq!(summary["route"], "waiting");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn dropped_stream_teardown_uses_its_original_logger(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let output = logger.instrument(streamed()).await.unwrap();
|
||||
Logger::default().scope(|| drop(output));
|
||||
let emitted = records.try_iter().collect::<Vec<_>>();
|
||||
assert_eq!(emitted.len(), 2);
|
||||
assert_eq!(
|
||||
emitted[0],
|
||||
json!({"route":"fixture", "stream":true, "event":"drop"})
|
||||
);
|
||||
assert_eq!(emitted[1]["outcome"], "cancelled");
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "disabled", level = "debug", skip_all, fields(outcome))]
|
||||
async fn disabled_child() {
|
||||
let _: Result<(), ()> = operation("disabled", async { Err(()) }).await;
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_filtered_operation_does_not_overwrite_its_parent_outcome(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
logger
|
||||
.instrument(async {
|
||||
let parent = tracing::info_span!("parent", outcome = "original");
|
||||
tracing::Instrument::instrument(disabled_child(), parent).await;
|
||||
})
|
||||
.await;
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["span_name"], "parent");
|
||||
assert_eq!(summary["outcome"], "original");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(stream = true, outcome))]
|
||||
async fn completed() -> Result<CallOutput<(), (), u8, ()>, ()> {
|
||||
call(async { Ok(CallOutput::Complete(())) }).await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn streaming_mode_reflects_the_returned_output(logger: (Logger, mpsc::Receiver<Value>)) {
|
||||
let (logger, records) = logger;
|
||||
logger.instrument(completed()).await.unwrap();
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["stream"], false);
|
||||
assert_eq!(summary["outcome"], "success");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
}
|
||||
|
|
@ -37,34 +37,23 @@ pub enum RouteError {
|
|||
Http(#[from] litellm_http::Error),
|
||||
#[error(transparent)]
|
||||
Secret(#[from] SecretError),
|
||||
#[error("post-call hook failed: {0}")]
|
||||
PostCallHook(#[source] Arc<RouteError>),
|
||||
}
|
||||
|
||||
/// Whether the provider had already been called when the route failed. Before the send, a
|
||||
/// host may retry on another path; after it, the provider has done the work and billed for it.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Phase {
|
||||
BeforeSend,
|
||||
AfterSend,
|
||||
impl From<litellm_host::machine::MachineFault> for RouteError {
|
||||
fn from(fault: litellm_host::machine::MachineFault) -> Self {
|
||||
use litellm_host::machine::MachineFault;
|
||||
Self::InvalidRequest(match fault {
|
||||
MachineFault::Abandoned => "host driver was abandoned".into(),
|
||||
MachineFault::Protocol(message) => format!("host {message}").into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl RouteError {
|
||||
pub fn phase(&self) -> Phase {
|
||||
match self {
|
||||
Self::InvalidResponse(_)
|
||||
| Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => {
|
||||
Phase::AfterSend
|
||||
}
|
||||
Self::Transport(TransportError::Connect(_))
|
||||
| Self::InvalidType { .. }
|
||||
| Self::MissingField(_)
|
||||
| Self::InvalidProvider(_)
|
||||
| Self::InvalidRequest(_)
|
||||
| Self::Unsupported(_)
|
||||
| Self::Auth(_)
|
||||
| Self::Headers(_)
|
||||
| Self::Http(_)
|
||||
| Self::Secret(_) => Phase::BeforeSend,
|
||||
}
|
||||
pub(crate) fn post_call(error: Self) -> Self {
|
||||
Self::PostCallHook(Arc::new(error))
|
||||
}
|
||||
|
||||
/// The caller's request is what is wrong, as opposed to the environment, the wire, or
|
||||
|
|
@ -78,9 +67,11 @@ impl RouteError {
|
|||
| Self::Unsupported(_)
|
||||
| Self::Headers(_) => true,
|
||||
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
|
||||
Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => {
|
||||
false
|
||||
}
|
||||
Self::InvalidResponse(_)
|
||||
| Self::Transport(_)
|
||||
| Self::Http(_)
|
||||
| Self::Secret(_)
|
||||
| Self::PostCallHook(_) => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -124,34 +115,10 @@ impl Eq for SecretError {}
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{Phase, RouteError};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use super::RouteError;
|
||||
use litellm_llms::{Error as LlmError, ErrorDetail};
|
||||
use rstest::rstest;
|
||||
|
||||
#[test]
|
||||
fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() {
|
||||
let after = [
|
||||
RouteError::InvalidResponse("bad json".into()),
|
||||
RouteError::Transport(TransportError::Http {
|
||||
status: 500,
|
||||
body: "boom".into(),
|
||||
}),
|
||||
RouteError::Transport(TransportError::Network("reset".into())),
|
||||
];
|
||||
for error in after {
|
||||
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
|
||||
}
|
||||
let before = [
|
||||
RouteError::Transport(TransportError::Connect("refused".into())),
|
||||
RouteError::Unsupported("streaming"),
|
||||
RouteError::Auth(litellm_auth::Error::InvalidHeader),
|
||||
];
|
||||
for error in before {
|
||||
assert_eq!(error.phase(), Phase::BeforeSend, "{error:?}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_missing_api_key_is_the_environment_not_the_request() {
|
||||
assert!(
|
||||
|
|
@ -166,12 +133,9 @@ mod tests {
|
|||
assert!(!RouteError::InvalidResponse("bad json".into()).is_request());
|
||||
}
|
||||
#[rstest]
|
||||
#[case::request(true, Phase::BeforeSend)]
|
||||
#[case::response(false, Phase::AfterSend)]
|
||||
fn contextual_errors_preserve_sources_and_route_classification(
|
||||
#[case] request: bool,
|
||||
#[case] phase: Phase,
|
||||
) {
|
||||
#[case::request(true)]
|
||||
#[case::response(false)]
|
||||
fn contextual_errors_preserve_sources_and_route_classification(#[case] request: bool) {
|
||||
let source = serde_json::from_str::<serde_json::Value>("{").unwrap_err();
|
||||
let source_message = source.to_string();
|
||||
let detail = ErrorDetail::invalid("test payload", source);
|
||||
|
|
@ -180,7 +144,6 @@ mod tests {
|
|||
} else {
|
||||
LlmError::InvalidResponse(detail)
|
||||
});
|
||||
assert_eq!(error.phase(), phase);
|
||||
assert_eq!(error.is_request(), request);
|
||||
let category = if request { "request" } else { "response" };
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -1,11 +1,40 @@
|
|||
mod context;
|
||||
mod diagnostic;
|
||||
|
||||
pub mod audio_transcription;
|
||||
pub mod caching;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
pub mod error;
|
||||
pub mod messages;
|
||||
pub mod ocr;
|
||||
mod outbound;
|
||||
mod provider;
|
||||
pub mod resources;
|
||||
pub mod responses;
|
||||
|
||||
pub use error::{Phase, RouteError};
|
||||
pub use error::RouteError;
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct CallOptions {
|
||||
pub cache: Option<litellm_cache_response::CachePolicy>,
|
||||
pub observers: Option<litellm_host::observation::ObservationSender>,
|
||||
}
|
||||
|
||||
impl From<Option<litellm_host::observation::ObservationSender>> for CallOptions {
|
||||
fn from(observers: Option<litellm_host::observation::ObservationSender>) -> Self {
|
||||
Self {
|
||||
cache: None,
|
||||
observers,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<litellm_cache_response::CachePolicy> for CallOptions {
|
||||
fn from(cache: litellm_cache_response::CachePolicy) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
observers: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
This directory owns provider-independent Messages call orchestration: the entrypoint, call envelopes, provider selection, credential resolution, transport coordination, hooks, and stream lifecycle. Shared API data contracts belong in `litellm-types::messages`, adapter contracts and execution inputs in `llms/src/base_llm/messages`, and provider implementations in `llms/src/<provider>/messages`
|
||||
This directory owns provider-independent Messages call orchestration: the entrypoint, call envelopes, provider selection, credential resolution, transport coordination, hooks, and stream lifecycle. Shared API data contracts belong in `litellm-llms-types::formats::messages`, adapter contracts and execution inputs in `llms/src/base_llm/messages`, and provider implementations in `llms/src/<provider>/messages`
|
||||
|
||||
Select concrete provider adapters and invoke their contracts. Delegate authentication policy, beta selection, payload rewriting, and response interpretation to those adapters. Keep provider policy out of request preparation and transport handlers. Calling a concrete provider helper for every provider is still a policy dependency
|
||||
|
||||
|
|
|
|||
|
|
@ -3,18 +3,17 @@ pub(super) use litellm_http::request::truncate_error_body;
|
|||
use litellm_llms::{
|
||||
anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::messages::transformation::BaseAnthropicMessagesConfig,
|
||||
base_llm::messages::transformation::BaseMessagesConfig,
|
||||
bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
use super::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
const HEADER_CONTEXT: &str = "messages";
|
||||
|
||||
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum MessagesProvider {
|
||||
Anthropic,
|
||||
AzureAi,
|
||||
|
|
@ -23,10 +22,15 @@ pub(crate) enum MessagesProvider {
|
|||
|
||||
impl MessagesProvider {
|
||||
pub(crate) fn as_str(self) -> &'static str {
|
||||
self.into()
|
||||
match self {
|
||||
Self::Anthropic => LlmProviders::Anthropic,
|
||||
Self::AzureAi => LlmProviders::AzureAi,
|
||||
Self::Bedrock => LlmProviders::Bedrock,
|
||||
}
|
||||
.into()
|
||||
}
|
||||
|
||||
pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig {
|
||||
pub(crate) fn config(self) -> &'static dyn BaseMessagesConfig {
|
||||
match self {
|
||||
Self::Anthropic => &ANTHROPIC_MESSAGES_CONFIG,
|
||||
Self::AzureAi => &AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
|
|
@ -35,6 +39,21 @@ impl MessagesProvider {
|
|||
}
|
||||
}
|
||||
|
||||
pub(crate) fn messages_provider(provider: LlmProviders) -> Option<MessagesProvider> {
|
||||
match provider {
|
||||
LlmProviders::Anthropic => Some(MessagesProvider::Anthropic),
|
||||
LlmProviders::AzureAi => Some(MessagesProvider::AzureAi),
|
||||
LlmProviders::Bedrock => Some(MessagesProvider::Bedrock),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn string_headers(
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
|
|
@ -47,8 +66,9 @@ mod tests {
|
|||
|
||||
use rstest::rstest;
|
||||
|
||||
use super::{MessagesProvider, string_headers, truncate_error_body};
|
||||
use super::{MessagesProvider, messages_provider, string_headers, truncate_error_body};
|
||||
use crate::messages::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic", MessagesProvider::Anthropic)]
|
||||
|
|
@ -58,13 +78,16 @@ mod tests {
|
|||
#[case] name: &str,
|
||||
#[case] provider: MessagesProvider,
|
||||
) {
|
||||
assert_eq!(name.parse::<MessagesProvider>(), Ok(provider));
|
||||
assert_eq!(
|
||||
messages_provider(name.parse::<LlmProviders>().unwrap()),
|
||||
Some(provider)
|
||||
);
|
||||
assert_eq!(provider.as_str(), name);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_without_a_messages_config_is_rejected() {
|
||||
assert!("openai".parse::<MessagesProvider>().is_err());
|
||||
assert_eq!(messages_provider(LlmProviders::Openai), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -2,99 +2,144 @@ use std::time::Duration;
|
|||
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_host::interceptors::{Interceptors, ProviderIdentity, RequestContext, WireRequest};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{Authenticated, resolve_auth},
|
||||
messages::{
|
||||
streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
|
||||
transformation::BaseAnthropicMessagesConfig,
|
||||
transformation::BaseMessagesConfig,
|
||||
},
|
||||
};
|
||||
use litellm_tracing::{ByteChunk, debug};
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use litellm_llms_types::formats::messages::MessagesResponse;
|
||||
use litellm_tracing::ByteChunk;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error, MessagesResponse, common_utils::truncate_error_body, prepare::ProviderMessagesRequest,
|
||||
Error, MessagesCallResponse, MessagesRoute, common_utils::truncate_error_body,
|
||||
prepare::ProviderMessagesRequest,
|
||||
};
|
||||
use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request};
|
||||
use crate::{constants::MESSAGES_TIMEOUT_SECS, context::CallContext, outbound::outbound_request};
|
||||
|
||||
pub(super) async fn execute(
|
||||
http: &litellm_http::Client,
|
||||
auth: &AuthServices,
|
||||
request: ProviderMessagesRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let ProviderMessagesRequest {
|
||||
provider,
|
||||
url,
|
||||
body,
|
||||
environment,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
let stream = body.params.stream == Some(true);
|
||||
let context = RequestContext {
|
||||
model: body.model.clone(),
|
||||
custom_llm_provider: provider.as_str().to_string(),
|
||||
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
|
||||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
body: serde_json::to_value(&body).map_err(serialize_failure)?,
|
||||
pub(super) struct ProviderCall {
|
||||
pub identity: ProviderIdentity,
|
||||
pub wire: WireRequest,
|
||||
provider: super::common_utils::MessagesProvider,
|
||||
signer: Option<litellm_auth_aws::SigV4Signer>,
|
||||
timeout: Option<Duration>,
|
||||
stream: bool,
|
||||
}
|
||||
|
||||
impl ProviderCall {
|
||||
pub fn cacheable(&self) -> bool {
|
||||
self.signer.is_none()
|
||||
}
|
||||
}
|
||||
|
||||
impl MessagesRoute {
|
||||
pub(super) async fn prepare_outbound(
|
||||
&self,
|
||||
request: ProviderMessagesRequest,
|
||||
context: &CallContext<'_, impl Interceptors<Error>>,
|
||||
) -> Result<ProviderCall, Error> {
|
||||
let ProviderMessagesRequest {
|
||||
provider,
|
||||
url,
|
||||
body,
|
||||
environment,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
let request_context = RequestContext {
|
||||
model: body.model.clone(),
|
||||
custom_llm_provider: provider.as_str().to_string(),
|
||||
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
|
||||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated =
|
||||
resolve_auth(&self.auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let identity = ProviderIdentity {
|
||||
model: request_context.model.clone(),
|
||||
provider: request_context.custom_llm_provider.clone(),
|
||||
};
|
||||
let wire = context
|
||||
.interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
body: serde_json::to_value(&body).map_err(serialize_failure)?,
|
||||
},
|
||||
request_context,
|
||||
)
|
||||
.await?;
|
||||
let stream = match wire.body.get("stream") {
|
||||
None | Some(Value::Null) => false,
|
||||
Some(Value::Bool(stream)) => *stream,
|
||||
Some(value) => {
|
||||
return Err(Error::InvalidRequest(
|
||||
litellm_llms::ErrorDetail::InvalidValue {
|
||||
field: "stream",
|
||||
expected: "a boolean",
|
||||
actual: value.clone(),
|
||||
},
|
||||
));
|
||||
}
|
||||
};
|
||||
Ok(ProviderCall {
|
||||
identity,
|
||||
wire,
|
||||
provider,
|
||||
signer: authenticated.signer,
|
||||
timeout,
|
||||
stream,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn call_provider(
|
||||
&self,
|
||||
request: ProviderCall,
|
||||
context: &CallContext<'_, impl Interceptors<Error>>,
|
||||
) -> Result<MessagesCallResponse, Error> {
|
||||
let ProviderCall {
|
||||
identity,
|
||||
wire,
|
||||
provider,
|
||||
signer,
|
||||
timeout,
|
||||
stream,
|
||||
} = request;
|
||||
let provider_name = provider.as_str();
|
||||
log_request_body(provider_name, stream, &wire.body);
|
||||
let response = send(
|
||||
&self.http,
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer,
|
||||
},
|
||||
context,
|
||||
&wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)
|
||||
.await?;
|
||||
let provider_name = provider.as_str();
|
||||
debug!(provider = provider_name, stream, body = %wire.body, "provider request");
|
||||
let response = send(
|
||||
http,
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
&wire.url,
|
||||
&wire.body,
|
||||
timeout,
|
||||
)
|
||||
.await?;
|
||||
debug!(
|
||||
provider = provider_name,
|
||||
status = response.status().as_u16(),
|
||||
"provider response headers"
|
||||
);
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
let config = provider.config();
|
||||
if stream {
|
||||
return Ok(streaming_response(
|
||||
response,
|
||||
config.stream_decoder(),
|
||||
provider_name,
|
||||
));
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
log_response_body(&text);
|
||||
context.response_received(&text).await?;
|
||||
decode_response(config, &identity.model, &text)
|
||||
.map(|message| MessagesCallResponse::Complete(Box::new(message)))
|
||||
}
|
||||
let config = provider.config();
|
||||
if stream {
|
||||
return Ok(streaming_response(
|
||||
response,
|
||||
config.stream_decoder(),
|
||||
provider_name,
|
||||
));
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
debug!(body = text.as_str(), "provider response body");
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
decode_response(config, &body.model, &text)
|
||||
.map(|message| MessagesResponse::Message(Box::new(message)))
|
||||
}
|
||||
|
||||
fn serialize_failure(err: serde_json::Error) -> Error {
|
||||
|
|
@ -121,14 +166,14 @@ async fn send(
|
|||
body,
|
||||
Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
|
||||
)?;
|
||||
request.send(http).await.map_err(network)
|
||||
crate::outbound::send(request, http).await.map_err(network)
|
||||
}
|
||||
|
||||
async fn provider_error(response: reqwest::Response) -> Error {
|
||||
let status = response.status().as_u16();
|
||||
match response.text().await {
|
||||
Ok(text) => {
|
||||
litellm_tracing::debug!(status, body = text.as_str(), "provider error body");
|
||||
log_error_body(status, &text);
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
|
|
@ -139,10 +184,10 @@ async fn provider_error(response: reqwest::Response) -> Error {
|
|||
}
|
||||
|
||||
fn decode_response(
|
||||
config: &dyn BaseAnthropicMessagesConfig,
|
||||
config: &dyn BaseMessagesConfig,
|
||||
model: &str,
|
||||
text: &str,
|
||||
) -> Result<AnthropicMessagesResponse, Error> {
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let response = serde_json::from_str(text).map_err(|err| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"messages response JSON",
|
||||
|
|
@ -158,7 +203,7 @@ fn streaming_response(
|
|||
response: reqwest::Response,
|
||||
decoder: Option<StreamDecoder>,
|
||||
provider: &'static str,
|
||||
) -> MessagesResponse {
|
||||
) -> MessagesCallResponse {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
|
|
@ -175,7 +220,10 @@ fn streaming_response(
|
|||
.boxed(),
|
||||
Some(decode) => decoded_chunks(response, decode, provider),
|
||||
};
|
||||
MessagesResponse::Stream { headers, chunks }
|
||||
MessagesCallResponse::Stream {
|
||||
head: super::route::MessagesStreamHead { headers },
|
||||
chunks,
|
||||
}
|
||||
}
|
||||
|
||||
fn decoded_chunks(
|
||||
|
|
@ -199,9 +247,21 @@ fn decoded_chunks(
|
|||
.boxed()
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
|
||||
fn log_request_body(provider: &str, stream: bool, body: &serde_json::Value) {
|
||||
tracing::debug!(provider, stream, body = %body, "provider request");
|
||||
}
|
||||
|
||||
fn log_response_body(body: &str) {
|
||||
tracing::debug!(body, "provider response body");
|
||||
}
|
||||
|
||||
fn log_error_body(status: u16, body: &str) {
|
||||
tracing::debug!(status, body, "provider error body");
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &bytes::Bytes) {
|
||||
let chunk = ByteChunk::new(data);
|
||||
debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
|
||||
tracing::debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -217,6 +277,7 @@ mod tests {
|
|||
"data: {\"type\":\"ping\"}\n\n",
|
||||
Some("event: ping\ndata: {\"type\":\"ping\"}\n\n")
|
||||
)]
|
||||
#[rstest::rstest]
|
||||
#[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)]
|
||||
#[tokio::test]
|
||||
async fn decoded_streams_encode_events_and_stop_at_the_first_error(
|
||||
|
|
@ -233,7 +294,7 @@ mod tests {
|
|||
.send()
|
||||
.await
|
||||
.unwrap();
|
||||
let MessagesResponse::Stream { mut chunks, .. } =
|
||||
let MessagesCallResponse::Stream { mut chunks, .. } =
|
||||
streaming_response(response, Some(anthropic_sse_event_stream), "test")
|
||||
else {
|
||||
panic!("a streaming response returns chunks");
|
||||
|
|
|
|||
|
|
@ -1,27 +1,100 @@
|
|||
//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`.
|
||||
//!
|
||||
//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the
|
||||
//! same two steps as a machine for a host that answers the call's operations itself.
|
||||
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
mod types;
|
||||
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use futures_util::FutureExt;
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::interceptors::{ExecutionFacts, Interceptors, ResultSource};
|
||||
|
||||
use crate::{caching::CallCache, context::CallContext};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub use crate::error::RouteError as Error;
|
||||
pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body};
|
||||
pub use types::{MessagesCall, MessagesCallResponse, MessagesShaping, messages_body};
|
||||
|
||||
pub async fn messages(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: &dyn SecretSource,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let request = prepare::prepare(call, secrets).await?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
#[derive(Clone)]
|
||||
pub struct MessagesRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
}
|
||||
|
||||
impl MessagesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
cache: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> Result<MessagesCallResponse, Error> {
|
||||
let context = CallContext::new(interceptors, options.into());
|
||||
litellm_host::lifecycle::observe_call(context.observers.clone(), self.run(call, context))
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "messages",
|
||||
model = %call.body.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = call.body.params.stream == Some(true),
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
context: CallContext<'_, impl Interceptors<Error>>,
|
||||
) -> Result<MessagesCallResponse, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
let prepared = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&prepared.body.model, prepared.provider.as_str());
|
||||
let request = self.prepare_outbound(prepared, &context).boxed().await?;
|
||||
let cache = CallCache::<route::Messages>::from_wire(
|
||||
self.cache.as_ref().filter(|_| request.cacheable()),
|
||||
context.cache,
|
||||
&request.identity,
|
||||
&request.wire,
|
||||
);
|
||||
let identity = request.identity.clone();
|
||||
let (output, source) = match cache.lookup().await {
|
||||
Some(hit) => hit,
|
||||
None => (
|
||||
self.call_provider(request, &context).await?,
|
||||
ResultSource::Provider,
|
||||
),
|
||||
};
|
||||
context
|
||||
.result_ready(ExecutionFacts {
|
||||
provider: identity,
|
||||
source: source.clone(),
|
||||
})
|
||||
.await?;
|
||||
Ok(cache.finish(output, &source).await)
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,22 +3,21 @@ use std::time::Duration;
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::{
|
||||
dot_notation_indexing::delete_nested_value,
|
||||
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
|
||||
get_provider_specific_headers::get_provider_specific_headers,
|
||||
settings::Lookup,
|
||||
get_provider_specific_headers::get_provider_specific_headers, settings::Lookup,
|
||||
};
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::ValidatedEnvironment, messages::context::MessagesTransformContext,
|
||||
};
|
||||
use litellm_llms_types::formats::messages::MessagesRequest;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
|
||||
use super::{
|
||||
Error, MessagesCall,
|
||||
common_utils::{MessagesProvider, string_headers},
|
||||
common_utils::{MessagesProvider, messages_provider, string_headers},
|
||||
types::invalid_request,
|
||||
};
|
||||
use crate::provider::resolve_llm_provider;
|
||||
|
||||
struct ResolvedProvider {
|
||||
model: String,
|
||||
|
|
@ -28,13 +27,14 @@ struct ResolvedProvider {
|
|||
pub(super) struct ProviderMessagesRequest {
|
||||
pub(super) provider: MessagesProvider,
|
||||
pub(super) url: String,
|
||||
pub(super) body: AnthropicMessagesRequest,
|
||||
pub(super) body: MessagesRequest,
|
||||
pub(super) environment: ValidatedEnvironment,
|
||||
pub(super) timeout: Option<Duration>,
|
||||
/// The caller's own credential, reported to the host beside the wire request.
|
||||
pub(super) api_key: Option<SecretValue>,
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) async fn prepare(
|
||||
call: MessagesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
|
|
@ -50,26 +50,11 @@ fn resolve_provider(
|
|||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for messages request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let provider = provider
|
||||
.parse()
|
||||
.map_err(|_| Error::InvalidProvider(provider.to_string()))?;
|
||||
let resolved = resolve_llm_provider(model, custom_llm_provider, "messages")?;
|
||||
let provider = messages_provider(resolved.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(resolved.provider).to_string()))?;
|
||||
Ok(ResolvedProvider {
|
||||
model: model.to_string(),
|
||||
model: resolved.model.to_string(),
|
||||
provider,
|
||||
})
|
||||
}
|
||||
|
|
@ -94,7 +79,7 @@ fn prepare_provider_request(
|
|||
let env_lookup = |key: &str| secrets.get(key);
|
||||
|
||||
let sanitized = config.shape_request(
|
||||
AnthropicMessagesRequest { model, ..body },
|
||||
MessagesRequest { model, ..body },
|
||||
shaping.reasoning_auto_summary,
|
||||
)?;
|
||||
let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?;
|
||||
|
|
@ -139,9 +124,9 @@ fn prepare_provider_request(
|
|||
}
|
||||
|
||||
fn without_additional_drop_params(
|
||||
request: AnthropicMessagesRequest,
|
||||
request: MessagesRequest,
|
||||
paths: &[String],
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
) -> Result<MessagesRequest, Error> {
|
||||
if paths.is_empty() {
|
||||
return Ok(request);
|
||||
}
|
||||
|
|
@ -149,7 +134,7 @@ fn without_additional_drop_params(
|
|||
let trimmed = paths
|
||||
.iter()
|
||||
.fold(params, |params, path| delete_nested_value(params, path));
|
||||
Ok(AnthropicMessagesRequest {
|
||||
Ok(MessagesRequest {
|
||||
params: serde_json::from_value(trimmed).map_err(invalid_request)?,
|
||||
..request
|
||||
})
|
||||
|
|
@ -158,7 +143,7 @@ fn without_additional_drop_params(
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::base_llm::auth::resolve_auth;
|
||||
use litellm_types::utils::ProviderSpecificHeaders;
|
||||
use litellm_llms_types::headers::ProviderSpecificHeaders;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
|
|
@ -170,7 +155,7 @@ mod tests {
|
|||
MessagesShaping::default()
|
||||
}
|
||||
|
||||
fn body(value: Value) -> AnthropicMessagesRequest {
|
||||
fn body(value: Value) -> MessagesRequest {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,26 +1,15 @@
|
|||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_host::{
|
||||
host::{Demand, Host},
|
||||
machine::{CallMachine, HostChannel, MachineFault},
|
||||
call::{HostedCompletion, HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_http::{Client, ClientVariant, HttpClientConfig};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use litellm_llms_types::formats::messages::MessagesResponse;
|
||||
|
||||
use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare};
|
||||
use super::{Error, MessagesCall};
|
||||
|
||||
pub enum MessagesOutput {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
/// Every chunk already reached the host through `Deliver`.
|
||||
Streamed,
|
||||
}
|
||||
pub type MessagesOutput = HostedCompletion<Box<MessagesResponse>>;
|
||||
|
||||
/// The upstream response as the caller sees it at stream hand-off, before any chunk.
|
||||
pub struct MessagesStreamHead {
|
||||
|
|
@ -30,91 +19,60 @@ pub struct MessagesStreamHead {
|
|||
pub struct Messages;
|
||||
|
||||
impl Protocol for Messages {
|
||||
type Response = MessagesOutput;
|
||||
type Response = Box<MessagesResponse>;
|
||||
type Error = Error;
|
||||
type Projection = MessagesCall;
|
||||
type Op = Infallible;
|
||||
type Request = MessagesCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = MessagesStreamHead;
|
||||
}
|
||||
|
||||
impl From<MachineFault> for Error {
|
||||
fn from(fault: MachineFault) -> Self {
|
||||
Self::InvalidRequest(match fault {
|
||||
MachineFault::Abandoned => "messages host driver was abandoned".into(),
|
||||
MachineFault::Protocol(message) => format!("messages {message}").into(),
|
||||
pub type MessagesMachine = HostedMachine<Messages>;
|
||||
|
||||
impl super::MessagesRoute {
|
||||
pub fn machine(
|
||||
self,
|
||||
request: super::MessagesCall,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> MessagesMachine {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
hosted_call(
|
||||
request,
|
||||
observers,
|
||||
move |call, _, interceptors, observers| async move {
|
||||
let context = crate::context::CallContext::new(
|
||||
&interceptors,
|
||||
crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
},
|
||||
);
|
||||
self.run(call, context).await
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::Cachable for Messages {
|
||||
const SURFACE: &'static str = "messages";
|
||||
}
|
||||
|
||||
impl crate::caching::StreamCachable for Messages {
|
||||
const TERMINAL_EVENT: &'static str = "message_stop";
|
||||
|
||||
fn replay(data: bytes::Bytes) -> Option<litellm_host::call::OutputOf<Self>> {
|
||||
Some(litellm_host::call::CallOutput::Stream {
|
||||
head: MessagesStreamHead {
|
||||
headers: Vec::new(),
|
||||
},
|
||||
chunks: Box::pin(futures_util::stream::iter([Ok(data)])),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub type MessagesHost = HostChannel<Messages>;
|
||||
pub type MessagesMachine = CallMachine<Messages>;
|
||||
|
||||
/// The in-process host for a request already in hand. It answers projection once and
|
||||
/// observes nothing.
|
||||
pub struct LocalMessagesHost {
|
||||
call: Mutex<Option<MessagesCall>>,
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
pub fn new(call: MessagesCall) -> Self {
|
||||
Self {
|
||||
call: Mutex::new(Some(call)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for LocalMessagesHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn messages_machine(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesMachine, litellm_http::Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let auth = resources.auth.clone();
|
||||
Ok(CallMachine::new(move |host| {
|
||||
Box::pin(drive(host, http, auth, secrets))
|
||||
}))
|
||||
}
|
||||
|
||||
/// The call as its host sees it: projection first, then the same prepare and execute as
|
||||
/// [`super::messages`], with each chunk of a stream handed over as it arrives.
|
||||
async fn drive(
|
||||
host: MessagesHost,
|
||||
http: Client,
|
||||
auth: Arc<litellm_auth::AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let call = host.project().await?;
|
||||
let request = prepare(call, secrets.as_ref()).await?;
|
||||
match execute(&http, &auth, request, &host).await? {
|
||||
MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)),
|
||||
MessagesResponse::Stream {
|
||||
headers,
|
||||
mut chunks,
|
||||
} => {
|
||||
if host.open(MessagesStreamHead { headers }).await? == Demand::Detached {
|
||||
return Ok(MessagesOutput::Streamed);
|
||||
}
|
||||
while let Some(chunk) = chunks.try_next().await? {
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(MessagesOutput::Streamed)
|
||||
}
|
||||
fn bytes(chunk: &Self::Chunk) -> &[u8] {
|
||||
chunk.as_ref()
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,13 +1,11 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream::BoxStream;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
},
|
||||
utils::ProviderSpecificHeaders,
|
||||
use litellm_llms_types::{
|
||||
formats::messages::{MessagesRequest, MessagesResponse},
|
||||
headers::ProviderSpecificHeaders,
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
|
@ -15,7 +13,7 @@ use serde_json::{Map, Value};
|
|||
use super::Error;
|
||||
|
||||
pub struct MessagesCall {
|
||||
pub body: AnthropicMessagesRequest,
|
||||
pub body: MessagesRequest,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
|
|
@ -25,24 +23,16 @@ pub struct MessagesCall {
|
|||
pub shaping: MessagesShaping,
|
||||
}
|
||||
|
||||
pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesRequest, Error> {
|
||||
pub fn messages_body(body: Map<String, Value>) -> Result<MessagesRequest, Error> {
|
||||
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
|
||||
}
|
||||
|
||||
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(litellm_llms::ErrorDetail::invalid(
|
||||
"Anthropic messages request",
|
||||
err,
|
||||
))
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into())
|
||||
}
|
||||
|
||||
pub enum MessagesResponse {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
Stream {
|
||||
headers: Vec<(String, String)>,
|
||||
chunks: BoxStream<'static, Result<Bytes, Error>>,
|
||||
},
|
||||
}
|
||||
pub type MessagesCallResponse =
|
||||
CallOutput<Box<MessagesResponse>, super::route::MessagesStreamHead, Bytes, Error>;
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
|
|
|
|||
|
|
@ -1,15 +1,81 @@
|
|||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_host::interceptors::Interceptors;
|
||||
use litellm_llms::base_llm::ocr::{error::Error, handler::OcrClient};
|
||||
use litellm_llms_types::formats::ocr::LiteLLMOcrResponse;
|
||||
|
||||
use super::{
|
||||
handler::perform_ocr_request,
|
||||
types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest},
|
||||
};
|
||||
|
||||
use crate::ocr::{
|
||||
route::{LocalOcrHost, ocr_machine},
|
||||
types::LiteLLMOcrRequest,
|
||||
};
|
||||
|
||||
pub async fn perform(
|
||||
client: &OcrClient,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_host::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await
|
||||
#[derive(Clone)]
|
||||
pub struct OcrRoute {
|
||||
client: OcrClient,
|
||||
}
|
||||
|
||||
impl OcrRoute {
|
||||
pub fn new(client: OcrClient) -> Self {
|
||||
Self { client }
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
request: LiteLLMOcrRequest,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<ObservationSender>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_host::lifecycle::observe_unary(
|
||||
observers.clone(),
|
||||
self.run(request, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "ocr",
|
||||
model = %request.model,
|
||||
resolved_model = %request.model,
|
||||
provider = <&str>::from(request.config.provider()),
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub(super) async fn run(
|
||||
&self,
|
||||
request: LiteLLMOcrRequest,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let caller_document = matches!(&request.document, OcrDocumentInput::Document(_));
|
||||
let prepared = prepare_request_document(request).await?;
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<LiteLLMOcrResponse, Error>> =
|
||||
Box::pin(perform_ocr_request(
|
||||
&self.client,
|
||||
prepared,
|
||||
interceptors,
|
||||
caller_document,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
async fn prepare_request_document(
|
||||
request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
) -> Result<ResolvedOcrRequest, Error> {
|
||||
if let OcrDocumentInput::Document(_) = &request.document {
|
||||
return request.map_document(super::document::prepare_document);
|
||||
}
|
||||
let logger = litellm_tracing::Logger::current();
|
||||
let span = tracing::Span::current();
|
||||
tokio::task::spawn_blocking(move || {
|
||||
logger.scope(|| span.in_scope(|| request.map_document(super::document::prepare_document)))
|
||||
})
|
||||
.await
|
||||
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,10 +1,8 @@
|
|||
use std::{collections::BTreeMap as Map, io::Read, path::Path};
|
||||
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
transformation::{OCR_INLINE_MAX_BYTES, OcrDocument},
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{error::Error, transformation::OCR_INLINE_MAX_BYTES};
|
||||
use litellm_llms_types::formats::ocr::OcrDocument;
|
||||
|
||||
use crate::ocr::types::OcrDocumentInput;
|
||||
|
||||
|
|
|
|||
|
|
@ -1,24 +1,23 @@
|
|||
use futures_util::future::BoxFuture;
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::{CallHooks, OcrClient},
|
||||
transformation::{LiteLLMOcrResponse, PreparedOcrRequest},
|
||||
transformation::PreparedOcrRequest,
|
||||
};
|
||||
use litellm_llms_types::formats::ocr::LiteLLMOcrResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind,
|
||||
route::OcrHost,
|
||||
};
|
||||
use super::{arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind};
|
||||
use crate::ocr::types::ResolvedOcrRequest;
|
||||
|
||||
pub(crate) async fn perform_ocr_request(
|
||||
client: &OcrClient,
|
||||
request: ResolvedOcrRequest,
|
||||
host: &OcrHost,
|
||||
host: &impl Interceptors<Error>,
|
||||
caller_document: bool,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
request.response_format()?;
|
||||
let config = request.config;
|
||||
|
|
@ -28,56 +27,62 @@ pub(crate) async fn perform_ocr_request(
|
|||
.await
|
||||
.map_err(|error| Error::Secret(std::sync::Arc::new(error)))?;
|
||||
let request = prepare_request(request, caller_document, client, secrets);
|
||||
let hooks = OcrCallHooks::new(host.clone(), &request, config);
|
||||
config.ocr(client, &request, &hooks).await
|
||||
let interceptors = OcrCallHooks::new(host, &request, config, observers);
|
||||
config.ocr(client, &request, &interceptors).await
|
||||
}
|
||||
|
||||
/// Lets provider code reach the host mid-call, filling in the request context only the
|
||||
/// route knows.
|
||||
pub(crate) struct OcrCallHooks {
|
||||
host: OcrHost,
|
||||
model: String,
|
||||
custom_llm_provider: &'static str,
|
||||
optional_params: Value,
|
||||
secret_fields: Vec<String>,
|
||||
api_key: Option<SecretValue>,
|
||||
struct OcrCallHooks<'a, H> {
|
||||
interceptors: &'a H,
|
||||
context: RequestContext,
|
||||
observers: Option<&'a ObservationSender>,
|
||||
}
|
||||
|
||||
impl OcrCallHooks {
|
||||
pub(crate) fn new(host: OcrHost, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self {
|
||||
impl<'a, H> OcrCallHooks<'a, H> {
|
||||
fn new(
|
||||
interceptors: &'a H,
|
||||
request: &PreparedOcrRequest,
|
||||
config: OcrConfigKind,
|
||||
observers: Option<&'a ObservationSender>,
|
||||
) -> Self {
|
||||
Self {
|
||||
host,
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: config.provider().into(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
secret_fields: request
|
||||
.optional_params
|
||||
.keys()
|
||||
.filter(|name| is_secret_param(name))
|
||||
.cloned()
|
||||
.collect(),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
interceptors,
|
||||
observers,
|
||||
context: RequestContext {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: <&str>::from(config.provider()).to_owned(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
secret_fields: request
|
||||
.optional_params
|
||||
.keys()
|
||||
.filter(|name| is_secret_param(name))
|
||||
.cloned()
|
||||
.collect(),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl CallHooks<Error> for OcrCallHooks {
|
||||
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
let context = RequestContext {
|
||||
model: self.model.clone(),
|
||||
custom_llm_provider: self.custom_llm_provider.into(),
|
||||
optional_params: self.optional_params.clone(),
|
||||
secret_fields: self.secret_fields.clone(),
|
||||
api_key: self.api_key.clone(),
|
||||
};
|
||||
Box::pin(self.host.before_send(wire, context))
|
||||
impl<H: Interceptors<Error>> CallHooks<Error> for OcrCallHooks<'_, H> {
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
Box::pin(
|
||||
self.interceptors
|
||||
.before_provider_request(wire, self.context.clone()),
|
||||
)
|
||||
}
|
||||
|
||||
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
|
||||
Box::pin(self.host.emit(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse {
|
||||
body: String::from_utf8_lossy(body).into_owned(),
|
||||
},
|
||||
}))
|
||||
let raw = RawResponse {
|
||||
body: String::from_utf8_lossy(body).into_owned(),
|
||||
};
|
||||
if let Some(observers) = self.observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
Box::pin(self.interceptors.after_provider_response(raw))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod arguments;
|
||||
pub mod client;
|
||||
mod client;
|
||||
pub use client::OcrRoute;
|
||||
pub mod document;
|
||||
pub(crate) mod handler;
|
||||
pub(crate) mod prepare;
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ use litellm_llms::base_llm::ocr::{
|
|||
};
|
||||
use litellm_secrets::source::Secrets;
|
||||
|
||||
use super::provider_config::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest};
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
pub(crate) fn prepare_request(
|
||||
request: ResolvedOcrRequest,
|
||||
|
|
@ -16,15 +16,19 @@ pub(crate) fn prepare_request(
|
|||
) -> PreparedOcrRequest {
|
||||
let credentials = request.credentials.clone();
|
||||
let (preferred_api_key_env, api_base_env) = match request.config.provider() {
|
||||
OcrProvider::Mistral => (
|
||||
LlmProviders::Mistral => (
|
||||
Some("MISTRAL_AZURE_API_KEY"),
|
||||
Some("MISTRAL_AZURE_API_BASE"),
|
||||
),
|
||||
OcrProvider::AzureAi => (None, Some("AZURE_AI_API_BASE")),
|
||||
OcrProvider::AwsTextract
|
||||
| OcrProvider::Cohere
|
||||
| OcrProvider::Reducto
|
||||
| OcrProvider::VertexAi => (None, None),
|
||||
LlmProviders::AzureAi => (None, Some("AZURE_AI_API_BASE")),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => (None, None),
|
||||
};
|
||||
let secret = |name: &str| secrets.truthy(name);
|
||||
let dynamic_api_key = credentials.dynamic_api_key.or_else(|| {
|
||||
|
|
@ -76,17 +80,18 @@ mod tests {
|
|||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_core_utils::call_arguments::{CallArguments, compose_body, parse_options};
|
||||
use litellm_host::event::WireRequest;
|
||||
use litellm_host::interceptors::WireRequest;
|
||||
use litellm_llms::{
|
||||
base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::{CallHooks, OcrClient},
|
||||
transformation::{BaseOcrConfig, OcrResponseFormat},
|
||||
transformation::BaseOcrConfig,
|
||||
},
|
||||
cohere::ocr::transformation::CohereParseConfig,
|
||||
mistral::ocr::transformation::MistralOcrConfig,
|
||||
vertex_ai::ocr::transformation::VertexAiOcrConfig,
|
||||
};
|
||||
use litellm_llms_types::formats::ocr::OcrResponseFormat;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
|
@ -96,11 +101,14 @@ mod tests {
|
|||
wire::{OcrWireRequest, decode_request},
|
||||
};
|
||||
|
||||
/// Stands in for a host with no hooks registered.
|
||||
/// Stands in for a host with no interceptors registered.
|
||||
struct NoHooks;
|
||||
|
||||
impl CallHooks<Error> for NoHooks {
|
||||
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
Box::pin(async move { Ok(wire) })
|
||||
}
|
||||
|
||||
|
|
@ -144,6 +152,7 @@ mod tests {
|
|||
json!({"type": "image_url", "image_url": url})
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn cohere_body_keeps_native_document_fields_and_untyped_overrides() {
|
||||
let request = request(
|
||||
|
|
@ -176,6 +185,7 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn explicit_null_options_use_defaults_before_http() {
|
||||
let request = request(
|
||||
|
|
@ -199,6 +209,7 @@ mod tests {
|
|||
assert!(body.get("req_format").is_none());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn direct_and_vertex_mistral_build_the_same_request_and_share_normalization() {
|
||||
let options = json!({
|
||||
|
|
@ -276,7 +287,7 @@ mod tests {
|
|||
pages: Option<Vec<i64>>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn parsed_provider_params_separates_known_and_extra_params() {
|
||||
let arguments: CallArguments = serde_json::from_value(json!({
|
||||
"pages": [0, 2],
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::provider::LlmProviders;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::{
|
||||
aws_textract::ocr::{
|
||||
|
|
@ -13,8 +14,7 @@ use litellm_llms::{
|
|||
error::Error,
|
||||
handler::{self, CallHooks, OcrClient},
|
||||
transformation::{
|
||||
BaseOcrConfig, LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument, OcrResponseFormat,
|
||||
PreparedOcrRequest, ResolvedOcrCredentials,
|
||||
BaseOcrConfig, OcrCredentialInputs, PreparedOcrRequest, ResolvedOcrCredentials,
|
||||
},
|
||||
},
|
||||
cohere::ocr::transformation::CohereParseConfig,
|
||||
|
|
@ -24,7 +24,7 @@ use litellm_llms::{
|
|||
deepseek_transformation::VertexAIDeepSeekOCRConfig, transformation::VertexAiOcrConfig,
|
||||
},
|
||||
};
|
||||
use strum::{EnumString, IntoStaticStr};
|
||||
use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrResponseFormat};
|
||||
|
||||
macro_rules! with_config {
|
||||
($kind:expr, $config:ident => $body:expr) => {
|
||||
|
|
@ -93,16 +93,16 @@ pub(crate) enum OcrConfigKind {
|
|||
}
|
||||
|
||||
impl OcrConfigKind {
|
||||
pub(crate) const fn provider(self) -> OcrProvider {
|
||||
pub(crate) const fn provider(self) -> LlmProviders {
|
||||
match self {
|
||||
Self::AwsTextract | Self::AwsTextractAnalyze => OcrProvider::AwsTextract,
|
||||
Self::Cohere => OcrProvider::Cohere,
|
||||
Self::Mistral => OcrProvider::Mistral,
|
||||
Self::AwsTextract | Self::AwsTextractAnalyze => LlmProviders::AwsTextract,
|
||||
Self::Cohere => LlmProviders::Cohere,
|
||||
Self::Mistral => LlmProviders::Mistral,
|
||||
Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => {
|
||||
OcrProvider::AzureAi
|
||||
LlmProviders::AzureAi
|
||||
}
|
||||
Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto,
|
||||
Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi,
|
||||
Self::ReductoLegacy | Self::ReductoV3 => LlmProviders::Reducto,
|
||||
Self::VertexAi | Self::VertexDeepSeek => LlmProviders::VertexAi,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -133,9 +133,9 @@ impl OcrConfigKind {
|
|||
self,
|
||||
client: &OcrClient,
|
||||
request: &PreparedOcrRequest,
|
||||
hooks: &dyn CallHooks<Error>,
|
||||
interceptors: &dyn CallHooks<Error>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
with_config!(self, config => handler::ocr(&config, client, request, hooks).await)
|
||||
with_config!(self, config => handler::ocr(&config, client, request, interceptors).await)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -187,17 +187,6 @@ pub fn passthrough_response(
|
|||
.map(Some)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub(crate) enum OcrProvider {
|
||||
AwsTextract,
|
||||
Cohere,
|
||||
Mistral,
|
||||
AzureAi,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_provider_config(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
|
|
@ -205,37 +194,45 @@ pub(crate) fn resolve_provider_config(
|
|||
let provider =
|
||||
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: OcrProvider::Mistral.into(),
|
||||
custom_llm_provider: LlmProviders::Mistral.into(),
|
||||
});
|
||||
let ocr_provider = provider
|
||||
let llm_provider = provider
|
||||
.custom_llm_provider
|
||||
.parse::<OcrProvider>()
|
||||
.parse::<LlmProviders>()
|
||||
.map_err(|_| Error::InvalidProvider(provider.custom_llm_provider.to_string()))?;
|
||||
let config = match ocr_provider {
|
||||
OcrProvider::AwsTextract => match TextractOperation::from_model(provider.model)? {
|
||||
let config = match llm_provider {
|
||||
LlmProviders::AwsTextract => match TextractOperation::from_model(provider.model)? {
|
||||
TextractOperation::DetectDocumentText => OcrConfigKind::AwsTextract,
|
||||
TextractOperation::AnalyzeDocument => OcrConfigKind::AwsTextractAnalyze,
|
||||
},
|
||||
OcrProvider::Cohere => OcrConfigKind::Cohere,
|
||||
OcrProvider::Mistral => OcrConfigKind::Mistral,
|
||||
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
LlmProviders::Cohere => OcrConfigKind::Cohere,
|
||||
LlmProviders::Mistral => OcrConfigKind::Mistral,
|
||||
LlmProviders::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
OcrConfigKind::AzureDocumentIntelligence
|
||||
}
|
||||
OcrProvider::AzureAi
|
||||
LlmProviders::AzureAi
|
||||
if provider.model.to_ascii_lowercase().contains("cohere")
|
||||
&& provider.model.to_ascii_lowercase().contains("parse") =>
|
||||
{
|
||||
OcrConfigKind::AzureCohere
|
||||
}
|
||||
OcrProvider::AzureAi => OcrConfigKind::AzureAi,
|
||||
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
LlmProviders::AzureAi => OcrConfigKind::AzureAi,
|
||||
LlmProviders::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
OcrConfigKind::ReductoLegacy
|
||||
}
|
||||
OcrProvider::Reducto => OcrConfigKind::ReductoV3,
|
||||
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
LlmProviders::Reducto => OcrConfigKind::ReductoV3,
|
||||
LlmProviders::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
OcrConfigKind::VertexDeepSeek
|
||||
}
|
||||
OcrProvider::VertexAi => OcrConfigKind::VertexAi,
|
||||
LlmProviders::VertexAi => OcrConfigKind::VertexAi,
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike => {
|
||||
return Err(Error::InvalidProvider(
|
||||
provider.custom_llm_provider.to_string(),
|
||||
));
|
||||
}
|
||||
};
|
||||
Ok((provider.model.to_string(), config))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,25 +1,22 @@
|
|||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use litellm_auth::ResolvedCredential;
|
||||
use litellm_auth::{ResolvedCredential, TokenProviderHandle};
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use litellm_host::{
|
||||
event::{CallEvent, RequestContext, WireRequest},
|
||||
host::Reply,
|
||||
machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol},
|
||||
call::{CallOutput, HostedMachine, hosted_call},
|
||||
machine::HostServices,
|
||||
protocol::Protocol,
|
||||
protocol::Reply,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::error::Error;
|
||||
use litellm_llms_types::formats::ocr::LiteLLMOcrResponse;
|
||||
|
||||
use super::handler::perform_ocr_request;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput};
|
||||
|
||||
pub enum OcrOp {
|
||||
AcquireAzureAdToken(Reply<ResolvedCredential>),
|
||||
}
|
||||
|
||||
/// The caller's request as the host projects it.
|
||||
pub struct OcrProjection {
|
||||
pub struct OcrCall {
|
||||
pub request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
/// The caller passed its own Azure AD token provider, which the host keeps.
|
||||
pub caller_token: bool,
|
||||
|
|
@ -30,134 +27,45 @@ pub struct Ocr;
|
|||
impl Protocol for Ocr {
|
||||
type Response = LiteLLMOcrResponse;
|
||||
type Error = Error;
|
||||
type Projection = OcrProjection;
|
||||
type Op = OcrOp;
|
||||
type Request = OcrCall;
|
||||
type HostCall = OcrOp;
|
||||
type Chunk = std::convert::Infallible;
|
||||
type StreamHead = std::convert::Infallible;
|
||||
}
|
||||
|
||||
impl TokenProtocol for Ocr {
|
||||
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> OcrOp {
|
||||
OcrOp::AcquireAzureAdToken(reply)
|
||||
}
|
||||
}
|
||||
pub type OcrMachine = HostedMachine<Ocr>;
|
||||
|
||||
pub type OcrHost = HostChannel<Ocr>;
|
||||
pub type OcrMachine = CallMachine<Ocr>;
|
||||
|
||||
/// The OCR call as a machine: projection and token acquisition are host operations;
|
||||
/// everything else runs in Rust.
|
||||
pub fn ocr_machine(client: OcrClient) -> OcrMachine {
|
||||
CallMachine::new(move |host| Box::pin(execute(client, host)))
|
||||
}
|
||||
|
||||
async fn execute(client: OcrClient, host: OcrHost) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let OcrProjection {
|
||||
request,
|
||||
caller_token,
|
||||
} = host.project().await?;
|
||||
let request = LiteLLMOcrRequest {
|
||||
azure_ad_token_provider: caller_token
|
||||
.then(|| HostTokenProvider::handle(host.clone()))
|
||||
.or(request.azure_ad_token_provider),
|
||||
..request
|
||||
};
|
||||
let caller_document = matches!(request.document, OcrDocumentInput::Document(_));
|
||||
let request = prepare_request_document(request).await?;
|
||||
perform_ocr_request(&client, request, &host, caller_document).await
|
||||
}
|
||||
|
||||
async fn prepare_request_document(
|
||||
request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
) -> Result<ResolvedOcrRequest, Error> {
|
||||
if let OcrDocumentInput::Document(_) = &request.document {
|
||||
return request.map_document(super::document::prepare_document);
|
||||
}
|
||||
tokio::task::spawn_blocking(move || request.map_document(super::document::prepare_document))
|
||||
.await
|
||||
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
|
||||
}
|
||||
|
||||
type BeforeSend =
|
||||
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
|
||||
|
||||
/// The in-process host for a request that is already in hand: the request answers
|
||||
/// projection, and the optional observer sees and may rewrite the wire request.
|
||||
pub struct LocalOcrHost {
|
||||
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
|
||||
before_send: Option<BeforeSend>,
|
||||
observer: Option<Observer>,
|
||||
}
|
||||
|
||||
impl LocalOcrHost {
|
||||
pub fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
|
||||
Self {
|
||||
request: Mutex::new(Some(request)),
|
||||
before_send: None,
|
||||
observer: None,
|
||||
fn caller_token_provider(services: HostServices<Ocr>) -> TokenProviderHandle {
|
||||
TokenProviderHandle::from_callback(move || {
|
||||
let host_services = services.clone();
|
||||
async move {
|
||||
host_services
|
||||
.call(OcrOp::AcquireAzureAdToken)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
litellm_auth::Error::CredentialAcquisition(error.to_string().into())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_before_send(
|
||||
self,
|
||||
before_send: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
) -> Self {
|
||||
Self {
|
||||
before_send: Some(Box::new(before_send)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self {
|
||||
Self {
|
||||
observer: Some(Box::new(observer)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
impl litellm_host::host::Host<Ocr> for LocalOcrHost {
|
||||
async fn project(&self) -> Result<OcrProjection, Error> {
|
||||
self.request
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.map(|request| OcrProjection {
|
||||
request,
|
||||
caller_token: false,
|
||||
})
|
||||
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into()))
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
|
||||
match op {
|
||||
OcrOp::AcquireAzureAdToken(_) => {
|
||||
Err(Error::Auth(litellm_auth::Error::CredentialAcquisition(
|
||||
"OCR host has no Azure AD token provider".into(),
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
match &self.before_send {
|
||||
Some(before_send) => before_send(wire, context),
|
||||
None => Ok(wire),
|
||||
}
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
if let Some(observer) = &self.observer {
|
||||
observer(event);
|
||||
}
|
||||
Ok(())
|
||||
impl crate::ocr::OcrRoute {
|
||||
pub fn machine(self, request: OcrCall, observers: Option<ObservationSender>) -> OcrMachine {
|
||||
hosted_call(
|
||||
request,
|
||||
observers,
|
||||
move |projection: OcrCall, services, interceptors, observers| async move {
|
||||
let request = LiteLLMOcrRequest {
|
||||
azure_ad_token_provider: projection
|
||||
.caller_token
|
||||
.then(|| caller_token_provider(services))
|
||||
.or(projection.request.azure_ad_token_provider),
|
||||
..projection.request
|
||||
};
|
||||
self.run(request, &interceptors, observers.as_ref())
|
||||
.await
|
||||
.map(CallOutput::Complete)
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,10 +5,9 @@ use litellm_auth::{InputSource, SecretValue, TokenProviderHandle};
|
|||
use litellm_core_utils::call_arguments::CallArguments;
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
transformation::{
|
||||
OcrCredentialInputs, OcrDocument, OcrResponseFormat, OcrTransportConfig, response_format,
|
||||
},
|
||||
transformation::{OcrCredentialInputs, OcrTransportConfig, response_format},
|
||||
};
|
||||
use litellm_llms_types::formats::ocr::{OcrDocument, OcrResponseFormat};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::provider_config::{OcrConfigKind, resolve_provider_config};
|
||||
|
|
@ -222,7 +221,7 @@ mod tests {
|
|||
use super::*;
|
||||
|
||||
fn document() -> OcrDocument {
|
||||
OcrDocument::try_from(
|
||||
serde_json::from_value(
|
||||
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
|
||||
)
|
||||
.unwrap()
|
||||
|
|
|
|||
|
|
@ -1,10 +1,8 @@
|
|||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use litellm_auth::{InputSource, SecretValue};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
transformation::{OcrDocument, decode_request_value},
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{error::Error, transformation::decode_request_value};
|
||||
use litellm_llms_types::formats::ocr::OcrDocument;
|
||||
use serde::Deserialize;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,21 @@ use litellm_http::outbound::OutboundRequest;
|
|||
use litellm_llms::base_llm::auth::Authenticated;
|
||||
use serde_json::Value;
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.provider.send",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(status)
|
||||
)]
|
||||
pub(crate) async fn send(
|
||||
request: OutboundRequest,
|
||||
client: &litellm_http::Client,
|
||||
) -> Result<reqwest::Response, reqwest::Error> {
|
||||
request.send(client).await.inspect(|response| {
|
||||
tracing::Span::current().record("status", response.status().as_u16());
|
||||
})
|
||||
}
|
||||
|
||||
/// Header credentials are already in `headers`; SigV4 is applied here, over the
|
||||
/// bytes that are sent.
|
||||
pub(crate) fn outbound_request(
|
||||
|
|
|
|||
36
litellm-rust/crates/core/src/provider.rs
Normal file
36
litellm-rust/crates/core/src/provider.rs
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
pub use litellm_core_utils::get_llm_provider_logic::LlmProviders;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use crate::error::RouteError as Error;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) struct ResolvedProvider<'a> {
|
||||
pub(crate) model: &'a str,
|
||||
pub(crate) provider: LlmProviders,
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_llm_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
route: &'static str,
|
||||
) -> Result<ResolvedProvider<'a>, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(format!(
|
||||
"unable to resolve custom_llm_provider for {route} request"
|
||||
))
|
||||
})?;
|
||||
let provider = custom_llm_provider
|
||||
.parse()
|
||||
.map_err(|_| Error::InvalidProvider(custom_llm_provider.to_string()))?;
|
||||
Ok(ResolvedProvider { model, provider })
|
||||
}
|
||||
|
|
@ -1,9 +1,7 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy};
|
||||
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_http::HttpClientPool;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CoreResources {
|
||||
|
|
@ -18,21 +16,4 @@ impl CoreResources {
|
|||
auth: Arc::new(AuthServices::default()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn ocr_client(
|
||||
&self,
|
||||
config: &HttpClientConfig,
|
||||
url_policy: UrlPolicy,
|
||||
settings: OcrSettings,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<OcrClient, litellm_http::Error> {
|
||||
OcrClient::new(
|
||||
&self.pool,
|
||||
config,
|
||||
url_policy,
|
||||
self.auth.clone(),
|
||||
settings,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
116
litellm-rust/crates/core/src/responses/handler.rs
Normal file
116
litellm-rust/crates/core/src/responses/handler.rs
Normal file
|
|
@ -0,0 +1,116 @@
|
|||
use litellm_host::lifecycle::ExecutionEvent;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::StreamExt;
|
||||
use litellm_host::interceptors::{Interceptors, RawResponse, WireRequest};
|
||||
use litellm_llms::base_llm::auth::{Authenticated, resolve_auth};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesOutput, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub(super) async fn execute(
|
||||
http: &litellm_http::Client,
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderResponsesRequest,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
cache_options: Option<litellm_cache_response::CachePolicy>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let authenticated = resolve_auth(auth, request.environment, &|_| None).await?;
|
||||
let identity = litellm_host::interceptors::ProviderIdentity {
|
||||
model: request.context.model.clone(),
|
||||
provider: request.context.custom_llm_provider.clone(),
|
||||
};
|
||||
let wire = interceptors
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url: request.url,
|
||||
headers: authenticated.headers,
|
||||
body: request.body,
|
||||
},
|
||||
request.context,
|
||||
)
|
||||
.await?;
|
||||
let cache = cache.filter(|_| authenticated.signer.is_none());
|
||||
let cache_request =
|
||||
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
|
||||
crate::caching::execute_streaming::<super::route::Responses, _, _>(
|
||||
cache_request,
|
||||
cache.as_ref().map(|cache| cache.service.clone()),
|
||||
cache.as_ref().map(|cache| cache.options(cache_options)),
|
||||
interceptors,
|
||||
observers,
|
||||
|| async move {
|
||||
let stream = match wire.body.get("stream") {
|
||||
None => false,
|
||||
Some(serde_json::Value::Bool(value)) => *value,
|
||||
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
|
||||
};
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(network)?;
|
||||
let status = response.status().as_u16();
|
||||
if !response.status().is_success() {
|
||||
let body = response.text().await.map_err(network)?;
|
||||
return Err(litellm_http::transport::Error::Http {
|
||||
status,
|
||||
body: litellm_http::request::truncate_error_body(&body),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
if stream {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| {
|
||||
Some((name.to_string(), value.to_str().ok()?.to_owned()))
|
||||
})
|
||||
.collect();
|
||||
let chunks = response
|
||||
.bytes_stream()
|
||||
.map(|chunk| chunk.map_err(network))
|
||||
.boxed();
|
||||
return Ok(ResponsesOutput::Stream {
|
||||
head: ResponsesStreamHead { headers },
|
||||
chunks,
|
||||
});
|
||||
}
|
||||
let body = response.text().await.map_err(network)?;
|
||||
let raw = RawResponse { body: body.clone() };
|
||||
if let Some(observers) = observers {
|
||||
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
|
||||
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
|
||||
));
|
||||
}
|
||||
interceptors
|
||||
.after_provider_response(raw)
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
let value = serde_json::from_str(&body)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into()))?;
|
||||
request
|
||||
.config
|
||||
.transform_response_api_response(value)
|
||||
.map(ResponsesOutput::Complete)
|
||||
.map_err(Error::from)
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn network(error: reqwest::Error) -> Error {
|
||||
litellm_http::transport::Error::Network(error.to_string()).into()
|
||||
}
|
||||
|
|
@ -1,2 +1,106 @@
|
|||
pub use crate::error::RouteError as Error;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
pub mod websocket;
|
||||
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::interceptors::Interceptors;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use types::{ResponsesCall, ResponsesOutput};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
cache: Option<litellm_cache_response::ScopedCache>,
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
cache: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
|
||||
Self {
|
||||
cache: Some(cache),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
litellm_host::lifecycle::observe_call(
|
||||
observers.clone(),
|
||||
self.run(call, cache_options, interceptors, observers.as_ref()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "responses",
|
||||
model = %call.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = call.optional_params.get("stream").and_then(serde_json::Value::as_bool).unwrap_or(false),
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
cache_options: Option<litellm_cache_response::CachePolicy>,
|
||||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
self.run_provider(call, cache_options, interceptors, observers)
|
||||
.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn run_provider(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
cache_options: Option<litellm_cache_response::CachePolicy>,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.context.model, &request.context.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
&self.auth,
|
||||
request,
|
||||
self.cache.clone(),
|
||||
cache_options,
|
||||
interceptors,
|
||||
observers,
|
||||
));
|
||||
execute.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
71
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
71
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
use litellm_host::interceptors::RequestContext;
|
||||
use litellm_llms::{
|
||||
base_llm::responses::transformation::BaseResponsesApiConfig,
|
||||
openai::responses::transformation::OpenAiResponsesApiConfig,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesCall},
|
||||
};
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) async fn prepare(
|
||||
call: ResponsesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderResponsesRequest, Error> {
|
||||
let identity = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
|
||||
let provider = identity.provider.as_str();
|
||||
let model = identity.model.as_str();
|
||||
let config: &'static dyn BaseResponsesApiConfig = &OpenAiResponsesApiConfig;
|
||||
let snapshot = secrets
|
||||
.resolve(config.secret_names(call.api_key.as_deref(), call.api_base.as_deref()))
|
||||
.await?;
|
||||
let lookup = |name: &str| snapshot.get(name);
|
||||
let environment = config.validate_environment(
|
||||
litellm_http::request::string_headers("responses", call.extra_headers)?,
|
||||
call.api_key.as_deref(),
|
||||
&lookup,
|
||||
)?;
|
||||
let context = RequestContext {
|
||||
model: model.into(),
|
||||
custom_llm_provider: provider.into(),
|
||||
optional_params: Value::Object(call.optional_params.clone()),
|
||||
secret_fields: Vec::new(),
|
||||
api_key: match &environment.auth {
|
||||
litellm_llms::base_llm::auth::AuthScheme::Credential { secret, .. } => {
|
||||
Some(secret.clone())
|
||||
}
|
||||
_ => None,
|
||||
},
|
||||
};
|
||||
let body = config.transform_responses_api_request(model, call.input, call.optional_params)?;
|
||||
Ok(ProviderResponsesRequest {
|
||||
url: config.get_complete_url(call.api_base.as_deref(), &lookup),
|
||||
config,
|
||||
environment,
|
||||
body,
|
||||
context,
|
||||
timeout: call.timeout,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn resolve_provider(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<litellm_host::interceptors::ProviderIdentity, Error> {
|
||||
let provider = custom_llm_provider.unwrap_or("openai");
|
||||
if provider != "openai" {
|
||||
return Err(Error::Unsupported("native HTTP responses provider"));
|
||||
}
|
||||
let resolved = model.strip_prefix("openai/").unwrap_or(model);
|
||||
if resolved.is_empty() || resolved.contains('/') {
|
||||
return Err(Error::InvalidProvider(model.into()));
|
||||
}
|
||||
Ok(litellm_host::interceptors::ProviderIdentity {
|
||||
model: resolved.into(),
|
||||
provider: provider.into(),
|
||||
})
|
||||
}
|
||||
74
litellm-rust/crates/core/src/responses/route.rs
Normal file
74
litellm-rust/crates/core/src/responses/route.rs
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::{
|
||||
call::{HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_llms_types::formats::responses::ResponsesApiResponse;
|
||||
|
||||
use super::{
|
||||
Error, ResponsesRoute,
|
||||
types::{ResponsesCall, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub struct Responses;
|
||||
|
||||
impl Protocol for Responses {
|
||||
type Response = ResponsesApiResponse;
|
||||
type Error = Error;
|
||||
type Request = ResponsesCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = ResponsesStreamHead;
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn machine(
|
||||
self,
|
||||
call: ResponsesCall,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
) -> HostedMachine<Responses> {
|
||||
let crate::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
hosted_call(
|
||||
call,
|
||||
observers,
|
||||
move |call, _, interceptors, observers| async move {
|
||||
self.run(call, cache_options, &interceptors, observers.as_ref())
|
||||
.await
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::Cachable for Responses {
|
||||
const SURFACE: &'static str = "responses";
|
||||
|
||||
fn reusable(response: &Self::Response) -> bool {
|
||||
response
|
||||
.extra
|
||||
.get("status")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
== Some("completed")
|
||||
}
|
||||
}
|
||||
|
||||
impl crate::caching::StreamCachable for Responses {
|
||||
const TERMINAL_EVENT: &'static str = "response.completed";
|
||||
|
||||
fn replay(data: bytes::Bytes) -> Option<litellm_host::call::OutputOf<Self>> {
|
||||
Some(litellm_host::call::CallOutput::Stream {
|
||||
head: ResponsesStreamHead {
|
||||
headers: Vec::new(),
|
||||
},
|
||||
chunks: Box::pin(futures_util::stream::iter([Ok(data)])),
|
||||
})
|
||||
}
|
||||
|
||||
fn bytes(chunk: &Self::Chunk) -> &[u8] {
|
||||
chunk.as_ref()
|
||||
}
|
||||
}
|
||||
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::ValidatedEnvironment, responses::transformation::BaseResponsesApiConfig,
|
||||
};
|
||||
use litellm_llms_types::formats::responses::ResponsesApiResponse;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
|
||||
pub struct ResponsesCall {
|
||||
pub model: String,
|
||||
pub input: Value,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub struct ResponsesStreamHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub type ResponsesOutput = CallOutput<ResponsesApiResponse, ResponsesStreamHead, Bytes, Error>;
|
||||
|
||||
pub(super) struct ProviderResponsesRequest {
|
||||
pub config: &'static dyn BaseResponsesApiConfig,
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
pub context: litellm_host::interceptors::RequestContext,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
|
@ -2,7 +2,7 @@ use std::{collections::HashMap, sync::Arc, time::Duration};
|
|||
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use litellm_http::websocket::{UpstreamWebSocket, connect_upstream};
|
||||
use litellm_types::responses::streaming_websocket::ResponsesWsEventType;
|
||||
use litellm_llms_types::formats::responses::streaming_websocket::ResponsesWsEventType;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_tungstenite::tungstenite::{
|
||||
Message,
|
||||
|
|
@ -29,83 +29,121 @@ pub struct ResponsesWebSocketConnection {
|
|||
}
|
||||
|
||||
impl ResponsesWebSocketConnection {
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.connect_url",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn connect_url(
|
||||
url: &str,
|
||||
headers: &HashMap<String, String>,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<Self, Error> {
|
||||
let mut request = url.into_client_request().map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
for (name, value) in headers {
|
||||
let header_name = name
|
||||
.parse::<HeaderName>()
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
let header_value = HeaderValue::from_str(value)
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
request.headers_mut().insert(header_name, header_value);
|
||||
}
|
||||
let connect = connect_upstream(request);
|
||||
let result = match timeout {
|
||||
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket connection timed out".into(),
|
||||
))
|
||||
})?,
|
||||
None => connect.await,
|
||||
};
|
||||
let (socket, _) = result.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => {
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
})
|
||||
}
|
||||
other => Error::Transport(litellm_http::transport::Error::Network(other.to_string())),
|
||||
})?;
|
||||
Ok(Self {
|
||||
socket: Arc::new(Mutex::new(Some(socket))),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn send_text(&self, text: String) -> Result<(), Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket is closed".into(),
|
||||
)));
|
||||
};
|
||||
socket.send(Message::Text(text)).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Ok(None);
|
||||
};
|
||||
match socket.next().await {
|
||||
Some(Ok(Message::Text(text))) => Ok(Some(text)),
|
||||
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into())),
|
||||
Some(Ok(Message::Close(_))) | None => Ok(None),
|
||||
Some(Ok(_)) => Ok(None),
|
||||
Some(Err(error)) => Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
error.to_string(),
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn close(&self) -> Result<(), Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
if let Some(socket) = socket.as_mut() {
|
||||
socket.close(None).await.map_err(|error| {
|
||||
crate::diagnostic::operation("litellm.websocket.connect_url", async {
|
||||
let mut request = url.into_client_request().map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
}
|
||||
*socket = None;
|
||||
Ok(())
|
||||
for (name, value) in headers {
|
||||
let header_name = name
|
||||
.parse::<HeaderName>()
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
let header_value = HeaderValue::from_str(value)
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
request.headers_mut().insert(header_name, header_value);
|
||||
}
|
||||
let connect = connect_upstream(request);
|
||||
let result = match timeout {
|
||||
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket connection timed out".into(),
|
||||
))
|
||||
})?,
|
||||
None => connect.await,
|
||||
};
|
||||
let (socket, _) = result.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => {
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
})
|
||||
}
|
||||
other => {
|
||||
Error::Transport(litellm_http::transport::Error::Network(other.to_string()))
|
||||
}
|
||||
})?;
|
||||
Ok(Self {
|
||||
socket: Arc::new(Mutex::new(Some(socket))),
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.send_text",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn send_text(&self, text: String) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.send_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket is closed".into(),
|
||||
)));
|
||||
};
|
||||
socket.send(Message::Text(text)).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.recv_text",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.recv_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Ok(None);
|
||||
};
|
||||
match socket.next().await {
|
||||
Some(Ok(Message::Text(text))) => Ok(Some(text)),
|
||||
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into())),
|
||||
Some(Ok(Message::Close(_))) | None => Ok(None),
|
||||
Some(Ok(_)) => Ok(None),
|
||||
Some(Err(error)) => Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
error.to_string(),
|
||||
))),
|
||||
}
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.close",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn close(&self) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.close", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
if let Some(socket) = socket.as_mut() {
|
||||
socket.close(None).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
}
|
||||
*socket = None;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,4 @@
|
|||
use litellm_core::audio_transcription::{
|
||||
Error, audio_transcription, types::AudioTranscriptionRequest,
|
||||
};
|
||||
use litellm_core::audio_transcription::{Error, types::AudioTranscriptionRequest};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
|
@ -11,19 +9,11 @@ use support::*;
|
|||
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
|
||||
|
||||
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
audio_transcription(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
request,
|
||||
)
|
||||
.await
|
||||
audio_transcription_route().execute(request).await
|
||||
}
|
||||
|
||||
fn transcript_response(text: &str) -> ResponseTemplate {
|
||||
json_response(
|
||||
json!({"output": {"message": {"content": [{"text": text}]}}, "usage": {"inputTokens": 1, "outputTokens": 1}}),
|
||||
)
|
||||
json_response(json!({"output": {"message": {"content": [{"text": text}]}}}))
|
||||
}
|
||||
|
||||
fn aws_params(region: &str) -> Map<String, Value> {
|
||||
|
|
@ -70,7 +60,6 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region(
|
|||
assert_eq!(response, json!({"text": "hello"}));
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.method.as_str(), "POST");
|
||||
assert_eq!(sent.header("content-type"), Some("application/json"));
|
||||
assert_eq!(sent.url.path(), format!("/model/{MODEL}/converse"));
|
||||
let authorization = sent.header("authorization").expect("request is signed");
|
||||
assert!(
|
||||
|
|
@ -264,58 +253,28 @@ async fn an_unreadable_success_body_is_an_invalid_response(
|
|||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn injected_secrets_supply_signing_credentials_and_region(
|
||||
async fn transcription_records_route_and_resolved_provider(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
let secrets = RecordingSecrets::new([
|
||||
("AWS_ACCESS_KEY_ID", "injected-access-key"),
|
||||
("AWS_SECRET_ACCESS_KEY", "injected-secret-key"),
|
||||
("AWS_REGION_NAME", "eu-west-1"),
|
||||
("AWS_SESSION_TOKEN", "injected-session-token"),
|
||||
]);
|
||||
let response = audio_transcription(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&secrets,
|
||||
AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
optional_params: Map::new(),
|
||||
..request
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response, json!({"text": "hello"}));
|
||||
let sent = only_request(&upstream).await;
|
||||
let authorization = sent.header("authorization").unwrap();
|
||||
assert!(authorization.contains("Credential=injected-access-key/"));
|
||||
assert!(authorization.contains("/eu-west-1/bedrock/aws4_request"));
|
||||
assert_eq!(
|
||||
sent.header("x-amz-security-token"),
|
||||
Some("injected-session-token")
|
||||
);
|
||||
assert!(!sent.body_text().contains("injected-secret-key"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn secret_resolution_failure_prevents_transcription(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
let result = audio_transcription(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::failing(),
|
||||
AudioTranscriptionRequest {
|
||||
let model = request.model;
|
||||
traces
|
||||
.logger()
|
||||
.instrument(transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
},
|
||||
)
|
||||
.await;
|
||||
assert!(matches!(result, Err(Error::Secret(_))));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "audio_transcription");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(summaries[0]["resolved_model"], model);
|
||||
assert_eq!(summaries[0]["provider"], "bedrock");
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
assert!(!format!("{:?}", traces.records()).contains("secret-key"));
|
||||
}
|
||||
|
|
|
|||
1217
litellm-rust/crates/core/tests/caching.rs
Normal file
1217
litellm-rust/crates/core/tests/caching.rs
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -1,10 +1,13 @@
|
|||
use litellm_host::interceptors::RawResponse;
|
||||
use litellm_host::{
|
||||
interceptors::{ExecutionFacts, ResultSource},
|
||||
lifecycle::ExecutionEvent,
|
||||
};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_core::chat_completions::{
|
||||
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
|
||||
};
|
||||
use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
|
@ -15,13 +18,7 @@ use support::*;
|
|||
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
|
||||
|
||||
async fn complete(request: ChatCompletionsRequest<'_>) -> Result<ChatCompletionsResponse, Error> {
|
||||
chat_completions(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
request,
|
||||
)
|
||||
.await
|
||||
chat_completions_route().execute(request, &(), None).await
|
||||
}
|
||||
|
||||
fn object(value: Value) -> Map<String, Value> {
|
||||
|
|
@ -162,8 +159,6 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq
|
|||
assert_eq!(response.usage.total_tokens, 15);
|
||||
}
|
||||
|
||||
/// The provider already answered and billed these, so the host must not retry them on
|
||||
/// its own path: they surface as `InvalidResponse`, never as a pre-send decline.
|
||||
#[rstest]
|
||||
#[case::missing_usage(
|
||||
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#
|
||||
|
|
@ -215,10 +210,9 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
|
|||
);
|
||||
}
|
||||
|
||||
/// Nothing was sent, so nothing was billed and the host can still serve the request.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
||||
async fn a_connection_that_is_never_established_returns_a_connect_error(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
|
|
@ -236,9 +230,7 @@ async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
|||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
async fn a_timeout_after_sending_returns_a_network_error(request: ChatCompletionsRequest<'static>) {
|
||||
let upstream =
|
||||
upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await;
|
||||
let base = upstream.uri();
|
||||
|
|
@ -258,235 +250,147 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::accepted("anthropic/claude-sonnet-4-5", None, hi(), json!({"max_tokens": 16}), None)]
|
||||
#[case::accepted_bedrock("bedrock/anthropic.claude-sonnet-4-5", None, hi(), json!({}), None)]
|
||||
#[case::unknown_provider(
|
||||
"gpt-4o",
|
||||
Some("openai"),
|
||||
hi(),
|
||||
json!({}),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
)]
|
||||
#[case::unreadable_messages(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("hi"),
|
||||
json!({}),
|
||||
Some("unreadable message list")
|
||||
)]
|
||||
#[case::empty_messages("anthropic/claude-sonnet-4-5", None, json!([]), json!({}), Some("empty message list"))]
|
||||
#[case::streaming(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"stream": true}),
|
||||
Some("streaming")
|
||||
)]
|
||||
#[case::unrecognized_param(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"not_a_param": 1}),
|
||||
Some("unrecognized request parameter")
|
||||
)]
|
||||
#[case::opens_on_assistant_turn(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "assistant", "content": "hi"}]),
|
||||
json!({}),
|
||||
Some("conversation does not open on a user turn")
|
||||
)]
|
||||
fn decline_reason_names_why_the_core_would_not_serve_the_request(
|
||||
#[case] model: &str,
|
||||
#[case] provider: Option<&str>,
|
||||
#[case] messages: Value,
|
||||
#[case] params: Value,
|
||||
#[case] reason: Option<&str>,
|
||||
) {
|
||||
assert_eq!(
|
||||
chat_completions_decline_reason(model, provider, messages, &object(params)),
|
||||
reason
|
||||
);
|
||||
}
|
||||
|
||||
/// A request the decline check accepts must not be declined by the call itself.
|
||||
#[rstest]
|
||||
#[case::direct(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn a_declined_request_fails_the_call_before_sending(
|
||||
async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] hosted: bool,
|
||||
) {
|
||||
use litellm_core::chat_completions::route::ChatCompletions;
|
||||
use litellm_host::{call::HostedCompletion, lifecycle::CallEvent};
|
||||
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
optional_params: object(json!({"stream": true})),
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("streaming is declined");
|
||||
|
||||
assert_eq!(error, Error::Unsupported("streaming"));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::source_key(None, "source-key")]
|
||||
#[case::explicit_key(Some("explicit-key"), "explicit-key")]
|
||||
#[tokio::test]
|
||||
async fn injected_secrets_supply_credentials_and_endpoint(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] api_key: Option<&'static str>,
|
||||
#[case] expected_key: &str,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let secrets = RecordingSecrets::new([
|
||||
("ANTHROPIC_API_KEY", "source-key"),
|
||||
("ANTHROPIC_API_BASE", upstream.uri().as_str()),
|
||||
]);
|
||||
let response = chat_completions(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&secrets,
|
||||
let host = RecordingCall::<ChatCompletions>::new(
|
||||
ChatCompletionsRequest {
|
||||
api_key,
|
||||
api_base: None,
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.header("x-api-key"), Some(expected_key));
|
||||
assert_eq!(sent.url.path(), "/v1/messages");
|
||||
}
|
||||
.into(),
|
||||
);
|
||||
let response = if hosted {
|
||||
let result = litellm_host_native::in_process::run_hosted(
|
||||
chat_completions_route()
|
||||
.machine(host.request().unwrap(), Some(host.events.0.sender.clone())),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let HostedCompletion::Complete(response) = result else {
|
||||
panic!("expected a complete response")
|
||||
};
|
||||
response
|
||||
} else {
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
chat_completions_route()
|
||||
.execute(
|
||||
ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers,
|
||||
timeout: call.timeout,
|
||||
},
|
||||
&host,
|
||||
Some(host.events.0.sender.clone()),
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
};
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::accepted(false)]
|
||||
#[case::declined(true)]
|
||||
#[tokio::test]
|
||||
async fn secret_failure_stops_before_sending_and_declines_skip_resolution(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] declined: bool,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
let secrets = RecordingSecrets::failing();
|
||||
let result = chat_completions(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&secrets,
|
||||
ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
optional_params: if declined {
|
||||
object(json!({"stream": true}))
|
||||
} else {
|
||||
request.optional_params.clone()
|
||||
},
|
||||
..request
|
||||
},
|
||||
)
|
||||
.await;
|
||||
if declined {
|
||||
assert!(matches!(result, Err(Error::Unsupported(_))));
|
||||
assert!(secrets.requested().is_empty());
|
||||
} else {
|
||||
assert!(matches!(result, Err(Error::Secret(_))));
|
||||
assert!(!secrets.requested().is_empty());
|
||||
}
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::bearer(true)]
|
||||
#[case::signed(false)]
|
||||
#[tokio::test]
|
||||
async fn bedrock_chat_uses_the_injected_credential_source(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] bearer: bool,
|
||||
) {
|
||||
let upstream = upstream([json_response(json!({
|
||||
"output": {"message": {"content": [{"text": "hello"}]}},
|
||||
"usage": {"inputTokens": 1, "outputTokens": 1}
|
||||
}))])
|
||||
.await;
|
||||
let base = upstream.uri();
|
||||
let secrets = RecordingSecrets::new(
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("x-hook"),
|
||||
Some("called")
|
||||
);
|
||||
let events = host.events.0.lock().unwrap();
|
||||
assert!(matches!(
|
||||
&events[..],
|
||||
[
|
||||
("AWS_ACCESS_KEY_ID", "injected-access-key"),
|
||||
("AWS_SECRET_ACCESS_KEY", "injected-secret-key"),
|
||||
("AWS_REGION_NAME", "eu-west-1"),
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
.into_iter()
|
||||
.chain(bearer.then_some(("AWS_BEARER_TOKEN_BEDROCK", "injected-bearer"))),
|
||||
);
|
||||
let response = chat_completions(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&secrets,
|
||||
ChatCompletionsRequest {
|
||||
model: "test-model",
|
||||
custom_llm_provider: Some("bedrock"),
|
||||
api_key: None,
|
||||
api_base: Some(&base),
|
||||
optional_params: Map::new(),
|
||||
..request
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
let sent = only_request(&upstream).await;
|
||||
let authorization = sent.header("authorization").unwrap();
|
||||
if bearer {
|
||||
assert_eq!(authorization, "Bearer injected-bearer");
|
||||
} else {
|
||||
assert!(authorization.contains("Credential=injected-access-key/"));
|
||||
assert!(authorization.contains("/eu-west-1/bedrock/aws4_request"));
|
||||
}
|
||||
assert!(!sent.body_text().contains("injected-secret-key"));
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn openai_compatible_chat_resolves_its_injected_endpoint_and_key(
|
||||
async fn a_post_call_hook_failure_never_looks_safe_to_retry(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let upstream = upstream([json_response(json!({
|
||||
"id": "test-response",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
|
||||
}))]).await;
|
||||
let secrets = RecordingSecrets::new([
|
||||
("OPENAI_LIKE_API_BASE", upstream.uri().as_str()),
|
||||
("OPENAI_LIKE_API_KEY", "injected-key"),
|
||||
]);
|
||||
let response = chat_completions(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&secrets,
|
||||
ChatCompletionsRequest {
|
||||
model: "test-model",
|
||||
custom_llm_provider: Some("openai_like"),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
..request
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.url.path(), "/chat/completions");
|
||||
assert_eq!(sent.header("authorization"), Some("Bearer injected-key"));
|
||||
use litellm_host::interceptors::{Interceptors, RequestContext, WireRequest};
|
||||
struct FailingHook;
|
||||
impl Interceptors<Error> for FailingHook {
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
_: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
Ok(wire)
|
||||
}
|
||||
async fn after_provider_response(&self, _: RawResponse) -> Result<(), Error> {
|
||||
Err(Error::InvalidRequest("callback rejected".into()))
|
||||
}
|
||||
}
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
let error = chat_completions_route()
|
||||
.execute(
|
||||
ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
},
|
||||
&FailingHook,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
let Error::PostCallHook(source) = error else {
|
||||
panic!("expected retained callback error")
|
||||
};
|
||||
assert_eq!(*source, Error::InvalidRequest("callback rejected".into()));
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn completed_chat_records_route_and_resolved_provider(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
let model = request.model;
|
||||
traces
|
||||
.logger()
|
||||
.instrument(complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "chat_completions");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(summaries[0]["provider"], "anthropic");
|
||||
assert_eq!(
|
||||
summaries[0]["resolved_model"],
|
||||
only_request(&upstream).await.json()["model"]
|
||||
);
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,24 +1,27 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
use litellm_host::lifecycle::ExecutionEvent;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use litellm_core::messages::route::Messages;
|
||||
use litellm_core::messages::{MessagesCallResponse, route::Messages};
|
||||
use litellm_host::{
|
||||
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
|
||||
host::Host,
|
||||
interceptors::{ExecutionFacts, RequestContext, ResultSource, WireRequest},
|
||||
lifecycle::CallEvent,
|
||||
};
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
|
||||
/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps
|
||||
/// Projects like `LocalMessagesHost`, answers `before_provider_request` through `rewrite`, and keeps
|
||||
/// every event the driver emits.
|
||||
struct RecordingHost {
|
||||
call: LocalMessagesHost,
|
||||
rewrite: Rewrite,
|
||||
events: Mutex<Vec<CallEvent>>,
|
||||
events: super::support::Observations,
|
||||
optional_params: Mutex<Vec<Value>>,
|
||||
facts: Mutex<Vec<ExecutionFacts>>,
|
||||
reject_result: bool,
|
||||
}
|
||||
|
||||
impl RecordingHost {
|
||||
|
|
@ -26,8 +29,10 @@ impl RecordingHost {
|
|||
Self {
|
||||
call: LocalMessagesHost::new(call),
|
||||
rewrite,
|
||||
events: Mutex::new(Vec::new()),
|
||||
events: super::support::Observations::default(),
|
||||
optional_params: Mutex::new(Vec::new()),
|
||||
facts: Mutex::new(Vec::new()),
|
||||
reject_result: false,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -41,7 +46,7 @@ impl RecordingHost {
|
|||
.unwrap()
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { raw }) => {
|
||||
Some(raw.body.clone())
|
||||
}
|
||||
_ => None,
|
||||
|
|
@ -50,19 +55,40 @@ impl RecordingHost {
|
|||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for RecordingHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.project().await
|
||||
impl RecordingHost {
|
||||
pub fn request(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.request()
|
||||
}
|
||||
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, ()> {
|
||||
litellm_host_native::in_process::Host {
|
||||
services: &(),
|
||||
interceptors: self,
|
||||
stream: &(),
|
||||
observers: Some(&self.events.sender),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_host::lifecycle::CallObserver for RecordingHost {
|
||||
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
|
||||
self.events.sender.emit(event);
|
||||
}
|
||||
}
|
||||
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
|
||||
for RecordingHost
|
||||
{
|
||||
async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), Error> {
|
||||
self.facts.lock().unwrap().push(facts);
|
||||
if self.reject_result {
|
||||
return Err(Error::Unsupported("result rejected"));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
context: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
self.optional_params
|
||||
.lock()
|
||||
|
|
@ -70,15 +96,111 @@ impl Host<Messages> for RecordingHost {
|
|||
.push(context.optional_params.clone());
|
||||
(self.rewrite)(wire)
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
self.events.lock().unwrap().push(event.clone());
|
||||
async fn after_provider_response(
|
||||
&self,
|
||||
raw: litellm_host::interceptors::RawResponse,
|
||||
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::lifecycle::CallEvent::Execution(
|
||||
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
|
||||
),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::native_unary(false, false)]
|
||||
#[case::native_stream(true, false)]
|
||||
#[case::hosted_unary(false, true)]
|
||||
#[case::hosted_stream(true, true)]
|
||||
#[tokio::test]
|
||||
async fn rejected_results_are_not_delivered_or_cached(
|
||||
call: MessagesCall,
|
||||
#[case] streaming: bool,
|
||||
#[case] hosted: bool,
|
||||
) {
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{CacheScope, ResponseCache, ScopedCache};
|
||||
|
||||
let response = if streaming {
|
||||
ResponseTemplate::new(200).set_body_raw(
|
||||
"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n",
|
||||
"text/event-stream",
|
||||
)
|
||||
} else {
|
||||
message_response()
|
||||
};
|
||||
let upstream = upstream([response.clone(), response]).await;
|
||||
let route = messages_route(no_secrets()).with_cache(ScopedCache::new(
|
||||
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
|
||||
Some(100),
|
||||
Some(Duration::from_secs(60)),
|
||||
)))),
|
||||
CacheScope::Shared,
|
||||
));
|
||||
for (reject, expected_requests, cached) in [
|
||||
(true, 1, false),
|
||||
(false, 2, false),
|
||||
(true, 2, true),
|
||||
(false, 2, true),
|
||||
] {
|
||||
let request = authenticated(
|
||||
with_fields(
|
||||
MessagesCall {
|
||||
body: call.body.clone(),
|
||||
..super::call()
|
||||
},
|
||||
json!({"stream": streaming}),
|
||||
),
|
||||
upstream.uri(),
|
||||
);
|
||||
let host = RecordingHost {
|
||||
reject_result: reject,
|
||||
..RecordingHost::passthrough(request)
|
||||
};
|
||||
let result = if hosted {
|
||||
litellm_host_native::in_process::run_hosted(
|
||||
route.clone().machine(host.request().unwrap(), None),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.map(|_| ())
|
||||
} else {
|
||||
match route.execute(host.request().unwrap(), &host, None).await {
|
||||
Ok(MessagesCallResponse::Complete(_)) => Ok(()),
|
||||
Ok(MessagesCallResponse::Stream { chunks, .. }) => {
|
||||
chunks.try_collect::<Vec<_>>().await.map(|_| ())
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
};
|
||||
assert_eq!(
|
||||
result,
|
||||
if reject {
|
||||
Err(Error::Unsupported("result rejected"))
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
);
|
||||
assert_eq!(received(&upstream).await.len(), expected_requests);
|
||||
let facts = host.facts.lock().unwrap();
|
||||
assert_eq!(facts.len(), 1);
|
||||
assert_eq!(
|
||||
matches!(facts[0].source, ResultSource::Cache { .. }),
|
||||
cached
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
litellm_host_native::in_process::run_hosted(
|
||||
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
|
|
@ -118,6 +240,81 @@ async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCa
|
|||
assert_eq!(request.header("x-api-key"), Some("sk-ant"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::enable(false, json!(true), Some(true))]
|
||||
#[case::disable(true, json!(false), Some(false))]
|
||||
#[case::null(true, Value::Null, Some(false))]
|
||||
#[case::invalid(false, json!("true"), None)]
|
||||
#[tokio::test]
|
||||
async fn response_mode_follows_the_intercepted_request(
|
||||
call: MessagesCall,
|
||||
traces: TraceCapture,
|
||||
#[case] original_stream: bool,
|
||||
#[case] rewritten_stream: Value,
|
||||
#[case] expected_stream: Option<bool>,
|
||||
) {
|
||||
use futures_util::TryStreamExt;
|
||||
|
||||
let sse = "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
|
||||
let response = if expected_stream == Some(true) {
|
||||
ResponseTemplate::new(200).set_body_raw(sse, "text/event-stream")
|
||||
} else {
|
||||
message_response()
|
||||
};
|
||||
let upstream = upstream([response]).await;
|
||||
let rewrite = rewritten_stream.clone();
|
||||
let host = RecordingHost::new(
|
||||
authenticated(
|
||||
with_fields(call, json!({"stream": original_stream})),
|
||||
upstream.uri(),
|
||||
),
|
||||
Box::new(move |wire| {
|
||||
let mut body = wire.body;
|
||||
body["stream"] = rewrite.clone();
|
||||
Ok(WireRequest { body, ..wire })
|
||||
}),
|
||||
);
|
||||
let result = traces
|
||||
.logger()
|
||||
.instrument(async {
|
||||
let output = messages_route(no_secrets())
|
||||
.execute(host.request()?, &host, None)
|
||||
.await?;
|
||||
match output {
|
||||
MessagesCallResponse::Stream { chunks, .. } => {
|
||||
assert_eq!(expected_stream, Some(true));
|
||||
assert_eq!(
|
||||
chunks.try_collect::<Vec<_>>().await?.concat(),
|
||||
sse.as_bytes()
|
||||
);
|
||||
}
|
||||
MessagesCallResponse::Complete(message) => {
|
||||
assert_eq!(expected_stream, Some(false));
|
||||
assert_eq!(*message, serde_json::from_value(message_body()).unwrap());
|
||||
}
|
||||
}
|
||||
Ok::<_, Error>(())
|
||||
})
|
||||
.await;
|
||||
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
let Some(expected_stream) = expected_stream else {
|
||||
assert!(matches!(result, Err(Error::InvalidRequest(_))));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
assert_eq!(summaries[0]["outcome"], "failure");
|
||||
return;
|
||||
};
|
||||
result.unwrap();
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.json()["stream"],
|
||||
rewritten_stream
|
||||
);
|
||||
assert_eq!(host.raw_responses().len(), usize::from(!expected_stream));
|
||||
assert_eq!(summaries[0]["stream"], expected_stream);
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_before_send_failure_never_sends(call: MessagesCall) {
|
||||
|
|
@ -129,8 +326,7 @@ async fn a_before_send_failure_never_sends(call: MessagesCall) {
|
|||
|
||||
let error = run_through(&host)
|
||||
.await
|
||||
.err()
|
||||
.expect("the host failure fails the call");
|
||||
.expect_err("the host failure fails the call");
|
||||
|
||||
assert_eq!(error, Error::InvalidRequest("vetoed by the host".into()));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
|
@ -146,7 +342,7 @@ async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall)
|
|||
|
||||
let output = run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
assert!(matches!(output, MessagesOutput::Complete(_)));
|
||||
let [emitted] = <[String; 1]>::try_from(host.raw_responses())
|
||||
.unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len()));
|
||||
assert_eq!(serde_json::from_str::<Value>(&emitted).unwrap(), raw);
|
||||
|
|
@ -183,9 +379,9 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
|
|||
let host = RecordingHost::passthrough(authenticated(
|
||||
MessagesCall {
|
||||
shaping: MessagesShaping {
|
||||
capabilities: MessagesModelCapabilities {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..MessagesModelCapabilities::default()
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
drop_params: true,
|
||||
..MessagesShaping::default()
|
||||
|
|
@ -198,6 +394,6 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
|
|||
run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
.unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len()));
|
||||
assert_eq!(optional_params, json!({"max_tokens": 16}));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,14 +1,15 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use litellm_core::messages::{
|
||||
Error, MessagesCall, MessagesShaping,
|
||||
route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine},
|
||||
route::{Messages, MessagesMachine, MessagesOutput},
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
use litellm_llms_types::formats::messages::{MessagesRequest, MessagesResponse};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
};
|
||||
use rstest::fixture;
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
|
@ -32,7 +33,7 @@ fn object(value: Value) -> Map<String, Value> {
|
|||
map
|
||||
}
|
||||
|
||||
fn body(value: Value) -> AnthropicMessagesRequest {
|
||||
fn body(value: Value) -> MessagesRequest {
|
||||
serde_json::from_value(value).unwrap()
|
||||
}
|
||||
|
||||
|
|
@ -95,16 +96,17 @@ fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Ma
|
|||
)
|
||||
}
|
||||
|
||||
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
messages_machine(&support::resources(), &http_config(), secrets)
|
||||
.expect("default HTTP settings build a client")
|
||||
fn machine(secrets: Arc<dyn SecretSource>) -> impl FnOnce(MessagesCall) -> MessagesMachine {
|
||||
move |request| messages_route(secrets).machine(request, None)
|
||||
}
|
||||
|
||||
async fn run_with(
|
||||
secrets: Arc<RecordingSecrets>,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
|
||||
let host = LocalMessagesHost::new(call);
|
||||
litellm_host_native::in_process::run_hosted(machine(secrets)(host.request()?), host.runtime())
|
||||
.await
|
||||
}
|
||||
|
||||
/// Runs the route with a secret source that knows nothing, so no environment leaks in.
|
||||
|
|
@ -112,9 +114,71 @@ async fn run(call: MessagesCall) -> Result<MessagesOutput, Error> {
|
|||
run_with(Arc::new(RecordingSecrets::empty()), call).await
|
||||
}
|
||||
|
||||
async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse {
|
||||
async fn run_message(call: MessagesCall) -> MessagesResponse {
|
||||
match run(call).await.expect("messages call succeeds") {
|
||||
MessagesOutput::Message(message) => *message,
|
||||
MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"),
|
||||
MessagesOutput::Complete(message) => *message,
|
||||
MessagesOutput::StreamEnded | MessagesOutput::Detached => {
|
||||
panic!("a non-streaming call returned a stream")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct LocalMessagesHost {
|
||||
call: Mutex<Option<MessagesCall>>,
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
fn new(call: MessagesCall) -> Self {
|
||||
Self {
|
||||
call: Mutex::new(Some(call)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
pub fn request(&self) -> Result<MessagesCall, Error> {
|
||||
self.call
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
|
||||
}
|
||||
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, ()> {
|
||||
litellm_host_native::in_process::Host {
|
||||
services: &(),
|
||||
interceptors: self,
|
||||
stream: &(),
|
||||
observers: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_host::lifecycle::CallObserver for LocalMessagesHost {
|
||||
fn observe(&self, _: litellm_host::lifecycle::CallEvent) {}
|
||||
}
|
||||
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
|
||||
for LocalMessagesHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: litellm_host::interceptors::WireRequest,
|
||||
_: litellm_host::interceptors::RequestContext,
|
||||
) -> Result<
|
||||
litellm_host::interceptors::WireRequest,
|
||||
<Messages as litellm_host::protocol::Protocol>::Error,
|
||||
> {
|
||||
Ok(wire)
|
||||
}
|
||||
async fn after_provider_response(
|
||||
&self,
|
||||
raw: litellm_host::interceptors::RawResponse,
|
||||
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::lifecycle::CallEvent::Execution(
|
||||
litellm_host::lifecycle::ExecutionEvent::ProviderResponseReceived { raw },
|
||||
),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
use litellm_llms::base_llm::messages::context::{MessagesModelCapabilities, SupportedEffortTiers};
|
||||
use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet};
|
||||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use litellm_llms_types::{
|
||||
headers::{ProviderSpecificHeader, ProviderSpecificHeaders},
|
||||
providers::anthropic::{AnthropicBeta, BetaSet},
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
|
@ -87,8 +89,7 @@ async fn a_call_without_credentials_fails_before_sending(
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("a call without credentials fails");
|
||||
.expect_err("a call without credentials fails");
|
||||
|
||||
assert!(
|
||||
matches!(
|
||||
|
|
@ -158,8 +159,7 @@ async fn unsupported_providers_are_rejected_before_sending(
|
|||
..with_model(call, model)
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("unsupported provider errors");
|
||||
.expect_err("unsupported provider errors");
|
||||
|
||||
assert_eq!(error, Error::InvalidProvider(reported.into()));
|
||||
}
|
||||
|
|
@ -405,8 +405,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
|
|||
|
||||
let error = run(shaped(false))
|
||||
.await
|
||||
.err()
|
||||
.expect("an unsupported param is rejected without drop_params");
|
||||
.expect_err("an unsupported param is rejected without drop_params");
|
||||
assert!(
|
||||
matches!(&error, Error::InvalidRequest(message) if message.to_string().contains(rejected_as)),
|
||||
"{error:?}"
|
||||
|
|
@ -601,8 +600,7 @@ async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fie
|
|||
fields,
|
||||
))
|
||||
.await
|
||||
.err()
|
||||
.expect("the request is rejected");
|
||||
.expect_err("the request is rejected");
|
||||
|
||||
assert!(error.is_request(), "{error:?}");
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
|
@ -705,7 +703,7 @@ async fn provider_validation_runs_before_caller_parameter_removal(
|
|||
json!({"metadata": {"user_id": 7}}),
|
||||
))
|
||||
.await;
|
||||
let error = result.err().expect("metadata is validated before removal");
|
||||
let error = result.expect_err("metadata is validated before removal");
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
|
|
|
|||
|
|
@ -1,12 +1,97 @@
|
|||
use litellm_core::{
|
||||
Phase,
|
||||
messages::{MessagesResponse, messages, messages_body},
|
||||
use litellm_core::messages::{MessagesCallResponse, messages_body};
|
||||
use litellm_host::{
|
||||
interceptors::{ExecutionFacts, ResultSource},
|
||||
lifecycle::ExecutionEvent,
|
||||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::neither(false, false)]
|
||||
#[case::hooks_only(true, false)]
|
||||
#[case::observer_only(false, true)]
|
||||
#[case::both(true, true)]
|
||||
#[tokio::test]
|
||||
async fn calls_defer_execution_until_polled(
|
||||
call: MessagesCall,
|
||||
#[case] with_hooks: bool,
|
||||
#[case] with_observer: bool,
|
||||
) {
|
||||
use futures_util::future::BoxFuture;
|
||||
|
||||
use litellm_host::lifecycle::CallEvent;
|
||||
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let secrets = Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "test-key")]));
|
||||
let route = messages_route(secrets.clone());
|
||||
let host = RecordingCall::<Messages>::new(MessagesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
});
|
||||
let request = host.request().unwrap();
|
||||
let observer: Option<litellm_host::observation::ObservationSender> =
|
||||
with_observer.then(|| host.events.0.sender.clone());
|
||||
let future: BoxFuture<'_, Result<MessagesCallResponse, Error>> = if with_hooks {
|
||||
Box::pin(route.execute(request, &host, observer))
|
||||
} else {
|
||||
Box::pin(route.execute(request, &(), observer))
|
||||
};
|
||||
|
||||
assert!(secrets.requested().is_empty());
|
||||
assert!(host.events.0.lock().unwrap().is_empty());
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
||||
let MessagesCallResponse::Complete(response) = future.await.unwrap() else {
|
||||
panic!("expected a completed message");
|
||||
};
|
||||
assert_eq!(
|
||||
response.content,
|
||||
message_body()["content"].as_array().unwrap().as_slice()
|
||||
);
|
||||
assert!(secrets.requested().contains(&"ANTHROPIC_API_KEY".into()));
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.header("x-api-key"), Some("test-key"));
|
||||
assert_eq!(sent.header("x-hook"), with_hooks.then_some("called"));
|
||||
let events = host.events.0.lock().unwrap();
|
||||
assert!(matches!(
|
||||
(with_hooks, with_observer, events.as_slice()),
|
||||
(false, false, [])
|
||||
| (true, false, [])
|
||||
| (
|
||||
false,
|
||||
true,
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
)
|
||||
| (
|
||||
true,
|
||||
true,
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Execution(ExecutionEvent::ProviderResponseReceived { .. }),
|
||||
CallEvent::Execution(ExecutionEvent::ResultReady {
|
||||
facts: ExecutionFacts {
|
||||
source: ResultSource::Provider,
|
||||
..
|
||||
}
|
||||
}),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
)
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
|
|
@ -75,8 +160,7 @@ async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) {
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
let Error::Transport(TransportError::Http { status, body }) = error else {
|
||||
panic!("{error:?}");
|
||||
|
|
@ -97,8 +181,7 @@ async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
|
|
@ -126,8 +209,7 @@ async fn an_upstream_error_keeps_its_status_and_body(call: MessagesCall, #[case]
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
|
|
@ -154,10 +236,9 @@ async fn an_unreadable_success_body_is_an_invalid_response(
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("an unreadable body fails");
|
||||
.expect_err("an unreadable body fails");
|
||||
|
||||
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
|
||||
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -172,8 +253,7 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) {
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("the call times out");
|
||||
.expect_err("the call times out");
|
||||
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
}
|
||||
|
|
@ -188,20 +268,25 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
|
|||
..HttpSettings::default()
|
||||
};
|
||||
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&Resolution::from(&settings).config,
|
||||
&RecordingSecrets::empty(),
|
||||
let resources = support::resources();
|
||||
let response = litellm_core::messages::MessagesRoute::new(
|
||||
provider_http(&resources, &Resolution::from(&settings).config),
|
||||
resources.auth,
|
||||
no_secrets(),
|
||||
)
|
||||
.execute(
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(base),
|
||||
..call
|
||||
},
|
||||
&(),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
let MessagesResponse::Message(message) = response else {
|
||||
let MessagesCallResponse::Complete(message) = response else {
|
||||
panic!("a non-streaming request returns a message");
|
||||
};
|
||||
assert_eq!(message.id, "msg_1");
|
||||
|
|
@ -221,3 +306,151 @@ fn a_body_that_does_not_parse_is_an_invalid_request(#[case] raw: Value) {
|
|||
"{error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn message_route_summary_excludes_payload_diagnostics(
|
||||
call: MessagesCall,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let model = call.body.model.clone();
|
||||
traces
|
||||
.logger()
|
||||
.instrument(run_message(MessagesCall {
|
||||
api_key: Some("private-key-sentinel".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
}))
|
||||
.await;
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "messages");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(
|
||||
summaries[0]["resolved_model"],
|
||||
only_request(&upstream).await.json()["model"]
|
||||
);
|
||||
assert_eq!(summaries[0]["provider"], "anthropic");
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
assert!(summaries[0].get("body").is_none());
|
||||
assert!(!format!("{:?}", traces.records()).contains("private-key-sentinel"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::uncached(false, 2)]
|
||||
#[case::cached(true, 1)]
|
||||
#[tokio::test]
|
||||
async fn route_uses_injected_dependencies_and_optional_cache(
|
||||
#[case] caching: bool,
|
||||
#[case] expected_requests: usize,
|
||||
) {
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{CacheScope, ResponseCache, ScopedCache};
|
||||
use litellm_core::messages::MessagesRoute;
|
||||
|
||||
let upstream = upstream([message_response(), message_response()]).await;
|
||||
let resources = resources();
|
||||
let route = MessagesRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth.clone(),
|
||||
Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "route-key")])),
|
||||
);
|
||||
let route = if caching {
|
||||
route.with_cache(ScopedCache::new(
|
||||
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
|
||||
Some(100),
|
||||
Some(Duration::from_secs(60)),
|
||||
)))),
|
||||
CacheScope::Shared,
|
||||
))
|
||||
} else {
|
||||
route
|
||||
};
|
||||
for _ in 0..2 {
|
||||
let request = MessagesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..super::call()
|
||||
};
|
||||
let MessagesCallResponse::Complete(response) =
|
||||
route.execute(request, &(), None).await.unwrap()
|
||||
else {
|
||||
panic!("expected a completed message");
|
||||
};
|
||||
assert_eq!(
|
||||
response.content,
|
||||
message_body()["content"].as_array().unwrap().as_slice()
|
||||
);
|
||||
}
|
||||
let requests = received(&upstream).await;
|
||||
assert_eq!(requests.len(), expected_requests);
|
||||
assert_eq!(requests[0].header("x-api-key"), Some("route-key"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn cache_overrides_preserve_the_routes_isolated_scope(call: MessagesCall) {
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{CachePolicy, CacheScope, ResponseCache, ScopedCache};
|
||||
|
||||
let first_body = message_body();
|
||||
let second_body = Value::Object(
|
||||
first_body
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.map(|(key, value)| {
|
||||
(
|
||||
key.clone(),
|
||||
if key == "id" {
|
||||
json!("msg_second")
|
||||
} else {
|
||||
value.clone()
|
||||
},
|
||||
)
|
||||
})
|
||||
.collect(),
|
||||
);
|
||||
let upstream = upstream([
|
||||
json_response(first_body.clone()),
|
||||
json_response(second_body.clone()),
|
||||
])
|
||||
.await;
|
||||
let service = Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
|
||||
Some(100),
|
||||
Some(Duration::from_secs(60)),
|
||||
))));
|
||||
let first = messages_route(no_secrets()).with_cache(ScopedCache::new(
|
||||
service.clone(),
|
||||
CacheScope::Isolated("first".into()),
|
||||
));
|
||||
let second = messages_route(no_secrets()).with_cache(ScopedCache::new(
|
||||
service,
|
||||
CacheScope::Isolated("second".into()),
|
||||
));
|
||||
for (route, expected) in [
|
||||
(&first, &first_body),
|
||||
(&second, &second_body),
|
||||
(&first, &first_body),
|
||||
(&second, &second_body),
|
||||
] {
|
||||
let request = MessagesCall {
|
||||
body: call.body.clone(),
|
||||
api_key: Some("same-key".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..super::call()
|
||||
};
|
||||
let override_options = CachePolicy {
|
||||
ttl: Some(Duration::from_secs(30)),
|
||||
..CachePolicy::default()
|
||||
};
|
||||
let MessagesCallResponse::Complete(response) =
|
||||
route.execute(request, &(), override_options).await.unwrap()
|
||||
else {
|
||||
panic!("expected a completed message");
|
||||
};
|
||||
assert_eq!(response.id, expected["id"].as_str().unwrap());
|
||||
}
|
||||
assert_eq!(received(&upstream).await.len(), 2);
|
||||
}
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue