Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_vector_store_request_embedding_resolution

This commit is contained in:
mateo-berri 2026-09-02 15:24:21 -07:00
commit 61ae06d4d9
266 changed files with 14557 additions and 1694 deletions

View file

@ -112,10 +112,10 @@ commands:
node --version
npm --version
install_rust:
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.97.1) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself."
description: "Install pinned rustup (1.28.2) and Rust toolchain (1.98.0) with checksum verification. Adds ~/.cargo/bin to PATH. Run this before any `uv sync` or `uv build` of the workspace: the root package builds litellm-rust through maturin, and on an image without cargo maturin fetches an unpinned rustup and a floating toolchain by itself."
steps:
- run:
name: Install Rust (rustup 1.28.2, toolchain 1.97.1)
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
command: |
case "$(uname -m)" in
x86_64)
@ -135,7 +135,7 @@ commands:
"https://static.rust-lang.org/rustup/archive/1.28.2/${RUSTUP_TRIPLE}/rustup-init"
echo "${RUSTUP_SHA256} /tmp/rustup-init" | sha256sum -c -
chmod +x /tmp/rustup-init
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.97.1
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
rm -f /tmp/rustup-init
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.cargo/bin:$PATH"
@ -300,7 +300,7 @@ jobs:
if ($rustupActual -ne $rustupExpected) {
throw "rustup installer hash mismatch: expected $rustupExpected got $rustupActual"
}
& $rustupInit -y --profile minimal --default-toolchain stable
& $rustupInit -y --profile minimal --default-toolchain 1.98.0
if ($LASTEXITCODE -ne 0) {
exit $LASTEXITCODE
}

View file

@ -0,0 +1,68 @@
from __future__ import annotations
import subprocess
import sys
import tempfile
import zipfile
from pathlib import Path
from typing import Final
CHILD_SCRIPT: Final = """
from importlib.util import module_from_spec, spec_from_file_location
from pathlib import Path
import sys
native_path = Path(sys.argv[1])
spec = spec_from_file_location("litellm.rust_bridge._native", native_path)
if spec is None or spec.loader is None:
raise RuntimeError("cannot create native extension import specification")
module = module_from_spec(spec)
spec.loader.exec_module(module)
before = module.gil_stats()
if not isinstance(before.get("releases"), int):
raise AssertionError(f"unexpected gil_stats result: {before!r}")
try:
module._panic_for_test()
except BaseException as error:
if type(error).__name__ != "PanicException":
raise AssertionError(f"expected PanicException, got {type(error).__name__}") from error
else:
raise AssertionError("Rust panic returned without raising")
after = module.gil_stats()
if not isinstance(after.get("releases"), int):
raise AssertionError(f"native module unusable after panic: {after!r}")
"""
def main() -> int:
if len(sys.argv) != 2:
sys.stderr.write(f"usage: {Path(sys.argv[0]).name} WHEEL\n")
return 2
wheel: Final = Path(sys.argv[1])
with tempfile.TemporaryDirectory() as temporary_directory, zipfile.ZipFile(wheel) as archive:
native_members: Final = tuple(
member
for member in archive.infolist()
if member.filename.startswith("litellm/rust_bridge/_native.") and member.filename.endswith(".so")
)
if len(native_members) != 1:
sys.stderr.write(f"expected one native extension, found {len(native_members)}\n")
return 1
native_path: Final = Path(temporary_directory) / Path(native_members[0].filename).name
native_path.write_bytes(archive.read(native_members[0]))
result: Final = subprocess.run((sys.executable, "-c", CHILD_SCRIPT, str(native_path)), check=False)
if result.returncode != 0:
sys.stderr.write(f"native wheel smoke test exited with status {result.returncode}\n")
return 1
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -0,0 +1,282 @@
from __future__ import annotations
import importlib.util
import os
import re
import subprocess
import sys
import zipfile
from collections.abc import Callable, Mapping, Sequence
from itertools import product
from pathlib import Path, PurePosixPath
from types import MappingProxyType, ModuleType
from typing import Final, Protocol
EXPECTED_PYTHON_TAG: Final = "cp310"
EXPECTED_ABI_TAG: Final = "abi3"
EXPECTED_PLATFORM_TAG: Final = "linux_x86_64"
class CommandRunner(Protocol):
def __call__(
self,
command: tuple[str, ...],
*,
check: bool,
capture_output: bool,
text: bool,
) -> subprocess.CompletedProcess[str]: ...
def _run_command(
command: tuple[str, ...],
*,
check: bool,
capture_output: bool,
text: bool,
) -> subprocess.CompletedProcess[str]:
return subprocess.run(command, check=check, capture_output=capture_output, text=text)
def _dist_info_directory(member: zipfile.ZipInfo) -> str | None:
parts: Final = PurePosixPath(member.filename).parts
if not parts or not parts[0].endswith(".dist-info"):
return None
return parts[0]
def _wheel_metadata_tags(archive: zipfile.ZipFile, members: tuple[zipfile.ZipInfo, ...]) -> tuple[str, ...]:
if len(members) != 1:
return ()
lines: Final = archive.read(members[0]).splitlines()
return tuple(line.removeprefix(b"Tag:").strip().decode("ascii") for line in lines if line.startswith(b"Tag:"))
def _load_native_module(native_path: Path) -> ModuleType | None:
module_spec: Final = importlib.util.spec_from_file_location("litellm.rust_bridge._native", native_path)
if module_spec is None or module_spec.loader is None:
return None
try:
native_module: Final = importlib.util.module_from_spec(module_spec)
module_spec.loader.exec_module(native_module)
except Exception as error: # noqa: BLE001 # native module initialization can raise arbitrary exceptions
sys.stderr.write(f"native module load failed: {error}\n")
return None
return native_module
def main(
argv: Sequence[str] | None = None,
environment: Mapping[str, str] | None = None,
load_native_module: Callable[[Path], ModuleType | None] = _load_native_module,
run_command: CommandRunner = _run_command,
) -> int:
arguments: Final = tuple(sys.argv if argv is None else argv)
resolved_environment: Final = os.environ if environment is None else environment
if len(arguments) != 2:
sys.stderr.write(f"usage: {Path(arguments[0]).name} WHEEL\n")
return 2
wheel: Final = Path(arguments[1])
wheel_tags: Final = wheel.stem.rsplit("-", maxsplit=3)
if len(wheel_tags) != 4:
sys.stderr.write(f"cannot parse wheel tags from {wheel.name}\n")
return 1
wheel_identity: Final = wheel_tags[0].split("-")
if len(wheel_identity) != 2 or wheel_identity[0] != "litellm" or not wheel_identity[1]:
sys.stderr.write(f"unexpected wheel identity: {wheel_tags[0]}\n")
return 1
expected_dist_info_directory: Final = f"{wheel_tags[0]}.dist-info"
expected_dist_info_directories: Final = frozenset((expected_dist_info_directory,))
python_tag: Final = wheel_tags[1]
abi_tag: Final = wheel_tags[2]
platform_tag: Final = wheel_tags[3]
expanded_filename_tags: Final = frozenset(
"-".join(tag) for tag in product(python_tag.split("."), abi_tag.split("."), platform_tag.split("."))
)
with zipfile.ZipFile(wheel) as archive:
wheel_members: Final = archive.infolist()
dist_info_directories: Final = frozenset(
directory for member in wheel_members if (directory := _dist_info_directory(member)) is not None
)
required_dist_info_files: Final = ("METADATA", "RECORD", "WHEEL")
dist_info_file_counts: Final = MappingProxyType(
{
filename: sum(
member.filename == f"{expected_dist_info_directory}/{filename}" for member in wheel_members
)
for filename in required_dist_info_files
}
)
wheel_metadata_members: Final = tuple(
member for member in wheel_members if member.filename == f"{expected_dist_info_directory}/WHEEL"
)
wheel_metadata_tags: Final = _wheel_metadata_tags(archive, wheel_metadata_members)
native_members: Final = tuple(
member
for member in wheel_members
if member.filename.startswith("litellm/rust_bridge/_native.") and member.filename.endswith(".so")
)
if len(native_members) != 1:
sys.stderr.write(f"expected one native extension, found {len(native_members)}\n")
return 1
unexpected_members: Final = tuple(
member.filename
for member in wheel_members
if member.filename.endswith((".pdb", ".dwp", ".rlib", ".rmeta", "Cargo.toml", "Cargo.lock"))
or any(part.endswith(".dSYM") for part in PurePosixPath(member.filename).parts)
)
native_member: Final = native_members[0]
uncompressed_wheel_size: Final = sum(member.file_size for member in wheel_members)
native_path: Final = wheel.parent / "native" / Path(native_member.filename).name
native_path.parent.mkdir(parents=True, exist_ok=True)
native_path.write_bytes(archive.read(native_member))
wheel_metadata_tags_match: Final = (
len(wheel_metadata_tags) == len(expanded_filename_tags)
and frozenset(wheel_metadata_tags) == expanded_filename_tags
)
commit_sha: Final = resolved_environment.get(
"RELEASE_WHEEL_COMMIT_SHA", resolved_environment.get("GITHUB_SHA", "unknown")
)
rustc_version: Final = run_command(
("rustc", "--version"),
check=True,
capture_output=True,
text=True,
).stdout.strip()
pyproject: Final = (Path(__file__).parents[2] / "pyproject.toml").read_text()
maturin_match: Final = re.search(r'"maturin==([^";]+)', pyproject)
if maturin_match is None:
sys.stderr.write("build-system does not pin an exact Maturin version\n")
return 1
maturin_version: Final = maturin_match.group(1)
native_percentage: Final = native_member.file_size / uncompressed_wheel_size * 100
size_report: Final = "\n".join(
(
"## Native wheel build report",
"",
"| Build | Value |",
"| --- | --- |",
f"| Commit | `{commit_sha}` |",
f"| Platform | `{platform_tag}` |",
f"| Python ABI | `{python_tag}-{abi_tag}` |",
f"| Rust compiler | `{rustc_version}` |",
f"| Maturin | `{maturin_version}` |",
"| Cargo profile | `release` |",
"",
"| Artifact | Size |",
"| --- | ---: |",
f"| Compressed wheel | {wheel.stat().st_size / 1_000_000:.2f} MB |",
f"| Uncompressed wheel | {uncompressed_wheel_size / 1_000_000:.2f} MB |",
f"| Native extension | {native_member.file_size / 1_000_000:.2f} MB |",
f"| Native share | {native_percentage:.2f}% |",
"",
)
)
summary_path: Final = resolved_environment.get("GITHUB_STEP_SUMMARY")
if summary_path is None:
sys.stdout.write(size_report)
else:
Path(summary_path).write_text(size_report)
sections: Final = run_command(
("readelf", "--sections", "--wide", str(native_path)),
check=True,
capture_output=True,
text=True,
).stdout
debug_sections: Final = tuple(section for section in (".debug_", ".zdebug_") if section in sections)
debug_sections_absent: Final = not debug_sections
static_symbol_table_absent: Final = ".symtab" not in sections
dynamic_symbols: Final = run_command(
("readelf", "--dyn-syms", "--wide", str(native_path)),
check=True,
capture_output=True,
text=True,
).stdout
extension_entry_point_present: Final = "PyInit__native" in dynamic_symbols
native_module: Final = load_native_module(native_path)
native_module_loads: Final = native_module is not None
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
native_size_limit: Final = 20_000_000
native_size_within_limit: Final = native_member.file_size <= native_size_limit
validations: Final = (
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
(f"ABI tag is {EXPECTED_ABI_TAG}", abi_tag == EXPECTED_ABI_TAG),
(f"Platform tag is {EXPECTED_PLATFORM_TAG}", platform_tag == EXPECTED_PLATFORM_TAG),
("Wheel dist-info directory matches the filename", dist_info_directories == expected_dist_info_directories),
(
"Required dist-info files are present exactly once",
all(count == 1 for count in dist_info_file_counts.values()),
),
("Wheel metadata tags match the filename", wheel_metadata_tags_match),
("Debug sections are absent", debug_sections_absent),
("Static symbol table is absent", static_symbol_table_absent),
("Python extension entry point is present", extension_entry_point_present),
("Native module loads", native_module_loads),
("Production module omits the panic test hook", panic_test_hook_absent),
("Native extension does not exceed 20 MB", native_size_within_limit),
("Wheel contents are valid", not unexpected_members),
)
verified_report: Final = size_report + "\n".join(
("", "| Validation | Expected | Result |", "| --- | --- | :---: |")
+ tuple(f"| {label} | Yes | {'O' if passed else 'X'} |" for label, passed in validations)
+ ("",)
)
if summary_path is not None:
Path(summary_path).write_text(verified_report)
invalid_dist_info_files: Final = any(count != 1 for count in dist_info_file_counts.values())
validation_errors: Final = tuple(
message
for failed, message in (
(bool(debug_sections), f"{native_member.filename} contains debug sections: {', '.join(debug_sections)}"),
(not static_symbol_table_absent, f"{native_member.filename} contains a static symbol table"),
(not extension_entry_point_present, "native extension does not export PyInit__native"),
(
python_tag != EXPECTED_PYTHON_TAG,
f"unexpected Python tag: expected {EXPECTED_PYTHON_TAG}, found {python_tag}",
),
(abi_tag != EXPECTED_ABI_TAG, f"unexpected ABI tag: expected {EXPECTED_ABI_TAG}, found {abi_tag}"),
(
platform_tag != EXPECTED_PLATFORM_TAG,
f"unexpected platform tag: expected {EXPECTED_PLATFORM_TAG}, found {platform_tag}",
),
(
dist_info_directories != expected_dist_info_directories,
f"unexpected dist-info directories: expected {expected_dist_info_directory}, "
f"found {', '.join(sorted(dist_info_directories))}",
),
(invalid_dist_info_files, f"required dist-info file counts are invalid: {dist_info_file_counts}"),
(
not invalid_dist_info_files and not wheel_metadata_tags_match,
f"WHEEL tags do not match filename: expected {', '.join(sorted(expanded_filename_tags))}, "
f"found {', '.join(sorted(wheel_metadata_tags))}",
),
(
native_module is not None and not panic_test_hook_absent,
"production native module exposes _panic_for_test",
),
(
not native_size_within_limit,
f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB",
),
(bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"),
)
if failed
)
sys.stderr.write("".join(f"{message}\n" for message in validation_errors))
return 0 if all(passed for _, passed in validations) else 1
if __name__ == "__main__":
sys.exit(main())

View file

@ -80,7 +80,7 @@ jobs:
LITELLM_IMAGE: litellm-image-scan:${{ github.sha }}
run: |
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
# Scans the whole shipped artifact: OS/apk plus every language package
# baked into the image, including ones no lockfile declares (e.g. prisma's
@ -124,7 +124,7 @@ jobs:
LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }}
run: |
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
migrations-image:
name: migrations-image
@ -185,7 +185,7 @@ jobs:
LITELLM_COMPONENT_PORT: "4000"
run: |
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py -v
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
ui-image:
name: ui-image

View file

@ -0,0 +1,130 @@
name: Report LiteLLM Rust release wheel
on: # zizmor: ignore[dangerous-triggers] reporter executes no PR code and consumes no PR artifacts or outputs
workflow_run:
workflows:
- LiteLLM Rust
types:
- completed
permissions: {}
concurrency:
group: ${{ github.workflow }}-${{ github.event.workflow_run.pull_requests[0].number || github.event.workflow_run.id }}
cancel-in-progress: false
jobs:
report-release-wheel:
name: report release wheel
if: >-
github.event.workflow_run.event == 'pull_request' &&
github.event.workflow_run.path == '.github/workflows/test-rust.yml' &&
github.event.workflow_run.head_repository.full_name == github.repository &&
github.event.workflow_run.pull_requests[0].number != null
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
issues: write # PR comments use the issues API
pull-requests: read # Current-head validation rejects stale workflow runs
steps:
- name: Link release wheel report on PR
uses: actions/github-script@f28e40c7f34bde8b3046d885e986cb6290c5673b # v7.1.0
env:
COMMENT_MARKER: "<!-- litellm-release-wheel-size -->"
with:
script: |
const marker = process.env.COMMENT_MARKER;
const workflowRun = context.payload.workflow_run;
const allowedConclusions = new Set([
"action_required",
"cancelled",
"failure",
"neutral",
"skipped",
"stale",
"startup_failure",
"success",
"timed_out",
]);
if (
!allowedConclusions.has(workflowRun.conclusion) ||
workflowRun.event !== "pull_request" ||
workflowRun.path !== ".github/workflows/test-rust.yml" ||
workflowRun.head_repository?.full_name !==
`${context.repo.owner}/${context.repo.repo}` ||
workflowRun.pull_requests?.length !== 1
) {
throw new Error("unexpected source workflow");
}
const pullRequest = workflowRun.pull_requests[0];
const pullRequestNumber = pullRequest.number;
const headSha = workflowRun.head_sha;
const runId = workflowRun.id;
if (
!Number.isSafeInteger(pullRequestNumber) ||
pullRequestNumber <= 0 ||
!Number.isSafeInteger(runId) ||
runId <= 0 ||
!/^[0-9a-f]{40}$/.test(headSha) ||
pullRequest.head?.sha !== headSha
) {
throw new Error("invalid source workflow metadata");
}
const runUrl =
`${context.serverUrl}/${context.repo.owner}/${context.repo.repo}` +
`/actions/runs/${runId}`;
const result =
workflowRun.conclusion === "success"
? "successfully"
: `with \`${workflowRun.conclusion}\``;
const body = [
marker,
"## LiteLLM Rust workflow",
"",
`Workflow completed ${result} for \`${headSha}\``,
"",
`[View workflow run](${runUrl})`,
].join("\n");
const comments = await github.paginate(github.rest.issues.listComments, {
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: pullRequestNumber,
per_page: 100,
});
const existing = comments.find(
(comment) =>
comment.user?.login === "github-actions[bot]" &&
comment.body?.startsWith(marker),
);
const currentPullRequest = (
await github.rest.pulls.get({
owner: context.repo.owner,
repo: context.repo.repo,
pull_number: pullRequestNumber,
})
).data;
if (
currentPullRequest.state !== "open" ||
currentPullRequest.head.repo?.full_name !==
`${context.repo.owner}/${context.repo.repo}` ||
currentPullRequest.head.sha !== headSha
) {
core.info("source workflow no longer matches the current pull request head");
return;
}
if (existing) {
await github.rest.issues.updateComment({
owner: context.repo.owner,
repo: context.repo.repo,
comment_id: existing.id,
body,
});
} else {
await github.rest.issues.createComment({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: pullRequestNumber,
body,
});
}

View file

@ -4,6 +4,11 @@ on:
push:
paths:
- "litellm-rust/**"
- ".cargo/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- ".github/workflows/test-rust.yml"
pull_request:
branches:
@ -13,6 +18,11 @@ on:
- "litellm_**"
paths:
- "litellm-rust/**"
- ".cargo/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- ".github/workflows/test-rust.yml"
permissions:
@ -40,9 +50,7 @@ jobs:
persist-credentials: false
- name: Set up Rust
run: |
rustup toolchain install stable --profile minimal --component clippy,rustfmt
rustup default stable
run: rustup toolchain install
- name: Cache Cargo registry and target
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
@ -51,7 +59,7 @@ jobs:
~/.cargo/registry
~/.cargo/git
litellm-rust/target
key: ${{ runner.os }}-cargo-${{ hashFiles('litellm-rust/Cargo.lock') }}
key: ${{ runner.os }}-cargo-${{ hashFiles('rust-toolchain.toml', 'litellm-rust/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-cargo-
@ -69,3 +77,47 @@ jobs:
- name: Run core tests with Bedrock auth
run: cargo test -p litellm-core --features bedrock-auth --locked
release-wheel:
name: release wheel
runs-on: ubuntu-latest
timeout-minutes: 20
permissions:
contents: read
env:
CARGO_TERM_COLOR: always
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Set up Rust
run: rustup toolchain install
- name: Build release wheel
run: uv build --wheel --out-dir dist
- name: Build panic contract wheel
run: >-
uv build --wheel --out-dir panic-dist
--config-setting "maturin.build-args=--features panic-test,extension-module"
- name: Smoke-test native panic unwinding
run: python .github/scripts/smoke_test_native_wheel.py panic-dist/*.whl
- name: Verify stripped native extension
env:
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl

View file

@ -66,6 +66,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
# Copy full source tree
@ -87,6 +88,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -46,6 +46,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3.13
# Stage 2 — copy source and install the project + workspace members.
@ -57,6 +58,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -117,7 +117,7 @@
"limit": 111
},
"reportUnnecessaryComparison": {
"limit": 695
"limit": 692
},
"reportUnnecessaryContains": {
"limit": 5

View file

@ -64,6 +64,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
# Copy full source tree
@ -85,6 +86,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -70,6 +70,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13
# Copy full source tree
@ -97,6 +98,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13 \
--no-sources-package litellm-proxy-extras; \
else \
@ -106,6 +108,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--extra bedrock-realtime \
--python python3.13; \
fi

View file

@ -18,7 +18,7 @@ type: application
# This is the chart version. This version number should be incremented each time you make changes
# to the chart and its templates, including the app version.
# Versions are expected to follow Semantic Versioning (https://semver.org/)
version: 1.1.2
version: 1.1.3
# This is the version number of the application being deployed. This version number should be
# incremented each time you make changes to the application. Versions are not expected to

View file

@ -26,7 +26,7 @@ If `db.useStackgresOperator` is used (not yet implemented):
| `replicaCount` | The number of LiteLLM Proxy pods to be deployed | `1` |
| `masterkeySecretName` | The name of the Kubernetes Secret that contains the Master API Key for LiteLLM. If not specified, use the generated secret name. | N/A |
| `masterkeySecretKey` | The key within the Kubernetes Secret that contains the Master API Key for LiteLLM. If not specified, use `masterkey` as the key. | N/A |
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated. | N/A |
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated on first install and reused on upgrades. | N/A |
| `environmentSecrets` | An optional array of Secret object names. The keys and values in these secrets will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `environmentConfigMaps` | An optional array of ConfigMap object names. The keys and values in these configmaps will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `image.repository` | LiteLLM Proxy image repository | `ghcr.io/berriai/litellm` |
@ -212,6 +212,8 @@ service, the **Proxy Endpoint** should be set to `http://<RELEASE>-litellm:4000`
The **Proxy Key** is the value specified for `masterkey` or, if a `masterkey`
was not provided to the helm command line, the `masterkey` is a randomly
generated string in the `sk-...` format stored in the `<RELEASE>-litellm-masterkey` Kubernetes Secret.
The key is generated once on the first install; later `helm upgrade` runs reuse the
value already in that Secret, so upgrading never rotates the master key.
```bash
kubectl -n litellm get secret <RELEASE>-litellm-masterkey -o jsonpath="{.data.masterkey}"

View file

@ -1,9 +1,11 @@
{{- if not .Values.masterkeySecretName }}
{{ $masterkey := (.Values.masterkey | default (printf "sk-%s" (randAlphaNum 18))) }}
{{- $secretName := printf "%s-masterkey" (include "litellm.fullname" .) }}
{{- $existing := lookup "v1" "Secret" .Release.Namespace $secretName }}
{{- $masterkey := .Values.masterkey | default (dig "data" "masterkey" "" $existing | b64dec) | default (printf "sk-%s" (randAlphaNum 18)) }}
apiVersion: v1
kind: Secret
metadata:
name: {{ include "litellm.fullname" . }}-masterkey
name: {{ $secretName }}
data:
masterkey: {{ $masterkey | b64enc }}
type: Opaque

View file

@ -1,4 +1,4 @@
suite: "hpa with behavior"
suite: "hpa"
templates:
- hpa.yaml
tests:
@ -23,14 +23,44 @@ tests:
- equal: { path: spec.behavior.scaleUp.stabilizationWindowSeconds, value: 60 }
- equal: { path: spec.behavior.scaleDown.stabilizationWindowSeconds, value: 90 }
---
suite: "hpa without behavior"
templates:
- hpa.yaml
tests:
- it: "does not render behavior when not set"
set:
autoscaling.enabled: true
asserts:
- isKind: { of: HorizontalPodAutoscaler }
- isNull: { path: spec.behavior }
- it: "scales on cpu at the documented 60 percent by default"
set:
autoscaling.enabled: true
asserts:
- isKind: { of: HorizontalPodAutoscaler }
- equal: { path: "spec.metrics[0].resource.name", value: cpu }
- equal: { path: "spec.metrics[0].resource.target.type", value: Utilization }
- equal: { path: "spec.metrics[0].resource.target.averageUtilization", value: 60 }
- it: "does not scale on memory by default"
set:
autoscaling.enabled: true
asserts:
- lengthEqual: { path: spec.metrics, count: 1 }
- it: "honours an explicit cpu target override"
set:
autoscaling.enabled: true
autoscaling.targetCPUUtilizationPercentage: 75
asserts:
- equal: { path: "spec.metrics[0].resource.target.averageUtilization", value: 75 }
- it: "renders a memory metric only when a memory target is set"
set:
autoscaling.enabled: true
autoscaling.targetMemoryUtilizationPercentage: 80
asserts:
- lengthEqual: { path: spec.metrics, count: 2 }
- equal: { path: "spec.metrics[1].resource.name", value: memory }
- equal: { path: "spec.metrics[1].resource.target.averageUtilization", value: 80 }
- it: "renders no hpa when autoscaling is disabled"
asserts:
- hasDocuments: { count: 0 }

View file

@ -15,6 +15,53 @@ tests:
# Note: The masterkey is generated as "sk-<18-random-chars>" in plain text,
# but stored as base64 encoded in Kubernetes secret (requirement).
# "sk-" base64 encodes to "c2st", so we check for "^c2st" pattern.
- it: should reuse the master key already stored in the cluster instead of generating a new one on upgrade
template: secret-masterkey.yaml
set:
masterkeySecretName: ""
kubernetesProvider:
scheme:
"v1/Secret":
gvr:
version: "v1"
resource: "secrets"
namespaced: true
objects:
- kind: Secret
apiVersion: v1
metadata:
name: RELEASE-NAME-litellm-masterkey
namespace: NAMESPACE
data:
masterkey: c2stZXhpc3Rpbmcta2V5
asserts:
- equal:
path: data.masterkey
value: c2stZXhpc3Rpbmcta2V5
- it: should let an explicit masterkey value override the one already stored in the cluster
template: secret-masterkey.yaml
set:
masterkeySecretName: ""
masterkey: sk-explicit
kubernetesProvider:
scheme:
"v1/Secret":
gvr:
version: "v1"
resource: "secrets"
namespaced: true
objects:
- kind: Secret
apiVersion: v1
metadata:
name: RELEASE-NAME-litellm-masterkey
namespace: NAMESPACE
data:
masterkey: c2stZXhpc3Rpbmcta2V5
asserts:
- equal:
path: data.masterkey
value: c2stZXhwbGljaXQ=
- it: should not create a secret if masterkeySecretName is set
template: secret-masterkey.yaml
set:

View file

@ -200,7 +200,16 @@ autoscaling:
enabled: false
minReplicas: 1
maxReplicas: 100
targetCPUUtilizationPercentage: 80
# 60 is the documented recommendation. See "Recommended Machine Specifications"
# in https://docs.litellm.ai/docs/proxy/prod. A new replica clears the startupProbe
# above only after up to failureThreshold x periodSeconds = 300 seconds, so a target
# high enough to trip near saturation adds capacity minutes after it was needed.
targetCPUUtilizationPercentage: 60
# Deliberately left unset rather than given a value. The prisma query engine's
# resident memory is a high-water mark that ratchets to the pod's worst-ever write
# and is never returned, so a memory target reads the largest write a pod ever did
# rather than what it is doing now, and replicas ratchet up without scaling back in.
# Memory is a floor to provision under 'resources', not a signal to scale on.
# targetMemoryUtilizationPercentage: 80
# behavior: {}

View file

@ -14,9 +14,22 @@ then fails on a Node binary that was never written. Deleting a cache directory
that exists without a Node binary is what turns a killed bootstrap back into a
recoverable one.
Both budgets are overridable so an operator can widen them without a release:
``LITELLM_PRISMA_BOOTSTRAP_TIMEOUT`` for the toolchain install and
``LITELLM_PRISMA_COMMAND_TIMEOUT`` for every individual Prisma command.
``prisma migrate deploy`` is the other command whose runtime is not a
constant: it grows with the number of pending migrations, so a fresh database
that has to replay every migration this package ships overruns a per-command
budget sized for the short bookkeeping commands, on a laptop as much as on a
slow CI runner. The Python ``prisma`` wrapper spawns Node and the schema engine
as separate children, so killing the wrapper on timeout leaves them running:
the retry then contends with that orphan for Prisma's advisory lock and cannot
finish any sooner. Migrate deploy therefore runs under its own budget.
All three budgets are overridable so an operator can widen them without a
release: ``LITELLM_PRISMA_BOOTSTRAP_TIMEOUT`` for the toolchain install,
``LITELLM_PRISMA_MIGRATE_DEPLOY_TIMEOUT`` for ``prisma migrate deploy`` and
``LITELLM_PRISMA_COMMAND_TIMEOUT`` for every other Prisma command. The
per-command budget used to bound migrate deploy as well, so a deployment that
raised it above the deploy default keeps that larger budget for deploy unless
the deploy override says otherwise.
"""
import math
@ -36,10 +49,12 @@ except ImportError:
PRISMA_COMMAND_TIMEOUT_ENV_VAR = "LITELLM_PRISMA_COMMAND_TIMEOUT"
PRISMA_BOOTSTRAP_TIMEOUT_ENV_VAR = "LITELLM_PRISMA_BOOTSTRAP_TIMEOUT"
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR = "LITELLM_PRISMA_MIGRATE_DEPLOY_TIMEOUT"
NODEENV_CACHE_DIR_ENV_VAR = "PRISMA_NODEENV_CACHE_DIR"
DEFAULT_PRISMA_COMMAND_TIMEOUT = 60.0
DEFAULT_PRISMA_BOOTSTRAP_TIMEOUT = 600.0
DEFAULT_PRISMA_MIGRATE_DEPLOY_TIMEOUT = 600.0
BOOTSTRAP_ARG = "--version"
@ -88,6 +103,15 @@ def prisma_bootstrap_timeout() -> float:
)
def prisma_migrate_deploy_timeout() -> float:
"""Seconds one ``prisma migrate deploy`` may run for, however many migrations are pending."""
if os.getenv(PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR) is not None:
return _timeout_from_env(
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR, DEFAULT_PRISMA_MIGRATE_DEPLOY_TIMEOUT
)
return max(DEFAULT_PRISMA_MIGRATE_DEPLOY_TIMEOUT, prisma_command_timeout())
def nodeenv_cache_dir() -> Optional[Path]:
"""Where Prisma installs its private Node runtime, or None if unknowable."""
override = os.getenv(NODEENV_CACHE_DIR_ENV_VAR)

View file

@ -15,8 +15,11 @@ from litellm_proxy_extras.replica_identity import (
apply_replica_identity_full,
)
from litellm_proxy_extras.prisma_toolchain import (
PRISMA_COMMAND_TIMEOUT_ENV_VAR,
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR,
ensure_prisma_toolchain,
prisma_command_timeout,
prisma_migrate_deploy_timeout,
)
@ -698,12 +701,13 @@ class ProxyExtrasDBManager:
original_dir = os.getcwd()
os.chdir(migrations_dir)
deploy_timeout = prisma_migrate_deploy_timeout()
try:
for attempt in range(4):
try:
result = subprocess.run(
[_get_prisma_command(), "migrate", "deploy"],
timeout=prisma_command_timeout(),
timeout=deploy_timeout,
check=True,
capture_output=True,
text=True,
@ -713,8 +717,12 @@ class ProxyExtrasDBManager:
return True
except subprocess.TimeoutExpired:
logger.info(
f"prisma migrate deploy attempt {attempt + 1} timed out, retrying"
logger.warning(
"prisma migrate deploy attempt %s timed out after %ss, retrying. "
"Raise %s if this database needs longer to apply its pending migrations.",
attempt + 1,
deploy_timeout,
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR,
)
time.sleep(random.randrange(5, 15))
continue
@ -823,7 +831,8 @@ class ProxyExtrasDBManager:
"Database migration failed after 4 attempts (retry loop "
"exhausted by timeouts or repeated idempotent-recovery "
"continues). Check database connectivity, load, and "
"_prisma_migrations ledger state."
"_prisma_migrations ledger state, and raise "
f"{PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR} if the attempts timed out."
)
finally:
os.chdir(original_dir)
@ -908,7 +917,7 @@ class ProxyExtrasDBManager:
# Set migrations directory for Prisma
result = subprocess.run(
[_get_prisma_command(), "migrate", "deploy"],
timeout=prisma_command_timeout(),
timeout=prisma_migrate_deploy_timeout(),
check=True,
capture_output=True,
text=True,
@ -1126,7 +1135,11 @@ class ProxyExtrasDBManager:
)
return True
except subprocess.TimeoutExpired:
logger.info(f"Attempt {attempt + 1} timed out")
logger.warning(
"Attempt %s timed out. Raise %s if this database needs longer to apply its schema.",
attempt + 1,
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR if use_migrate else PRISMA_COMMAND_TIMEOUT_ENV_VAR,
)
time.sleep(random.randrange(5, 15))
except subprocess.CalledProcessError as e:
attempts_left = 3 - attempt

View file

@ -1,6 +1,6 @@
# AGENTS.md
litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are MODULES inside the layers.
litellm-rust has four crates. A crate is a layer or shared foundation, not a route. Routes (ocr, realtime, chat) and providers (mistral, openai) are modules inside the layers.
## Crates
@ -8,9 +8,10 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes (
|-------|------|
| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. |
| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. Owns API registration, domain wiring, and Python exception mapping. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
Dependency direction is acyclic: `litellm-python-bridge` depends on the domain layers and `litellm-python-interop`; the interop foundation depends on no LiteLLM domain crate.
## Where a route lives
@ -28,7 +29,7 @@ core/src/messages/
Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched.
Adding a crate: default to a MODULE. New crate ONLY on a real trigger — separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
Adding a crate: default to a module. A new crate requires a real trigger: separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional.

View file

@ -21,12 +21,13 @@ variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
## Crates (exactly three — see AGENTS.md)
## Crates (see AGENTS.md)
`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call.
`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and
`litellm-python-bridge` exposes it to the Python SDK. A crate is a **layer**, not
a route — add modules, not crates.
`litellm-python-bridge` exposes it to the Python SDK. `litellm-python-interop`
holds domain-neutral PyO3 primitives shared by Python-facing Rust code. A crate
is a layer or shared foundation, not a route; add modules, not crates.
## Core Boundary
@ -175,7 +176,7 @@ cd litellm-rust
cargo fmt --check
# the ai-gateway binary + server code is behind the `server` feature
cargo clippy -p litellm-ai-gateway --all-targets --features server -- -D warnings
cargo clippy -p litellm-core -p litellm-python-bridge --all-targets -- -D warnings
cargo clippy -p litellm-core -p litellm-python-interop -p litellm-python-bridge --all-targets -- -D warnings
cargo test --workspace
```

109
litellm-rust/Cargo.lock generated
View file

@ -919,6 +919,12 @@ version = "0.3.33"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b231ed28831efb4a61a08580c4bc233ec56bc009f4cd8f52da2c3cb97df0c109"
[[package]]
name = "futures-timer"
version = "3.0.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "af43fadb8a98512d547e37b4e92e0ced13e205c061b87b4623eff01d918d6968"
[[package]]
name = "futures-util"
version = "0.3.33"
@ -972,6 +978,12 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "glob"
version = "0.3.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e4eba85ea1d0a966a983acd07deee566e67395d2d96b6fb39e62b5a833f1eb0b"
[[package]]
name = "h2"
version = "0.3.27"
@ -1432,14 +1444,24 @@ dependencies = [
"criterion",
"litellm-ai-gateway",
"litellm-core",
"litellm-python-interop",
"pyo3",
"pyo3-async-runtimes",
"pythonize",
"serde",
"serde_json",
"tokio",
]
[[package]]
name = "litellm-python-interop"
version = "0.1.0"
dependencies = [
"pyo3",
"pythonize",
"rstest",
"serde",
"serde_json",
]
[[package]]
name = "litemap"
version = "0.8.2"
@ -1627,6 +1649,15 @@ dependencies = [
"zerocopy",
]
[[package]]
name = "proc-macro-crate"
version = "3.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f"
dependencies = [
"toml_edit",
]
[[package]]
name = "proc-macro2"
version = "1.0.107"
@ -1899,6 +1930,12 @@ version = "0.8.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4"
[[package]]
name = "relative-path"
version = "1.9.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ba39f3699c378cd8970968dcbff9c43159ea4cfbd88d43c00b22f2ef10a435d2"
[[package]]
name = "reqwest"
version = "0.12.28"
@ -1956,6 +1993,35 @@ dependencies = [
"windows-sys 0.52.0",
]
[[package]]
name = "rstest"
version = "0.26.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f5a3193c063baaa2a95a33f03035c8a72b83d97a54916055ba22d35ed3839d49"
dependencies = [
"futures-timer",
"futures-util",
"rstest_macros",
]
[[package]]
name = "rstest_macros"
version = "0.26.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9c845311f0ff7951c5506121a9ad75aec44d083c31583b2ea5a30bcb0b0abba0"
dependencies = [
"cfg-if",
"glob",
"proc-macro-crate",
"proc-macro2",
"quote",
"regex",
"relative-path",
"rustc_version",
"syn 2.0.119",
"unicode-ident",
]
[[package]]
name = "rustc-hash"
version = "2.1.3"
@ -2488,6 +2554,36 @@ dependencies = [
"tokio",
]
[[package]]
name = "toml_datetime"
version = "1.1.1+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3165f65f62e28e0115a00b2ebdd37eb6f3b641855f9d636d3cd4103767159ad7"
dependencies = [
"serde_core",
]
[[package]]
name = "toml_edit"
version = "0.25.13+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b"
dependencies = [
"indexmap",
"toml_datetime",
"toml_parser",
"winnow",
]
[[package]]
name = "toml_parser"
version = "1.1.3+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1d38ac1cf9b95face32296c0a3ede1fdc270627c9d9c02a7274dd6d960dc4d56"
dependencies = [
"winnow",
]
[[package]]
name = "tower"
version = "0.5.3"
@ -2903,6 +2999,15 @@ version = "0.52.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
[[package]]
name = "winnow"
version = "1.0.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "23b97319f7b8343df12cc98938e5c3eb436064524c8d2b4e30a1d3a36eecdf81"
dependencies = [
"memchr",
]
[[package]]
name = "writeable"
version = "0.6.3"

View file

@ -2,6 +2,7 @@
members = [
"crates/core",
"crates/ai-gateway",
"crates/python-interop",
"crates/python-bridge",
]
resolver = "2"
@ -15,12 +16,14 @@ repository = "https://github.com/BerriAI/litellm"
[workspace.dependencies]
litellm-core = { path = "crates/core" }
litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false }
litellm-python-interop = { path = "crates/python-interop" }
axum = "0.7"
pyo3 = "0.29.0"
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
pythonize = "0.29.0"
rand = "0.8"
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls", "http2", "stream"] }
rstest = "0.26.1"
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
sha2 = "0.10"

View file

@ -26,9 +26,10 @@ coverage and production evidence.
|-------|------|
| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. |
| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. Owns API registration, domain wiring, and Python exception mapping. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
Dependency direction is acyclic: `litellm-python-bridge` depends on the domain layers and `litellm-python-interop`; the interop foundation depends on no LiteLLM domain crate.
## Layout
@ -38,7 +39,8 @@ crates/
src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client
src/providers/anthropic/messages/transformation.rs
ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints.
python-bridge/ PyO3 bridge for Python LiteLLM.
python-interop/ Domain-neutral PyO3 conversion and GIL primitives.
python-bridge/ PyO3 API adapter for Python LiteLLM.
```
The folder shape follows the Python provider tree:

View file

@ -54,6 +54,6 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages`
cd litellm-rust
cargo fmt --check
cargo clippy -p litellm-ai-gateway --all-targets --features server -- -D warnings
cargo clippy -p litellm-core -p litellm-python-bridge --all-targets -- -D warnings
cargo clippy -p litellm-core -p litellm-python-interop -p litellm-python-bridge --all-targets -- -D warnings
cargo test --workspace
```

View file

@ -6,15 +6,16 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame.
## Crates
`litellm-rust` is exactly three crates (a crate is a **layer**, not a route):
`litellm-rust` has four crates. A crate is a layer or shared foundation, not a route:
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust — per-route entrypoints (`messages::messages()`) that resolve the provider, transform, and make the call; plus types, provider transforms, and the router. |
| litellm-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
| litellm-python-interop | Domain-neutral PyO3 foundation for GIL handling and typed Python/Serde conversion. |
| litellm-python-bridge | PyO3 cdylib exposing LiteLLM Rust APIs to the Python SDK. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
Dependency direction is acyclic: `litellm-python-bridge` depends on the domain layers and `litellm-python-interop`; the interop foundation depends on no LiteLLM domain crate.
- **Client endpoint:** `wss://<host>/v1/realtime?model=<model>` (WebSocket)
- **Auth:** `Authorization: Bearer $LITELLM_MASTER_KEY` (fails closed if unset)

View file

@ -1,10 +1,8 @@
use std::collections::BTreeMap;
use litellm_core::CoreResult;
use litellm_core::audio_transcription::transformation::AudioTranscriptionProviderConfig;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG;
use serde_json::{Map, Value};
use std::collections::BTreeMap;
pub(super) fn audio_transcription_provider_config(
provider: &str,
@ -17,7 +15,7 @@ pub(super) fn audio_transcription_provider_config(
pub(super) fn string_headers(
headers: Option<Map<String, Value>>,
) -> CoreResult<BTreeMap<String, String>> {
) -> Result<BTreeMap<String, String>, Error> {
headers
.unwrap_or_default()
.into_iter()
@ -26,7 +24,7 @@ pub(super) fn string_headers(
.as_str()
.map(|value| (key.clone(), value.to_string()))
.ok_or_else(|| {
CoreError::InvalidRequest(format!(
Error::InvalidRequest(format!(
"audio transcription extra_headers.{key} must be a string"
))
})

View file

@ -1,11 +1,9 @@
use std::time::SystemTime;
use litellm_core::CoreResult;
use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::providers::bedrock::audio_transcription::aws_auth_config;
use litellm_core::providers::bedrock::aws_base::{resolve_credentials, sign_bedrock_post};
use serde_json::Value;
use std::time::SystemTime;
use super::common_utils::truncate_error_body;
use super::types::ProviderAudioTranscriptionRequest;
@ -13,10 +11,9 @@ use crate::client::http_client;
pub(crate) async fn execute_audio_transcription_provider_call(
request: ProviderAudioTranscriptionRequest,
) -> CoreResult<Value> {
let body = serde_json::to_vec(&request.body).map_err(|error| {
CoreError::InvalidRequest(format!("invalid audio request body: {error}"))
})?;
) -> Result<Value, Error> {
let body = serde_json::to_vec(&request.body)
.map_err(|error| Error::InvalidRequest(format!("invalid audio request body: {error}")))?;
let mut request_builder = http_client().post(&request.url).body(body.clone());
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
@ -27,21 +24,20 @@ pub(crate) async fn execute_audio_transcription_provider_call(
let response = request_builder
.send()
.await
.map_err(|error| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
let status = response.status();
let text = response
.text()
.await
.map_err(|error| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response_json: Value = serde_json::from_str(&text).map_err(|error| {
CoreError::InvalidResponse(format!("invalid audio response JSON: {error}"))
})?;
let response_json: Value = serde_json::from_str(&text)
.map_err(|error| Error::InvalidResponse(format!("invalid audio response JSON: {error}")))?;
Ok(request
.config
.transform_transcription_response(&request.model, response_json)?
@ -51,14 +47,13 @@ pub(crate) async fn execute_audio_transcription_provider_call(
pub(crate) async fn sign_request(
request: &ProviderAudioTranscriptionRequest,
optional_params: &serde_json::Map<String, Value>,
) -> CoreResult<ProviderAudioTranscriptionRequest> {
) -> Result<ProviderAudioTranscriptionRequest, Error> {
let env_lookup = environment_lookup;
let auth = request
.config
.auth_strategy(&request.model, optional_params, &env_lookup)?;
let body = serde_json::to_vec(&request.body).map_err(|error| {
CoreError::InvalidRequest(format!("invalid audio request body: {error}"))
})?;
let body = serde_json::to_vec(&request.body)
.map_err(|error| Error::InvalidRequest(format!("invalid audio request body: {error}")))?;
let mut headers = super::common_utils::string_headers(None)?;
headers.insert("Content-Type".to_string(), "application/json".to_string());
headers.extend(request.upstream_headers.iter().cloned());

View file

@ -1,11 +1,9 @@
use std::future::Future;
use std::pin::Pin;
use litellm_core::CoreResult;
use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use serde_json::{Map, Value, json};
use std::future::Future;
use std::pin::Pin;
use super::common_utils::{audio_transcription_provider_config, has_header, string_headers};
use super::handler::sign_request;
@ -26,7 +24,7 @@ pub(crate) struct AudioTranscriptionLifecycleHooks {
request_metadata: RequestMetadata,
}
type AudioFuture<'a, T> = Pin<Box<dyn Future<Output = CoreResult<T>> + Send + 'a>>;
type AudioFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
type AudioLogFuture<'a> = Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
impl AudioTranscriptionLifecycleHooks {
@ -45,7 +43,7 @@ impl AudioTranscriptionLifecycleHooks {
async fn run_pre_call_guardrails(
&self,
request: PreparedAudioTranscriptionRequest,
) -> CoreResult<PreparedAudioTranscriptionRequest> {
) -> Result<PreparedAudioTranscriptionRequest, Error> {
if self.guardrail_runner.is_empty() {
return Ok(request);
}
@ -63,17 +61,17 @@ impl AudioTranscriptionLifecycleHooks {
.await
.map_err(guardrail_error_to_core_error)?;
let Value::Object(mut data) = guardrail_request.data else {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"audio transcription pre_call guardrail must return an object".to_string(),
));
};
let audio = data.remove("audio").ok_or_else(|| {
CoreError::InvalidRequest("audio transcription guardrail removed audio".to_string())
Error::InvalidRequest("audio transcription guardrail removed audio".to_string())
})?;
let optional_params = match data.remove("optional_params") {
Some(Value::Object(value)) => value,
Some(_) => {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"audio transcription optional_params must be an object".to_string(),
));
}
@ -89,9 +87,9 @@ impl AudioTranscriptionLifecycleHooks {
async fn prepare_provider_request(
&self,
request: PreparedAudioTranscriptionRequest,
) -> CoreResult<ProviderAudioTranscriptionRequest> {
) -> Result<ProviderAudioTranscriptionRequest, Error> {
let config = audio_transcription_provider_config(&request.custom_llm_provider)
.ok_or_else(|| CoreError::InvalidProvider(request.custom_llm_provider.clone()))?;
.ok_or_else(|| Error::InvalidProvider(request.custom_llm_provider.clone()))?;
let env_lookup = super::handler::environment_lookup;
let headers = string_headers(request.extra_headers)?;
let url = config.complete_url(
@ -135,7 +133,7 @@ impl AudioTranscriptionLifecycleHooks {
async fn run_during_call_guardrails(
&self,
request: ProviderAudioTranscriptionRequest,
) -> CoreResult<ProviderAudioTranscriptionRequest> {
) -> Result<ProviderAudioTranscriptionRequest, Error> {
if self.guardrail_runner.is_empty() {
return Ok(request);
}
@ -153,12 +151,12 @@ impl AudioTranscriptionLifecycleHooks {
.await
.map_err(guardrail_error_to_core_error)?;
let Value::Object(mut data) = guardrail_request.data else {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"audio transcription during_call guardrail must return an object".to_string(),
));
};
let body = data.remove("body").ok_or_else(|| {
CoreError::InvalidRequest("audio transcription guardrail removed body".to_string())
Error::InvalidRequest("audio transcription guardrail removed body".to_string())
})?;
Ok(ProviderAudioTranscriptionRequest { body, ..request })
}
@ -241,7 +239,7 @@ impl CallLifecycleHooks<PreparedAudioTranscriptionRequest, ProviderAudioTranscri
fn async_log_failure_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
error: &'a CoreError,
error: &'a Error,
timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
Box::pin(async move {
@ -281,22 +279,22 @@ fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
}
}
fn guardrail_error_to_core_error(error: GuardrailError) -> CoreError {
CoreError::InvalidRequest(format!("{}: {}", error.kind, error.message))
fn guardrail_error_to_core_error(error: GuardrailError) -> Error {
Error::InvalidRequest(format!("{}: {}", error.kind, error.message))
}
fn core_error_kind(error: &CoreError) -> &'static str {
fn core_error_kind(error: &Error) -> &'static str {
match error {
CoreError::Auth(_) => "AuthError",
CoreError::InvalidProvider(_) => "InvalidProvider",
CoreError::InvalidRequest(_) => "InvalidRequest",
CoreError::InvalidType { .. } => "InvalidType",
CoreError::MissingField(_) => "MissingField",
CoreError::Http { .. } => "HttpError",
CoreError::InvalidResponse(_) => "InvalidResponse",
CoreError::Network(_) => "NetworkError",
CoreError::Connect(_) => "ConnectError",
CoreError::Routing(_) => "RoutingError",
CoreError::Unsupported(_) => "UnsupportedRequest",
Error::Auth(_) => "AuthError",
Error::InvalidProvider(_) => "InvalidProvider",
Error::InvalidRequest(_) => "InvalidRequest",
Error::InvalidType { .. } => "InvalidType",
Error::MissingField(_) => "MissingField",
Error::Http { .. } => "HttpError",
Error::InvalidResponse(_) => "InvalidResponse",
Error::Network(_) => "NetworkError",
Error::Connect(_) => "ConnectError",
Error::Routing(_) => "RoutingError",
Error::Unsupported(_) => "UnsupportedRequest",
}
}

View file

@ -1,4 +1,4 @@
use litellm_core::CoreResult;
use litellm_core::Error;
use litellm_core::call_lifecycle::CallLifecycle;
use serde_json::Value;
@ -13,7 +13,7 @@ pub use types::AudioTranscriptionRequest;
use handler::execute_audio_transcription_provider_call;
use prepare::{PreparedAudioTranscriptionCall, prepare_audio_transcription_call};
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> CoreResult<Value> {
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
let PreparedAudioTranscriptionCall { request, hooks } =
prepare_audio_transcription_call(request);
CallLifecycle::default()

View file

@ -15,8 +15,7 @@ use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::realtime::transformation::RealtimeProviderConfig;
use litellm_core::realtime::types::RealtimeEvent;
use tokio::net::TcpStream;
@ -48,7 +47,7 @@ pub(crate) type UpstreamRx = SplitStream<UpstreamWs>;
/// Resolve the OpenAI API key from the explicit param or the environment.
///
/// Blank/whitespace values are treated as absent (guard at resolution time).
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> CoreResult<String> {
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result<String, Error> {
api_key
.map(str::trim)
.filter(|key| !key.is_empty())
@ -58,7 +57,7 @@ pub(crate) fn resolve_api_key(api_key: Option<&str>) -> CoreResult<String> {
.ok()
.filter(|key| !key.trim().is_empty())
})
.ok_or_else(|| CoreError::Auth(MISSING_KEY_MESSAGE.to_string()))
.ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string()))
}
/// Open the upstream WebSocket to OpenAI for `(model, api_key, api_base)`.
@ -70,24 +69,24 @@ pub(crate) async fn dial_upstream(
model: &str,
api_key: &str,
api_base: Option<&str>,
) -> CoreResult<UpstreamWs> {
) -> Result<UpstreamWs, Error> {
let url = OPENAI_REALTIME_CONFIG.complete_url(api_base, model);
let mut request = url
.as_str()
.into_client_request()
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
// GA realtime: only Authorization. The legacy OpenAI-Beta header triggers
// beta_api_shape_disabled, so we do not send it.
request.headers_mut().insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {api_key}"))
.map_err(|err| CoreError::Auth(err.to_string()))?,
.map_err(|err| Error::Auth(err.to_string()))?,
);
let (upstream, _response) = connect_async(request)
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
Ok(upstream)
}
@ -96,22 +95,22 @@ pub(crate) async fn dial_upstream(
/// Used by the pool to pre-read OpenAI's unprompted `session.created`. Returns an
/// error on a non-text frame, a closed socket, or undecodable JSON so the pool can
/// discard a misbehaving socket rather than warm it.
pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> CoreResult<RealtimeEvent> {
pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> Result<RealtimeEvent, Error> {
loop {
let message = upstream_rx
.next()
.await
.ok_or_else(|| CoreError::Network("upstream closed before first event".to_string()))?
.map_err(|err| CoreError::Network(err.to_string()))?;
.ok_or_else(|| Error::Network("upstream closed before first event".to_string()))?
.map_err(|err| Error::Network(err.to_string()))?;
match message {
Message::Text(text) => {
return serde_json::from_str(&text)
.map_err(|err| CoreError::InvalidResponse(err.to_string()));
.map_err(|err| Error::InvalidResponse(err.to_string()));
}
// Ignore protocol frames (ping/pong) while waiting for the first event.
Message::Ping(_) | Message::Pong(_) => continue,
Message::Close(_) => {
return Err(CoreError::Network(
return Err(Error::Network(
"upstream closed before first event".to_string(),
));
}
@ -139,7 +138,7 @@ pub(crate) async fn splice<In, Out>(
mut observe: impl FnMut(&RealtimeEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
@ -154,7 +153,7 @@ where
client_out
.send(outbound)
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
}
}
@ -175,26 +174,26 @@ where
// inflate its own spend log. Logging observes upstream events only.
for outbound in config.transform_realtime_request(&event, model)?.events {
let payload = serde_json::to_string(&outbound)
.map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
.map_err(|err| Error::InvalidResponse(err.to_string()))?;
upstream_tx
.send(Message::Text(payload))
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
}
}
// upstream -> client
upstream_message = upstream_rx.next() => {
let Some(message) = upstream_message else { break }; // upstream closed
match message.map_err(|err| CoreError::Network(err.to_string()))? {
match message.map_err(|err| Error::Network(err.to_string()))? {
Message::Text(text) => {
let event: RealtimeEvent = serde_json::from_str(&text)
.map_err(|err| CoreError::InvalidResponse(err.to_string()))?;
.map_err(|err| Error::InvalidResponse(err.to_string()))?;
observe(&event);
for outbound in config.transform_realtime_response(&event, model)?.events {
client_out
.send(outbound)
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
}
}
Message::Close(_) => break,
@ -225,7 +224,7 @@ pub async fn realtime<In, Out>(
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
@ -258,7 +257,7 @@ pub async fn realtime_warm<In, Out>(
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,

View file

@ -28,7 +28,7 @@ use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use futures_util::StreamExt;
use litellm_core::CoreResult;
use litellm_core::Error;
use litellm_core::realtime::types::RealtimeEvent;
use crate::io::realtime::{
@ -438,7 +438,7 @@ impl RealtimePool {
///
/// `key.api_key` is already resolved (non-blank). The first frame OpenAI sends
/// unprompted is `session.created`; we buffer exactly that and read nothing more.
async fn warm_one(key: &UpstreamKey) -> CoreResult<WarmConnection> {
async fn warm_one(key: &UpstreamKey) -> Result<WarmConnection, Error> {
let upstream: UpstreamWs =
dial_upstream(&key.model, &key.api_key, key.api_base.as_deref()).await?;
let (tx, mut rx) = upstream.split();

View file

@ -4,10 +4,10 @@ use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use litellm_core::Error;
use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG;
use litellm_core::responses::types::ResponsesWsEvent;
use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig;
use litellm_core::{CoreError, CoreResult};
use tokio::net::TcpStream;
use tokio::sync::Mutex;
use tokio_tungstenite::tungstenite::Message;
@ -37,51 +37,49 @@ impl ResponsesWebSocketConnection {
url: &str,
headers: &HashMap<String, String>,
timeout: Option<Duration>,
) -> CoreResult<Self> {
) -> Result<Self, Error> {
let mut request = url
.into_client_request()
.map_err(|error| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
for (name, value) in headers {
let header_name = name
.parse::<HeaderName>()
.map_err(|error| CoreError::InvalidRequest(error.to_string()))?;
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
let header_value = HeaderValue::from_str(value)
.map_err(|error| CoreError::InvalidRequest(error.to_string()))?;
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
request.headers_mut().insert(header_name, header_value);
}
let connect = connect_async(request);
let result = match timeout {
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
CoreError::Network("Responses WebSocket connection timed out".to_string())
Error::Network("Responses WebSocket connection timed out".to_string())
})?,
None => connect.await,
};
let (socket, _) = result.map_err(|error| match error {
tokio_tungstenite::tungstenite::Error::Http(response) => CoreError::Http {
tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http {
status: response.status().as_u16(),
body: String::new(),
},
other => CoreError::Network(other.to_string()),
other => Error::Network(other.to_string()),
})?;
Ok(Self {
socket: Arc::new(Mutex::new(Some(socket))),
})
}
pub async fn send_text(&self, text: String) -> CoreResult<()> {
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(CoreError::Network(
"Responses WebSocket is closed".to_string(),
));
return Err(Error::Network("Responses WebSocket is closed".to_string()));
};
socket
.send(Message::Text(text))
.await
.map_err(|error| CoreError::Network(error.to_string()))
.map_err(|error| Error::Network(error.to_string()))
}
pub async fn recv_text(&self) -> CoreResult<Option<String>> {
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
let mut socket_guard = self.socket.lock().await;
let Some(socket) = socket_guard.as_mut() else {
return Ok(None);
@ -90,27 +88,27 @@ impl ResponsesWebSocketConnection {
Some(Ok(Message::Text(text))) => Ok(Some(text)),
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
.map(Some)
.map_err(|error| CoreError::InvalidResponse(error.to_string())),
.map_err(|error| Error::InvalidResponse(error.to_string())),
Some(Ok(Message::Close(_))) | None => Ok(None),
Some(Ok(_)) => Ok(None),
Some(Err(error)) => Err(CoreError::Network(error.to_string())),
Some(Err(error)) => Err(Error::Network(error.to_string())),
}
}
pub async fn close(&self) -> CoreResult<()> {
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| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
}
*socket = None;
Ok(())
}
}
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> CoreResult<String> {
pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result<String, Error> {
api_key
.map(str::trim)
.filter(|value| !value.is_empty())
@ -120,38 +118,38 @@ pub(crate) fn resolve_api_key(api_key: Option<&str>) -> CoreResult<String> {
.ok()
.filter(|value| !value.trim().is_empty())
})
.ok_or_else(|| CoreError::Auth(MISSING_KEY_MESSAGE.to_string()))
.ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string()))
}
async fn dial_upstream(
model: &str,
api_key: &str,
api_base: Option<&str>,
) -> CoreResult<ResponsesUpstreamWs> {
) -> Result<ResponsesUpstreamWs, Error> {
let url = OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, model);
let mut request = url
.as_str()
.into_client_request()
.map_err(|error| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
request.headers_mut().insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {api_key}"))
.map_err(|error| CoreError::Auth(error.to_string()))?,
.map_err(|error| Error::Auth(error.to_string()))?,
);
let result = tokio::time::timeout(
Duration::from_secs(DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS),
connect_async(request),
)
.await
.map_err(|_| CoreError::Network("Responses WebSocket connection timed out".to_string()))?;
.map_err(|_| Error::Network("Responses WebSocket connection timed out".to_string()))?;
result
.map(|(socket, _)| socket)
.map_err(|error| match error {
tokio_tungstenite::tungstenite::Error::Http(response) => CoreError::Http {
tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http {
status: response.status().as_u16(),
body: String::new(),
},
other => CoreError::Network(other.to_string()),
other => Error::Network(other.to_string()),
})
}
@ -166,7 +164,7 @@ impl ResponsesWebSocketStreaming {
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
@ -193,7 +191,7 @@ pub(crate) async fn splice<In, Out>(
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
@ -210,18 +208,18 @@ where
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| CoreError::InvalidResponse(error.to_string()))?;
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx.send(Message::Text(payload))
.await
.map_err(|error| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
}
}
message = upstream_rx.next() => {
let Some(message) = message else { break };
match message.map_err(|error| CoreError::Network(error.to_string()))? {
match message.map_err(|error| Error::Network(error.to_string()))? {
Message::Text(text) => {
let event = serde_json::from_str::<ResponsesWsEvent>(&text)
.map_err(|error| CoreError::InvalidResponse(error.to_string()))?;
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
observe(&event);
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_response(&event, model)?
@ -229,7 +227,7 @@ where
{
client_out.send(outbound)
.await
.map_err(|error| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
}
}
Message::Close(_) => break,
@ -252,7 +250,7 @@ pub async fn async_responses_websocket<In, Out>(
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
@ -267,11 +265,11 @@ where
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| CoreError::InvalidResponse(error.to_string()))?;
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx
.send(Message::Text(payload))
.await
.map_err(|error| CoreError::Network(error.to_string()))?;
.map_err(|error| Error::Network(error.to_string()))?;
}
}
ResponsesWebSocketStreaming::bidirectional_forward(
@ -296,7 +294,7 @@ pub async fn responses_ws<In, Out>(
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
@ -514,7 +512,7 @@ mod tests {
)
.await
.expect_err("status error");
assert!(matches!(error, CoreError::Http { status: 401, .. }));
assert!(matches!(error, Error::Http { status: 401, .. }));
server.await.expect("server task");
}
@ -543,7 +541,7 @@ mod tests {
)
.await
.expect_err("status error");
assert!(matches!(error, CoreError::Http { status: 500, .. }));
assert!(matches!(error, Error::Http { status: 500, .. }));
server.await.expect("server task");
}
}

View file

@ -3,8 +3,7 @@ use std::time::{Duration, Instant};
use base64::Engine;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::ocr::transformation::OcrProviderConfig;
use reqwest::Url;
use serde_json::{Map, Value};
@ -56,7 +55,7 @@ fn is_azure_document_intelligence_model(model: &str) -> bool {
pub(super) fn string_headers(
extra_headers: Option<Map<String, Value>>,
) -> CoreResult<Vec<(String, String)>> {
) -> Result<Vec<(String, String)>, Error> {
extra_headers
.unwrap_or_default()
.into_iter()
@ -65,7 +64,7 @@ pub(super) fn string_headers(
.as_str()
.map(|value| (key.clone(), value.to_string()))
.ok_or_else(|| {
CoreError::InvalidRequest(format!(
Error::InvalidRequest(format!(
"OCR extra_headers.{key} must be a string, got {}",
litellm_core::error::json_type_name(&value)
))
@ -80,7 +79,7 @@ pub(super) fn has_header(headers: &[(String, String)], name: &str) -> bool {
.any(|(key, _)| key.eq_ignore_ascii_case(name))
}
fn document_url_field(document: &Value) -> CoreResult<Option<(&str, &str)>> {
fn document_url_field(document: &Value) -> Result<Option<(&str, &str)>, Error> {
let Some(object) = document.as_object() else {
return Ok(None);
};
@ -138,13 +137,13 @@ fn is_blocked_ip(ip: IpAddr) -> bool {
}
}
fn blocked_url_error(url: &Url) -> CoreError {
CoreError::InvalidRequest(format!(
fn blocked_url_error(url: &Url) -> Error {
Error::InvalidRequest(format!(
"OCR document URL rejected by SSRF protection: {url}"
))
}
async fn validate_safe_fetch_url(url: &Url) -> CoreResult<()> {
async fn validate_safe_fetch_url(url: &Url) -> Result<(), Error> {
if !matches!(url.scheme(), "http" | "https") {
return Err(blocked_url_error(url));
}
@ -162,7 +161,7 @@ async fn validate_safe_fetch_url(url: &Url) -> CoreResult<()> {
.ok_or_else(|| blocked_url_error(url))?;
let addresses = tokio::net::lookup_host((host, port))
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
let mut saw_address = false;
for address in addresses {
saw_address = true;
@ -176,25 +175,25 @@ async fn validate_safe_fetch_url(url: &Url) -> CoreResult<()> {
Ok(())
}
fn redirect_location(response: &reqwest::Response, url: &Url) -> CoreResult<Url> {
fn redirect_location(response: &reqwest::Response, url: &Url) -> Result<Url, Error> {
let location = response
.headers()
.get(reqwest::header::LOCATION)
.and_then(|value| value.to_str().ok())
.ok_or_else(|| {
CoreError::InvalidResponse("OCR document redirect missing Location header".to_string())
Error::InvalidResponse("OCR document redirect missing Location header".to_string())
})?;
url.join(location)
.map_err(|err| CoreError::InvalidResponse(format!("invalid OCR document redirect: {err}")))
.map_err(|err| Error::InvalidResponse(format!("invalid OCR document redirect: {err}")))
}
async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response)> {
async fn safe_get_document_url(url: &str) -> Result<(Url, reqwest::Response), Error> {
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
let mut current_url = Url::parse(url)
.map_err(|err| CoreError::InvalidRequest(format!("invalid OCR document URL: {err}")))?;
.map_err(|err| Error::InvalidRequest(format!("invalid OCR document URL: {err}")))?;
for _ in 0..MAX_SAFE_FETCH_REDIRECTS {
validate_safe_fetch_url(&current_url).await?;
@ -202,28 +201,28 @@ async fn safe_get_document_url(url: &str) -> CoreResult<(Url, reqwest::Response)
.get(current_url.clone())
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
if !response.status().is_redirection() {
return Ok((current_url, response));
}
current_url = redirect_location(&response, &current_url)?;
}
Err(CoreError::InvalidRequest(
Err(Error::InvalidRequest(
"Too many redirects while fetching OCR document URL".to_string(),
))
}
fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> CoreResult<()> {
fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> Result<(), Error> {
if max_bytes == 0 {
return Err(CoreError::InvalidRequest(format!(
return Err(Error::InvalidRequest(format!(
"OCR document URL download is disabled (MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0). url={url}"
)));
}
if content_length > max_bytes {
let size_mb = content_length as f64 / (1024.0 * 1024.0);
let max_size_mb = max_bytes as f64 / (1024.0 * 1024.0);
return Err(CoreError::InvalidRequest(format!(
return Err(Error::InvalidRequest(format!(
"OCR document size ({size_mb:.2}MB) exceeds maximum allowed size ({max_size_mb:.2}MB). url={url}"
)));
}
@ -233,7 +232,7 @@ fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> Core
async fn read_response_with_limit(
mut response: reqwest::Response,
url: &Url,
) -> CoreResult<Vec<u8>> {
) -> Result<Vec<u8>, Error> {
let max_bytes = max_document_download_bytes();
if let Some(content_length) = response.content_length() {
enforce_download_size(content_length, max_bytes, url)?;
@ -246,7 +245,7 @@ async fn read_response_with_limit(
while let Some(chunk) = response
.chunk()
.await
.map_err(|err| CoreError::Network(err.to_string()))?
.map_err(|err| Error::Network(err.to_string()))?
{
bytes_downloaded += chunk.len() as u64;
enforce_download_size(bytes_downloaded, max_bytes, url)?;
@ -255,7 +254,7 @@ async fn read_response_with_limit(
Ok(bytes)
}
pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreResult<Value> {
pub(super) async fn convert_document_url_to_data_uri(document: Value) -> Result<Value, Error> {
let Some((field, url)) = document_url_field(&document)? else {
return Ok(document);
};
@ -267,7 +266,7 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
return Err(CoreError::Http {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&body),
});
@ -290,7 +289,7 @@ pub(super) async fn convert_document_url_to_data_uri(document: Value) -> CoreRes
let mut transformed = document
.as_object()
.cloned()
.ok_or_else(|| CoreError::InvalidRequest("OCR document must be an object".to_string()))?;
.ok_or_else(|| Error::InvalidRequest("OCR document must be an object".to_string()))?;
transformed.insert(field.to_string(), Value::String(data_uri));
Ok(Value::Object(transformed))
}
@ -316,11 +315,11 @@ fn retry_after_secs(response: &reqwest::Response) -> u64 {
.unwrap_or(2)
}
fn operation_status(response_json: &Value) -> CoreResult<&str> {
fn operation_status(response_json: &Value) -> Result<&str, Error> {
let status = response_json
.get("status")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("status"))?;
.ok_or(Error::MissingField("status"))?;
match status {
"succeeded" => Ok("succeeded"),
"running" | "notStarted" => Ok("running"),
@ -330,11 +329,11 @@ fn operation_status(response_json: &Value) -> CoreResult<&str> {
.and_then(|error| error.get("message"))
.and_then(Value::as_str)
.unwrap_or("Unknown error");
Err(CoreError::InvalidResponse(format!(
Err(Error::InvalidResponse(format!(
"Azure Document Intelligence analysis failed: {message}"
)))
}
other => Err(CoreError::InvalidResponse(format!(
other => Err(Error::InvalidResponse(format!(
"Unknown operation status: {other}"
))),
}
@ -345,9 +344,9 @@ pub(super) async fn poll_document_intelligence(
original_url: &str,
headers: &[(String, String)],
timeout: Option<Duration>,
) -> CoreResult<Value> {
) -> Result<Value, Error> {
if !same_origin(operation_url, original_url) {
return Err(CoreError::InvalidResponse(
return Err(Error::InvalidResponse(
"Azure Document Intelligence: rejected cross-origin polling URL".to_string(),
));
}
@ -358,7 +357,7 @@ pub(super) async fn poll_document_intelligence(
));
loop {
if start.elapsed() > timeout {
return Err(CoreError::Network(format!(
return Err(Error::Network(format!(
"Azure Document Intelligence operation polling timed out after {} seconds",
timeout.as_secs()
)));
@ -373,21 +372,21 @@ pub(super) async fn poll_document_intelligence(
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
let retry_after = retry_after_secs(&response);
let status = response.status();
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response_json: Value = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid Azure DI poll response JSON: {err}"))
Error::InvalidResponse(format!("invalid Azure DI poll response JSON: {err}"))
})?;
if operation_status(&response_json)? == "succeeded" {
return Ok(response_json);
@ -426,7 +425,7 @@ mod tests {
assert!(matches!(
error,
CoreError::InvalidRequest(message)
Error::InvalidRequest(message)
if message.contains("SSRF protection")
));
}

View file

@ -1,5 +1,4 @@
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::ocr::transformation::OcrResponseHandling;
use serde_json::Value;
@ -7,7 +6,7 @@ use super::common_utils::{poll_document_intelligence, truncate_error_body};
use super::types::ProviderOcrRequest;
use crate::client::http_client;
pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> CoreResult<Value> {
pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Result<Value, Error> {
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
@ -19,7 +18,7 @@ pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Co
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
let status = response.status();
if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll
@ -31,7 +30,7 @@ pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Co
.and_then(|value| value.to_str().ok())
.map(str::to_string)
.ok_or_else(|| {
CoreError::InvalidResponse(
Error::InvalidResponse(
"Azure Document Intelligence returned 202 but no Operation-Location header found"
.to_string(),
)
@ -52,17 +51,17 @@ pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Co
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response_json: Value = serde_json::from_str(&text)
.map_err(|err| CoreError::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
.map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
Ok(request
.config

View file

@ -1,11 +1,9 @@
use std::future::Future;
use std::pin::Pin;
use litellm_core::CoreResult;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::ocr::transformation::OcrAuthStrategy;
use serde_json::{Map, Value, json};
use std::future::Future;
use std::pin::Pin;
use super::common_utils::{
convert_document_url_to_data_uri, has_header, ocr_provider_config, string_headers,
@ -27,7 +25,7 @@ pub(crate) struct OcrLifecycleHooks {
request_metadata: RequestMetadata,
}
type OcrFuture<'a, T> = Pin<Box<dyn Future<Output = CoreResult<T>> + Send + 'a>>;
type OcrFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
type OcrLogFuture<'a> = Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
impl OcrLifecycleHooks {
@ -46,7 +44,7 @@ impl OcrLifecycleHooks {
async fn run_pre_call_guardrails(
&self,
request: PreparedOcrRequest,
) -> CoreResult<PreparedOcrRequest> {
) -> Result<PreparedOcrRequest, Error> {
if self.guardrail_runner.is_empty() {
return Ok(request);
}
@ -74,9 +72,9 @@ impl OcrLifecycleHooks {
async fn prepare_provider_request(
&self,
request: PreparedOcrRequest,
) -> CoreResult<ProviderOcrRequest> {
) -> Result<ProviderOcrRequest, Error> {
let config = ocr_provider_config(&request.custom_llm_provider, &request.model)
.ok_or_else(|| CoreError::InvalidProvider(request.custom_llm_provider.clone()))?;
.ok_or_else(|| Error::InvalidProvider(request.custom_llm_provider.clone()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
let headers = string_headers(request.extra_headers)?;
let auth_strategy = config.auth_strategy();
@ -120,7 +118,7 @@ impl OcrLifecycleHooks {
custom_llm_provider: &str,
url: &str,
body: Value,
) -> CoreResult<Value> {
) -> Result<Value, Error> {
if self.guardrail_runner.is_empty() {
return Ok(body);
}
@ -217,7 +215,7 @@ impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> for OcrLi
fn async_log_failure_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
error: &'a CoreError,
error: &'a Error,
timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
Box::pin(async move {
@ -278,19 +276,19 @@ fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
fn parse_ocr_pre_call_guardrail_request(
request: GuardrailRequest,
) -> CoreResult<(Value, Map<String, Value>)> {
) -> Result<(Value, Map<String, Value>), Error> {
let Value::Object(mut data) = request.data else {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"OCR pre_call guardrail must return an object".to_string(),
));
};
let document = data.remove("document").ok_or_else(|| {
CoreError::InvalidRequest("OCR pre_call guardrail removed document".to_string())
Error::InvalidRequest("OCR pre_call guardrail removed document".to_string())
})?;
let optional_params = match data.remove("optional_params") {
Some(Value::Object(params)) => params,
Some(_) => {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"OCR pre_call guardrail optional_params must be an object".to_string(),
));
}
@ -299,33 +297,32 @@ fn parse_ocr_pre_call_guardrail_request(
Ok((document, optional_params))
}
fn parse_ocr_during_call_guardrail_request(request: GuardrailRequest) -> CoreResult<Value> {
fn parse_ocr_during_call_guardrail_request(request: GuardrailRequest) -> Result<Value, Error> {
let Value::Object(mut data) = request.data else {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"OCR during_call guardrail must return an object".to_string(),
));
};
data.remove("body").ok_or_else(|| {
CoreError::InvalidRequest("OCR during_call guardrail removed body".to_string())
})
data.remove("body")
.ok_or_else(|| Error::InvalidRequest("OCR during_call guardrail removed body".to_string()))
}
fn guardrail_error_to_core_error(error: GuardrailError) -> CoreError {
CoreError::InvalidRequest(format!("{}: {}", error.kind, error.message))
fn guardrail_error_to_core_error(error: GuardrailError) -> Error {
Error::InvalidRequest(format!("{}: {}", error.kind, error.message))
}
fn core_error_kind(error: &CoreError) -> &'static str {
fn core_error_kind(error: &Error) -> &'static str {
match error {
CoreError::Auth(_) => "AuthError",
CoreError::InvalidProvider(_) => "InvalidProvider",
CoreError::InvalidRequest(_) => "InvalidRequest",
CoreError::InvalidType { .. } => "InvalidType",
CoreError::MissingField(_) => "MissingField",
CoreError::Http { .. } => "HttpError",
CoreError::InvalidResponse(_) => "InvalidResponse",
CoreError::Network(_) => "NetworkError",
CoreError::Connect(_) => "ConnectError",
CoreError::Routing(_) => "RoutingError",
CoreError::Unsupported(_) => "UnsupportedRequest",
Error::Auth(_) => "AuthError",
Error::InvalidProvider(_) => "InvalidProvider",
Error::InvalidRequest(_) => "InvalidRequest",
Error::InvalidType { .. } => "InvalidType",
Error::MissingField(_) => "MissingField",
Error::Http { .. } => "HttpError",
Error::InvalidResponse(_) => "InvalidResponse",
Error::Network(_) => "NetworkError",
Error::Connect(_) => "ConnectError",
Error::Routing(_) => "RoutingError",
Error::Unsupported(_) => "UnsupportedRequest",
}
}

View file

@ -1,4 +1,4 @@
use litellm_core::CoreResult;
use litellm_core::Error;
use litellm_core::call_lifecycle::CallLifecycle;
use serde_json::Value;
@ -13,7 +13,7 @@ pub use types::OcrRequest;
use handler::execute_ocr_provider_call;
use prepare::{PreparedOcrCall, prepare_ocr_call};
pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
pub async fn ocr(request: OcrRequest<'_>) -> Result<Value, Error> {
let PreparedOcrCall { request, hooks } = prepare_ocr_call(request);
CallLifecycle::default()
.run_request(request, &hooks, execute_ocr_provider_call)

View file

@ -1,7 +1,7 @@
use std::sync::{Arc, Mutex};
use std::time::Duration;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::ocr::transformation::OcrResponseHandling;
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
@ -395,7 +395,7 @@ async fn ocr_lifecycle_runs_failure_hook_on_provider_error() {
.await
.expect_err("provider error propagates");
assert!(matches!(err, CoreError::Http { status: 500, .. }));
assert!(matches!(err, Error::Http { status: 500, .. }));
server.await.expect("server task completes");
assert_eq!(
logger.events(),
@ -439,7 +439,7 @@ async fn ocr_lifecycle_pre_call_block_skips_provider_socket() {
.await
.expect_err("guardrail blocks request");
assert!(matches!(err, CoreError::InvalidRequest(_)));
assert!(matches!(err, Error::InvalidRequest(_)));
assert_eq!(guardrail.events(), vec!["async_pre_call_hook"]);
assert_eq!(
logger.events(),
@ -607,7 +607,7 @@ fn string_headers_rejects_non_string_values() {
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
CoreError::InvalidRequest(
Error::InvalidRequest(
"OCR extra_headers.x-retry-count must be a string, got number".to_string()
)
);

View file

@ -6,33 +6,31 @@
//! (and recorded in [`crate::gil`]); the realtime hot path never touches Python.
//!
//! Compiled only under the `python-config` feature.
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::router::{Deployment, Router};
use pyo3::prelude::*;
use crate::gil;
/// Load the router's `model_list` from `config_path` via the Python reader.
pub fn load_router_from_config(config_path: &str) -> CoreResult<Router> {
pub fn load_router_from_config(config_path: &str) -> Result<Router, Error> {
gil::record_acquisition();
Python::attach(|py| {
let model_list = py
.import("litellm.proxy.read_model_list")
.and_then(|module| module.getattr("read_model_list"))
.and_then(|reader| reader.call1((config_path,)))
.map_err(|err| CoreError::Routing(format!("read_model_list failed: {err}")))?;
.map_err(|err| Error::Routing(format!("read_model_list failed: {err}")))?;
let model_list_json: String = py
.import("json")
.and_then(|json| json.getattr("dumps"))
.and_then(|dumps| dumps.call1((model_list,)))
.and_then(|encoded| encoded.extract())
.map_err(|err| CoreError::Routing(format!("serializing model_list failed: {err}")))?;
.map_err(|err| Error::Routing(format!("serializing model_list failed: {err}")))?;
let deployments: Vec<Deployment> = serde_json::from_str(&model_list_json)
.map_err(|err| CoreError::Routing(format!("parsing model_list failed: {err}")))?;
.map_err(|err| Error::Routing(format!("parsing model_list failed: {err}")))?;
Ok(Router::new(deployments))
})

View file

@ -9,7 +9,7 @@ use axum::http::StatusCode;
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue};
use axum::response::{IntoResponse, Response};
use axum::routing::post;
use litellm_core::CoreError;
use litellm_core::Error;
use serde_json::{Map, Value};
use crate::auth::RequireMasterKey;
@ -46,7 +46,7 @@ fn stream_response(upstream: reqwest::Response) -> Result<Response, MessagesRout
let mut response = Response::builder()
.status(
StatusCode::from_u16(upstream.status().as_u16()).map_err(|error| {
MessagesRouteError(CoreError::InvalidResponse(format!(
MessagesRouteError(Error::InvalidResponse(format!(
"invalid upstream response status: {error}"
)))
})?,
@ -58,13 +58,13 @@ fn stream_response(upstream: reqwest::Response) -> Result<Response, MessagesRout
response
.body(Body::from_stream(upstream.bytes_stream()))
.map_err(|error| {
MessagesRouteError(CoreError::InvalidResponse(format!(
MessagesRouteError(Error::InvalidResponse(format!(
"failed to build streaming response: {error}"
)))
})
}
fn forwarded_headers(headers: &HeaderMap) -> Result<Option<Map<String, Value>>, CoreError> {
fn forwarded_headers(headers: &HeaderMap) -> Result<Option<Map<String, Value>>, Error> {
let forwarded = headers
.iter()
.filter(|(name, _)| {
@ -74,19 +74,19 @@ fn forwarded_headers(headers: &HeaderMap) -> Result<Option<Map<String, Value>>,
})
.map(|(name, value)| {
let value = value.to_str().map_err(|_| {
CoreError::InvalidRequest(format!("invalid value for header {}", name.as_str()))
Error::InvalidRequest(format!("invalid value for header {}", name.as_str()))
})?;
Ok((name.to_string(), Value::String(value.to_string())))
})
.collect::<Result<Map<_, _>, CoreError>>()?;
.collect::<Result<Map<_, _>, Error>>()?;
Ok((!forwarded.is_empty()).then_some(forwarded))
}
#[derive(Debug)]
struct MessagesRouteError(CoreError);
struct MessagesRouteError(Error);
impl From<CoreError> for MessagesRouteError {
fn from(error: CoreError) -> Self {
impl From<Error> for MessagesRouteError {
fn from(error: Error) -> Self {
Self(error)
}
}
@ -94,28 +94,28 @@ impl From<CoreError> for MessagesRouteError {
impl IntoResponse for MessagesRouteError {
fn into_response(self) -> Response {
let (status, message) = match self.0 {
CoreError::InvalidRequest(message) => (StatusCode::BAD_REQUEST, message),
CoreError::InvalidProvider(_) | CoreError::Routing(_) => (
Error::InvalidRequest(message) => (StatusCode::BAD_REQUEST, message),
Error::InvalidProvider(_) | Error::Routing(_) => (
StatusCode::NOT_FOUND,
"no messages deployment is configured for this model".to_string(),
),
CoreError::Auth(_) => (
Error::Auth(_) => (
StatusCode::BAD_GATEWAY,
"messages provider authentication failed".to_string(),
),
CoreError::Http { .. }
| CoreError::Network(_)
| CoreError::Connect(_)
| CoreError::InvalidResponse(_)
| CoreError::InvalidType { .. }
| CoreError::MissingField(_) => (
Error::Http { .. }
| Error::Network(_)
| Error::Connect(_)
| Error::InvalidResponse(_)
| Error::InvalidType { .. }
| Error::MissingField(_) => (
StatusCode::BAD_GATEWAY,
"messages provider request failed".to_string(),
),
// The gateway has no Python implementation to decline to, so a
// request the core cannot serve is reported to the caller. The
// reason is a fixed internal string, never provider content.
CoreError::Unsupported(reason) => (
Error::Unsupported(reason) => (
StatusCode::BAD_REQUEST,
format!("messages request is not supported: {reason}"),
),

View file

@ -1,10 +1,10 @@
use std::sync::Arc;
use litellm_core::Error;
use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER;
use litellm_core::messages::types::MessagesRequest;
use litellm_core::messages::{messages, messages_stream};
use litellm_core::router::Router;
use litellm_core::{CoreError, CoreResult};
use serde_json::{Map, Value};
pub(crate) enum MessagesResponse {
@ -16,16 +16,16 @@ pub async fn run(
router: &Arc<Router>,
body: Value,
extra_headers: Option<Map<String, Value>>,
) -> CoreResult<MessagesResponse> {
) -> Result<MessagesResponse, Error> {
let model = body
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|model| !model.is_empty())
.ok_or_else(|| CoreError::InvalidRequest("messages body requires a model".to_string()))?;
let deployment = router.get_available_deployment(model).ok_or_else(|| {
CoreError::Routing(format!("no deployment available for model '{model}'"))
})?;
.ok_or_else(|| Error::InvalidRequest("messages body requires a model".to_string()))?;
let deployment = router
.get_available_deployment(model)
.ok_or_else(|| Error::Routing(format!("no deployment available for model '{model}'")))?;
let provider_model = deployment.litellm_params.model.as_str();
let upstream_model = provider_model
.split_once('/')
@ -37,7 +37,7 @@ pub async fn run(
};
let mut body = body;
body.as_object_mut()
.ok_or_else(|| CoreError::InvalidRequest("messages body must be an object".to_string()))?
.ok_or_else(|| Error::InvalidRequest("messages body must be an object".to_string()))?
.insert(
"model".to_string(),
Value::String(upstream_model.to_string()),
@ -60,6 +60,6 @@ pub async fn run(
serde_json::to_value(response)
.map(MessagesResponse::Json)
.map_err(|err| {
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
Error::InvalidResponse(format!("failed to serialize messages response: {err}"))
})
}

View file

@ -11,8 +11,7 @@ use std::time::Duration;
use crate::io::realtime_pool::{RealtimePool, upstream_key};
use futures_util::{Sink, Stream};
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::realtime::types::RealtimeEvent;
use litellm_core::router::Router;
@ -29,15 +28,15 @@ pub async fn run<In, Out>(
observe: impl FnMut(&RealtimeEvent) + Send,
client_in: In,
client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = RealtimeEvent> + Unpin + Send,
Out: Sink<RealtimeEvent> + Unpin + Send,
<Out as Sink<RealtimeEvent>>::Error: std::fmt::Display,
{
let deployment = router.get_available_deployment(model).ok_or_else(|| {
CoreError::Routing(format!("no deployment available for model '{model}'"))
})?;
let deployment = router
.get_available_deployment(model)
.ok_or_else(|| Error::Routing(format!("no deployment available for model '{model}'")))?;
let params = &deployment.litellm_params;
// Strip a leading `openai/` so the OpenAI-only realtime fn gets the bare model.
let provider_model = params

View file

@ -2,13 +2,13 @@ use std::sync::Arc;
use std::time::Duration;
use futures_util::{Sink, Stream};
use litellm_core::Error;
use litellm_core::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use litellm_core::responses::instrumentation::{
ResponsesWsCallbackPayload, ResponsesWsInstrumentation, ResponsesWsLogOutcome,
ResponsesWsMetadata,
};
use litellm_core::responses::types::ResponsesWsEvent;
use litellm_core::{CoreError, CoreResult};
use crate::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails,
@ -26,22 +26,22 @@ pub async fn run<In, Out>(
metadata: RequestMetadata,
client_in: In,
client_out: Out,
) -> CoreResult<()>
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let deployment = router.get_available_deployment(model).ok_or_else(|| {
CoreError::Routing(format!("no deployment available for model '{model}'"))
})?;
let deployment = router
.get_available_deployment(model)
.ok_or_else(|| Error::Routing(format!("no deployment available for model '{model}'")))?;
let params = &deployment.litellm_params;
let provider_model = params
.model
.strip_prefix("openai/")
.unwrap_or(&params.model);
if params.model.contains('/') && !params.model.starts_with("openai/") {
return Err(CoreError::InvalidProvider(
return Err(Error::InvalidProvider(
"Responses WebSocket route supports OpenAI deployments only".to_string(),
));
}

View file

@ -1,7 +1,6 @@
use crate::Error;
use serde_json::{Map, Value};
use crate::CoreResult;
use super::types::{AudioTranscriptionRequestData, AudioTranscriptionResponseData};
#[derive(Clone, Debug, PartialEq, Eq)]
@ -32,13 +31,13 @@ pub trait AudioTranscriptionProviderConfig: Sync {
model: &str,
audio: Value,
optional_params: Map<String, Value>,
) -> CoreResult<AudioTranscriptionRequestData>;
) -> Result<AudioTranscriptionRequestData, Error>;
fn transform_transcription_response(
&self,
model: &str,
response_json: Value,
) -> CoreResult<AudioTranscriptionResponseData>;
) -> Result<AudioTranscriptionResponseData, Error>;
fn complete_url(
&self,
@ -46,12 +45,12 @@ pub trait AudioTranscriptionProviderConfig: Sync {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
) -> Result<String, Error>;
fn auth_strategy(
&self,
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<AudioTranscriptionAuth>;
) -> Result<AudioTranscriptionAuth, Error>;
}

View file

@ -1,7 +1,7 @@
use std::future::Future;
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use crate::{CoreError, CoreResult};
use crate::Error;
pub mod types;
@ -11,14 +11,14 @@ pub use types::{
};
pub trait CallLifecycleHooks<InitialReq, ProviderReq, Resp>: Send + Sync {
type PreCallFuture<'a>: Future<Output = CoreResult<InitialReq>> + Send + 'a
type PreCallFuture<'a>: Future<Output = Result<InitialReq, Error>> + Send + 'a
where
Self: 'a,
InitialReq: 'a,
ProviderReq: 'a,
Resp: 'a;
type DuringCallFuture<'a>: Future<Output = CoreResult<ProviderReq>> + Send + 'a
type DuringCallFuture<'a>: Future<Output = Result<ProviderReq, Error>> + Send + 'a
where
Self: 'a,
InitialReq: 'a,
@ -56,7 +56,7 @@ pub trait CallLifecycleHooks<InitialReq, ProviderReq, Resp>: Send + Sync {
fn async_log_failure_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
error: &'a CoreError,
error: &'a Error,
timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a>;
}
@ -86,12 +86,12 @@ impl<'a> CallLifecycle<'a> {
request: InitialReq,
hooks: &Hooks,
provider_call: ProviderCall,
) -> CoreResult<Resp>
) -> Result<Resp, Error>
where
InitialReq: CallLifecycleRequest,
Hooks: CallLifecycleHooks<InitialReq, ProviderReq, Resp>,
ProviderCall: FnOnce(ProviderReq) -> ProviderFuture,
ProviderFuture: Future<Output = CoreResult<Resp>>,
ProviderFuture: Future<Output = Result<Resp, Error>>,
{
let context = request.lifecycle_context();
self.run(context, request, hooks, provider_call).await
@ -103,11 +103,11 @@ impl<'a> CallLifecycle<'a> {
request: InitialReq,
hooks: &Hooks,
provider_call: ProviderCall,
) -> CoreResult<Resp>
) -> Result<Resp, Error>
where
Hooks: CallLifecycleHooks<InitialReq, ProviderReq, Resp>,
ProviderCall: FnOnce(ProviderReq) -> ProviderFuture,
ProviderFuture: Future<Output = CoreResult<Resp>>,
ProviderFuture: Future<Output = Result<Resp, Error>>,
{
let call_start = epoch_seconds();
let mut phases = Vec::new();
@ -166,7 +166,7 @@ impl<'a> CallLifecycle<'a> {
&self,
context: &CallLifecycleContext,
hooks: &Hooks,
error: &CoreError,
error: &Error,
call_start: f64,
phases: &mut Vec<CallLifecyclePhaseTiming>,
) where
@ -251,8 +251,8 @@ mod tests {
}
impl CallLifecycleHooks<String, String, String> for RecordingHooks {
type PreCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
type DuringCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
type PreCallFuture<'a> = BoxFuture<'a, Result<String, Error>>;
type DuringCallFuture<'a> = BoxFuture<'a, Result<String, Error>>;
type SuccessFuture<'a> = BoxFuture<'a, ()>;
type FailureFuture<'a> = BoxFuture<'a, ()>;
@ -294,7 +294,7 @@ mod tests {
fn async_log_failure_event<'a>(
&'a self,
_context: &'a CallLifecycleContext,
_error: &'a CoreError,
_error: &'a Error,
_timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
Box::pin(async move {
@ -304,8 +304,8 @@ mod tests {
}
impl CallLifecycleHooks<RecordingRequest, String, String> for RecordingHooks {
type PreCallFuture<'a> = BoxFuture<'a, CoreResult<RecordingRequest>>;
type DuringCallFuture<'a> = BoxFuture<'a, CoreResult<String>>;
type PreCallFuture<'a> = BoxFuture<'a, Result<RecordingRequest, Error>>;
type DuringCallFuture<'a> = BoxFuture<'a, Result<String, Error>>;
type SuccessFuture<'a> = BoxFuture<'a, ()>;
type FailureFuture<'a> = BoxFuture<'a, ()>;
@ -345,7 +345,7 @@ mod tests {
fn async_log_failure_event<'a>(
&'a self,
_context: &'a CallLifecycleContext,
_error: &'a CoreError,
_error: &'a Error,
_timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
Box::pin(async move {
@ -383,13 +383,13 @@ mod tests {
"request".to_string(),
&hooks,
|_request| async move {
Err::<String, CoreError>(CoreError::Network("provider down".to_string()))
Err::<String, Error>(Error::Network("provider down".to_string()))
},
)
.await
.expect_err("call fails");
assert_eq!(error, CoreError::Network("provider down".to_string()));
assert_eq!(error, Error::Network("provider down".to_string()));
assert_eq!(hooks.events(), vec!["pre_call", "during_call", "failure"]);
}

View file

@ -1,8 +1,7 @@
use serde_json::{Map, Value};
use crate::error::CoreResult;
use crate::Error;
use crate::http_utils::string_headers as shared_string_headers;
use crate::providers::anthropic::chat_completions::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG;
use serde_json::{Map, Value};
use super::transformation::ChatCompletionsProviderConfig;
@ -23,6 +22,6 @@ pub(super) fn chat_completions_provider_config(
pub(super) fn string_headers(
extra_headers: Option<Map<String, Value>>,
) -> CoreResult<Vec<(String, String)>> {
) -> Result<Vec<(String, String)>, Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers)
}

View file

@ -1,6 +1,6 @@
use serde_json::Value;
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use crate::http_utils::truncate_error_body;
use super::client::http_client;
@ -11,9 +11,9 @@ use super::types::{
pub(super) async fn execute_chat_completions_provider_call(
request: ProviderChatCompletionsRequest,
) -> CoreResult<ChatCompletionsResponse> {
) -> Result<ChatCompletionsResponse, Error> {
let body = serde_json::to_vec(&request.body).map_err(|err| {
CoreError::InvalidRequest(format!(
Error::InvalidRequest(format!(
"failed to serialize chat completions request: {err}"
))
})?;
@ -32,9 +32,9 @@ pub(super) async fn execute_chat_completions_provider_call(
// 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() {
CoreError::Connect(err.to_string())
Error::Connect(err.to_string())
} else {
CoreError::Network(err.to_string())
Error::Network(err.to_string())
}
})?;
@ -42,17 +42,17 @@ pub(super) async fn execute_chat_completions_provider_call(
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let body: Value = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
Error::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
})?;
request
.config
@ -69,10 +69,10 @@ pub(super) async fn execute_chat_completions_provider_call(
/// second kind has already been billed, and a host that keeps a reference
/// implementation must not retry those, so collapse them to one variant that
/// can only mean the provider was already called.
pub(super) fn as_response_error(err: CoreError) -> CoreError {
pub(super) fn as_response_error(err: Error) -> Error {
match err {
already @ (CoreError::InvalidResponse(_) | CoreError::Http { .. }) => already,
other => CoreError::InvalidResponse(other.to_string()),
already @ (Error::InvalidResponse(_) | Error::Http { .. }) => already,
other => Error::InvalidResponse(other.to_string()),
}
}
@ -80,7 +80,7 @@ pub(super) fn as_response_error(err: CoreError) -> CoreError {
pub(super) async fn signed_headers(
request: &ProviderChatCompletionsRequest,
body: &[u8],
) -> CoreResult<Vec<(String, String)>> {
) -> Result<Vec<(String, String)>, Error> {
use std::collections::BTreeMap;
use std::time::SystemTime;
@ -101,7 +101,7 @@ pub(super) async fn signed_headers(
.iter()
.any(|(name, _)| is_sigv4_computed_header(name))
{
return Err(CoreError::Unsupported(
return Err(Error::Unsupported(
"request forwards a header AWS SigV4 computes",
));
}
@ -137,9 +137,9 @@ pub(super) async fn signed_headers(
pub(super) async fn signed_headers(
request: &ProviderChatCompletionsRequest,
_body: &[u8],
) -> CoreResult<Vec<(String, String)>> {
) -> Result<Vec<(String, String)>, Error> {
match &request.auth {
ChatCompletionsAuth::AwsSigV4 { .. } => Err(CoreError::Unsupported(
ChatCompletionsAuth::AwsSigV4 { .. } => Err(Error::Unsupported(
"AWS SigV4 requires the bedrock-auth feature",
)),
_ => Ok(request.upstream_headers.clone()),

View file

@ -6,6 +6,7 @@
//! credentials, and it resolves the provider, translates the conversation,
//! calls the provider, and returns a typed OpenAI-shaped response.
use crate::Error;
mod client;
mod common_utils;
pub mod conversation;
@ -17,15 +18,13 @@ pub mod types;
use serde_json::{Map, Value};
use crate::error::CoreResult;
use handler::execute_chat_completions_provider_call;
use prepare::{parse_messages, prepare_chat_completions_call, resolve_provider_config};
use types::{ChatCompletionsRequest, ChatCompletionsResponse};
pub async fn chat_completions(
request: ChatCompletionsRequest<'_>,
) -> CoreResult<ChatCompletionsResponse> {
) -> Result<ChatCompletionsResponse, Error> {
execute_chat_completions_provider_call(prepare_chat_completions_call(request)?).await
}

View file

@ -1,6 +1,6 @@
use serde_json::Value;
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use crate::http_utils::has_header;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
@ -11,7 +11,7 @@ use super::types::{ChatCompletionsRequest, ChatMessage, ProviderChatCompletionsR
pub(super) fn resolve_provider_config<'a>(
model: &'a str,
custom_llm_provider: Option<&'a str>,
) -> CoreResult<(String, &'static dyn ChatCompletionsProviderConfig)> {
) -> Result<(String, &'static dyn ChatCompletionsProviderConfig), Error> {
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
.or_else(|| {
custom_llm_provider.map(|provider| CustomLlmProvider {
@ -20,35 +20,34 @@ pub(super) fn resolve_provider_config<'a>(
})
})
.ok_or_else(|| {
CoreError::InvalidProvider(
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(|| CoreError::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
Ok((provider_info.model.to_string(), config))
}
pub(super) fn parse_messages(messages: Value) -> CoreResult<Vec<ChatMessage>> {
serde_json::from_value(messages).map_err(|err| {
CoreError::InvalidRequest(format!("invalid chat completions messages: {err}"))
})
pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error> {
serde_json::from_value(messages)
.map_err(|err| Error::InvalidRequest(format!("invalid chat completions messages: {err}")))
}
pub(super) fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> CoreResult<ProviderChatCompletionsRequest> {
) -> Result<ProviderChatCompletionsRequest, Error> {
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?;
let env_lookup = |key: &str| std::env::var(key).ok();
let messages = parse_messages(request.messages)?;
if messages.is_empty() {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"chat completions requires at least one message".to_string(),
));
}
if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) {
return Err(CoreError::Unsupported(reason.0));
return Err(Error::Unsupported(reason.0));
}
let mut headers = string_headers(request.extra_headers)?;

View file

@ -1,6 +1,6 @@
use serde_json::{Map, Value, json};
use crate::error::CoreError;
use crate::error::Error;
use super::prepare::prepare_chat_completions_call;
use super::transformation::ChatCompletionsAuth;
@ -29,7 +29,7 @@ fn request<'a>(
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
/// carry resolved credentials), so unwrap the failure case by hand.
fn decline(request: ChatCompletionsRequest<'_>) -> CoreError {
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),
@ -196,7 +196,7 @@ fn declines_an_unsupported_request_before_resolving_credentials() {
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), CoreError::Unsupported("streaming"));
assert_eq!(decline(call), Error::Unsupported("streaming"));
}
#[test]
@ -208,7 +208,7 @@ fn rejects_an_unknown_provider() {
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
CoreError::InvalidProvider("openai".to_string())
Error::InvalidProvider("openai".to_string())
);
}
@ -221,7 +221,7 @@ fn rejects_a_model_with_no_resolvable_provider() {
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
CoreError::InvalidProvider(_)
Error::InvalidProvider(_)
));
}
@ -234,7 +234,7 @@ fn rejects_an_empty_or_malformed_message_list() {
json!([]),
json!({}),
)),
CoreError::InvalidRequest("chat completions requires at least one message".to_string())
Error::InvalidRequest("chat completions requires at least one message".to_string())
);
assert!(matches!(
decline(request(
@ -243,7 +243,7 @@ fn rejects_an_empty_or_malformed_message_list() {
json!("not a list"),
json!({}),
)),
CoreError::InvalidRequest(_)
Error::InvalidRequest(_)
));
}
@ -258,7 +258,7 @@ fn rejects_non_string_extra_headers() {
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
assert_eq!(
decline(call),
CoreError::InvalidRequest(
Error::InvalidRequest(
"chat completions extra_headers.x-trace must be a string, got number".to_string()
)
);
@ -374,7 +374,7 @@ async fn a_forwarded_header_the_signer_computes_declines_to_python() {
.await
.expect_err("{forwarded} should decline instead of being signed");
assert!(
matches!(error, CoreError::Unsupported(_)),
matches!(error, Error::Unsupported(_)),
"{forwarded} declined as {error:?}, which the host would not fall back on"
);
}
@ -727,7 +727,7 @@ mod round_trip {
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, CoreError::InvalidResponse(_)),
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
@ -745,7 +745,7 @@ mod round_trip {
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, CoreError::InvalidResponse(_)),
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
@ -763,7 +763,7 @@ mod round_trip {
.expect_err("upstream rejects");
handle.await.expect("server task");
assert!(
matches!(err, CoreError::Http { status: 429, .. }),
matches!(err, Error::Http { status: 429, .. }),
"expected a 429, got {err:?}"
);
}
@ -787,7 +787,7 @@ mod round_trip {
.await
.expect_err("nothing is listening");
assert!(
matches!(err, CoreError::Connect(_)),
matches!(err, Error::Connect(_)),
"expected a pre-send connect failure, got {err:?}"
);
}
@ -797,24 +797,24 @@ mod round_trip {
use crate::chat_completions::handler::as_response_error;
for original in [
CoreError::MissingField("usage"),
CoreError::Unsupported("non-text response content block"),
CoreError::InvalidRequest("whatever".to_string()),
CoreError::Auth("whatever".to_string()),
Error::MissingField("usage"),
Error::Unsupported("non-text response content block"),
Error::InvalidRequest("whatever".to_string()),
Error::Auth("whatever".to_string()),
] {
let label = format!("{original:?}");
assert!(
matches!(as_response_error(original), CoreError::InvalidResponse(_)),
matches!(as_response_error(original), Error::InvalidResponse(_)),
"{label} must not stay retryable once the provider has answered"
);
}
// An upstream status is already unambiguous, so it survives intact.
assert!(matches!(
as_response_error(CoreError::Http {
as_response_error(Error::Http {
status: 500,
body: "boom".to_string()
}),
CoreError::Http { status: 500, .. }
Error::Http { status: 500, .. }
));
}
}

View file

@ -1,7 +1,6 @@
use crate::Error;
use serde_json::{Map, Value};
use crate::error::CoreResult;
use super::types::{
ChatCompletionsResponse, ChatMessage, ChatMessageContent, ProviderChatRequestData,
ProviderChatResponseData,
@ -39,7 +38,7 @@ pub trait ChatCompletionsProviderConfig: Sync {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
) -> Result<String, Error>;
fn auth(
&self,
@ -47,7 +46,7 @@ pub trait ChatCompletionsProviderConfig: Sync {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<ChatCompletionsAuth>;
) -> Result<ChatCompletionsAuth, Error>;
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
&[("content-type", "application/json")]
@ -91,13 +90,13 @@ pub trait ChatCompletionsProviderConfig: Sync {
model: &str,
messages: Vec<ChatMessage>,
optional_params: Map<String, Value>,
) -> CoreResult<ProviderChatRequestData>;
) -> Result<ProviderChatRequestData, Error>;
fn transform_response(
&self,
model: &str,
response: ProviderChatResponseData,
) -> CoreResult<ChatCompletionsResponse>;
) -> Result<ChatCompletionsResponse, Error>;
}
pub fn unsupported_param(

View file

@ -1,9 +1,7 @@
use thiserror::Error;
use thiserror::Error as ThisError;
pub type CoreResult<T> = Result<T, CoreError>;
#[derive(Debug, Error, PartialEq, Eq)]
pub enum CoreError {
#[derive(Debug, ThisError, PartialEq, Eq)]
pub enum Error {
#[error("expected {expected}, got {actual}")]
InvalidType {
expected: &'static str,

View file

@ -3,7 +3,7 @@
use serde_json::{Map, Value};
use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS;
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::error::{Error, json_type_name};
/// Bound an upstream error body before it crosses a host boundary, so provider
/// bodies stay data-minimized.
@ -18,7 +18,7 @@ pub fn truncate_error_body(body: &str) -> String {
pub fn string_headers(
context: &'static str,
extra_headers: Option<Map<String, Value>>,
) -> CoreResult<Vec<(String, String)>> {
) -> Result<Vec<(String, String)>, Error> {
extra_headers
.unwrap_or_default()
.into_iter()
@ -27,7 +27,7 @@ pub fn string_headers(
.as_str()
.map(|value| (key.clone(), value.to_string()))
.ok_or_else(|| {
CoreError::InvalidRequest(format!(
Error::InvalidRequest(format!(
"{context} extra_headers.{key} must be a string, got {}",
json_type_name(&value)
))
@ -81,7 +81,7 @@ mod tests {
let err = string_headers("chat completions", Some(headers)).expect_err("non-string value");
assert_eq!(
err,
CoreError::InvalidRequest(
Error::InvalidRequest(
"chat completions extra_headers.x-trace must be a string, got number".to_string()
)
);

View file

@ -13,4 +13,4 @@ pub mod responses;
pub mod router;
pub mod routing_utils;
pub use error::{CoreError, CoreResult};
pub use error::Error;

View file

@ -1,9 +1,8 @@
use serde_json::{Map, Value};
use crate::error::CoreResult;
use crate::Error;
use crate::http_utils::string_headers as shared_string_headers;
use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
use serde_json::{Map, Value};
use super::transformation::AnthropicMessagesProviderConfig;
@ -23,6 +22,6 @@ pub(super) fn messages_provider_config(
pub(super) fn string_headers(
extra_headers: Option<Map<String, Value>>,
) -> CoreResult<Vec<(String, String)>> {
) -> Result<Vec<(String, String)>, Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers)
}

View file

@ -1,5 +1,5 @@
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use super::client::http_client;
use super::common_utils::truncate_error_body;
@ -7,7 +7,7 @@ use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
pub(super) async fn execute_messages_provider_call(
request: ProviderMessagesRequest,
) -> CoreResult<AnthropicMessagesResponse> {
) -> Result<AnthropicMessagesResponse, Error> {
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
@ -19,32 +19,31 @@ pub(super) async fn execute_messages_provider_call(
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
let status = response.status();
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
})?;
let response = serde_json::from_str(&text)
.map_err(|err| Error::InvalidResponse(format!("invalid messages response JSON: {err}")))?;
request.config.transform_response(&request.model, response)
}
pub(super) async fn execute_messages_provider_stream(
request: ProviderMessagesRequest,
) -> CoreResult<reqwest::Response> {
) -> Result<reqwest::Response, Error> {
if request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"streaming messages is not supported for this provider".to_string(),
));
}
@ -60,14 +59,14 @@ pub(super) async fn execute_messages_provider_stream(
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
.map_err(|err| Error::Network(err.to_string()))?;
let status = response.status();
if !status.is_success() {
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
return Err(CoreError::Http {
.map_err(|err| Error::Network(err.to_string()))?;
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});

View file

@ -7,6 +7,7 @@
//! is the streaming variant; it hands the raw upstream response back so a host
//! can splice the event stream to its own caller.
use crate::Error;
mod client;
mod common_utils;
mod handler;
@ -14,17 +15,15 @@ mod prepare;
pub mod transformation;
pub mod types;
use crate::error::CoreResult;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use prepare::prepare_messages_call;
use types::{AnthropicMessagesResponse, MessagesRequest};
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<AnthropicMessagesResponse> {
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
execute_messages_provider_call(prepare_messages_call(request)?).await
}
pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult<reqwest::Response> {
pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Response, Error> {
execute_messages_provider_stream(prepare_messages_call(request)?).await
}

View file

@ -1,4 +1,4 @@
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
@ -7,7 +7,7 @@ use super::types::{MessagesRequest, ProviderMessagesRequest};
pub(super) fn prepare_messages_call(
request: MessagesRequest<'_>,
) -> CoreResult<ProviderMessagesRequest> {
) -> Result<ProviderMessagesRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.or_else(|| {
request
@ -18,7 +18,7 @@ pub(super) fn prepare_messages_call(
})
})
.ok_or_else(|| {
CoreError::InvalidProvider(
Error::InvalidProvider(
"unable to resolve custom_llm_provider for messages request".to_string(),
)
})?;
@ -26,7 +26,7 @@ pub(super) fn prepare_messages_call(
let provider = provider_info.custom_llm_provider;
let config = messages_provider_config(provider)
.ok_or_else(|| CoreError::InvalidProvider(provider.to_string()))?;
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
let mut headers = string_headers(request.extra_headers)?;
@ -53,11 +53,11 @@ pub(super) fn prepare_messages_call(
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
let typed_request = serde_json::from_value(request.body).map_err(|err| {
CoreError::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
})?;
let transformed = config.transform_request(typed_request)?;
let body = serde_json::to_value(transformed).map_err(|err| {
CoreError::InvalidRequest(format!(
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"
))
})?;

View file

@ -4,7 +4,7 @@ use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use crate::error::CoreError;
use crate::error::Error;
use super::common_utils::{
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
@ -77,7 +77,7 @@ fn truncate_error_body_caps_long_payloads() {
fn string_headers_rejects_non_string_values() {
let headers = json!({"x-count": 3}).as_object().unwrap().clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert!(matches!(err, CoreError::InvalidRequest(_)));
assert!(matches!(err, Error::InvalidRequest(_)));
}
#[test]
@ -341,7 +341,7 @@ async fn messages_requires_auth_when_no_key_and_no_header() {
.await
.expect_err("missing auth errors");
assert!(matches!(err, CoreError::Auth(_)));
assert!(matches!(err, Error::Auth(_)));
}
#[tokio::test]
@ -420,7 +420,7 @@ async fn messages_maps_provider_error_status_to_http_error() {
.await
.expect_err("provider error propagates");
assert!(matches!(err, CoreError::Http { status: 401, .. }));
assert!(matches!(err, Error::Http { status: 401, .. }));
}
#[tokio::test]
@ -437,5 +437,5 @@ async fn messages_rejects_unsupported_provider() {
.await
.expect_err("unsupported provider errors");
assert!(matches!(err, CoreError::InvalidProvider(provider) if provider == "openai"));
assert!(matches!(err, Error::InvalidProvider(provider) if provider == "openai"));
}

View file

@ -1,6 +1,5 @@
use crate::error::CoreResult;
use super::types::{AnthropicMessagesRequest, AnthropicMessagesResponse};
use crate::Error;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesAuthStrategy {
@ -23,13 +22,13 @@ pub trait AnthropicMessagesProviderConfig: Sync {
api_base: Option<&str>,
model: &str,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
) -> Result<String, Error>;
fn resolve_api_key(
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
) -> Result<String, Error>;
fn auth_strategy(&self) -> MessagesAuthStrategy {
MessagesAuthStrategy::Header("x-api-key")
@ -49,7 +48,7 @@ pub trait AnthropicMessagesProviderConfig: Sync {
fn transform_request(
&self,
request: AnthropicMessagesRequest,
) -> CoreResult<AnthropicMessagesRequest> {
) -> Result<AnthropicMessagesRequest, Error> {
Ok(request)
}
@ -57,7 +56,7 @@ pub trait AnthropicMessagesProviderConfig: Sync {
&self,
_model: &str,
response: AnthropicMessagesResponse,
) -> CoreResult<AnthropicMessagesResponse> {
) -> Result<AnthropicMessagesResponse, Error> {
Ok(response)
}
}

View file

@ -1,7 +1,6 @@
use crate::Error;
use serde_json::{Map, Value};
use crate::CoreResult;
use super::types::{OcrRequestData, OcrResponseData};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
@ -43,13 +42,13 @@ pub trait OcrProviderConfig: Sync {
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData>;
) -> Result<OcrRequestData, Error>;
fn transform_ocr_response(
&self,
model: &str,
response_json: Value,
) -> CoreResult<OcrResponseData>;
) -> Result<OcrResponseData, Error>;
fn complete_url(
&self,
@ -57,13 +56,13 @@ pub trait OcrProviderConfig: Sync {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
) -> Result<String, Error>;
fn resolve_api_key(
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
) -> Result<String, Error>;
fn auth_strategy(&self) -> OcrAuthStrategy {
OcrAuthStrategy::Bearer

View file

@ -1,4 +1,5 @@
use super::*;
use crate::Error;
use serde_json::json;
fn messages(value: Value) -> Vec<ChatMessage> {
@ -19,7 +20,7 @@ fn transform(model: &str, msgs: Value, opts: Value) -> Value {
.body
}
fn transform_response(body: Value) -> CoreResult<ChatCompletionsResponse> {
fn transform_response(body: Value) -> Result<ChatCompletionsResponse, Error> {
ANTHROPIC_CHAT_COMPLETIONS_CONFIG
.transform_response("claude-sonnet-4-5", ProviderChatResponseData { body })
}
@ -390,29 +391,26 @@ fn declines_a_response_carrying_a_non_text_block() {
"usage": {"input_tokens": 1, "output_tokens": 1}
}))
.expect_err("non-text block");
assert_eq!(
err,
CoreError::Unsupported("non-text response content block")
);
assert_eq!(err, Error::Unsupported("non-text response content block"));
}
#[test]
fn errors_on_a_response_missing_required_fields() {
assert_eq!(
transform_response(json!("nope")).expect_err("not an object"),
CoreError::InvalidResponse("messages response is not an object".to_string())
Error::InvalidResponse("messages response is not an object".to_string())
);
assert_eq!(
transform_response(json!({"model": "m", "usage": {}})).expect_err("no content"),
CoreError::MissingField("content")
Error::MissingField("content")
);
assert_eq!(
transform_response(json!({"model": "m", "content": []})).expect_err("no usage"),
CoreError::MissingField("usage")
Error::MissingField("usage")
);
assert_eq!(
transform_response(json!({"content": [], "usage": {}})).expect_err("no model"),
CoreError::MissingField("model")
Error::MissingField("model")
);
}

View file

@ -10,7 +10,7 @@ use crate::chat_completions::types::{
ProviderChatRequestData, ProviderChatResponseData,
};
use crate::constants::ANTHROPIC_OAUTH_TOKEN_PREFIX;
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use crate::providers::anthropic::messages::transformation::{
complete_anthropic_url, resolve_anthropic_api_key,
};
@ -74,7 +74,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
_model: &str,
_optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
Ok(complete_anthropic_url(api_base, env_lookup))
}
@ -84,7 +84,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
_model: &str,
_optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<ChatCompletionsAuth> {
) -> Result<ChatCompletionsAuth, Error> {
Ok(ChatCompletionsAuth::Header {
name: "x-api-key",
value: resolve_anthropic_api_key(api_key, env_lookup)?,
@ -137,7 +137,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
model: &str,
messages: Vec<ChatMessage>,
optional_params: Map<String, Value>,
) -> CoreResult<ProviderChatRequestData> {
) -> Result<ProviderChatRequestData, Error> {
Ok(ProviderChatRequestData {
body: anthropic_body(model, &build_conversation(&messages), optional_params),
})
@ -147,15 +147,16 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
&self,
_model: &str,
response: ProviderChatResponseData,
) -> CoreResult<ChatCompletionsResponse> {
let body = response.body.as_object().ok_or_else(|| {
CoreError::InvalidResponse("messages response is not an object".into())
})?;
) -> Result<ChatCompletionsResponse, Error> {
let body = response
.body
.as_object()
.ok_or_else(|| Error::InvalidResponse("messages response is not an object".into()))?;
let content = body
.get("content")
.and_then(Value::as_array)
.ok_or(CoreError::MissingField("content"))?;
.ok_or(Error::MissingField("content"))?;
// The route declines tool and thinking requests, so a non-text block
// means the response carries something this path never asked for.
// Decline rather than silently dropping it; the host falls back.
@ -163,7 +164,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
.iter()
.any(|block| block.get("type").and_then(Value::as_str) != Some("text"))
{
return Err(CoreError::Unsupported("non-text response content block"));
return Err(Error::Unsupported("non-text response content block"));
}
let text: String = content
.iter()
@ -173,7 +174,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
let usage = body
.get("usage")
.and_then(Value::as_object)
.ok_or(CoreError::MissingField("usage"))?;
.ok_or(Error::MissingField("usage"))?;
let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0);
Ok(ChatCompletionsResponse {
@ -181,7 +182,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
model: body
.get("model")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("model"))?
.ok_or(Error::MissingField("model"))?
.to_string(),
choices: vec![ChatCompletionsChoice {
index: 0,

View file

@ -1,4 +1,4 @@
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy};
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
@ -17,12 +17,12 @@ pub fn non_empty(value: Option<&str>) -> Option<&str> {
pub fn resolve_anthropic_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
non_empty(api_key)
.map(str::to_string)
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| {
CoreError::Auth(
Error::Auth(
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY \
environment variable"
.to_string(),
@ -52,7 +52,7 @@ impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig {
api_base: Option<&str>,
_model: &str,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
Ok(complete_anthropic_url(api_base, env_lookup))
}
@ -60,7 +60,7 @@ impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig {
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_anthropic_api_key(api_key, env_lookup)
}
@ -121,7 +121,7 @@ mod tests {
);
assert!(matches!(
resolve_anthropic_api_key(None, &|_| None).expect_err("missing key"),
CoreError::Auth(_)
Error::Auth(_)
));
}

View file

@ -1,4 +1,4 @@
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy};
use crate::messages::types::{
AnthropicMessage, AnthropicMessagesRequest, AnthropicMessagesResponse, ContentBlock,
@ -28,12 +28,12 @@ pub const AZURE_ANTHROPIC_MESSAGES_CONFIG: AzureAnthropicMessagesConfig =
pub fn resolve_azure_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
non_empty(api_key)
.map(str::to_string)
.or_else(|| env_lookup(AZURE_API_KEY_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| {
CoreError::Auth(
Error::Auth(
"Missing Azure API Key - Set `api_key` or the AZURE_API_KEY environment variable"
.to_string(),
)
@ -43,12 +43,12 @@ pub fn resolve_azure_api_key(
pub fn complete_azure_anthropic_url(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
let api_base = non_empty(api_base)
.map(str::to_string)
.or_else(|| env_lookup(AZURE_API_BASE_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| {
CoreError::Auth(
Error::Auth(
"Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. \
Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
.to_string(),
@ -147,7 +147,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
api_base: Option<&str>,
_model: &str,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
complete_azure_anthropic_url(api_base, env_lookup)
}
@ -155,7 +155,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_azure_api_key(api_key, env_lookup)
}
@ -174,7 +174,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
fn transform_request(
&self,
request: AnthropicMessagesRequest,
) -> CoreResult<AnthropicMessagesRequest> {
) -> Result<AnthropicMessagesRequest, Error> {
let mut request = fold_system_role_messages(request);
if let Some(system) = request.system.as_mut() {
strip_scope_from_system(system);
@ -190,7 +190,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
&self,
model: &str,
response: AnthropicMessagesResponse,
) -> CoreResult<AnthropicMessagesResponse> {
) -> Result<AnthropicMessagesResponse, Error> {
self.anthropic.transform_response(model, response)
}
}
@ -268,7 +268,7 @@ mod tests {
"https://env.services.ai.azure.com/anthropic/v1/messages"
);
let err = complete_azure_anthropic_url(Some(" "), &|_| None).expect_err("missing base");
assert!(matches!(err, CoreError::Auth(_)));
assert!(matches!(err, Error::Auth(_)));
}
#[test]
@ -284,7 +284,7 @@ mod tests {
);
assert!(matches!(
resolve_azure_api_key(None, &|_| None).expect_err("missing key"),
CoreError::Auth(_)
Error::Auth(_)
));
}

View file

@ -1,6 +1,6 @@
use std::collections::BTreeSet;
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::error::{Error, json_type_name};
use crate::ocr::transformation::{OcrAuthStrategy, OcrProviderConfig, OcrResponseHandling};
use crate::ocr::types::{OcrRequestData, OcrResponseData};
use serde_json::{Map, Value, json};
@ -32,17 +32,17 @@ fn resolve_value(
env_name: &str,
env_lookup: &dyn Fn(&str) -> Option<String>,
missing_message: &str,
) -> CoreResult<String> {
) -> Result<String, Error> {
non_empty(explicit)
.map(str::to_string)
.or_else(|| env_lookup(env_name).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| CoreError::Auth(missing_message.to_string()))
.ok_or_else(|| Error::Auth(missing_message.to_string()))
}
pub fn resolve_azure_ai_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_value(
api_key,
AZURE_AI_API_KEY_ENV,
@ -54,7 +54,7 @@ pub fn resolve_azure_ai_api_key(
pub fn resolve_azure_ai_api_base(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_value(
api_base,
AZURE_AI_API_BASE_ENV,
@ -66,7 +66,7 @@ pub fn resolve_azure_ai_api_base(
pub fn complete_azure_ai_url(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
let base = resolve_azure_ai_api_base(api_base, env_lookup)?;
Ok(format!(
"{}/providers/mistral/azure/ocr",
@ -77,7 +77,7 @@ pub fn complete_azure_ai_url(
pub fn resolve_document_intelligence_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_value(
api_key,
AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV,
@ -89,7 +89,7 @@ pub fn resolve_document_intelligence_api_key(
pub fn resolve_document_intelligence_endpoint(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_value(
api_base,
AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV,
@ -127,7 +127,7 @@ fn pages_token_is_valid(token: &str) -> bool {
}
}
fn normalize_pages_param(pages: &Value) -> CoreResult<Option<String>> {
fn normalize_pages_param(pages: &Value) -> Result<Option<String>, Error> {
match pages {
Value::String(value) => {
let normalized = value
@ -138,7 +138,7 @@ fn normalize_pages_param(pages: &Value) -> CoreResult<Option<String>> {
if normalized.split(',').all(pages_token_is_valid) {
Ok(Some(normalized))
} else {
Err(CoreError::InvalidRequest(format!(
Err(Error::InvalidRequest(format!(
"Invalid `pages` string for Azure Document Intelligence: {value:?}. Expected format like '1-3,5,7-9'."
)))
}
@ -152,7 +152,7 @@ fn normalize_pages_param(pages: &Value) -> CoreResult<Option<String>> {
for value in values {
let page = value.as_i64().expect("checked is_i64");
if page < 0 {
return Err(CoreError::InvalidRequest(
return Err(Error::InvalidRequest(
"`pages` integers must be >= 0 (Mistral 0-based indices)".to_string(),
));
}
@ -176,16 +176,16 @@ fn normalize_pages_param(pages: &Value) -> CoreResult<Option<String>> {
if normalized.split(',').all(pages_token_is_valid) {
return Ok(Some(normalized));
}
return Err(CoreError::InvalidRequest(format!(
return Err(Error::InvalidRequest(format!(
"Invalid `pages` list for Azure Document Intelligence: {values:?}. Expected tokens like '1' or '3-5'."
)));
}
Err(CoreError::InvalidRequest(
Err(Error::InvalidRequest(
"`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'."
.to_string(),
))
}
_ => Err(CoreError::InvalidRequest(
_ => Err(Error::InvalidRequest(
"`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'."
.to_string(),
)),
@ -197,7 +197,7 @@ pub fn complete_document_intelligence_url(
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
let endpoint = resolve_document_intelligence_endpoint(api_base, env_lookup)?;
let mut url = format!(
"{}/documentintelligence/documentModels/{}:analyze?api-version={}",
@ -216,20 +216,20 @@ pub fn complete_document_intelligence_url(
Ok(url)
}
fn document_url_from_mistral_document(document: &Value) -> CoreResult<&str> {
let object = document.as_object().ok_or_else(|| CoreError::InvalidType {
fn document_url_from_mistral_document(document: &Value) -> Result<&str, Error> {
let object = document.as_object().ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(document),
})?;
let doc_type = object
.get("type")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("document.type"))?;
.ok_or(Error::MissingField("document.type"))?;
let field_name = match doc_type {
"document_url" => "document_url",
"image_url" => "image_url",
other => {
return Err(CoreError::InvalidRequest(format!(
return Err(Error::InvalidRequest(format!(
"Invalid document type: {other}. Must be 'document_url' or 'image_url'"
)));
}
@ -238,7 +238,7 @@ fn document_url_from_mistral_document(document: &Value) -> CoreResult<&str> {
.get(field_name)
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.ok_or(CoreError::MissingField(field_name))
.ok_or(Error::MissingField(field_name))
}
fn extract_base64_from_data_uri(data_uri: &str) -> &str {
@ -290,7 +290,7 @@ impl OcrProviderConfig for AzureAiOcrConfig {
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData> {
) -> Result<OcrRequestData, Error> {
MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params)
}
@ -298,7 +298,7 @@ impl OcrProviderConfig for AzureAiOcrConfig {
&self,
model: &str,
response_json: Value,
) -> CoreResult<OcrResponseData> {
) -> Result<OcrResponseData, Error> {
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
}
@ -308,7 +308,7 @@ impl OcrProviderConfig for AzureAiOcrConfig {
_model: &str,
_optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
complete_azure_ai_url(api_base, env_lookup)
}
@ -316,7 +316,7 @@ impl OcrProviderConfig for AzureAiOcrConfig {
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_azure_ai_api_key(api_key, env_lookup)
}
@ -335,7 +335,7 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
_model: &str,
document: Value,
_optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData> {
) -> Result<OcrRequestData, Error> {
let document_url = document_url_from_mistral_document(&document)?;
let mut data = Map::new();
if document_url.starts_with("data:") {
@ -359,19 +359,19 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
&self,
model: &str,
response_json: Value,
) -> CoreResult<OcrResponseData> {
) -> Result<OcrResponseData, Error> {
let response = response_json
.as_object()
.ok_or_else(|| CoreError::InvalidType {
.ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
let status = response
.get("status")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("status"))?;
.ok_or(Error::MissingField("status"))?;
if status != "succeeded" {
return Err(CoreError::InvalidResponse(format!(
return Err(Error::InvalidResponse(format!(
"Azure Document Intelligence analysis failed with status: {status}"
)));
}
@ -414,7 +414,7 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
complete_document_intelligence_url(api_base, model, optional_params, env_lookup)
}
@ -422,7 +422,7 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_document_intelligence_api_key(api_key, env_lookup)
}

View file

@ -6,7 +6,7 @@ use crate::audio_transcription::transformation::{
use crate::audio_transcription::types::{
AudioTranscriptionRequestData, AudioTranscriptionResponseData,
};
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::error::{Error, json_type_name};
pub use super::aws_base::{aws_auth_config, bedrock_model_id_and_region, resolve_bedrock_region};
use super::constants::{BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE};
@ -18,8 +18,8 @@ pub static BEDROCK_AUDIO_TRANSCRIPTION_CONFIG: BedrockAudioTranscriptionConfig =
pub struct BedrockAudioTranscriptionConfig;
fn audio_fields(audio: Value) -> CoreResult<(String, String)> {
let object = audio.as_object().ok_or_else(|| CoreError::InvalidType {
fn audio_fields(audio: Value) -> Result<(String, String), Error> {
let object = audio.as_object().ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(&audio),
})?;
@ -27,13 +27,13 @@ fn audio_fields(audio: Value) -> CoreResult<(String, String)> {
.get("data")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.ok_or(CoreError::MissingField("audio.data"))?;
.ok_or(Error::MissingField("audio.data"))?;
let format = object
.get("format")
.and_then(Value::as_str)
.filter(|value| matches!(*value, "wav" | "mp3" | "flac" | "ogg"))
.ok_or_else(|| {
CoreError::InvalidRequest("audio.format must be wav, mp3, flac, or ogg".to_string())
Error::InvalidRequest("audio.format must be wav, mp3, flac, or ogg".to_string())
})?;
Ok((data.to_string(), format.to_string()))
}
@ -55,7 +55,7 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
_model: &str,
audio: Value,
optional_params: Map<String, Value>,
) -> CoreResult<AudioTranscriptionRequestData> {
) -> Result<AudioTranscriptionRequestData, Error> {
let (data, format) = audio_fields(audio)?;
let mut instruction = "Transcribe the audio. Respond with only the transcript.".to_string();
if let Some(language) = optional_string(&optional_params, "language") {
@ -87,14 +87,14 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
&self,
_model: &str,
response_json: Value,
) -> CoreResult<AudioTranscriptionResponseData> {
) -> Result<AudioTranscriptionResponseData, Error> {
let content = response_json
.get("output")
.and_then(|value| value.get("message"))
.and_then(|value| value.get("content"))
.and_then(Value::as_array)
.ok_or_else(|| {
CoreError::InvalidResponse("Bedrock response has no output content".to_string())
Error::InvalidResponse("Bedrock response has no output content".to_string())
})?;
let mut text = String::new();
for block in content {
@ -111,7 +111,7 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
let (model_id, model_region) = bedrock_model_id_and_region(model);
let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup);
let endpoint = optional_params
@ -133,7 +133,7 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<AudioTranscriptionAuth> {
) -> Result<AudioTranscriptionAuth, Error> {
let (_, model_region) = bedrock_model_id_and_region(model);
Ok(AudioTranscriptionAuth::AwsSigV4 {
region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup),

View file

@ -4,7 +4,7 @@ use std::time::Duration;
use std::time::{SystemTime, UNIX_EPOCH};
use crate::caching::in_memory_cache::InMemoryCache;
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use aws_credential_types::Credentials;
use aws_credential_types::provider::ProvideCredentials;
use aws_sigv4::http_request::{
@ -197,7 +197,7 @@ pub fn classify_auth(
pub async fn resolve_credentials(
config: AwsAuthConfig,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> CoreResult<Credentials> {
) -> Result<Credentials, Error> {
let resolved = config.clone().with_environment(env_lookup);
let flow = classify_auth(config, env_lookup);
match flow {
@ -244,9 +244,10 @@ pub async fn resolve_credentials(
let provider = aws_config::profile::ProfileFileCredentialsProvider::builder()
.profile_name(name)
.build();
provider.provide_credentials().await.map_err(|error| {
CoreError::Auth(format!("AWS profile credentials failed: {error}"))
})
provider
.provide_credentials()
.await
.map_err(|error| Error::Auth(format!("AWS profile credentials failed: {error}")))
}
AwsAuthFlow::AssumeRole { role, session_name } => {
if is_already_running_as_role(&role, &resolved).await? {
@ -260,7 +261,7 @@ pub async fn resolve_credentials(
.build()
.await;
let credentials = provider.provide_credentials().await.map_err(|error| {
CoreError::Auth(format!("AWS default credentials failed: {error}"))
Error::Auth(format!("AWS default credentials failed: {error}"))
})?;
set_cached_credentials(
key,
@ -301,7 +302,7 @@ pub async fn resolve_credentials(
provider
.provide_credentials()
.await
.map_err(|error| CoreError::Auth(format!("AWS role credentials failed: {error}")))
.map_err(|error| Error::Auth(format!("AWS role credentials failed: {error}")))
}
AwsAuthFlow::WebIdentity {
token,
@ -325,13 +326,13 @@ pub async fn resolve_credentials(
.send()
.await
.map_err(|error| {
CoreError::Auth(format!("AWS web identity credentials failed: {error}"))
Error::Auth(format!("AWS web identity credentials failed: {error}"))
})?;
let credentials = response.credentials().ok_or_else(|| {
CoreError::Auth("AWS web identity response had no credentials".to_string())
Error::Auth("AWS web identity response had no credentials".to_string())
})?;
let expiration = SystemTime::try_from(*credentials.expiration()).map_err(|error| {
CoreError::Auth(format!("AWS web identity expiration was invalid: {error}"))
Error::Auth(format!("AWS web identity expiration was invalid: {error}"))
})?;
Ok(Credentials::new(
credentials.access_key_id(),
@ -350,9 +351,10 @@ pub async fn resolve_credentials(
aws_config::default_provider::credentials::DefaultCredentialsChain::builder()
.build()
.await;
let credentials = provider.provide_credentials().await.map_err(|error| {
CoreError::Auth(format!("AWS default credentials failed: {error}"))
})?;
let credentials = provider
.provide_credentials()
.await
.map_err(|error| Error::Auth(format!("AWS default credentials failed: {error}")))?;
set_cached_credentials(
key,
credentials.clone(),
@ -363,7 +365,7 @@ pub async fn resolve_credentials(
}
}
async fn is_already_running_as_role(role: &str, config: &AwsAuthConfig) -> CoreResult<bool> {
async fn is_already_running_as_role(role: &str, config: &AwsAuthConfig) -> Result<bool, Error> {
if role_identity(role).is_none() {
return Ok(false);
}
@ -437,7 +439,7 @@ pub fn sign_bedrock_post(
region: &str,
credentials: &Credentials,
signing_time: SystemTime,
) -> CoreResult<BTreeMap<String, String>> {
) -> Result<BTreeMap<String, String>, Error> {
let identity: Identity = credentials.clone().into();
let params = v4::SigningParams::builder()
.identity(&identity)
@ -447,14 +449,14 @@ pub fn sign_bedrock_post(
.settings(SigningSettings::default())
.build()
.map(SigningParams::from)
.map_err(|error| CoreError::Auth(format!("AWS signing parameters failed: {error}")))?;
.map_err(|error| Error::Auth(format!("AWS signing parameters failed: {error}")))?;
let header_refs = headers
.iter()
.map(|(name, value)| (name.as_str(), value.as_str()));
let request = SignableRequest::new("POST", url, header_refs, SignableBody::Bytes(body))
.map_err(|error| CoreError::Auth(format!("AWS signable request failed: {error}")))?;
.map_err(|error| Error::Auth(format!("AWS signable request failed: {error}")))?;
let (instructions, _) = sign(request, &params)
.map_err(|error| CoreError::Auth(format!("AWS request signing failed: {error}")))?
.map_err(|error| Error::Auth(format!("AWS request signing failed: {error}")))?
.into_parts();
Ok(instructions
.headers()

View file

@ -1,4 +1,5 @@
use super::*;
use crate::Error;
use serde_json::json;
fn messages(value: Value) -> Vec<ChatMessage> {
@ -23,7 +24,7 @@ fn transform(msgs: Value, opts: Value) -> Value {
.body
}
fn transform_response(body: Value) -> CoreResult<ChatCompletionsResponse> {
fn transform_response(body: Value) -> Result<ChatCompletionsResponse, Error> {
BEDROCK_CHAT_COMPLETIONS_CONFIG.transform_response(
"anthropic.claude-sonnet-4-5-v1:0",
ProviderChatResponseData { body },
@ -478,25 +479,22 @@ fn declines_a_response_carrying_a_tool_use_block() {
"usage": {"inputTokens": 1, "outputTokens": 1}
}))
.expect_err("tool use block");
assert_eq!(
err,
CoreError::Unsupported("non-text response content block")
);
assert_eq!(err, Error::Unsupported("non-text response content block"));
}
#[test]
fn errors_on_a_response_missing_required_fields() {
assert_eq!(
transform_response(json!("nope")).expect_err("not an object"),
CoreError::InvalidResponse("converse response is not an object".to_string())
Error::InvalidResponse("converse response is not an object".to_string())
);
assert_eq!(
transform_response(json!({"usage": {}})).expect_err("no output"),
CoreError::MissingField("output.message.content")
Error::MissingField("output.message.content")
);
assert_eq!(
transform_response(json!({"output": {"message": {"content": []}}})).expect_err("no usage"),
CoreError::MissingField("usage")
Error::MissingField("usage")
);
}

View file

@ -11,7 +11,7 @@ use crate::chat_completions::types::{
ChatCompletionsUsage, ChatMessage, ChatMessageContent, ProviderChatRequestData,
ProviderChatResponseData,
};
use crate::error::{CoreError, CoreResult};
use crate::error::Error;
use super::super::aws_base::{bedrock_model_id_and_region, resolve_bedrock_region};
use super::super::constants::{AWS_BEARER_TOKEN_BEDROCK, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE};
@ -110,7 +110,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
let (model_id, model_region) = bedrock_model_id_and_region(model);
let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup);
let endpoint = optional_params
@ -137,7 +137,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<ChatCompletionsAuth> {
) -> Result<ChatCompletionsAuth, Error> {
// Python reads `api_key` as the Bedrock bearer token and consults the
// env only when the caller passed none, so a caller-supplied empty key
// falls through to SigV4 without reaching for the environment. An
@ -208,7 +208,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
_model: &str,
messages: Vec<ChatMessage>,
optional_params: Map<String, Value>,
) -> CoreResult<ProviderChatRequestData> {
) -> Result<ProviderChatRequestData, Error> {
Ok(ProviderChatRequestData {
body: converse_body(&build_conversation(&messages), &optional_params),
})
@ -218,17 +218,18 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
&self,
model: &str,
response: ProviderChatResponseData,
) -> CoreResult<ChatCompletionsResponse> {
let body = response.body.as_object().ok_or_else(|| {
CoreError::InvalidResponse("converse response is not an object".into())
})?;
) -> Result<ChatCompletionsResponse, Error> {
let body = response
.body
.as_object()
.ok_or_else(|| Error::InvalidResponse("converse response is not an object".into()))?;
let content = body
.get("output")
.and_then(|output| output.get("message"))
.and_then(|message| message.get("content"))
.and_then(Value::as_array)
.ok_or(CoreError::MissingField("output.message.content"))?;
.ok_or(Error::MissingField("output.message.content"))?;
// The route declines tool requests, so anything other than a text block
// is something this path never asked for. Decline; the host falls back.
if content.iter().any(|block| {
@ -236,7 +237,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
.as_object()
.is_none_or(|block| block.len() != 1 || !block.contains_key("text"))
}) {
return Err(CoreError::Unsupported("non-text response content block"));
return Err(Error::Unsupported("non-text response content block"));
}
let text: String = content
.iter()
@ -246,7 +247,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
let usage = body
.get("usage")
.and_then(Value::as_object)
.ok_or(CoreError::MissingField("usage"))?;
.ok_or(Error::MissingField("usage"))?;
let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0);
let computed = usage_from_parts(
field("inputTokens"),

View file

@ -1,4 +1,4 @@
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::error::{Error, json_type_name};
use crate::ocr::transformation::OcrProviderConfig;
use crate::ocr::types::{OcrRequestData, OcrResponseData};
use serde_json::{Map, Value};
@ -47,7 +47,7 @@ pub fn complete_url(api_base: Option<&str>) -> String {
/// Resolve the Mistral API key from the explicit param or the environment.
///
/// Blank/whitespace values are treated as absent. Returns `CoreError::Auth`
/// Blank/whitespace values are treated as absent. Returns `Error::Auth`
/// when no usable key is available.
///
/// Note: the env fallback only reads the process environment. Secret-manager
@ -56,13 +56,13 @@ pub fn complete_url(api_base: Option<&str>) -> String {
pub fn resolve_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
api_key
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or_else(|| CoreError::Auth(MISSING_KEY_MESSAGE.to_string()))
.ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string()))
}
pub struct MistralOcrConfig;
@ -79,9 +79,9 @@ impl OcrProviderConfig for MistralOcrConfig {
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData> {
) -> Result<OcrRequestData, Error> {
if !document.is_object() {
return Err(CoreError::InvalidType {
return Err(Error::InvalidType {
expected: "object",
actual: json_type_name(&document),
});
@ -104,10 +104,10 @@ impl OcrProviderConfig for MistralOcrConfig {
&self,
model: &str,
response_json: Value,
) -> CoreResult<OcrResponseData> {
) -> Result<OcrResponseData, Error> {
let response_object = response_json
.as_object()
.ok_or_else(|| CoreError::InvalidType {
.ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
@ -140,7 +140,7 @@ impl OcrProviderConfig for MistralOcrConfig {
_model: &str,
_optional_params: &Map<String, Value>,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
Ok(complete_url(api_base))
}
@ -148,7 +148,7 @@ impl OcrProviderConfig for MistralOcrConfig {
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_api_key(api_key, env_lookup)
}
}
@ -165,11 +165,11 @@ pub fn transform_ocr_request(
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData> {
) -> Result<OcrRequestData, Error> {
MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params)
}
pub fn transform_ocr_response(model: &str, response_json: Value) -> CoreResult<OcrResponseData> {
pub fn transform_ocr_response(model: &str, response_json: Value) -> Result<OcrResponseData, Error> {
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
}
@ -250,7 +250,7 @@ mod tests {
assert_eq!(
err,
CoreError::InvalidType {
Error::InvalidType {
expected: "object",
actual: "string",
}
@ -307,6 +307,6 @@ mod tests {
#[test]
fn resolve_api_key_errors_when_absent() {
let err = resolve_api_key(None, &|_| None).expect_err("missing key should error");
assert_eq!(err, CoreError::Auth(MISSING_KEY_MESSAGE.to_string()));
assert_eq!(err, Error::Auth(MISSING_KEY_MESSAGE.to_string()));
}
}

View file

@ -1,4 +1,4 @@
use crate::CoreResult;
use crate::Error;
use crate::realtime::transformation::RealtimeProviderConfig;
use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
@ -72,7 +72,7 @@ impl RealtimeProviderConfig for OpenAiRealtimeConfig {
&self,
event: &RealtimeEvent,
_model: &str,
) -> CoreResult<RealtimeTransformResult> {
) -> Result<RealtimeTransformResult, Error> {
Ok(RealtimeTransformResult::passthrough(event.clone()))
}
@ -80,7 +80,7 @@ impl RealtimeProviderConfig for OpenAiRealtimeConfig {
&self,
event: &RealtimeEvent,
_model: &str,
) -> CoreResult<RealtimeTransformResult> {
) -> Result<RealtimeTransformResult, Error> {
Ok(RealtimeTransformResult::passthrough(event.clone()))
}
}
@ -88,14 +88,14 @@ impl RealtimeProviderConfig for OpenAiRealtimeConfig {
pub fn transform_realtime_request(
event: &RealtimeEvent,
model: &str,
) -> CoreResult<RealtimeTransformResult> {
) -> Result<RealtimeTransformResult, Error> {
OPENAI_REALTIME_CONFIG.transform_realtime_request(event, model)
}
pub fn transform_realtime_response(
event: &RealtimeEvent,
model: &str,
) -> CoreResult<RealtimeTransformResult> {
) -> Result<RealtimeTransformResult, Error> {
OPENAI_REALTIME_CONFIG.transform_realtime_response(event, model)
}

View file

@ -1,4 +1,4 @@
use crate::CoreResult;
use crate::Error;
use crate::responses::types::{ResponsesWsEvent, ResponsesWsTransformResult};
use crate::responses::websocket::{ResponsesWebSocketProviderConfig, enforce_model};
@ -15,7 +15,7 @@ impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig {
&self,
event: &ResponsesWsEvent,
model: &str,
) -> CoreResult<ResponsesWsTransformResult> {
) -> Result<ResponsesWsTransformResult, Error> {
Ok(ResponsesWsTransformResult::passthrough(enforce_model(
event, model,
)))
@ -25,7 +25,7 @@ impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig {
&self,
event: &ResponsesWsEvent,
_model: &str,
) -> CoreResult<ResponsesWsTransformResult> {
) -> Result<ResponsesWsTransformResult, Error> {
Ok(ResponsesWsTransformResult::passthrough(event.clone()))
}
}

View file

@ -1,4 +1,4 @@
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::error::{Error, json_type_name};
use crate::ocr::transformation::OcrProviderConfig;
use crate::ocr::types::{OcrRequestData, OcrResponseData};
use serde_json::{Map, Value, json};
@ -43,7 +43,7 @@ pub fn is_deepseek_model(model: &str) -> bool {
pub fn resolve_vertex_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
api_key
.map(str::trim)
.filter(|key| !key.is_empty())
@ -51,7 +51,7 @@ pub fn resolve_vertex_api_key(
.or_else(|| env_lookup(VERTEX_AI_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.or_else(|| env_lookup(VERTEXAI_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or_else(|| {
CoreError::Auth(
Error::Auth(
"Missing Vertex AI access token - pass api_key or provide Authorization via extra_headers"
.to_string(),
)
@ -61,12 +61,12 @@ pub fn resolve_vertex_api_key(
fn vertex_project(
params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
string_param(params, &["vertex_project", "vertex_ai_project"])
.map(str::to_string)
.or_else(|| env_lookup(VERTEXAI_PROJECT_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| {
CoreError::InvalidRequest(
Error::InvalidRequest(
"Missing vertex_project - Set VERTEXAI_PROJECT environment variable or pass vertex_project parameter"
.to_string(),
)
@ -99,7 +99,7 @@ pub fn complete_vertex_mistral_url(
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
let project = vertex_project(optional_params, env_lookup)?;
let location = vertex_location(optional_params, env_lookup);
let base = vertex_mistral_api_base(api_base, &location);
@ -112,7 +112,7 @@ pub fn complete_vertex_deepseek_url(
api_base: Option<&str>,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
let project = vertex_project(optional_params, env_lookup)?;
let location = vertex_location(optional_params, env_lookup);
let base = api_base
@ -125,20 +125,20 @@ pub fn complete_vertex_deepseek_url(
))
}
fn document_content_item(document: &Value) -> CoreResult<Value> {
let object = document.as_object().ok_or_else(|| CoreError::InvalidType {
fn document_content_item(document: &Value) -> Result<Value, Error> {
let object = document.as_object().ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(document),
})?;
let doc_type = object
.get("type")
.and_then(Value::as_str)
.ok_or(CoreError::MissingField("document.type"))?;
.ok_or(Error::MissingField("document.type"))?;
let url_field = match doc_type {
"image_url" => "image_url",
"document_url" => "document_url",
other => {
return Err(CoreError::InvalidRequest(format!(
return Err(Error::InvalidRequest(format!(
"Unsupported document type: {other}. Expected 'image_url' or 'document_url'"
)));
}
@ -147,7 +147,7 @@ fn document_content_item(document: &Value) -> CoreResult<Value> {
.get(url_field)
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.ok_or(CoreError::MissingField(url_field))?;
.ok_or(Error::MissingField(url_field))?;
Ok(json!({
"type": "image_url",
@ -163,7 +163,7 @@ fn deepseek_model_name(model: &str) -> String {
}
}
fn first_choice_content(response: &Value) -> CoreResult<Value> {
fn first_choice_content(response: &Value) -> Result<Value, Error> {
response
.get("choices")
.and_then(Value::as_array)
@ -176,9 +176,7 @@ fn first_choice_content(response: &Value) -> CoreResult<Value> {
Value::Object(_) => true,
_ => false,
})
.ok_or_else(|| {
CoreError::InvalidResponse("No content in DeepSeek OCR response".to_string())
})
.ok_or_else(|| Error::InvalidResponse("No content in DeepSeek OCR response".to_string()))
}
fn ocr_data_from_content(content: Value, usage: Option<Value>, model: &str) -> Value {
@ -219,7 +217,7 @@ impl OcrProviderConfig for VertexAiOcrConfig {
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData> {
) -> Result<OcrRequestData, Error> {
MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params)
}
@ -227,7 +225,7 @@ impl OcrProviderConfig for VertexAiOcrConfig {
&self,
model: &str,
response_json: Value,
) -> CoreResult<OcrResponseData> {
) -> Result<OcrResponseData, Error> {
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
}
@ -237,7 +235,7 @@ impl OcrProviderConfig for VertexAiOcrConfig {
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
complete_vertex_mistral_url(api_base, model, optional_params, env_lookup)
}
@ -245,7 +243,7 @@ impl OcrProviderConfig for VertexAiOcrConfig {
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_vertex_api_key(api_key, env_lookup)
}
@ -264,7 +262,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
model: &str,
document: Value,
optional_params: Map<String, Value>,
) -> CoreResult<OcrRequestData> {
) -> Result<OcrRequestData, Error> {
let mut data = Map::new();
data.insert(
"model".to_string(),
@ -289,10 +287,10 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
&self,
model: &str,
response_json: Value,
) -> CoreResult<OcrResponseData> {
) -> Result<OcrResponseData, Error> {
let response = response_json
.as_object()
.ok_or_else(|| CoreError::InvalidType {
.ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
@ -314,7 +312,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
});
}
let object = ocr_data.as_object().ok_or_else(|| CoreError::InvalidType {
let object = ocr_data.as_object().ok_or_else(|| Error::InvalidType {
expected: "object",
actual: json_type_name(&ocr_data),
})?;
@ -346,7 +344,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
_model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
complete_vertex_deepseek_url(api_base, optional_params, env_lookup)
}
@ -354,7 +352,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
) -> Result<String, Error> {
resolve_vertex_api_key(api_key, env_lookup)
}
}

View file

@ -1,4 +1,4 @@
use crate::CoreResult;
use crate::Error;
use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
pub trait RealtimeProviderConfig {
@ -11,12 +11,12 @@ pub trait RealtimeProviderConfig {
&self,
event: &RealtimeEvent,
model: &str,
) -> CoreResult<RealtimeTransformResult>;
) -> Result<RealtimeTransformResult, Error>;
/// Transform a backend → client event before it is forwarded downstream.
fn transform_realtime_response(
&self,
event: &RealtimeEvent,
model: &str,
) -> CoreResult<RealtimeTransformResult>;
) -> Result<RealtimeTransformResult, Error>;
}

View file

@ -5,9 +5,9 @@ use std::time::{SystemTime, UNIX_EPOCH};
use serde_json::Value;
use crate::Error;
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType};
use crate::{CoreError, CoreResult};
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct ResponsesWsUsage {
@ -205,7 +205,7 @@ impl ResponsesWsInstrumentation {
}
}
type LifecycleFuture<'a, T> = Pin<Box<dyn Future<Output = CoreResult<T>> + Send + 'a>>;
type LifecycleFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation {
type PreCallFuture<'a> = LifecycleFuture<'a, ()>;
@ -246,7 +246,7 @@ impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation {
fn async_log_failure_event<'a>(
&'a self,
_context: &'a CallLifecycleContext,
_error: &'a CoreError,
_error: &'a Error,
_timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
Box::pin(async move {
@ -342,7 +342,7 @@ mod tests {
),
(),
&instrumentation,
|_| async { Ok::<(), CoreError>(()) },
|_| async { Ok::<(), Error>(()) },
)
.await;

View file

@ -1,4 +1,4 @@
use crate::CoreResult;
use crate::Error;
use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH};
use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult};
@ -19,13 +19,13 @@ pub trait ResponsesWebSocketProviderConfig: Sync {
&self,
event: &ResponsesWsEvent,
model: &str,
) -> CoreResult<ResponsesWsTransformResult>;
) -> Result<ResponsesWsTransformResult, Error>;
fn transform_ws_response(
&self,
event: &ResponsesWsEvent,
model: &str,
) -> CoreResult<ResponsesWsTransformResult>;
) -> Result<ResponsesWsTransformResult, Error>;
}
pub fn complete_websocket_url(

View file

@ -1,7 +1,8 @@
//! Enforcement: the litellm-rust workspace has exactly three crates.
//! Enforcement: the litellm-rust workspace has exactly four crates.
//!
//! `core` (pure translation), `ai-gateway` (routes + all network I/O), and
//! `python-bridge` (the PyO3 cdylib). Adding or removing a crate must be a
//! `core` (the Rust SDK), `ai-gateway` (the HTTP/WebSocket host),
//! `python-interop` (domain-neutral PyO3 primitives), and `python-bridge` (the
//! PyO3 cdylib). Adding or removing a crate must be a
//! deliberate act: this test fails until the allowlist here is updated, forcing
//! whoever changes the crate set to justify the new crate per the rule that a
//! crate is a layer needing independent compilation / its own deps / a separate
@ -16,10 +17,15 @@ use std::path::{Path, PathBuf};
/// The one true crate set. Update BOTH this and `litellm-rust/AGENTS.md` when the
/// workspace legitimately gains or loses a crate.
const EXPECTED_MEMBERS: &[&str] = &["crates/core", "crates/ai-gateway", "crates/python-bridge"];
const EXPECTED_MEMBERS: &[&str] = &[
"crates/core",
"crates/ai-gateway",
"crates/python-interop",
"crates/python-bridge",
];
/// The crate subdirectory names that must exist under `crates/`.
const EXPECTED_CRATE_DIRS: &[&str] = &["core", "ai-gateway", "python-bridge"];
const EXPECTED_CRATE_DIRS: &[&str] = &["core", "ai-gateway", "python-interop", "python-bridge"];
const MISMATCH: &str = "litellm-rust crate set changed — update this allowlist AND litellm-rust/AGENTS.md, and justify the crate per the rule (crate = layer needing independent compilation / its own deps / a separate artifact).";

View file

@ -1,3 +1,3 @@
litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over the litellm-core route entrypoints (e.g. `litellm_core::messages::messages`).
litellm-python-bridge is the PyO3 cdylib that exposes LiteLLM Rust APIs to the Python SDK. Keep API registration, domain dependency wiring, request assembly, and Python exception mapping here. Put domain-neutral Python/Serde conversion and GIL primitives in litellm-python-interop.
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call the core entrypoint.

View file

@ -5,8 +5,9 @@ Rules for `litellm-rust/crates/python-bridge`.
## Responsibility
`python-bridge` is the PyO3 boundary between Python LiteLLM and Rust transforms.
Keep this crate thin. It adapts Python objects to Rust payloads and returns
Python-compatible dictionaries.
Keep this crate thin. It exposes LiteLLM Rust APIs, assembles domain requests,
maps domain errors to Python exceptions, and delegates generic conversion and
GIL handling to `litellm-python-interop`.
## Bridge Shape

View file

@ -13,14 +13,14 @@ crate-type = ["cdylib"]
default = ["abi3"]
abi3 = ["pyo3/abi3-py310"]
extension-module = ["pyo3/extension-module"]
panic-test = []
[dependencies]
litellm-core = { workspace = true, features = ["bedrock-auth"] }
litellm-ai-gateway = { workspace = true, default-features = false }
litellm-python-interop.workspace = true
pyo3.workspace = true
pyo3-async-runtimes.workspace = true
pythonize.workspace = true
serde.workspace = true
serde_json.workspace = true
tokio.workspace = true

View file

@ -2,6 +2,7 @@ use std::hint::black_box;
use std::time::Duration;
use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main};
use litellm_python_interop::{from_py, to_py};
use pyo3::prelude::*;
use pyo3::types::PyDict;
use serde_json::{Value, json};
@ -25,7 +26,7 @@ fn former_json_roundtrip_from_py(py: Python<'_>, value: &Bound<'_, PyAny>) -> Va
}
fn pythonize_from_py(value: &Bound<'_, PyAny>) -> Value {
pythonize::depythonize(value).expect("payload should depythonize")
from_py(value).expect("payload should depythonize")
}
fn former_json_roundtrip_to_py(py: Python<'_>, value: &Value) -> Py<PyAny> {
@ -37,12 +38,10 @@ fn former_json_roundtrip_to_py(py: Python<'_>, value: &Value) -> Py<PyAny> {
}
fn pythonize_to_py(py: Python<'_>, value: &Value) -> Py<PyAny> {
pythonize::pythonize(py, value)
.expect("response should pythonize")
.unbind()
to_py(py, value).expect("response should pythonize")
}
fn serialization(c: &mut Criterion) {
fn bridge_serialization(c: &mut Criterion) {
Python::initialize();
Python::attach(|py| {
for &(label, payload_bytes) in PAYLOAD_SIZES {
@ -98,6 +97,6 @@ criterion_group! {
.sample_size(20)
.warm_up_time(Duration::from_secs(1))
.measurement_time(Duration::from_secs(4));
targets = serialization
targets = bridge_serialization
}
criterion_main!(benches);

View file

@ -1,32 +0,0 @@
//! GIL accounting.
//!
//! A single chokepoint for releasing the GIL around blocking work. Every
//! blocking call in the bridge goes through [`release_gil`] instead of calling
//! `Python::detach` directly, so the release count stays accurate and we
//! have one place to extend later (timing histograms, per-call labels, etc.).
use std::sync::atomic::{AtomicU64, Ordering};
use pyo3::prelude::*;
/// Number of times the bridge has released the GIL since process start.
static GIL_RELEASES: AtomicU64 = AtomicU64::new(0);
/// Release the GIL around `f`, recording the release.
///
/// `f` must not touch any Python state — that is what makes releasing the GIL
/// safe. Returning the value back to Python re-acquires the GIL at the call
/// site, after `f` has finished.
pub fn release_gil<T, F>(py: Python<'_>, f: F) -> T
where
F: FnOnce() -> T + Send,
T: Send,
{
GIL_RELEASES.fetch_add(1, Ordering::Relaxed);
py.detach(f)
}
/// Total GIL releases performed by the bridge so far.
pub fn release_count() -> u64 {
GIL_RELEASES.load(Ordering::Relaxed)
}

View file

@ -10,19 +10,15 @@ use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompleti
use litellm_core::chat_completions::{
chat_completions as run_chat_completions, chat_completions_decline_reason,
};
use litellm_core::error::CoreError;
use litellm_core::error::Error;
use litellm_core::messages::messages as run_messages;
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
use litellm_python_interop::{from_py, release_count, release_gil, to_py};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::{PyAny, PyDict};
use serde_json::{Map, Value};
mod gil;
mod marshal;
use marshal::{from_py, to_py};
pyo3::create_exception!(
_native,
RustBridgeDeclined,
@ -58,13 +54,13 @@ fn chat_completions_response_to_py(
to_py(py, &response)
}
fn core_error_to_pyerr(err: CoreError) -> PyErr {
fn core_error_to_pyerr(err: Error) -> PyErr {
match err {
CoreError::Auth(message) => PyValueError::new_err(message),
CoreError::InvalidProvider(_)
| CoreError::InvalidRequest(_)
| CoreError::InvalidType { .. }
| CoreError::MissingField(_) => PyValueError::new_err(err.to_string()),
Error::Auth(message) => PyValueError::new_err(message),
Error::InvalidProvider(_)
| Error::InvalidRequest(_)
| Error::InvalidType { .. }
| Error::MissingField(_) => PyValueError::new_err(err.to_string()),
other => PyRuntimeError::new_err(other.to_string()),
}
}
@ -75,22 +71,22 @@ fn core_error_to_pyerr(err: CoreError) -> PyErr {
/// Everything raised before the request goes out is safe for the host to retry
/// on its own path; anything after it is not, because the provider has already
/// done the work and billed for it.
fn chat_completions_error_to_pyerr(err: CoreError) -> PyErr {
fn chat_completions_error_to_pyerr(err: Error) -> PyErr {
match err {
CoreError::Unsupported(_)
| CoreError::Auth(_)
| CoreError::InvalidProvider(_)
| CoreError::InvalidRequest(_)
| CoreError::InvalidType { .. }
| CoreError::MissingField(_)
| CoreError::Routing(_)
Error::Unsupported(_)
| Error::Auth(_)
| Error::InvalidProvider(_)
| Error::InvalidRequest(_)
| Error::InvalidType { .. }
| Error::MissingField(_)
| Error::Routing(_)
// Nothing reached the provider, so serving it on Python cannot double
// bill and is the only way the caller gets an answer at all.
| CoreError::Connect(_) => RustBridgeDeclined::new_err(err.to_string()),
CoreError::Http { status, body } => {
| Error::Connect(_) => RustBridgeDeclined::new_err(err.to_string()),
Error::Http { status, body } => {
RustUpstreamError::new_err((status, format!("{status}: {body}")))
}
CoreError::Network(message) | CoreError::InvalidResponse(message) => {
Error::Network(message) | Error::InvalidResponse(message) => {
RustUpstreamError::new_err((0u16, message))
}
}
@ -230,7 +226,7 @@ fn ocr(
timeout_seconds,
)?;
let result = gil::release_gil(py, || {
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_ocr(OcrRequest {
model: &model,
document,
@ -318,7 +314,7 @@ fn transcription(
};
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let timeout = optional_timeout(timeout_seconds);
let result = gil::release_gil(py, || {
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_audio_transcription(
AudioTranscriptionRequest {
model: &model,
@ -419,7 +415,7 @@ fn messages(
let (body, extra_headers, timeout) =
marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?;
let result = gil::release_gil(py, || {
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_messages(MessagesRequest {
model: &model,
body,
@ -546,7 +542,7 @@ fn chat_completions(
timeout_seconds,
)?;
let result = gil::release_gil(py, || {
let result = release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_chat_completions(
ChatCompletionsRequest {
model: &model,
@ -610,10 +606,16 @@ fn achat_completions(
#[pyfunction]
fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
let stats = PyDict::new(py);
stats.set_item("releases", gil::release_count())?;
stats.set_item("releases", release_count())?;
Ok(stats.into_any().unbind())
}
#[cfg(feature = "panic-test")]
#[pyfunction]
fn _panic_for_test() {
panic!("intentional PyO3 panic smoke test");
}
#[pymodule]
fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> {
let py = module.py();
@ -630,5 +632,7 @@ fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(achat_completions, module)?)?;
module.add_class::<ResponsesWebSocketConnection>()?;
module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
#[cfg(feature = "panic-test")]
module.add_function(wrap_pyfunction!(_panic_for_test, module)?)?;
Ok(())
}

View file

@ -1,7 +1,7 @@
use std::fs;
use std::path::{Path, PathBuf};
const DISALLOWED_OUTSIDE_MARSHAL: &[&str] = &[
const DISALLOWED_OUTSIDE_INTEROP: &[&str] = &[
"py.import(\"json\")",
"pythonize::",
"serde_json::to_string",
@ -33,18 +33,15 @@ fn rust_sources(directory: &Path) -> Vec<PathBuf> {
}
#[test]
fn serialization_is_centralized_in_marshal_module() {
fn serialization_uses_the_interop_boundary() {
let root = source_root();
for path in rust_sources(&root) {
if path == root.join("marshal.rs") {
continue;
}
let source = fs::read_to_string(&path).expect("bridge source should be readable");
for disallowed in DISALLOWED_OUTSIDE_MARSHAL {
for disallowed in DISALLOWED_OUTSIDE_INTEROP {
assert!(
!source.contains(disallowed),
"{} bypasses the typed marshal module with `{disallowed}`",
"{} bypasses litellm-python-interop with `{disallowed}`",
path.display()
);
}

View file

@ -0,0 +1 @@
litellm-python-interop is the domain-neutral PyO3 foundation. Keep generic Python/Serde conversion and interpreter primitives here. Do not add LiteLLM domain crates, route types, API registration, or cdylib build features.

View file

@ -0,0 +1,15 @@
[package]
name = "litellm-python-interop"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
pyo3.workspace = true
pythonize.workspace = true
serde.workspace = true
[dev-dependencies]
rstest.workspace = true
serde_json.workspace = true

View file

@ -0,0 +1,21 @@
use std::sync::atomic::{AtomicU64, Ordering};
use pyo3::prelude::*;
static GIL_RELEASES: AtomicU64 = AtomicU64::new(0);
/// Runs work detached from the interpreter and records the release.
///
/// `f` must not access Python state while the interpreter is detached.
pub fn release_gil<T, F>(py: Python<'_>, f: F) -> T
where
F: FnOnce() -> T + Send,
T: Send,
{
GIL_RELEASES.fetch_add(1, Ordering::Relaxed);
py.detach(f)
}
pub fn release_count() -> u64 {
GIL_RELEASES.load(Ordering::Relaxed)
}

View file

@ -0,0 +1,5 @@
mod gil;
mod marshal;
pub use gil::{release_count, release_gil};
pub use marshal::{from_py, to_py};

View file

@ -0,0 +1,44 @@
use pyo3::Python;
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use litellm_python_interop::{from_py, release_count, release_gil, to_py};
struct InitializedPython;
impl InitializedPython {
fn attach<F, R>(&self, f: F) -> R
where
F: for<'py> FnOnce(Python<'py>) -> R,
{
Python::attach(f)
}
}
#[fixture]
#[once]
fn initialized_python() -> InitializedPython {
Python::initialize();
InitializedPython
}
#[rstest]
fn serde_values_round_trip_through_python(#[from(initialized_python)] python: &InitializedPython) {
python.attach(|py| {
let expected = json!({"model": "test", "items": [1, true, null]});
let python_value = to_py(py, &expected).expect("value should convert to Python");
let actual: Value =
from_py(python_value.bind(py)).expect("Python value should convert to serde");
assert_eq!(actual, expected);
});
}
#[rstest]
fn release_gil_runs_work_and_records_it(#[from(initialized_python)] python: &InitializedPython) {
let before = release_count();
let result = python.attach(|py| release_gil(py, || 42));
assert_eq!(result, 42);
assert_eq!(release_count(), before + 1);
}

View file

@ -10,8 +10,19 @@ from collections.abc import AsyncIterator, Mapping
from typing import Any, Final
from litellm._logging import verbose_logger
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
from litellm.a2a_protocol.utils import (
get_session_id_from_a2a_params,
scope_session_to_principal,
)
from litellm.exceptions import BadRequestError
from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig
RUNTIME_SESSION_ID_MIN_LENGTH: Final = 33
RUNTIME_SESSION_ID_MAX_LENGTH: Final = 256
# Reserved outbound header names that must never be sourced from per-request
# ``agent_extra_headers`` for AgentCore requests. ``agent_extra_headers`` carries
# values rewritten from the client-controlled ``x-a2a-{agent}-*`` convention, so
@ -19,8 +30,9 @@ from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreCo
# request identity / SigV4 metadata by overwriting headers the proxy sets from
# trusted server-side config.
#
# The runtime headers (session / user id) are derived server-side from
# ``runtimeSessionId`` / ``runtimeUserId`` in the agent's ``litellm_params``;
# The runtime headers (session / user id) are derived server-side from the A2A
# ``message.contextId`` and ``runtimeSessionId`` / ``runtimeUserId`` in the
# agent's ``litellm_params``;
# ``authorization`` is set by the AgentCore signer (JWT or SigV4); ``host`` and
# the ``x-amz-*`` family are owned by SigV4 itself.
_RESERVED_EXACT_HEADERS: Final = frozenset(
@ -66,6 +78,31 @@ def _filter_reserved_headers(
return filtered or None
def _request_scoped_runtime_session_id(
params: Mapping[str, Any],
litellm_params: Mapping[str, Any],
) -> str | None:
context_id: Final = get_session_id_from_a2a_params(params)
if not isinstance(context_id, str) or not context_id:
return None
return scope_session_to_principal(context_id, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM))
def _validate_runtime_session_id(session_id: str, model: str) -> str:
if RUNTIME_SESSION_ID_MIN_LENGTH <= len(session_id) <= RUNTIME_SESSION_ID_MAX_LENGTH:
return session_id
raise BadRequestError(
message=(
f"Invalid AgentCore runtime session id {session_id!r}: AWS requires "
f"{RUNTIME_SESSION_ID_MIN_LENGTH}-{RUNTIME_SESSION_ID_MAX_LENGTH} characters. It is built from the A2A "
"message.contextId (prefixed with a 16-hex-char hash of the calling key and '-') when set, "
"otherwise from the agent's configured runtimeSessionId."
),
model=model,
llm_provider="bedrock",
)
class BedrockAgentCoreA2ATransformation:
"""
Request/response transformation for Bedrock AgentCore A2A agents.
@ -100,7 +137,9 @@ class BedrockAgentCoreA2ATransformation:
here to prevent a caller-controlled ``x-a2a-{agent}-*`` header from
spoofing the AgentCore runtime user id or other SigV4 metadata. Use
``api_key`` / ``runtimeUserId`` / ``runtimeSessionId`` in litellm_params
(not ``agent_extra_headers``) to override those values.
(not ``agent_extra_headers``) to override those values. The runtime
session id is taken from ``params["message"]["contextId"]`` (scoped to
the calling key) when present, then ``runtimeSessionId``, else generated.
Returns:
Tuple of (url, signed_headers, signed_body_bytes)
@ -139,7 +178,11 @@ class BedrockAgentCoreA2ATransformation:
# Set required AgentCore session headers (normally set by transform_request,
# which we skip because it also builds {"prompt": "..."})
headers: Final[dict] = {}
session_id: Final = agentcore_config._get_runtime_session_id(optional_params)
session_id: Final = _validate_runtime_session_id(
_request_scoped_runtime_session_id(params, litellm_params)
or agentcore_config._get_runtime_session_id(optional_params),
model=model,
)
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = session_id
runtime_user_id: Final = agentcore_config._get_runtime_user_id(optional_params)
if runtime_user_id:

View file

@ -2,6 +2,8 @@
Utility functions for A2A protocol.
"""
import hashlib
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import litellm
@ -140,6 +142,29 @@ class A2ARequestUtils:
return prompt_tokens, completion_tokens, total_tokens
def get_session_id_from_a2a_params(params: Mapping[str, Any]) -> str | None:
message: Final = params.get("message", {})
if isinstance(message, dict):
return message.get("contextId")
return getattr(message, "contextId", None)
def scope_session_to_principal(session_id: str, principal: str | None) -> str:
"""
Bind a client-supplied A2A contextId to the authenticated principal.
Without this, two distinct keys authorized for the same agent could set the
same contextId and read/append to each other's backend memory. The
principal is hashed (it is already a hashed token) so the raw value is never
sent to the agent backend, while the original contextId is kept as a suffix
for operator-side correlation.
"""
if not principal:
return session_id
principal_prefix: Final = hashlib.sha256(principal.encode("utf-8")).hexdigest()[:16]
return f"{principal_prefix}-{session_id}"
# Backwards compatibility aliases
def extract_text_from_a2a_message(message: Any) -> str:
return A2ARequestUtils.extract_text_from_message(message)

View file

@ -9,6 +9,38 @@ DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT"
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5))
ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS: Final = 2000
RUNTIME_UPDATABLE_ROUTER_SETTINGS: Final[frozenset[str]] = frozenset(
{
"routing_strategy_args",
"routing_strategy",
"routing_groups",
"allowed_fails",
"cooldown_time",
"num_retries",
"timeout",
"max_retries",
"retry_after",
"fallbacks",
"context_window_fallbacks",
"retry_policy",
"model_group_retry_policy",
"model_group_alias",
"enable_weighted_failover",
"enable_tag_filtering",
"tag_routing_prefix",
"optional_pre_call_checks",
}
)
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
{
"model_list",
"search_tools",
"assistants_config",
"router_general_settings",
"ignore_invalid_deployments",
"fallback_access_check",
}
)
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))

View file

@ -1,7 +1,7 @@
import asyncio
import contextvars
import json
from collections.abc import Coroutine, Mapping
from collections.abc import Callable, Coroutine, Mapping
from functools import partial
from typing import Final, Literal, overload
@ -47,6 +47,13 @@ __all__ = [
##### Container Create #######################
async def _encode_created_container_id(
pending: Coroutine[object, object, ContainerObject],
encode: Callable[[ContainerObject], ContainerObject],
) -> ContainerObject:
return encode(await pending)
@client
async def acreate_container(
name: str,
@ -256,16 +263,16 @@ def create_container(
_is_async=_is_async,
)
# Encode container_id with provider/model metadata for routing
encode: Final = partial(
ContainerRequestUtils.encode_container_id_in_response,
custom_llm_provider=custom_llm_provider,
litellm_metadata=kwargs.get("litellm_metadata"),
extra_body=extra_body,
)
if isinstance(container_obj, ContainerObject):
container_obj = ContainerRequestUtils.encode_container_id_in_response(
response_obj=container_obj,
custom_llm_provider=custom_llm_provider,
litellm_metadata=kwargs.get("litellm_metadata"),
extra_body=extra_body,
)
return encode(container_obj)
return container_obj
return _encode_created_container_id(pending=container_obj, encode=encode)
except Exception as e:
raise litellm.exception_type(

View file

@ -1,4 +1,5 @@
import contextvars
import copy
import hashlib
import os
import secrets
@ -39,6 +40,7 @@ except ImportError:
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
dc: Final = DualCache()
@ -852,6 +854,69 @@ class CustomGuardrail(CustomLogger):
return result
async def async_logging_hook(
self,
kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract
result: object,
call_type: str,
) -> tuple[dict, object]: # mutable-ok: CustomLogger.async_logging_hook contract
"""logging_only: run apply_guardrail on copies of the logged request/response and record the verdict."""
from litellm.llms import get_guardrail_translation_mapping
if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks:
return kwargs, result
try:
translation: Final = get_guardrail_translation_mapping(CallTypes(call_type))()
except ValueError:
verbose_logger.debug(
"Guardrail %s: no guardrail translation for call_type=%s, skipping logging_only scan",
self.guardrail_name,
call_type,
)
return kwargs, result
litellm_params: Final = kwargs.get("litellm_params") or {}
scratch_metadata: Final = {
key: value
for key, value in (litellm_params.get("metadata") or {}).items()
if key != "standard_logging_guardrail_information"
}
try:
await self._scan_logged_call(kwargs, result, translation, scratch_metadata)
except Exception as e:
verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e)
recorded: Final = scratch_metadata.get("standard_logging_guardrail_information")
standard_logging_object: Final = kwargs.get("standard_logging_object")
if not recorded or not isinstance(standard_logging_object, dict):
return kwargs, result
entries: Final = recorded if isinstance(recorded, list) else [recorded]
existing: Final = standard_logging_object.get("guardrail_information") or []
return {
**kwargs,
"standard_logging_object": {**standard_logging_object, "guardrail_information": [*existing, *entries]},
}, result
async def _scan_logged_call(
self,
kwargs: dict, # mutable-ok: CustomLogger.async_logging_hook contract
result: object,
translation: "BaseTranslation",
scratch_metadata: dict, # mutable-ok: apply_guardrail records its verdict into request metadata
) -> None:
optional_params: Final = kwargs.get("optional_params") or {}
scratch_input: Final = copy.deepcopy(kwargs.get("messages") or kwargs.get("input"))
scratch_request: Final = {
"model": kwargs.get("model"),
"messages": scratch_input,
"input": scratch_input,
"tools": copy.deepcopy(optional_params.get("tools")),
"litellm_call_id": kwargs.get("litellm_call_id"),
"metadata": scratch_metadata,
}
await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self)
await translation.process_output_response(
response=copy.deepcopy(result), guardrail_to_apply=self, request_data=scratch_request
)
def supports_scan_only_tool_results(self) -> bool:
"""Whether this guardrail can scan tool-result content.

View file

@ -11,6 +11,7 @@ import json
import os
from collections.abc import Mapping, Sequence
from datetime import datetime
from types import MappingProxyType
from typing import Any, Final, Literal
import httpx
@ -30,12 +31,16 @@ from litellm.integrations.datadog.datadog_mock_client import (
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
handle_any_messages_to_chat_completion_str_messages_conversion,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.spend_tracking.savings import extract_cache_creation_tokens, extract_cache_read_tokens
from litellm.types.integrations.datadog_llm_obs import *
from litellm.types.utils import (
CallTypes,
@ -44,6 +49,189 @@ from litellm.types.utils import (
StandardLoggingPayloadErrorInformation,
)
_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({})
_EMPTY_MESSAGE: Final[Message] = {"role": "", "content": ""}
_MAX_PARSED_TOOL_ARGUMENT_CHARS: Final = 256 * 1024
def _mapping_field(source: Mapping[str, Any], key: str) -> Mapping[str, Any]:
"""The value at `key` when it is a mapping, else an empty one."""
value: Final = source.get(key)
return value if isinstance(value, dict) else _EMPTY_MAPPING
def _content_blocks(message: Mapping[str, Any]) -> tuple[Mapping[str, Any], ...]:
content: Final = message.get("content")
if not isinstance(content, list):
return ()
return tuple(block for block in content if isinstance(block, dict))
def _to_dd_arguments(raw_arguments: object) -> dict[str, Any] | str:
"""
Arguments as the object LLM Obs types them as, or the raw string when they are not one.
Strings past the size bound ship unparsed: decoding multiplies memory on hostile compact
JSON, and the raw string is what the intake receives either way.
"""
if not isinstance(raw_arguments, str):
return raw_arguments if isinstance(raw_arguments, dict) else str(raw_arguments)
if len(raw_arguments) > _MAX_PARSED_TOOL_ARGUMENT_CHARS:
return raw_arguments
parsed: Final = safe_json_loads(raw_arguments)
return parsed if isinstance(parsed, dict) else raw_arguments
def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]:
"""
The tool calls a message carries, in LLM Obs' ToolCall schema, from either dialect.
OpenAI puts them in `tool_calls` with the callee nested under `function` and `arguments`
serialized; Anthropic puts them in `content` as `tool_use` blocks with `input` already an
object. LLM Obs reads `name` / `arguments` / `tool_id` either way.
"""
raw_tool_calls: Final = message.get("tool_calls")
openai_calls: Final = tuple(
ToolCall(
name=function.get("name", ""),
arguments=_to_dd_arguments(function.get("arguments", "")),
tool_id=tool_call.get("id", ""),
type=tool_call.get("type", "function"),
)
for tool_call in (raw_tool_calls if isinstance(raw_tool_calls, list) else ())
if isinstance(tool_call, dict)
for function in [_mapping_field(tool_call, "function")]
)
anthropic_calls: Final = tuple(
ToolCall(
name=block.get("name", ""),
arguments=_to_dd_arguments(block.get("input") or {}),
tool_id=block.get("id", ""),
type="tool_use",
)
for block in _content_blocks(message)
if block.get("type") == "tool_use"
)
return openai_calls + anthropic_calls
def _to_dd_tool_results(message: Mapping[str, Any], tool_call_names: Mapping[str, str]) -> tuple[ToolResult, ...]:
"""
The tool results a message carries, linked back to the call each answers.
OpenAI models a result as a whole `role: "tool"` message keyed by `tool_call_id`;
Anthropic nests `tool_result` blocks inside a user message, keyed by `tool_use_id`.
"""
def to_result(tool_id: str, result: object) -> ToolResult:
return ToolResult(
name=tool_call_names.get(tool_id, ""),
result=result if isinstance(result, str) else safe_dumps(result),
tool_id=tool_id,
type="function",
)
if message.get("role") == "tool":
return (to_result(str(message.get("tool_call_id", "")), message.get("content") or ""),)
return tuple(
to_result(str(block.get("tool_use_id", "")), block.get("content") or "")
for block in _content_blocks(message)
if block.get("type") == "tool_result"
)
def _tool_call_names_by_id(messages: Sequence[object]) -> Mapping[str, str]:
"""Ids to tool names for result linking; reads names structurally and parses nothing."""
openai_pairs: Final = tuple(
(tool_call.get("id"), function.get("name", ""))
for message in messages
if isinstance(message, dict) and isinstance(message.get("tool_calls"), list)
for tool_call in message["tool_calls"]
if isinstance(tool_call, dict)
for function in [_mapping_field(tool_call, "function")]
)
anthropic_pairs: Final = tuple(
(block.get("id"), block.get("name", ""))
for message in messages
if isinstance(message, dict)
for block in _content_blocks(message)
if block.get("type") == "tool_use"
)
return MappingProxyType({str(tool_id): str(name) for tool_id, name in openai_pairs + anthropic_pairs if tool_id})
def _to_dd_message(message: object, tool_call_names: Mapping[str, str]) -> Message:
"""
Map one chat message onto LLM Obs' Message schema, adding fields and never destroying content.
Content collapses to its text only when it has text; a content list with none (tool blocks,
images) rides along unchanged so nothing the caller logged is lost. Tool calls and results
move into the fields the LLM Obs Tools panel reads, from both the OpenAI and Anthropic shapes.
"""
if not isinstance(message, dict):
converted: Final = handle_any_messages_to_chat_completion_str_messages_conversion(message)
return converted[0] if converted else _EMPTY_MESSAGE
text: Final = convert_content_list_to_str(message) # pyright: ignore[reportArgumentType] # caller-supplied dict
original_content: Final = message.get("content")
content: Final = (
text if text or not isinstance(original_content, list) or not original_content else original_content
)
reasoning: Final = message.get("reasoning_content")
tool_calls: Final = _to_dd_tool_calls(message)
tool_results: Final = _to_dd_tool_results(message, tool_call_names)
dd_message: Final[Message] = {
"role": message.get("role", ""),
"content": content,
**({"reasoning_content": reasoning} if reasoning is not None else {}),
**({"tool_calls": tool_calls} if tool_calls else {}),
**({"tool_results": tool_results} if tool_results else {}),
}
return dd_message
def _to_dd_messages(messages: object) -> tuple[Message, ...]:
"""Map a whole conversation, resolving each tool result against the calls that precede it."""
if messages is None:
return ()
if not isinstance(messages, list):
return tuple(handle_any_messages_to_chat_completion_str_messages_conversion(messages))
tool_call_names: Final = _tool_call_names_by_id(messages)
return tuple(_to_dd_message(message, tool_call_names) for message in messages)
def _to_dd_tool_definition(entry: Mapping[str, Any]) -> ToolDefinition | None:
function: Final = entry.get("function")
declared: Final[Mapping[str, Any]] = function if isinstance(function, dict) else entry
name: Final = declared.get("name")
if not name:
return None
schema: Final = declared.get("parameters") or declared.get("input_schema")
description: Final = declared.get("description", "")
if not isinstance(schema, dict):
return ToolDefinition(name=name, description=description)
return ToolDefinition(name=name, description=description, schema=schema)
def _to_dd_tool_definitions(model_parameters: object) -> tuple[ToolDefinition, ...]:
"""
Map the request's declared tools onto LLM Obs' ToolDefinition schema.
Handles the wrapped chat-completions shape and the bare shape the Anthropic and
Responses surfaces use, since both reach this logger through `model_parameters`.
"""
if not isinstance(model_parameters, dict):
return ()
raw_tools: Final = model_parameters.get("tools") or model_parameters.get("functions")
if not isinstance(raw_tools, list):
return ()
return tuple(
definition
for entry in raw_tools
if isinstance(entry, dict)
if (definition := _to_dd_tool_definition(entry)) is not None
)
class DataDogLLMObsLogger(CustomBatchLogger):
def __init__(self, **kwargs):
@ -222,12 +410,9 @@ class DataDogLLMObsLogger(CustomBatchLogger):
if standard_logging_payload is None:
raise Exception("DataDogLLMObs: standard_logging_object is not set")
messages = standard_logging_payload["messages"]
messages = self._ensure_string_content(messages=messages)
metadata: Final = kwargs.get("litellm_params", {}).get("metadata", {})
input_meta: Final = InputMeta(messages=handle_any_messages_to_chat_completion_str_messages_conversion(messages))
input_meta: Final = InputMeta(messages=_to_dd_messages(standard_logging_payload["messages"]))
output_meta: Final = OutputMeta(
messages=self._get_response_messages(
standard_logging_payload=standard_logging_payload,
@ -241,22 +426,20 @@ class DataDogLLMObsLogger(CustomBatchLogger):
if isinstance(metadata, dict):
metadata_parent_id = metadata.get("parent_id")
meta: Final = Meta(
kind=self._get_datadog_span_kind(standard_logging_payload.get("call_type"), metadata_parent_id),
input=input_meta,
output=output_meta,
metadata=self._get_dd_llm_obs_payload_metadata(standard_logging_payload),
error=error_info,
)
tool_definitions: Final = _to_dd_tool_definitions(standard_logging_payload.get("model_parameters"))
span_kind: Final = self._get_datadog_span_kind(standard_logging_payload.get("call_type"), metadata_parent_id)
payload_metadata: Final = self._get_dd_llm_obs_payload_metadata(standard_logging_payload)
# Calculate metrics (you may need to adjust these based on available data)
metrics: Final = LLMMetrics(
input_tokens=float(standard_logging_payload.get("prompt_tokens", 0)),
output_tokens=float(standard_logging_payload.get("completion_tokens", 0)),
total_tokens=float(standard_logging_payload.get("total_tokens", 0)),
total_cost=float(standard_logging_payload.get("response_cost", 0)),
time_to_first_token=self._get_time_to_first_token_seconds(standard_logging_payload),
)
meta: Final[Meta] = {
"kind": span_kind,
"input": input_meta,
"output": output_meta,
"metadata": payload_metadata,
"error": error_info,
**({"tool_definitions": tool_definitions} if tool_definitions else {}),
}
metrics: Final = self._assemble_metrics(standard_logging_payload)
payload: Final[LLMObsPayload] = LLMObsPayload(
parent_id=metadata_parent_id if metadata_parent_id else "undefined",
@ -314,6 +497,45 @@ class DataDogLLMObsLogger(CustomBatchLogger):
)
return error_info
def _assemble_metrics(self, standard_logging_payload: StandardLoggingPayload) -> LLMMetrics:
"""
Build the span metrics, including the prompt-cache counts LLM Obs charts cache savings from.
Cache counts resolve through the same owners the savings dashboard uses, so every provider
spelling is covered, and `non_cached_input_tokens` subtracts BOTH cache categories because
litellm's normalized prompt count includes both (the invariant the cost calculator's custom
pricing helper documents). A zero residual on a fully cached request is real data and is
emitted; a zero read or write count is absence and is not.
"""
prompt_tokens: Final = float(standard_logging_payload.get("prompt_tokens", 0))
completion_tokens: Final = float(standard_logging_payload.get("completion_tokens", 0))
total_tokens: Final = float(standard_logging_payload.get("total_tokens", 0))
total_cost: Final = float(standard_logging_payload.get("response_cost", 0))
time_to_first_token: Final = self._get_time_to_first_token_seconds(standard_logging_payload)
raw_usage: Final = (standard_logging_payload.get("metadata") or {}).get("usage_object")
usage_object: Final = raw_usage if isinstance(raw_usage, dict) else None
cache_read: Final = float(extract_cache_read_tokens(usage_object))
cache_write: Final = float(extract_cache_creation_tokens(usage_object))
metrics: Final[LLMMetrics] = {
"input_tokens": prompt_tokens,
"output_tokens": completion_tokens,
"total_tokens": total_tokens,
"total_cost": total_cost,
"time_to_first_token": time_to_first_token,
**(
{
**({"cache_read_input_tokens": cache_read} if cache_read else {}),
**({"cache_write_input_tokens": cache_write} if cache_write else {}),
"non_cached_input_tokens": max(prompt_tokens - cache_read - cache_write, 0.0),
}
if cache_read or cache_write
else {}
),
}
return metrics
def _get_time_to_first_token_seconds(self, standard_logging_payload: StandardLoggingPayload) -> float:
"""
Get the time to first token in seconds
@ -335,7 +557,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
def _get_response_messages(
self, standard_logging_payload: StandardLoggingPayload, call_type: str | None
) -> list[object]:
) -> tuple[Message, ...]:
"""
Get the messages from the response object
@ -344,7 +566,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
response_obj = standard_logging_payload.get("response")
if response_obj is None:
return []
return ()
# edge case: handle response_obj is a string representation of a dict
if isinstance(response_obj, str):
@ -357,7 +579,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
# fallback to json parsing
response_obj = json.loads(str(response_obj))
except json.JSONDecodeError:
return []
return ()
if call_type in [
CallTypes.completion.value,
@ -375,12 +597,12 @@ class DataDogLLMObsLogger(CustomBatchLogger):
if isinstance(response_obj, dict) and "choices" in response_obj:
choices: Final = response_obj["choices"]
if choices and len(choices) > 0 and "message" in choices[0]:
return [choices[0]["message"]]
return []
return _to_dd_messages([choices[0]["message"]])
return ()
except (KeyError, IndexError, TypeError):
# In case of any error accessing the response structure, return empty list
return []
return []
return ()
return ()
def _get_datadog_span_kind(
self, call_type: str | None, parent_id: str | None = None
@ -485,17 +707,6 @@ class DataDogLLMObsLogger(CustomBatchLogger):
# Default fallback for unknown or passthrough operations
return "llm"
def _ensure_string_content(self, messages: str | Sequence[object] | Mapping[object, object] | None) -> list[object]:
if messages is None:
return []
if isinstance(messages, str):
return [messages]
elif isinstance(messages, list):
return [message for message in messages]
elif isinstance(messages, dict):
return [str(messages.get("content", ""))]
return []
def _get_dd_llm_obs_payload_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
"""
Fields to track in DD LLM Observability metadata from litellm standard logging payload
@ -524,10 +735,6 @@ class DataDogLLMObsLogger(CustomBatchLogger):
spend_metrics: Final = self._get_spend_metrics(standard_logging_payload)
_metadata.update({"spend_metrics": dict(spend_metrics)})
## extract tool calls and add to metadata
tool_call_metadata: Final = self._extract_tool_call_metadata(standard_logging_payload)
_metadata.update(tool_call_metadata)
_standard_logging_metadata: Final[dict] = dict(standard_logging_payload.get("metadata", {})) or {}
_metadata.update(_standard_logging_metadata)
return _metadata
@ -647,107 +854,3 @@ class DataDogLLMObsLogger(CustomBatchLogger):
verbose_logger.debug("Original value: %s", user_api_key_budget_reset_at)
return spend_metrics
def _process_input_messages_preserving_tool_calls(self, messages: Sequence[object]) -> list[dict[str, object]]:
"""
Process input messages while preserving tool_calls and tool message types.
This bypasses the lossy string conversion when tool calls are present,
allowing complex nested tool_calls objects to be preserved for Datadog.
"""
processed: Final = []
for msg in messages:
if isinstance(msg, dict):
# Preserve messages with tool_calls or tool role as-is
if "tool_calls" in msg or msg.get("role") == "tool":
processed.append(msg)
else:
# For regular messages, still apply string conversion
converted = handle_any_messages_to_chat_completion_str_messages_conversion([msg])
processed.extend(converted)
else:
# For non-dict messages, apply string conversion
converted = handle_any_messages_to_chat_completion_str_messages_conversion([msg])
processed.extend(converted)
return processed
@staticmethod
def _tool_calls_kv_pair(tool_calls: list[dict[str, Any]]) -> dict[str, object]:
"""
Extract tool call information into key-value pairs for Datadog metadata.
Similar to OpenTelemetry's implementation but adapted for Datadog's format.
"""
kv_pairs: Final[dict[str, object]] = {}
for idx, tool_call in enumerate(tool_calls):
try:
# Extract tool call ID
tool_id = tool_call.get("id")
if tool_id:
kv_pairs[f"tool_calls.{idx}.id"] = tool_id
# Extract tool call type
tool_type = tool_call.get("type")
if tool_type:
kv_pairs[f"tool_calls.{idx}.type"] = tool_type
# Extract function information
function = tool_call.get("function")
if function:
function_name = function.get("name")
if function_name:
kv_pairs[f"tool_calls.{idx}.function.name"] = function_name
function_arguments = function.get("arguments")
if function_arguments:
# Store arguments as JSON string for Datadog
if isinstance(function_arguments, str):
kv_pairs[f"tool_calls.{idx}.function.arguments"] = function_arguments
else:
import json
kv_pairs[f"tool_calls.{idx}.function.arguments"] = json.dumps(function_arguments)
except (KeyError, TypeError, ValueError) as e:
verbose_logger.debug("DataDogLLMObs: Error processing tool call %s: %s", idx, e)
continue
return kv_pairs
def _extract_tool_call_metadata(self, standard_logging_payload: StandardLoggingPayload) -> dict[str, object]:
"""
Extract tool call information from both input messages and response for Datadog metadata.
"""
tool_call_metadata: Final[dict[str, object]] = {}
try:
# Extract tool calls from input messages
messages: Final = standard_logging_payload.get("messages", [])
if messages and isinstance(messages, list):
for message in messages:
if isinstance(message, dict) and "tool_calls" in message:
tool_calls = message.get("tool_calls")
if tool_calls:
input_tool_calls_kv = self._tool_calls_kv_pair(tool_calls)
# Prefix with "input_" to distinguish from response tool calls
for key, value in input_tool_calls_kv.items():
tool_call_metadata[f"input_{key}"] = value
# Extract tool calls from response
response_obj: Final = standard_logging_payload.get("response")
if response_obj and isinstance(response_obj, dict):
choices: Final = response_obj.get("choices", [])
for choice in choices:
if isinstance(choice, dict):
message = choice.get("message")
if message and isinstance(message, dict):
tool_calls = message.get("tool_calls")
if tool_calls:
response_tool_calls_kv = self._tool_calls_kv_pair(tool_calls)
# Prefix with "output_" to distinguish from input tool calls
for key, value in response_tool_calls_kv.items():
tool_call_metadata[f"output_{key}"] = value
except Exception as e:
verbose_logger.debug("DataDogLLMObs: Error extracting tool call metadata: %s", e)
return tool_call_metadata

View file

@ -6,6 +6,7 @@ import json
from collections.abc import Mapping
from dataclasses import dataclass, field
from enum import Enum
from types import MappingProxyType
from typing import TYPE_CHECKING, ClassVar, Final, cast
from urllib.parse import urlsplit
@ -62,6 +63,31 @@ if TYPE_CHECKING:
# --- typed sub-structures ---------------------------------------------------- #
def _cache_token_value(*values: object) -> int | None:
explicit_zero = False
invalid_before_zero = False
for raw_value in values:
if raw_value is None:
continue
if isinstance(raw_value, bool):
parsed = None
else:
try:
parsed = as_int(raw_value)
except (OverflowError, ValueError):
parsed = None
if parsed is None:
if not explicit_zero:
invalid_before_zero = True
elif parsed > 0:
return parsed
elif parsed == 0:
explicit_zero = True
elif not explicit_zero:
invalid_before_zero = True
return 0 if explicit_zero and not invalid_before_zero else None
@dataclass(frozen=True)
class LLMRequestParams:
temperature: float | None = None
@ -104,12 +130,25 @@ class LLMUsage:
metadata: Final[Mapping[str, object]] = payload.get("metadata") or {}
raw_usage: Final = metadata.get("usage_object")
usage_object: Final[Mapping[str, object]] = raw_usage if isinstance(raw_usage, Mapping) else {}
raw_details: Final = usage_object.get("prompt_tokens_details")
prompt_details: Final[Mapping[str, object]] = (
raw_details if isinstance(raw_details, Mapping) else MappingProxyType({})
)
return cls(
input_tokens=as_int(payload.get("prompt_tokens")),
output_tokens=as_int(payload.get("completion_tokens")),
total_tokens=as_int(payload.get("total_tokens")),
cache_creation_input_tokens=as_int(usage_object.get("cache_creation_input_tokens")),
cache_read_input_tokens=as_int(usage_object.get("cache_read_input_tokens")),
cache_creation_input_tokens=_cache_token_value(
usage_object.get("cache_creation_input_tokens"),
prompt_details.get("cache_write_tokens"),
prompt_details.get("cache_creation_tokens"),
prompt_details.get("cache_creation_input_tokens"),
),
cache_read_input_tokens=_cache_token_value(
usage_object.get("cache_read_input_tokens"),
prompt_details.get("cached_tokens"),
usage_object.get("prompt_cache_hit_tokens"),
),
)

View file

@ -8,6 +8,7 @@ import math
import os
import sys
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import replace
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, cast
@ -58,6 +59,7 @@ from litellm.types.utils import (
if TYPE_CHECKING:
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from prometheus_client import Gauge
from prometheus_client.metrics import MetricWrapperBase
from litellm.router import Router
@ -476,6 +478,30 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_remaining_api_key_tokens_for_model"),
)
self.litellm_api_key_rate_limit_allowed_metric = self._gauge_factory(
"litellm_api_key_rate_limit_allowed_metric",
"Configured rate limit for the API Key in the current window (rpm_limit / tpm_limit), by rate_limit_type",
labelnames=self.get_labels_for_metric("litellm_api_key_rate_limit_allowed_metric"),
)
self.litellm_api_key_rate_limit_used_metric = self._gauge_factory(
"litellm_api_key_rate_limit_used_metric",
"Requests or tokens the API Key has consumed in the current rate limit window, by rate_limit_type",
labelnames=self.get_labels_for_metric("litellm_api_key_rate_limit_used_metric"),
)
self.litellm_team_rate_limit_allowed_metric = self._gauge_factory(
"litellm_team_rate_limit_allowed_metric",
"Configured rate limit for the Team in the current window (team rpm_limit / tpm_limit), by rate_limit_type",
labelnames=self.get_labels_for_metric("litellm_team_rate_limit_allowed_metric"),
)
self.litellm_team_rate_limit_used_metric = self._gauge_factory(
"litellm_team_rate_limit_used_metric",
"Requests or tokens the Team has consumed in the current rate limit window, by rate_limit_type",
labelnames=self.get_labels_for_metric("litellm_team_rate_limit_used_metric"),
)
########################################
# LLM API Deployment Metrics / analytics
########################################
@ -1475,6 +1501,11 @@ class PrometheusLogger(CustomLogger):
model_id=enum_values.model_id,
)
self._set_key_and_team_rate_limit_metrics(
standard_logging_payload=standard_logging_payload, # pyright: ignore[reportArgumentType] # isinstance(dict) above narrows the TypedDict to dict[Unknown, Unknown]
enum_values=enum_values,
)
# set latency metrics
self._set_latency_metrics(
kwargs=kwargs,
@ -2002,17 +2033,102 @@ class PrometheusLogger(CustomLogger):
"""
if standard_logging_payload is None:
return None
return PrometheusLogger._get_int_from_v3_rate_limit_headers(
standard_logging_payload=standard_logging_payload,
header_name=f"x-ratelimit-model_per_key-remaining-{rate_limit_type}",
)
@staticmethod
def _get_int_from_v3_rate_limit_headers(
standard_logging_payload: StandardLoggingPayload,
header_name: str,
) -> int | None:
hidden_params: Final = standard_logging_payload.get("hidden_params")
if hidden_params is None:
return None
additional_headers: Final = hidden_params.get("additional_headers")
additional_headers: Final[Mapping[str, object] | None] = hidden_params.get("additional_headers")
if additional_headers is None:
return None
value: Final = dict(additional_headers).get(f"x-ratelimit-model_per_key-remaining-{rate_limit_type}")
value: Final = additional_headers.get(header_name)
if isinstance(value, bool) or not isinstance(value, int):
return None
return value
def _set_key_and_team_rate_limit_metrics(
self,
standard_logging_payload: StandardLoggingPayload,
enum_values: UserAPIKeyLabelValues,
) -> None:
"""
Export the key-level and team-level RPM / TPM limit and current window
usage from the ``x-ratelimit-{api_key,team}-{limit,remaining}-*``
headers the v3 rate limiter mirrors into the logging payload. The
limiter already read these counters (from Redis when configured) on
the request path, so no extra store lookup happens here. Descriptors
without a configured limit emit no header, so their series is removed
rather than left at the value from before the limit was dropped.
"""
descriptor_gauges: Final[
tuple[tuple[Literal["api_key", "team"], DEFINED_PROMETHEUS_METRICS, Gauge, Gauge], ...]
] = (
(
"api_key",
"litellm_api_key_rate_limit_allowed_metric",
self.litellm_api_key_rate_limit_allowed_metric,
self.litellm_api_key_rate_limit_used_metric,
),
(
"team",
"litellm_team_rate_limit_allowed_metric",
self.litellm_team_rate_limit_allowed_metric,
self.litellm_team_rate_limit_used_metric,
),
)
for descriptor_key, metric_name, allowed_gauge, used_gauge in descriptor_gauges:
for rate_limit_type in ("requests", "tokens"):
self._set_rate_limit_allowed_and_used_gauges(
standard_logging_payload=standard_logging_payload,
enum_values=enum_values,
descriptor_key=descriptor_key,
metric_name=metric_name,
allowed_gauge=allowed_gauge,
used_gauge=used_gauge,
rate_limit_type=rate_limit_type,
)
def _set_rate_limit_allowed_and_used_gauges(
self,
standard_logging_payload: StandardLoggingPayload,
enum_values: UserAPIKeyLabelValues,
descriptor_key: Literal["api_key", "team"],
metric_name: DEFINED_PROMETHEUS_METRICS,
allowed_gauge: Gauge,
used_gauge: Gauge,
rate_limit_type: Literal["requests", "tokens"],
) -> None:
limit: Final = self._get_int_from_v3_rate_limit_headers(
standard_logging_payload=standard_logging_payload,
header_name=f"x-ratelimit-{descriptor_key}-limit-{rate_limit_type}",
)
remaining: Final = self._get_int_from_v3_rate_limit_headers(
standard_logging_payload=standard_logging_payload,
header_name=f"x-ratelimit-{descriptor_key}-remaining-{rate_limit_type}",
)
labelled_values: Final = replace(enum_values, rate_limit_type=rate_limit_type)
labelnames: Final = self.get_labels_for_metric(metric_name)
labels: Final = prometheus_label_factory(
supported_enum_labels=labelnames,
enum_values=labelled_values,
label_context=PrometheusLabelFactoryContext(labelled_values),
)
if limit is None or remaining is None:
label_values: Final = tuple(labels.get(label) for label in labelnames)
self._bounded_prometheus_series_tracker.remove_series(allowed_gauge, label_values)
self._bounded_prometheus_series_tracker.remove_series(used_gauge, label_values)
return
allowed_gauge.labels(**labels).set(limit)
used_gauge.labels(**labels).set(limit - remaining)
def _set_virtual_key_rate_limit_metrics(
self,
user_api_key: str | None,

View file

@ -60,6 +60,10 @@ class BoundedPrometheusSeriesTracker:
break
del series[tracked_label_values]
def remove_series(self, metric: object, label_values: tuple[str | None, ...]) -> bool:
"""Drop one child series, True when it is gone (removed or never existed)."""
return self._remove_metric_child(metric, label_values)
def _should_run_ttl_cleanup(
self,
metric_name: str,

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