mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Merge pull request #39593 from BerriAI/litellm_internal_staging
chore(ci): promote internal staging to main
This commit is contained in:
commit
5be7ed1307
967 changed files with 123513 additions and 7105 deletions
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
5
.github/ci-coverage-allowlist.yml
vendored
5
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -4,6 +4,11 @@ description: >-
|
|||
by a job nor listed here, so every entry below is a decision on the record.
|
||||
|
||||
test_paths:
|
||||
- reason: >-
|
||||
The Rust/Python parity harness is run manually through its local CLI. Recorded replay,
|
||||
fixture generation, and harness checks are intentionally outside pull request CI
|
||||
paths:
|
||||
- tests/rust-python-harness
|
||||
- reason: >-
|
||||
What is left of the caching suite in tests/local_testing that runs nowhere. Every job that
|
||||
globs that directory either deselects it (local_testing_part1 and part2 carry `-k "... and
|
||||
|
|
|
|||
68
.github/scripts/smoke_test_native_wheel.py
vendored
Normal file
68
.github/scripts/smoke_test_native_wheel.py
vendored
Normal 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())
|
||||
282
.github/scripts/verify_linux_native_wheel.py
vendored
Normal file
282
.github/scripts/verify_linux_native_wheel.py
vendored
Normal 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())
|
||||
6
.github/workflows/image-scan.yml
vendored
6
.github/workflows/image-scan.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
130
.github/workflows/report-rust-release-wheel.yml
vendored
Normal file
130
.github/workflows/report-rust-release-wheel.yml
vendored
Normal 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,
|
||||
});
|
||||
}
|
||||
33
.github/workflows/test-code-quality.yml
vendored
33
.github/workflows/test-code-quality.yml
vendored
|
|
@ -128,6 +128,9 @@ jobs:
|
|||
- name: check_fastuuid_usage
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_fastuuid_usage.py
|
||||
|
||||
- name: check_py310_typing_imports
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_py310_typing_imports.py
|
||||
|
||||
- name: check_e2e_no_raw_requests
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/check_e2e_no_raw_requests.py
|
||||
|
||||
|
|
@ -145,3 +148,33 @@ jobs:
|
|||
|
||||
- name: documentation_test_api_docs
|
||||
run: uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
|
||||
|
||||
python-310-import-smoke:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
|
||||
steps:
|
||||
- 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.10"
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Install dependencies
|
||||
run: uv sync --frozen --extra proxy --python 3.10
|
||||
|
||||
- run: uv run --no-sync python --version
|
||||
|
||||
- name: Import litellm
|
||||
run: uv run --no-sync python -c "import litellm"
|
||||
|
||||
- name: Check litellm CLI
|
||||
run: uv run --no-sync litellm --version
|
||||
|
|
|
|||
65
.github/workflows/test-rust.yml
vendored
65
.github/workflows/test-rust.yml
vendored
|
|
@ -4,6 +4,12 @@ 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"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
|
|
@ -13,6 +19,12 @@ 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"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
|
||||
permissions:
|
||||
|
|
@ -40,9 +52,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 +61,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 +79,50 @@ 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
|
||||
|
||||
- name: Test native route wheel
|
||||
run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -96,6 +96,7 @@ jobs:
|
|||
- shard: misc
|
||||
artifact-name: misc
|
||||
test-path: >-
|
||||
tests/sdk_function_trace
|
||||
tests/test_litellm/batches
|
||||
tests/test_litellm/secret_managers
|
||||
tests/test_litellm/a2a_protocol
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
@ -101,6 +103,12 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
|
|||
|
||||
USER root
|
||||
|
||||
# The base image only configures Chainguard's authenticated apk repo, which
|
||||
# requires an enterprise subscription. Add the public Wolfi repo so `apk add`
|
||||
# also works for anyone installing extra packages into a running container.
|
||||
# https://github.com/BerriAI/litellm/issues/33518
|
||||
RUN echo "https://packages.wolfi.dev/os" >> /etc/apk/repositories
|
||||
|
||||
# node (without npm) is required by the prisma CLI at runtime
|
||||
RUN apk add --no-cache bash openssl tzdata nodejs python-3.13 libsndfile
|
||||
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 14076
|
||||
"limit": 14074
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2216
|
||||
"limit": 2215
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 4128
|
||||
"limit": 4124
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 5601
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15306
|
||||
"limit": 15287
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,31 +99,31 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44364
|
||||
"limit": 44362
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38350
|
||||
"limit": 38323
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19626
|
||||
"limit": 19624
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29890
|
||||
"limit": 29861
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 111
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 695
|
||||
"limit": 687
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 5
|
||||
"limit": 4
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 826
|
||||
"limit": 819
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
@ -138,7 +138,7 @@
|
|||
"limit": 138
|
||||
},
|
||||
"reportUnusedImport": {
|
||||
"limit": 543
|
||||
"limit": 542
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"limit": 137
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -11,17 +11,35 @@ import sys
|
|||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
import functools
|
||||
import tempfile
|
||||
from typing import Optional
|
||||
from contextvars import ContextVar
|
||||
from typing import TYPE_CHECKING, ClassVar, Literal, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails._content_utils import walk_user_text
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
GUARDRAIL_NAME = "hide_secrets"
|
||||
|
||||
GUARDRAIL_PROVIDER = "hide-secrets"
|
||||
|
||||
# Per-invocation tally of redacted secrets by detect-secrets plugin type; None
|
||||
# means the guardrail did not run, so _process_response records nothing.
|
||||
_masked_entity_count: ContextVar[Optional[dict]] = ContextVar(
|
||||
"hide_secrets_masked_entity_count", default=None
|
||||
)
|
||||
|
||||
_custom_plugins_path = "file://" + os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)), "secrets_plugins"
|
||||
)
|
||||
|
|
@ -422,6 +440,10 @@ _default_detect_secrets_config = {
|
|||
|
||||
|
||||
class _ENTERPRISE_SecretDetection(CustomGuardrail):
|
||||
# Keeps proxied traffic on async_pre_call_hook (the unified apply_guardrail
|
||||
# path skips should_run_check and never sees data["prompt"]).
|
||||
use_native_lifecycle_hooks: ClassVar[bool] = True
|
||||
|
||||
def __init__(self, detect_secrets_config: Optional[dict] = None, **kwargs):
|
||||
self.user_defined_detect_secrets_config = detect_secrets_config
|
||||
super().__init__(**kwargs)
|
||||
|
|
@ -455,6 +477,26 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
|
|||
|
||||
return detected_secrets
|
||||
|
||||
def redact_text(self, text: str, source: str = "message") -> str:
|
||||
"""Replace every detected secret in ``text`` with ``[REDACTED]`` and
|
||||
tally the detected types into the per-invocation masked-entity count."""
|
||||
detected_secrets = self.scan_message_for_secrets(text)
|
||||
if not detected_secrets:
|
||||
return text
|
||||
counts = _masked_entity_count.get()
|
||||
if counts is not None:
|
||||
for secret in detected_secrets:
|
||||
counts[secret["type"]] = counts.get(secret["type"], 0) + 1
|
||||
secret_types = [secret["type"] for secret in detected_secrets]
|
||||
verbose_proxy_logger.warning(
|
||||
f"Detected and redacted secrets in {source}: {secret_types}"
|
||||
)
|
||||
return functools.reduce(
|
||||
lambda redacted, secret: redacted.replace(secret["value"], "[REDACTED]"),
|
||||
detected_secrets,
|
||||
text,
|
||||
)
|
||||
|
||||
async def should_run_check(self, user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
if user_api_key_dict.permissions is not None:
|
||||
if GUARDRAIL_NAME in user_api_key_dict.permissions:
|
||||
|
|
@ -463,7 +505,45 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
|
|||
|
||||
return True
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Unified-interface entrypoint, used by /guardrails/apply_guardrail
|
||||
(the UI test playground). Proxied traffic keeps using
|
||||
``async_pre_call_hook``, see ``use_native_lifecycle_hooks``."""
|
||||
texts = inputs.get("texts")
|
||||
if not texts or not any(texts):
|
||||
return inputs
|
||||
_masked_entity_count.set({})
|
||||
return {**inputs, "texts": [self.redact_text(text) for text in texts]}
|
||||
|
||||
def _redact_prompt(self, data: dict) -> int:
|
||||
"""Redact ``data["prompt"]`` (the text-completion shape, which
|
||||
``walk_user_text`` does not cover) and return how many non-empty
|
||||
strings were inspected."""
|
||||
prompt = data.get("prompt")
|
||||
if isinstance(prompt, str):
|
||||
if not prompt:
|
||||
return 0
|
||||
data["prompt"] = self.redact_text(prompt, source="prompt")
|
||||
return 1
|
||||
if isinstance(prompt, list):
|
||||
data["prompt"] = [ # mutable-ok: data["prompt"] is a list on the wire
|
||||
self.redact_text(item, source="prompt")
|
||||
if isinstance(item, str) and item
|
||||
else item
|
||||
for item in prompt
|
||||
]
|
||||
return sum(1 for item in prompt if isinstance(item, str) and item)
|
||||
return 0
|
||||
|
||||
#### CALL HOOKS - proxy only ####
|
||||
@log_guardrail_information
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -471,53 +551,84 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
|
|||
data: dict,
|
||||
call_type: str, # "completion", "embeddings", "image_generation", "moderation"
|
||||
):
|
||||
_masked_entity_count.set(None)
|
||||
if await self.should_run_check(user_api_key_dict) is False:
|
||||
return
|
||||
|
||||
_masked_entity_count.set({})
|
||||
|
||||
# Covers multimodal list content + Responses-API input.
|
||||
def _redact_message_text(text: str) -> str:
|
||||
detected_secrets = self.scan_message_for_secrets(text)
|
||||
for secret in detected_secrets:
|
||||
text = text.replace(secret["value"], "[REDACTED]")
|
||||
if detected_secrets:
|
||||
secret_types = [secret["type"] for secret in detected_secrets]
|
||||
verbose_proxy_logger.warning(
|
||||
f"Detected and redacted secrets in message: {secret_types}"
|
||||
)
|
||||
return text
|
||||
inspected = walk_user_text(data, self.redact_text) + self._redact_prompt(data)
|
||||
|
||||
walk_user_text(data, _redact_message_text)
|
||||
if inspected == 0:
|
||||
# Image-only, empty-text, and unsupported payloads inspected
|
||||
# nothing, so recording "allow" would count a run that never
|
||||
# looked at any content.
|
||||
_masked_entity_count.set(None)
|
||||
|
||||
if "prompt" in data:
|
||||
if isinstance(data["prompt"], str):
|
||||
detected_secrets = self.scan_message_for_secrets(data["prompt"])
|
||||
for secret in detected_secrets:
|
||||
data["prompt"] = data["prompt"].replace(
|
||||
secret["value"], "[REDACTED]"
|
||||
)
|
||||
if len(detected_secrets) > 0:
|
||||
secret_types = [secret["type"] for secret in detected_secrets]
|
||||
verbose_proxy_logger.warning(
|
||||
f"Detected and redacted secrets in prompt: {secret_types}"
|
||||
)
|
||||
elif isinstance(data["prompt"], list):
|
||||
# Index back into the list — assigning to ``item`` would only
|
||||
# rebind the loop variable and leave ``data["prompt"]``
|
||||
# carrying the unredacted secret.
|
||||
for idx, item in enumerate(data["prompt"]):
|
||||
if isinstance(item, str):
|
||||
detected_secrets = self.scan_message_for_secrets(item)
|
||||
for secret in detected_secrets:
|
||||
item = item.replace(secret["value"], "[REDACTED]")
|
||||
data["prompt"][idx] = item
|
||||
if len(detected_secrets) > 0:
|
||||
secret_types = [
|
||||
secret["type"] for secret in detected_secrets
|
||||
]
|
||||
verbose_proxy_logger.warning(
|
||||
f"Detected and redacted secrets in prompt: {secret_types}"
|
||||
)
|
||||
|
||||
# ``data["input"]`` (Responses API and embeddings/moderation) is
|
||||
# already covered by ``walk_user_text`` above.
|
||||
return
|
||||
|
||||
def _process_response(
|
||||
self,
|
||||
response: Optional[dict],
|
||||
request_data: dict,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
original_inputs: Optional[dict] = None,
|
||||
):
|
||||
"""Record allow/mask plus the masked-entity tally for a completed run.
|
||||
|
||||
Records nothing when the guardrail inspected nothing (opted-out key,
|
||||
empty inputs) or when the instance has no guardrail_name (legacy
|
||||
``litellm_settings.callbacks`` deployments, which predate guardrail
|
||||
telemetry and stay without it).
|
||||
"""
|
||||
counts = _masked_entity_count.get()
|
||||
_masked_entity_count.set(None)
|
||||
if counts is None or self.guardrail_name is None:
|
||||
return response
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response="mask" if counts else "allow",
|
||||
request_data=request_data,
|
||||
guardrail_status="success",
|
||||
duration=duration,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
event_type=event_type,
|
||||
guardrail_provider=GUARDRAIL_PROVIDER,
|
||||
masked_entity_count=counts,
|
||||
)
|
||||
return response
|
||||
|
||||
def _process_error(
|
||||
self,
|
||||
e: Exception,
|
||||
request_data: dict,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
event_type: Optional[GuardrailEventHooks] = None,
|
||||
):
|
||||
"""Label the failed run with this guardrail's provider so error rows
|
||||
group with the successful ones in the monitor. Nameless legacy
|
||||
instances record nothing, matching ``_process_response``."""
|
||||
_masked_entity_count.set(None)
|
||||
if self.guardrail_name is None:
|
||||
raise e
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=e,
|
||||
request_data=request_data,
|
||||
guardrail_status=(
|
||||
"guardrail_intervened"
|
||||
if self._is_guardrail_intervention(e)
|
||||
else "guardrail_failed_to_respond"
|
||||
),
|
||||
duration=duration,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
event_type=event_type,
|
||||
guardrail_provider=GUARDRAIL_PROVIDER,
|
||||
)
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -39,9 +39,9 @@ async def available_enterprise_users(
|
|||
|
||||
if not premium_user:
|
||||
# check if SSO is enabled - show 5 user limit
|
||||
from litellm.proxy.auth.auth_utils import _has_user_setup_sso
|
||||
from litellm.proxy.auth.auth_utils import has_user_setup_sso
|
||||
|
||||
if _has_user_setup_sso():
|
||||
if has_user_setup_sso():
|
||||
premium_user_data = EnterpriseLicenseData(
|
||||
max_users=5,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.63"
|
||||
version = "0.1.64"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.63"
|
||||
version = "0.1.64"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 }
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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: {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import shutil
|
|||
import subprocess
|
||||
import tempfile
|
||||
import time
|
||||
from dataclasses import dataclass, replace
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
|
|
@ -15,8 +16,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,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -40,6 +44,40 @@ def _get_prisma_env() -> dict:
|
|||
|
||||
_MIGRATION_TS_RE = re.compile(r"^(\d{14})_")
|
||||
|
||||
_MIGRATION_DEADLOCK_MARKER = "deadlock detected"
|
||||
|
||||
MAX_MIGRATE_DEPLOY_ATTEMPTS = 4
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _MigrateAttemptBudget:
|
||||
"""Retries left, and the recoveries already run.
|
||||
|
||||
A recovery that lands something new costs nothing, so a database full of
|
||||
objects `prisma db push` created works through them one per pass. Anything
|
||||
that made no progress spends an attempt, so a stuck run still gives up.
|
||||
"""
|
||||
|
||||
attempts_left: int
|
||||
recoveries: frozenset[str] = frozenset()
|
||||
|
||||
@property
|
||||
def exhausted(self) -> bool:
|
||||
return self.attempts_left <= 0
|
||||
|
||||
@property
|
||||
def attempt_number(self) -> int:
|
||||
return MAX_MIGRATE_DEPLOY_ATTEMPTS - self.attempts_left + 1
|
||||
|
||||
def spend(self) -> "_MigrateAttemptBudget":
|
||||
return replace(self, attempts_left=self.attempts_left - 1)
|
||||
|
||||
def after_recovery(self, recovery: str) -> "_MigrateAttemptBudget":
|
||||
if recovery in self.recoveries:
|
||||
return self.spend()
|
||||
return replace(self, recoveries=self.recoveries | {recovery})
|
||||
|
||||
|
||||
_SPEND_LOGS_ALTER_RE = re.compile(r'^ALTER\s+TABLE\s+"LiteLLM_SpendLogs"\s', re.IGNORECASE)
|
||||
_SPEND_LOGS_ARTIFACT_DROP_RE = re.compile(
|
||||
r'^DROP\s+TABLE\s+"LiteLLM_SpendLogs_[^"]*"', re.IGNORECASE
|
||||
|
|
@ -262,6 +300,50 @@ class ProxyExtrasDBManager:
|
|||
env=prisma_env,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _roll_back_migration_best_effort(migration_name: str) -> None:
|
||||
"""Mark a migration rolled back, tolerating a concurrent resolver
|
||||
having already done it."""
|
||||
try:
|
||||
ProxyExtrasDBManager._roll_back_migration(migration_name)
|
||||
except (subprocess.CalledProcessError, subprocess.TimeoutExpired):
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _failed_migration_logs(migration_name: str) -> Optional[str]:
|
||||
"""Return failed migration logs, or None if the ledger is unavailable."""
|
||||
database_url = os.getenv("DATABASE_URL")
|
||||
if not database_url:
|
||||
return None
|
||||
|
||||
try:
|
||||
import psycopg
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
cleaned_url = ProxyExtrasDBManager._strip_prisma_query_params(database_url)
|
||||
ledger_table = psycopg.sql.SQL("{}.{}").format(
|
||||
psycopg.sql.Identifier(
|
||||
ProxyExtrasDBManager._prisma_schema_param(database_url) or "public"
|
||||
),
|
||||
psycopg.sql.Identifier("_prisma_migrations"),
|
||||
)
|
||||
try:
|
||||
with psycopg.connect(
|
||||
cleaned_url, connect_timeout=10, autocommit=True
|
||||
) as conn:
|
||||
row = conn.execute(
|
||||
psycopg.sql.SQL(
|
||||
"SELECT logs FROM {} "
|
||||
"WHERE migration_name = %s AND finished_at IS NULL "
|
||||
"AND rolled_back_at IS NULL"
|
||||
).format(ledger_table),
|
||||
(migration_name,),
|
||||
).fetchone()
|
||||
except (psycopg.OperationalError, psycopg.DatabaseError):
|
||||
return None
|
||||
return (row[0] or "") if row else ""
|
||||
|
||||
@staticmethod
|
||||
def _resolve_specific_migration(migration_name: str):
|
||||
"""Mark a specific migration as applied"""
|
||||
|
|
@ -658,7 +740,8 @@ class ProxyExtrasDBManager:
|
|||
v2 migration resolver (opt-in via --use_v2_migration_resolver).
|
||||
|
||||
Runs `prisma migrate deploy` and handles standard recovery paths
|
||||
(P3005 baseline, P3009/P3018 idempotent errors). Critically, it does
|
||||
(P3005 baseline, P3009/P3018 idempotent errors, deadlocks against a
|
||||
concurrent migrate deploy). Critically, it does
|
||||
NOT call `_resolve_all_migrations` — the diff-and-force recovery that
|
||||
caused schema thrashing when two LiteLLM versions contended for the
|
||||
same DB during rolling deploys.
|
||||
|
|
@ -666,6 +749,9 @@ class ProxyExtrasDBManager:
|
|||
Ahead-of-HEAD state (DB has migrations newer than this build ships)
|
||||
is logged as a warning, not a fatal error — users whose DBs got into
|
||||
weird shapes from the old thrashing should still be able to start.
|
||||
|
||||
The retry budget only counts attempts that made no progress: see
|
||||
_MigrateAttemptBudget.
|
||||
"""
|
||||
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
|
||||
migrations_dir = ProxyExtrasDBManager._get_prisma_dir()
|
||||
|
|
@ -698,12 +784,14 @@ class ProxyExtrasDBManager:
|
|||
|
||||
original_dir = os.getcwd()
|
||||
os.chdir(migrations_dir)
|
||||
deploy_timeout = prisma_migrate_deploy_timeout()
|
||||
budget = _MigrateAttemptBudget(attempts_left=MAX_MIGRATE_DEPLOY_ATTEMPTS)
|
||||
try:
|
||||
for attempt in range(4):
|
||||
while not budget.exhausted:
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[_get_prisma_command(), "migrate", "deploy"],
|
||||
timeout=prisma_command_timeout(),
|
||||
timeout=deploy_timeout,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
|
|
@ -713,121 +801,158 @@ 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.",
|
||||
budget.attempt_number,
|
||||
deploy_timeout,
|
||||
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR,
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
continue
|
||||
next_budget = budget.spend()
|
||||
|
||||
except subprocess.CalledProcessError as e:
|
||||
stderr = e.stderr or ""
|
||||
next_budget = ProxyExtrasDBManager._budget_after_deploy_failure(
|
||||
e, budget, schema_path
|
||||
)
|
||||
|
||||
if "P3005" in stderr and "database schema is not empty" in stderr:
|
||||
logger.info(
|
||||
"Schema exists but no migrations ledger — creating baseline"
|
||||
)
|
||||
ProxyExtrasDBManager._create_baseline_migration(schema_path)
|
||||
continue
|
||||
|
||||
if "P3009" in stderr:
|
||||
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
|
||||
if (
|
||||
migration_match
|
||||
and ProxyExtrasDBManager._is_idempotent_error(stderr)
|
||||
):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} failed idempotently — marking applied and retrying"
|
||||
)
|
||||
try:
|
||||
ProxyExtrasDBManager._roll_back_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
):
|
||||
pass # may already be rolled-back
|
||||
try:
|
||||
ProxyExtrasDBManager._resolve_specific_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
) as resolve_err:
|
||||
# We're already inside the outer
|
||||
# `except CalledProcessError` handler —
|
||||
# re-raising CalledProcessError from here
|
||||
# would escape as itself, bypassing
|
||||
# proxy_cli.py's `except RuntimeError`.
|
||||
raise RuntimeError(
|
||||
f"Failed to mark migration {name} as applied "
|
||||
f"after idempotent recovery. Manual "
|
||||
f"intervention may be required.\n\n"
|
||||
f"Detail: {resolve_err}"
|
||||
) from resolve_err
|
||||
continue
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
if "P3018" in stderr:
|
||||
if ProxyExtrasDBManager._is_permission_error(stderr):
|
||||
raise RuntimeError(
|
||||
"Database migration failed due to insufficient "
|
||||
"permissions. Please grant the required privileges "
|
||||
f"and retry.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
migration_match = re.search(
|
||||
r"Migration name: (\d+_\S+)", stderr
|
||||
)
|
||||
if (
|
||||
migration_match
|
||||
and ProxyExtrasDBManager._is_idempotent_error(stderr)
|
||||
):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
|
||||
)
|
||||
try:
|
||||
ProxyExtrasDBManager._roll_back_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
):
|
||||
pass # may already be rolled-back
|
||||
try:
|
||||
ProxyExtrasDBManager._resolve_specific_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
) as resolve_err:
|
||||
raise RuntimeError(
|
||||
f"Failed to mark migration {name} as applied "
|
||||
f"after idempotent recovery. Manual "
|
||||
f"intervention may be required.\n\n"
|
||||
f"Detail: {resolve_err}"
|
||||
) from resolve_err
|
||||
continue
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
if next_budget.attempts_left < budget.attempts_left:
|
||||
time.sleep(random.randrange(5, 15))
|
||||
budget = next_budget # rebind-ok: the loop carries the budget from one migrate deploy pass to the next
|
||||
|
||||
raise RuntimeError(
|
||||
"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."
|
||||
f"Database migration failed after {MAX_MIGRATE_DEPLOY_ATTEMPTS} "
|
||||
"attempts that made no progress (timeouts, deadlock retries, or a "
|
||||
"recovery that had already run once). Check database connectivity, "
|
||||
"load, and _prisma_migrations ledger state, and raise "
|
||||
f"{PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR} if the attempts timed out."
|
||||
)
|
||||
finally:
|
||||
os.chdir(original_dir)
|
||||
|
||||
@staticmethod
|
||||
def _budget_after_deploy_failure(
|
||||
error: subprocess.CalledProcessError,
|
||||
budget: "_MigrateAttemptBudget",
|
||||
schema_path: str,
|
||||
) -> "_MigrateAttemptBudget":
|
||||
"""Recover from one failed `prisma migrate deploy`, and price the pass.
|
||||
|
||||
Returns the budget the next pass runs under, or raises when the failure
|
||||
is not one this resolver knows how to recover from.
|
||||
"""
|
||||
stderr = error.stderr or ""
|
||||
|
||||
if "P3005" in stderr and "database schema is not empty" in stderr:
|
||||
logger.info("Schema exists but no migrations ledger — creating baseline")
|
||||
if ProxyExtrasDBManager._create_baseline_migration(schema_path):
|
||||
return budget.after_recovery("baseline")
|
||||
return budget.spend()
|
||||
|
||||
if "P3009" in stderr:
|
||||
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
|
||||
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} failed idempotently — marking applied and retrying"
|
||||
)
|
||||
ProxyExtrasDBManager._mark_migration_applied(name)
|
||||
return budget.after_recovery(f"resolved:{name}")
|
||||
if migration_match:
|
||||
migration_name = migration_match.group(1)
|
||||
ledger_logs = ProxyExtrasDBManager._failed_migration_logs(migration_name)
|
||||
if ledger_logs is not None and (
|
||||
ledger_logs == "" or _MIGRATION_DEADLOCK_MARKER in ledger_logs
|
||||
):
|
||||
logger.info(
|
||||
"Migration %s failed in a concurrent migrate deploy "
|
||||
"deadlock race, rolling its ledger row back and retrying",
|
||||
migration_name,
|
||||
)
|
||||
ProxyExtrasDBManager._roll_back_migration_best_effort(migration_name)
|
||||
return budget.spend()
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
if "P3018" in stderr:
|
||||
if ProxyExtrasDBManager._is_permission_error(stderr):
|
||||
raise RuntimeError(
|
||||
"Database migration failed due to insufficient "
|
||||
"permissions. Please grant the required privileges "
|
||||
f"and retry.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
migration_match = re.search(r"Migration name: (\d+_\S+)", stderr)
|
||||
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
|
||||
)
|
||||
ProxyExtrasDBManager._mark_migration_applied(name)
|
||||
return budget.after_recovery(f"resolved:{name}")
|
||||
|
||||
if migration_match and _MIGRATION_DEADLOCK_MARKER in stderr:
|
||||
logger.info(
|
||||
"Migration %s deadlocked against a concurrent "
|
||||
"migrate deploy, rolling its ledger row back "
|
||||
"and retrying",
|
||||
migration_match.group(1),
|
||||
)
|
||||
ProxyExtrasDBManager._roll_back_migration_best_effort(
|
||||
migration_match.group(1)
|
||||
)
|
||||
return budget.spend()
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
if _MIGRATION_DEADLOCK_MARKER in stderr:
|
||||
logger.info(
|
||||
"prisma migrate deploy attempt %s deadlocked against "
|
||||
"a concurrent migrate deploy, retrying",
|
||||
budget.attempt_number,
|
||||
)
|
||||
return budget.spend()
|
||||
|
||||
if "P1002" in stderr and "advisory lock" in stderr:
|
||||
logger.info(
|
||||
"prisma migrate deploy attempt %s timed out waiting for "
|
||||
"the advisory lock a concurrent migrate deploy holds, retrying",
|
||||
budget.attempt_number,
|
||||
)
|
||||
return budget.spend()
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _mark_migration_applied(name: str) -> None:
|
||||
"""Roll a failed ledger row back if it is still there, then mark it applied."""
|
||||
try:
|
||||
ProxyExtrasDBManager._roll_back_migration(name)
|
||||
except (subprocess.CalledProcessError, subprocess.TimeoutExpired):
|
||||
pass # may already be rolled-back
|
||||
try:
|
||||
ProxyExtrasDBManager._resolve_specific_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
) as resolve_err:
|
||||
# We're called from inside an `except CalledProcessError` handler —
|
||||
# re-raising CalledProcessError from here would escape as itself,
|
||||
# bypassing proxy_cli.py's `except RuntimeError`.
|
||||
raise RuntimeError(
|
||||
f"Failed to mark migration {name} as applied "
|
||||
f"after idempotent recovery. Manual "
|
||||
f"intervention may be required.\n\n"
|
||||
f"Detail: {resolve_err}"
|
||||
) from resolve_err
|
||||
|
||||
@staticmethod
|
||||
def apply_replica_identity_full_if_requested() -> bool:
|
||||
"""
|
||||
|
|
@ -908,7 +1033,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 +1251,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
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.92"
|
||||
version = "0.4.93"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.92"
|
||||
version = "0.4.93"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -240,3 +240,223 @@ def test_v2_does_not_call_resolve_all_migrations(monkeypatch, tmp_path):
|
|||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
assert ok is True
|
||||
assert resolve_called["n"] == 0, "v2 must not invoke the diff-and-force recovery"
|
||||
|
||||
|
||||
_DEADLOCK_P3018_STDERR = (
|
||||
"Error: P3018\n"
|
||||
"Migration name: 20260415120000_health_check_latest_per_model_index\n"
|
||||
"Database error code: 40P01\n"
|
||||
"deadlock detected"
|
||||
)
|
||||
|
||||
|
||||
def _stub_v2_env(monkeypatch, tmp_path):
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
|
||||
)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
|
||||
(tmp_path / "schema.prisma").write_text("// stub")
|
||||
monkeypatch.setattr("time.sleep", lambda _: None)
|
||||
|
||||
|
||||
def _succeed_after(failures: int, stderr: str):
|
||||
calls = {"n": 0}
|
||||
|
||||
class _OkResult:
|
||||
stdout = "Applied migration.\n"
|
||||
stderr = ""
|
||||
|
||||
def _run(*args, **kwargs):
|
||||
if "deploy" not in args[0]:
|
||||
return _OkResult()
|
||||
calls["n"] += 1
|
||||
if calls["n"] <= failures:
|
||||
raise subprocess.CalledProcessError(
|
||||
returncode=1, cmd=args[0], stderr=stderr, output=""
|
||||
)
|
||||
return _OkResult()
|
||||
|
||||
return _run
|
||||
|
||||
|
||||
def test_v2_p3018_deadlock_rolls_back_and_retries(monkeypatch, tmp_path):
|
||||
"""v2: losing the migrate deploy deadlock race against a concurrent
|
||||
instance rolls the ledger row back and retries instead of dying."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
|
||||
rolled_back = []
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_roll_back_migration",
|
||||
lambda name: rolled_back.append(name),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_resolve_specific_migration",
|
||||
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
|
||||
)
|
||||
monkeypatch.setattr("subprocess.run", _succeed_after(1, _DEADLOCK_P3018_STDERR))
|
||||
|
||||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
assert ok is True
|
||||
assert rolled_back == ["20260415120000_health_check_latest_per_model_index"]
|
||||
|
||||
|
||||
def test_v2_p3018_persistent_deadlock_exhausts_attempts(monkeypatch, tmp_path):
|
||||
"""v2: a deadlock on every attempt still fails after the retry budget."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_roll_back_migration", lambda name: None)
|
||||
|
||||
with patch(
|
||||
"subprocess.run",
|
||||
side_effect=_fake_migrate_deploy_failure(1, _DEADLOCK_P3018_STDERR),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="after 4 attempts"):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
|
||||
|
||||
def test_v2_p3009_deadlocked_ledger_row_rolls_back_and_retries(monkeypatch, tmp_path):
|
||||
"""v2: the surviving instance sees the victim's failed ledger row as P3009.
|
||||
When that row's logs show a deadlock, roll it back and retry."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
|
||||
stderr = (
|
||||
"Error: P3009\n"
|
||||
"migrate found failed migrations in the target database\n"
|
||||
"The `20260415120000_health_check_latest_per_model_index` migration "
|
||||
"started at 2026-09-01 18:46:13 UTC failed"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_failed_migration_logs",
|
||||
lambda name: "ERROR: deadlock detected\nDETAIL: Process 72 waits for ShareLock",
|
||||
)
|
||||
rolled_back = []
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_roll_back_migration",
|
||||
lambda name: rolled_back.append(name),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_resolve_specific_migration",
|
||||
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
|
||||
)
|
||||
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
|
||||
|
||||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
assert ok is True
|
||||
assert rolled_back == ["20260415120000_health_check_latest_per_model_index"]
|
||||
|
||||
|
||||
def test_v2_p3009_empty_ledger_logs_rolls_back_and_retries(monkeypatch, tmp_path):
|
||||
"""v2: empty failed ledger logs mean a concurrent deploy moved it on."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
|
||||
stderr = (
|
||||
"Error: P3009\n"
|
||||
"migrate found failed migrations in the target database\n"
|
||||
"The `20260415120000_health_check_latest_per_model_index` migration "
|
||||
"started at 2026-09-01 18:46:13 UTC failed"
|
||||
)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_failed_migration_logs", lambda name: "")
|
||||
rolled_back = []
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_roll_back_migration",
|
||||
lambda name: rolled_back.append(name),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_resolve_specific_migration",
|
||||
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
|
||||
)
|
||||
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
|
||||
|
||||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
assert ok is True
|
||||
assert rolled_back == ["20260415120000_health_check_latest_per_model_index"]
|
||||
|
||||
|
||||
def test_v2_p3009_unreadable_ledger_still_raises(monkeypatch, tmp_path):
|
||||
"""v2: an unreadable ledger cannot establish that P3009 was a deadlock."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
|
||||
stderr = (
|
||||
"Error: P3009\n"
|
||||
"migrate found failed migrations in the target database\n"
|
||||
"The `20260415120000_health_check_latest_per_model_index` migration "
|
||||
"started at 2026-09-01 18:46:13 UTC failed"
|
||||
)
|
||||
monkeypatch.setattr(ProxyExtrasDBManager, "_failed_migration_logs", lambda name: None)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_roll_back_migration",
|
||||
lambda name: pytest.fail("an unreadable ledger must not trigger a retry"),
|
||||
)
|
||||
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
|
||||
|
||||
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
|
||||
|
||||
def test_v2_p3009_non_deadlock_ledger_row_still_raises(monkeypatch, tmp_path):
|
||||
"""v2: a failed ledger row whose logs show a real SQL error stays fatal."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
|
||||
stderr = (
|
||||
"Error: P3009\n"
|
||||
"migrate found failed migrations in the target database\n"
|
||||
"The `20260101000000_genuinely_broken` migration started at "
|
||||
"2026-09-01 18:46:13 UTC failed"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_failed_migration_logs",
|
||||
lambda name: 'ERROR: syntax error at or near "BRKN"',
|
||||
)
|
||||
|
||||
with patch("subprocess.run", side_effect=_fake_migrate_deploy_failure(1, stderr)):
|
||||
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
|
||||
|
||||
def test_v2_bare_deadlock_stderr_retries(monkeypatch, tmp_path):
|
||||
"""v2: a deadlock reported without a Prisma error code (the advisory-lock
|
||||
waiter as victim) is retried, not fatal."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
monkeypatch.setattr(
|
||||
"subprocess.run", _succeed_after(1, "Database error: deadlock detected")
|
||||
)
|
||||
|
||||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
assert ok is True
|
||||
|
||||
|
||||
_P1002_ADVISORY_LOCK_STDERR = (
|
||||
"Error: P1002\n\n"
|
||||
"The database server at `127.0.0.1`:`45743` was reached but timed out.\n\n"
|
||||
"Context: Timed out trying to acquire a postgres advisory lock "
|
||||
"(SELECT pg_advisory_lock(72707369)). Elapsed: 10000ms."
|
||||
)
|
||||
|
||||
|
||||
def test_v2_advisory_lock_timeout_retries(monkeypatch, tmp_path):
|
||||
"""v2: the advisory-lock waiter that times out while a peer's retry holds
|
||||
the lock retries instead of dying."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
monkeypatch.setattr("subprocess.run", _succeed_after(2, _P1002_ADVISORY_LOCK_STDERR))
|
||||
|
||||
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
assert ok is True
|
||||
|
||||
|
||||
def test_v2_p1002_without_advisory_lock_context_still_raises(monkeypatch, tmp_path):
|
||||
"""v2: a plain P1002 (database unreachable) stays fatal."""
|
||||
_stub_v2_env(monkeypatch, tmp_path)
|
||||
stderr = "Error: P1002\n\nThe database server at `db`:`5432` was reached but timed out."
|
||||
monkeypatch.setattr("subprocess.run", _succeed_after(1, stderr))
|
||||
|
||||
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
|
||||
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
```
|
||||
|
||||
|
|
|
|||
169
litellm-rust/Cargo.lock
generated
169
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
@ -1380,6 +1392,12 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lazy_static"
|
||||
version = "1.5.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe"
|
||||
|
||||
[[package]]
|
||||
name = "libc"
|
||||
version = "0.2.186"
|
||||
|
|
@ -1404,6 +1422,7 @@ dependencies = [
|
|||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"tower",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -1423,6 +1442,7 @@ dependencies = [
|
|||
"sha2 0.10.9",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -1430,14 +1450,29 @@ name = "litellm-python-bridge"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"criterion",
|
||||
"futures-util",
|
||||
"litellm-ai-gateway",
|
||||
"litellm-core",
|
||||
"litellm-python-interop",
|
||||
"pyo3",
|
||||
"pyo3-async-runtimes",
|
||||
"pythonize",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-python-interop"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"pyo3",
|
||||
"pythonize",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -1627,6 +1662,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"
|
||||
|
|
@ -1638,9 +1682,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3"
|
||||
version = "0.29.0"
|
||||
version = "0.29.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cd274650b21d4bfc26a0a47587962c1edb425f69287324355cd040c3ea66071c"
|
||||
checksum = "4688ddedf473e32662b9b067670129a8afb8c18e351482c70d62ba4a88171e8b"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"once_cell",
|
||||
|
|
@ -1666,18 +1710,18 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3-build-config"
|
||||
version = "0.29.0"
|
||||
version = "0.29.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c5e2a7d2f0d013342f295c048ad19237add5154a55b1c5a254c0ec93d4109078"
|
||||
checksum = "f41027e41b4bd03f6e60f9f417fe24a6341a6bb744edd62b6f709f2a52ea30e9"
|
||||
dependencies = [
|
||||
"target-lexicon",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pyo3-ffi"
|
||||
version = "0.29.0"
|
||||
version = "0.29.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ca85c467da1bbc8d866eea5deff9cf29ea5f7785054a17da36e65bda9c05845b"
|
||||
checksum = "e591a95526fead067432c3b3a33fc74770b87b1e04e73671090d9c2055a2b327"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"pyo3-build-config",
|
||||
|
|
@ -1685,9 +1729,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3-macros"
|
||||
version = "0.29.0"
|
||||
version = "0.29.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9ac53762fd065daa3194dd09337a38bd793a188100fd1a9304c4ab312d901771"
|
||||
checksum = "73225868fc1cd84eef2c3c230ddb91273bf1de46aeb8a4248da76d32a0924a1c"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"pyo3-macros-backend",
|
||||
|
|
@ -1697,9 +1741,9 @@ dependencies = [
|
|||
|
||||
[[package]]
|
||||
name = "pyo3-macros-backend"
|
||||
version = "0.29.0"
|
||||
version = "0.29.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4ca3a1557399783172dc5bf39cfca835157732532cba56b71d2292161e53b362"
|
||||
checksum = "571575aa3749fa6216757dd47d2a3e7ef360f329a40f0666a9fbd14889024952"
|
||||
dependencies = [
|
||||
"heck",
|
||||
"proc-macro2",
|
||||
|
|
@ -1899,6 +1943,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 +2006,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"
|
||||
|
|
@ -2210,6 +2289,15 @@ dependencies = [
|
|||
"digest 0.11.3",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sharded-slab"
|
||||
version = "0.1.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6"
|
||||
dependencies = [
|
||||
"lazy_static",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "shlex"
|
||||
version = "2.0.1"
|
||||
|
|
@ -2348,6 +2436,15 @@ dependencies = [
|
|||
"syn 3.0.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "thread_local"
|
||||
version = "1.1.10"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "time"
|
||||
version = "0.3.53"
|
||||
|
|
@ -2488,6 +2585,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"
|
||||
|
|
@ -2566,6 +2693,17 @@ dependencies = [
|
|||
"once_cell",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tracing-subscriber"
|
||||
version = "0.3.23"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319"
|
||||
dependencies = [
|
||||
"sharded-slab",
|
||||
"thread_local",
|
||||
"tracing-core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "try-lock"
|
||||
version = "0.2.5"
|
||||
|
|
@ -2903,6 +3041,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"
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
members = [
|
||||
"crates/core",
|
||||
"crates/ai-gateway",
|
||||
"crates/python-interop",
|
||||
"crates/python-bridge",
|
||||
]
|
||||
resolver = "2"
|
||||
|
|
@ -13,16 +14,20 @@ license = "MIT"
|
|||
repository = "https://github.com/BerriAI/litellm"
|
||||
|
||||
[workspace.dependencies]
|
||||
tracing = "0.1"
|
||||
tracing-subscriber = { version = "0.3", default-features = false, features = ["registry", "std"] }
|
||||
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 = "0.29.2"
|
||||
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"
|
||||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
sha2 = "0.10"
|
||||
subtle = "2"
|
||||
thiserror = "2.0"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
```
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ path = "src/main.rs"
|
|||
required-features = ["server"]
|
||||
|
||||
[dependencies]
|
||||
tracing.workspace = true
|
||||
litellm-core = { workspace = true, features = ["bedrock-auth"] }
|
||||
# reqwest (rustls + json) is used by io/ocr and ships realtime logs to the
|
||||
# Python proxy callbacks API.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -1,48 +0,0 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::audio_transcription::transformation::AudioTranscriptionProviderConfig;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub(super) fn audio_transcription_provider_config(
|
||||
provider: &str,
|
||||
) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
|
||||
match provider {
|
||||
"bedrock" => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn string_headers(
|
||||
headers: Option<Map<String, Value>>,
|
||||
) -> CoreResult<BTreeMap<String, String>> {
|
||||
headers
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(|(key, value)| {
|
||||
value
|
||||
.as_str()
|
||||
.map(|value| (key.clone(), value.to_string()))
|
||||
.ok_or_else(|| {
|
||||
CoreError::InvalidRequest(format!(
|
||||
"audio transcription extra_headers.{key} must be a string"
|
||||
))
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn has_header(headers: &BTreeMap<String, String>, name: &str) -> bool {
|
||||
headers.keys().any(|key| key.eq_ignore_ascii_case(name))
|
||||
}
|
||||
|
||||
pub(super) fn truncate_error_body(body: &str) -> String {
|
||||
let truncated: String = body.chars().take(256).collect();
|
||||
if truncated.chars().count() == body.chars().count() {
|
||||
truncated
|
||||
} else {
|
||||
format!("{truncated}... (truncated)")
|
||||
}
|
||||
}
|
||||
|
|
@ -1,89 +0,0 @@
|
|||
use std::time::SystemTime;
|
||||
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth;
|
||||
use litellm_core::error::CoreError;
|
||||
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 super::common_utils::truncate_error_body;
|
||||
use super::types::ProviderAudioTranscriptionRequest;
|
||||
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}"))
|
||||
})?;
|
||||
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);
|
||||
}
|
||||
if let Some(duration) = request.timeout {
|
||||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
let response = request_builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| CoreError::Network(error.to_string()))?;
|
||||
let status = response.status();
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|error| CoreError::Network(error.to_string()))?;
|
||||
if !status.is_success() {
|
||||
return Err(CoreError::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}"))
|
||||
})?;
|
||||
Ok(request
|
||||
.config
|
||||
.transform_transcription_response(&request.model, response_json)?
|
||||
.into_json())
|
||||
}
|
||||
|
||||
pub(crate) async fn sign_request(
|
||||
request: &ProviderAudioTranscriptionRequest,
|
||||
optional_params: &serde_json::Map<String, Value>,
|
||||
) -> CoreResult<ProviderAudioTranscriptionRequest> {
|
||||
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 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());
|
||||
match auth {
|
||||
AudioTranscriptionAuth::Bearer => {}
|
||||
AudioTranscriptionAuth::AwsSigV4 { region, .. } => {
|
||||
let credentials =
|
||||
resolve_credentials(aws_auth_config(optional_params, &env_lookup), &env_lookup)
|
||||
.await?;
|
||||
headers.extend(sign_bedrock_post(
|
||||
&request.url,
|
||||
&body,
|
||||
&headers,
|
||||
®ion,
|
||||
&credentials,
|
||||
SystemTime::now(),
|
||||
)?);
|
||||
}
|
||||
}
|
||||
Ok(ProviderAudioTranscriptionRequest {
|
||||
upstream_headers: headers.into_iter().collect(),
|
||||
..request.clone()
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn environment_lookup(key: &str) -> Option<String> {
|
||||
std::env::var(key).ok()
|
||||
}
|
||||
|
|
@ -1,15 +1,14 @@
|
|||
use litellm_core::audio_transcription::{
|
||||
AudioTranscriptionRequest as CoreAudioTranscriptionRequest, ProviderAudioTranscriptionRequest,
|
||||
prepare_audio_transcription_provider_call,
|
||||
};
|
||||
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
|
||||
use litellm_core::error::Error;
|
||||
use serde_json::{Map, Value, json};
|
||||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth;
|
||||
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
|
||||
use litellm_core::error::CoreError;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::common_utils::{audio_transcription_provider_config, has_header, string_headers};
|
||||
use super::handler::sign_request;
|
||||
use super::types::{PreparedAudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
|
||||
use super::types::PreparedAudioTranscriptionRequest;
|
||||
use crate::integrations::custom_guardrail::{
|
||||
CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest,
|
||||
};
|
||||
|
|
@ -26,7 +25,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 +44,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 +62,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,53 +88,36 @@ impl AudioTranscriptionLifecycleHooks {
|
|||
async fn prepare_provider_request(
|
||||
&self,
|
||||
request: PreparedAudioTranscriptionRequest,
|
||||
) -> CoreResult<ProviderAudioTranscriptionRequest> {
|
||||
let config = audio_transcription_provider_config(&request.custom_llm_provider)
|
||||
.ok_or_else(|| CoreError::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(
|
||||
request.api_base.as_deref(),
|
||||
&request.model,
|
||||
&request.optional_params,
|
||||
&env_lookup,
|
||||
)?;
|
||||
let filtered_params = config.map_transcription_params(&request.optional_params);
|
||||
let body = config.transform_transcription_request(
|
||||
&request.model,
|
||||
request.audio,
|
||||
filtered_params,
|
||||
)?;
|
||||
let auth = config.auth_strategy(&request.model, &request.optional_params, &env_lookup)?;
|
||||
let mut upstream_headers = headers.into_iter().collect::<Vec<_>>();
|
||||
if matches!(auth, AudioTranscriptionAuth::Bearer)
|
||||
&& !has_header(
|
||||
&upstream_headers
|
||||
.iter()
|
||||
.cloned()
|
||||
.collect::<std::collections::BTreeMap<_, _>>(),
|
||||
"authorization",
|
||||
)
|
||||
&& let Some(api_key) = request.api_key.as_deref()
|
||||
{
|
||||
upstream_headers.push(("Authorization".to_string(), format!("Bearer {api_key}")));
|
||||
}
|
||||
let provider_request = ProviderAudioTranscriptionRequest {
|
||||
model: request.model,
|
||||
config,
|
||||
url,
|
||||
body: body.body,
|
||||
upstream_headers,
|
||||
timeout: request.timeout,
|
||||
};
|
||||
let provider_request = self.run_during_call_guardrails(provider_request).await?;
|
||||
sign_request(&provider_request, &request.optional_params).await
|
||||
) -> Result<ProviderAudioTranscriptionRequest, Error> {
|
||||
let PreparedAudioTranscriptionRequest {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
audio,
|
||||
api_key,
|
||||
api_base,
|
||||
extra_headers,
|
||||
optional_params,
|
||||
timeout,
|
||||
..
|
||||
} = request;
|
||||
let provider_request =
|
||||
prepare_audio_transcription_provider_call(CoreAudioTranscriptionRequest {
|
||||
model: &model,
|
||||
audio,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: Some(&custom_llm_provider),
|
||||
extra_headers,
|
||||
optional_params,
|
||||
timeout,
|
||||
})?;
|
||||
self.run_during_call_guardrails(provider_request).await
|
||||
}
|
||||
|
||||
async fn run_during_call_guardrails(
|
||||
&self,
|
||||
request: ProviderAudioTranscriptionRequest,
|
||||
) -> CoreResult<ProviderAudioTranscriptionRequest> {
|
||||
) -> Result<ProviderAudioTranscriptionRequest, Error> {
|
||||
if self.guardrail_runner.is_empty() {
|
||||
return Ok(request);
|
||||
}
|
||||
|
|
@ -144,23 +126,23 @@ impl AudioTranscriptionLifecycleHooks {
|
|||
.run_during_call(
|
||||
&guardrail_context(&self.request_metadata),
|
||||
GuardrailRequest::new(json!({
|
||||
"model": request.model,
|
||||
"custom_llm_provider": "bedrock",
|
||||
"url": request.url,
|
||||
"body": request.body,
|
||||
"model": request.model(),
|
||||
"custom_llm_provider": request.custom_llm_provider(),
|
||||
"url": request.url(),
|
||||
"body": request.body(),
|
||||
})),
|
||||
)
|
||||
.await
|
||||
.map_err(guardrail_error_to_core_error)?;
|
||||
let Value::Object(mut data) = guardrail_request.data else {
|
||||
return Err(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 })
|
||||
Ok(request.with_body(body))
|
||||
}
|
||||
|
||||
fn logging_payload(
|
||||
|
|
@ -241,7 +223,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 +263,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",
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,19 +1,17 @@
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::Error;
|
||||
use litellm_core::audio_transcription::execute_audio_transcription_provider_call;
|
||||
use litellm_core::call_lifecycle::CallLifecycle;
|
||||
use serde_json::Value;
|
||||
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod hooks;
|
||||
mod prepare;
|
||||
mod types;
|
||||
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_core::audio_transcription::transformation::AudioTranscriptionProviderConfig;
|
||||
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
|
|
@ -46,13 +45,3 @@ impl CallLifecycleRequest for PreparedAudioTranscriptionRequest {
|
|||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct ProviderAudioTranscriptionRequest {
|
||||
pub(crate) model: String,
|
||||
pub(crate) config: &'static dyn AudioTranscriptionProviderConfig,
|
||||
pub(crate) url: String,
|
||||
pub(crate) body: Value,
|
||||
pub(crate) upstream_headers: Vec<(String, String)>,
|
||||
pub(crate) timeout: Option<Duration>,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
@ -33,6 +32,7 @@ pub(super) fn truncate_error_body(body: &str) -> String {
|
|||
format!("{truncated}... (truncated)")
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(super) fn ocr_provider_config(
|
||||
provider: &str,
|
||||
model: &str,
|
||||
|
|
@ -56,7 +56,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 +65,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)
|
||||
))
|
||||
|
|
@ -74,13 +74,7 @@ pub(super) fn string_headers(
|
|||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn has_header(headers: &[(String, String)], name: &str) -> bool {
|
||||
headers
|
||||
.iter()
|
||||
.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 +132,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 +156,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 +170,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(¤t_url).await?;
|
||||
|
|
@ -202,28 +196,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, ¤t_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 +227,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 +240,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 +249,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 +261,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 +284,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 +310,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 +324,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 +339,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 +352,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 +367,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 +420,7 @@ mod tests {
|
|||
|
||||
assert!(matches!(
|
||||
error,
|
||||
CoreError::InvalidRequest(message)
|
||||
Error::InvalidRequest(message)
|
||||
if message.contains("SSRF protection")
|
||||
));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,13 +1,19 @@
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::error::Error;
|
||||
use litellm_core::http_utils::http_request;
|
||||
use litellm_core::ocr::transformation::OcrResponseHandling;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::common_utils::{poll_document_intelligence, truncate_error_body};
|
||||
use super::types::ProviderOcrRequest;
|
||||
use super::hooks::OcrLifecycleHooks;
|
||||
use super::types::PreparedOcrRequest;
|
||||
use crate::client::http_client;
|
||||
|
||||
pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> CoreResult<Value> {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) async fn execute_ocr_provider_call(
|
||||
request: PreparedOcrRequest,
|
||||
hooks: &OcrLifecycleHooks,
|
||||
) -> Result<Value, Error> {
|
||||
let request = hooks.prepare_provider_request(request).await?;
|
||||
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);
|
||||
|
|
@ -16,10 +22,9 @@ pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> Co
|
|||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
let response = http_request(request_builder)
|
||||
.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 +36,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 +57,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
|
||||
|
|
|
|||
|
|
@ -1,15 +1,10 @@
|
|||
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
|
||||
use litellm_core::error::Error;
|
||||
use serde_json::{Map, Value, json};
|
||||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::ocr::transformation::OcrAuthStrategy;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::common_utils::{
|
||||
convert_document_url_to_data_uri, has_header, ocr_provider_config, string_headers,
|
||||
};
|
||||
use super::common_utils::{convert_document_url_to_data_uri, string_headers};
|
||||
use super::types::{PreparedOcrRequest, ProviderOcrRequest};
|
||||
use crate::integrations::custom_guardrail::{
|
||||
CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest,
|
||||
|
|
@ -27,7 +22,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 +41,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);
|
||||
}
|
||||
|
|
@ -64,6 +59,10 @@ impl OcrLifecycleHooks {
|
|||
.await
|
||||
.map_err(guardrail_error_to_core_error)?;
|
||||
let (document, optional_params) = parse_ocr_pre_call_guardrail_request(guardrail_request)?;
|
||||
let optional_params = match &request.config {
|
||||
Ok(config) => config.map_ocr_params(&optional_params),
|
||||
Err(_) => optional_params,
|
||||
};
|
||||
Ok(PreparedOcrRequest {
|
||||
document,
|
||||
optional_params,
|
||||
|
|
@ -71,25 +70,23 @@ impl OcrLifecycleHooks {
|
|||
})
|
||||
}
|
||||
|
||||
async fn prepare_provider_request(
|
||||
pub(crate) async fn prepare_provider_request(
|
||||
&self,
|
||||
request: PreparedOcrRequest,
|
||||
) -> CoreResult<ProviderOcrRequest> {
|
||||
let config = ocr_provider_config(&request.custom_llm_provider, &request.model)
|
||||
.ok_or_else(|| CoreError::InvalidProvider(request.custom_llm_provider.clone()))?;
|
||||
) -> Result<ProviderOcrRequest, Error> {
|
||||
let config = request.config?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let headers = string_headers(request.extra_headers)?;
|
||||
let auth_strategy = config.auth_strategy();
|
||||
let api_key = (!has_header(&headers, auth_strategy.header_name()))
|
||||
.then(|| config.resolve_api_key(request.api_key.as_deref(), &env_lookup))
|
||||
.transpose()?;
|
||||
let upstream_headers = config.validate_environment(
|
||||
string_headers(request.extra_headers)?,
|
||||
request.api_key.as_deref(),
|
||||
&env_lookup,
|
||||
)?;
|
||||
let url = config.complete_url(
|
||||
request.api_base.as_deref(),
|
||||
&request.model,
|
||||
&request.optional_params,
|
||||
&env_lookup,
|
||||
)?;
|
||||
let filtered_params = config.map_ocr_params(&request.optional_params);
|
||||
let model = request.model.clone();
|
||||
let custom_llm_provider = request.custom_llm_provider.clone();
|
||||
let document = if config.requires_data_uri_document() {
|
||||
|
|
@ -98,9 +95,8 @@ impl OcrLifecycleHooks {
|
|||
request.document
|
||||
};
|
||||
let body = config
|
||||
.transform_ocr_request(&request.model, document, filtered_params)?
|
||||
.transform_ocr_request(&request.model, document, request.optional_params)?
|
||||
.data;
|
||||
let upstream_headers = upstream_headers(&headers, auth_strategy, api_key.as_deref());
|
||||
let body = self
|
||||
.run_during_call_guardrails(&model, &custom_llm_provider, &url, body)
|
||||
.await?;
|
||||
|
|
@ -120,7 +116,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);
|
||||
}
|
||||
|
|
@ -169,9 +165,9 @@ impl OcrLifecycleHooks {
|
|||
}
|
||||
}
|
||||
|
||||
impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> for OcrLifecycleHooks {
|
||||
impl CallLifecycleHooks<PreparedOcrRequest, PreparedOcrRequest, Value> for OcrLifecycleHooks {
|
||||
type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>;
|
||||
type DuringCallFuture<'a> = OcrFuture<'a, ProviderOcrRequest>;
|
||||
type DuringCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>;
|
||||
type SuccessFuture<'a> = OcrLogFuture<'a>;
|
||||
type FailureFuture<'a> = OcrLogFuture<'a>;
|
||||
|
||||
|
|
@ -188,7 +184,7 @@ impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> for OcrLi
|
|||
_context: &'a CallLifecycleContext,
|
||||
request: PreparedOcrRequest,
|
||||
) -> Self::DuringCallFuture<'a> {
|
||||
Box::pin(async move { self.prepare_provider_request(request).await })
|
||||
Box::pin(async move { Ok(request) })
|
||||
}
|
||||
|
||||
fn async_log_success_event<'a>(
|
||||
|
|
@ -217,7 +213,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 {
|
||||
|
|
@ -249,21 +245,6 @@ impl CallLifecycleHooks<PreparedOcrRequest, ProviderOcrRequest, Value> for OcrLi
|
|||
}
|
||||
}
|
||||
|
||||
fn upstream_headers(
|
||||
headers: &[(String, String)],
|
||||
auth_strategy: OcrAuthStrategy,
|
||||
api_key: Option<&str>,
|
||||
) -> Vec<(String, String)> {
|
||||
api_key
|
||||
.map(|api_key| match auth_strategy {
|
||||
OcrAuthStrategy::Bearer => ("Authorization".to_string(), format!("Bearer {api_key}")),
|
||||
OcrAuthStrategy::Header(header_name) => (header_name.to_string(), api_key.to_string()),
|
||||
})
|
||||
.into_iter()
|
||||
.chain(headers.iter().cloned())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
|
||||
GuardrailContext {
|
||||
call_type: CallType::Ocr,
|
||||
|
|
@ -278,19 +259,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 +280,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",
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::Error;
|
||||
use litellm_core::call_lifecycle::CallLifecycle;
|
||||
use serde_json::Value;
|
||||
|
||||
|
|
@ -13,10 +13,13 @@ pub use types::OcrRequest;
|
|||
use handler::execute_ocr_provider_call;
|
||||
use prepare::{PreparedOcrCall, prepare_ocr_call};
|
||||
|
||||
pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
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)
|
||||
.run_request(request, &hooks, |request| {
|
||||
execute_ocr_provider_call(request, &hooks)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ use std::time::{SystemTime, UNIX_EPOCH};
|
|||
|
||||
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use super::common_utils::ocr_provider_config;
|
||||
use super::hooks::OcrLifecycleHooks;
|
||||
use super::types::{OcrRequest, PreparedOcrRequest};
|
||||
use crate::integrations::custom_guardrail::CustomGuardrailRunner;
|
||||
|
|
@ -13,6 +14,7 @@ pub(crate) struct PreparedOcrCall {
|
|||
pub(crate) hooks: OcrLifecycleHooks,
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
|
||||
let call_id = request
|
||||
.litellm_call_id
|
||||
|
|
@ -25,9 +27,25 @@ pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
|
|||
});
|
||||
let model = provider_info.model.to_string();
|
||||
let custom_llm_provider = provider_info.custom_llm_provider.to_string();
|
||||
let config = ocr_provider_config(&custom_llm_provider, &model)
|
||||
.ok_or_else(|| litellm_core::Error::InvalidProvider(custom_llm_provider.clone()));
|
||||
let optional_params = match &config {
|
||||
Ok(config) => {
|
||||
let supported = config.supported_ocr_params();
|
||||
config.map_ocr_params(
|
||||
&request
|
||||
.optional_params
|
||||
.into_iter()
|
||||
.filter(|(name, _)| supported.contains(&name.as_str()))
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
Err(_) => request.optional_params,
|
||||
};
|
||||
|
||||
PreparedOcrCall {
|
||||
request: PreparedOcrRequest {
|
||||
config,
|
||||
model,
|
||||
custom_llm_provider,
|
||||
litellm_call_id: call_id,
|
||||
|
|
@ -35,7 +53,7 @@ pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
|
|||
api_key: request.api_key.map(str::to_string),
|
||||
api_base: request.api_base.map(str::to_string),
|
||||
extra_headers: request.extra_headers,
|
||||
optional_params: request.optional_params,
|
||||
optional_params,
|
||||
timeout: request.timeout,
|
||||
},
|
||||
hooks: OcrLifecycleHooks::new(
|
||||
|
|
|
|||
|
|
@ -1,13 +1,14 @@
|
|||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::error::Error;
|
||||
use litellm_core::http_utils::has_header;
|
||||
use litellm_core::ocr::transformation::OcrResponseHandling;
|
||||
use serde_json::{Map, Value, json};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
use super::common_utils::{has_header, ocr_provider_config, string_headers, truncate_error_body};
|
||||
use super::common_utils::{ocr_provider_config, string_headers, truncate_error_body};
|
||||
use super::{OcrRequest, ocr};
|
||||
use crate::integrations::custom_guardrail::{
|
||||
CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook,
|
||||
|
|
@ -395,7 +396,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 +440,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 +608,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()
|
||||
)
|
||||
);
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ pub struct OcrRequest<'a> {
|
|||
}
|
||||
|
||||
pub(crate) struct PreparedOcrRequest {
|
||||
pub(crate) config: Result<&'static dyn OcrProviderConfig, litellm_core::Error>,
|
||||
pub(crate) model: String,
|
||||
pub(crate) custom_llm_provider: String,
|
||||
pub(crate) litellm_call_id: String,
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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}"),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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}"))
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(¶ms.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(),
|
||||
));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ reqwest.workspace = true
|
|||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
tracing.workspace = true
|
||||
sha2.workspace = true
|
||||
aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true }
|
||||
aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true }
|
||||
|
|
|
|||
14
litellm-rust/crates/core/src/audio_transcription/client.rs
Normal file
14
litellm-rust/crates/core/src/audio_transcription/client.rs
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
use std::sync::OnceLock;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS;
|
||||
|
||||
pub(super) fn http_client() -> &'static reqwest::Client {
|
||||
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
|
||||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new())
|
||||
})
|
||||
}
|
||||
91
litellm-rust/crates/core/src/audio_transcription/handler.rs
Normal file
91
litellm-rust/crates/core/src/audio_transcription/handler.rs
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
use serde_json::Value;
|
||||
|
||||
use crate::error::Error;
|
||||
use crate::http_utils::{http_request, truncate_error_body};
|
||||
|
||||
use super::client::http_client;
|
||||
use super::types::ProviderAudioTranscriptionRequest;
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn execute_audio_transcription_provider_call(
|
||||
request: ProviderAudioTranscriptionRequest,
|
||||
) -> Result<Value, Error> {
|
||||
let body = serde_json::to_vec(&request.body)
|
||||
.map_err(|error| Error::InvalidRequest(format!("invalid audio request body: {error}")))?;
|
||||
let headers = signed_headers(&request, &body).await?;
|
||||
let mut request_builder = http_client().post(&request.url).body(body);
|
||||
for (key, value) in headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
}
|
||||
if let Some(duration) = request.timeout {
|
||||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
let response = http_request(request_builder)
|
||||
.await
|
||||
.map_err(|error| Error::Network(error.to_string()))?;
|
||||
let status = response.status();
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|error| Error::Network(error.to_string()))?;
|
||||
if !status.is_success() {
|
||||
return Err(Error::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
let response_json = 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)?
|
||||
.into_json())
|
||||
}
|
||||
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
async fn signed_headers(
|
||||
request: &ProviderAudioTranscriptionRequest,
|
||||
body: &[u8],
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::SystemTime;
|
||||
|
||||
use crate::audio_transcription::transformation::AudioTranscriptionAuth;
|
||||
use crate::providers::bedrock::audio_transcription::aws_auth_config;
|
||||
use crate::providers::bedrock::aws_base::{resolve_credentials, sign_bedrock_post};
|
||||
|
||||
let AudioTranscriptionAuth::AwsSigV4 { region, .. } = &request.auth else {
|
||||
return Ok(request.upstream_headers.clone());
|
||||
};
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let credentials = resolve_credentials(
|
||||
aws_auth_config(&request.optional_params, &env_lookup),
|
||||
&env_lookup,
|
||||
)
|
||||
.await?;
|
||||
let unsigned: BTreeMap<String, String> = request.upstream_headers.iter().cloned().collect();
|
||||
let signature = sign_bedrock_post(
|
||||
&request.url,
|
||||
body,
|
||||
&unsigned,
|
||||
region,
|
||||
&credentials,
|
||||
SystemTime::now(),
|
||||
)?;
|
||||
Ok(unsigned.into_iter().chain(signature).collect())
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "bedrock-auth"))]
|
||||
async fn signed_headers(
|
||||
request: &ProviderAudioTranscriptionRequest,
|
||||
_body: &[u8],
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
use crate::audio_transcription::transformation::AudioTranscriptionAuth;
|
||||
|
||||
match request.auth {
|
||||
AudioTranscriptionAuth::AwsSigV4 { .. } => Err(Error::Unsupported(
|
||||
"AWS SigV4 requires the bedrock-auth feature",
|
||||
)),
|
||||
AudioTranscriptionAuth::Bearer => Ok(request.upstream_headers.clone()),
|
||||
}
|
||||
}
|
||||
|
|
@ -1,2 +1,21 @@
|
|||
use crate::Error;
|
||||
mod client;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod transformation;
|
||||
pub mod types;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
pub use handler::execute_audio_transcription_provider_call;
|
||||
pub use prepare::prepare_audio_transcription_provider_call;
|
||||
pub use types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
|
||||
.await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
74
litellm-rust/crates/core/src/audio_transcription/prepare.rs
Normal file
74
litellm-rust/crates/core/src/audio_transcription/prepare.rs
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
use crate::error::Error;
|
||||
use crate::http_utils::{has_header, string_headers};
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
use crate::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG;
|
||||
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig};
|
||||
use super::types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
if provider == "bedrock" {
|
||||
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);
|
||||
}
|
||||
let _ = provider;
|
||||
None
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub fn prepare_audio_transcription_provider_call(
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
) -> Result<ProviderAudioTranscriptionRequest, Error> {
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
.or_else(|| {
|
||||
request
|
||||
.custom_llm_provider
|
||||
.map(|provider| CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for audio transcription request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let model = provider_info.model.to_string();
|
||||
let config = provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let mut headers = string_headers("audio transcription", request.extra_headers)?;
|
||||
let auth = config.auth_strategy(&model, &request.optional_params, &env_lookup)?;
|
||||
if matches!(auth, AudioTranscriptionAuth::Bearer)
|
||||
&& !has_header(&headers, "authorization")
|
||||
&& let Some(api_key) = request.api_key
|
||||
{
|
||||
headers.push(("Authorization".to_string(), format!("Bearer {api_key}")));
|
||||
}
|
||||
if !has_header(&headers, "content-type") {
|
||||
headers.push(("Content-Type".to_string(), "application/json".to_string()));
|
||||
}
|
||||
let url = config.complete_url(
|
||||
request.api_base,
|
||||
&model,
|
||||
&request.optional_params,
|
||||
&env_lookup,
|
||||
)?;
|
||||
let filtered_params = config.map_transcription_params(&request.optional_params);
|
||||
let transformed =
|
||||
config.transform_transcription_request(&model, request.audio, filtered_params)?;
|
||||
Ok(ProviderAudioTranscriptionRequest {
|
||||
model,
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
config,
|
||||
url,
|
||||
body: transformed.body,
|
||||
upstream_headers: headers,
|
||||
auth,
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
optional_params: request.optional_params,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
}
|
||||
50
litellm-rust/crates/core/src/audio_transcription/tests.rs
Normal file
50
litellm-rust/crates/core/src/audio_transcription/tests.rs
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
use std::io::{Read, Write};
|
||||
use std::net::TcpListener;
|
||||
use std::thread;
|
||||
|
||||
use serde_json::{Map, json};
|
||||
|
||||
use super::audio_transcription;
|
||||
use super::types::AudioTranscriptionRequest;
|
||||
|
||||
#[tokio::test]
|
||||
async fn bedrock_request_is_signed_and_contains_audio() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").expect("listener");
|
||||
let address = listener.local_addr().expect("address");
|
||||
let server = thread::spawn(move || {
|
||||
let (mut stream, _) = listener.accept().expect("connection");
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0_u8; 16_384];
|
||||
let count = stream.read(&mut buffer).expect("request");
|
||||
request.extend_from_slice(&buffer[..count]);
|
||||
let request = String::from_utf8_lossy(&request);
|
||||
assert!(request.contains("POST /model/mistral.voxtral-mini-3b-2507/converse"));
|
||||
assert!(request.contains("authorization: AWS4-HMAC-SHA256"));
|
||||
assert!(request.contains("x-amz-date:"));
|
||||
assert!(request.contains("\"bytes\":\"AQI=\""));
|
||||
assert!(request.contains("Transcribe the audio. Respond with only the transcript."));
|
||||
let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}";
|
||||
stream.write_all(response).expect("response");
|
||||
});
|
||||
|
||||
let optional_params = Map::from_iter([
|
||||
("aws_access_key_id".to_string(), json!("access-key")),
|
||||
("aws_secret_access_key".to_string(), json!("secret-key")),
|
||||
("aws_region_name".to_string(), json!("us-east-1")),
|
||||
]);
|
||||
let api_base = format!("http://{address}");
|
||||
let response = audio_transcription(AudioTranscriptionRequest {
|
||||
model: "mistral.voxtral-mini-3b-2507",
|
||||
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
|
||||
api_key: None,
|
||||
api_base: Some(&api_base),
|
||||
custom_llm_provider: Some("bedrock"),
|
||||
extra_headers: None,
|
||||
optional_params,
|
||||
timeout: None,
|
||||
})
|
||||
.await
|
||||
.expect("transcription");
|
||||
assert_eq!(response, json!({"text": "hello"}));
|
||||
server.join().expect("server");
|
||||
}
|
||||
|
|
@ -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)]
|
||||
|
|
@ -16,6 +15,7 @@ pub enum AudioTranscriptionAuth {
|
|||
pub trait AudioTranscriptionProviderConfig: Sync {
|
||||
fn supported_transcription_params(&self) -> &'static [&'static str];
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn map_transcription_params(&self, params: &Map<String, Value>) -> Map<String, Value> {
|
||||
params
|
||||
.iter()
|
||||
|
|
@ -32,13 +32,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 +46,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>;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,56 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig};
|
||||
|
||||
pub struct AudioTranscriptionRequest<'a> {
|
||||
pub model: &'a str,
|
||||
pub audio: Value,
|
||||
pub api_key: Option<&'a str>,
|
||||
pub api_base: Option<&'a str>,
|
||||
pub custom_llm_provider: Option<&'a str>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ProviderAudioTranscriptionRequest {
|
||||
pub(super) model: String,
|
||||
pub(super) custom_llm_provider: String,
|
||||
pub(super) config: &'static dyn AudioTranscriptionProviderConfig,
|
||||
pub(super) url: String,
|
||||
pub(super) body: Value,
|
||||
pub(super) upstream_headers: Vec<(String, String)>,
|
||||
pub(super) auth: AudioTranscriptionAuth,
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
pub(super) optional_params: Map<String, Value>,
|
||||
pub(super) timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
impl ProviderAudioTranscriptionRequest {
|
||||
pub fn model(&self) -> &str {
|
||||
&self.model
|
||||
}
|
||||
|
||||
pub fn custom_llm_provider(&self) -> &str {
|
||||
&self.custom_llm_provider
|
||||
}
|
||||
|
||||
pub fn url(&self) -> &str {
|
||||
&self.url
|
||||
}
|
||||
|
||||
pub fn body(&self) -> &Value {
|
||||
&self.body
|
||||
}
|
||||
|
||||
pub fn with_body(self, body: Value) -> Self {
|
||||
Self { body, ..self }
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct AudioTranscriptionRequestData {
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,13 +1,13 @@
|
|||
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;
|
||||
|
||||
const HEADER_CONTEXT: &str = "chat completions";
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(super) fn chat_completions_provider_config(
|
||||
provider: &str,
|
||||
) -> Option<&'static dyn ChatCompletionsProviderConfig> {
|
||||
|
|
@ -23,6 +23,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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,19 +1,23 @@
|
|||
use serde_json::Value;
|
||||
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::http_utils::truncate_error_body;
|
||||
use crate::error::Error;
|
||||
use crate::http_utils::{http_request, truncate_error_body};
|
||||
|
||||
use super::client::http_client;
|
||||
use super::prepare::prepare_provider_request;
|
||||
use super::transformation::ChatCompletionsAuth;
|
||||
use super::types::{
|
||||
ChatCompletionsResponse, ProviderChatCompletionsRequest, ProviderChatResponseData,
|
||||
ResolvedChatCompletionsRequest,
|
||||
};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(super) async fn execute_chat_completions_provider_call(
|
||||
request: ProviderChatCompletionsRequest,
|
||||
) -> CoreResult<ChatCompletionsResponse> {
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let request = prepare_provider_request(request)?;
|
||||
let body = serde_json::to_vec(&request.body).map_err(|err| {
|
||||
CoreError::InvalidRequest(format!(
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize chat completions request: {err}"
|
||||
))
|
||||
})?;
|
||||
|
|
@ -27,14 +31,14 @@ pub(super) async fn execute_chat_completions_provider_call(
|
|||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder.send().await.map_err(|err| {
|
||||
let response = http_request(request_builder).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
if err.is_connect() || err.is_builder() {
|
||||
CoreError::Connect(err.to_string())
|
||||
Error::Connect(err.to_string())
|
||||
} else {
|
||||
CoreError::Network(err.to_string())
|
||||
Error::Network(err.to_string())
|
||||
}
|
||||
})?;
|
||||
|
||||
|
|
@ -42,17 +46,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 +73,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 +84,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 +105,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 +141,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()),
|
||||
|
|
|
|||
|
|
@ -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,16 +18,15 @@ 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 prepare::{parse_messages, resolve_provider_config, resolve_request};
|
||||
use types::{ChatCompletionsRequest, ChatCompletionsResponse};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn chat_completions(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> CoreResult<ChatCompletionsResponse> {
|
||||
execute_chat_completions_provider_call(prepare_chat_completions_call(request)?).await
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
execute_chat_completions_provider_call(resolve_request(request)?).await
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
|
|
|
|||
|
|
@ -1,17 +1,20 @@
|
|||
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};
|
||||
|
||||
use super::common_utils::{chat_completions_provider_config, string_headers};
|
||||
use super::transformation::{ChatCompletionsAuth, ChatCompletionsProviderConfig};
|
||||
use super::types::{ChatCompletionsRequest, ChatMessage, ProviderChatCompletionsRequest};
|
||||
use super::types::{
|
||||
ChatCompletionsRequest, ChatMessage, ProviderChatCompletionsRequest,
|
||||
ResolvedChatCompletionsRequest,
|
||||
};
|
||||
|
||||
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,41 +23,56 @@ 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(
|
||||
pub(super) fn resolve_request(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> CoreResult<ProviderChatCompletionsRequest> {
|
||||
) -> Result<ResolvedChatCompletionsRequest<'_>, 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));
|
||||
}
|
||||
Ok(ResolvedChatCompletionsRequest {
|
||||
model,
|
||||
config,
|
||||
messages,
|
||||
optional_params: request.optional_params,
|
||||
api_key: request.api_key,
|
||||
api_base: request.api_base,
|
||||
extra_headers: request.extra_headers,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
}
|
||||
|
||||
let mut headers = string_headers(request.extra_headers)?;
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn validate_environment(
|
||||
request: &ResolvedChatCompletionsRequest<'_>,
|
||||
model: &str,
|
||||
config: &dyn ChatCompletionsProviderConfig,
|
||||
) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let mut headers = string_headers(request.extra_headers.clone())?;
|
||||
let auth = config.auth(
|
||||
request.api_key,
|
||||
&model,
|
||||
model,
|
||||
&request.optional_params,
|
||||
&env_lookup,
|
||||
)?;
|
||||
|
|
@ -95,7 +113,16 @@ pub(super) fn prepare_chat_completions_call(
|
|||
headers.push(((*name).to_string(), (*value).to_string()));
|
||||
}
|
||||
}
|
||||
Ok((headers, auth))
|
||||
}
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
|
||||
let model = request.model;
|
||||
let config = request.config;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let url = config.complete_url(
|
||||
request.api_base,
|
||||
&model,
|
||||
|
|
@ -103,7 +130,7 @@ pub(super) fn prepare_chat_completions_call(
|
|||
&env_lookup,
|
||||
)?;
|
||||
let transformed =
|
||||
config.transform_request(&model, messages, request.optional_params.clone())?;
|
||||
config.transform_request(&model, request.messages, request.optional_params.clone())?;
|
||||
|
||||
Ok(ProviderChatCompletionsRequest {
|
||||
model,
|
||||
|
|
|
|||
|
|
@ -1,10 +1,16 @@
|
|||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::error::CoreError;
|
||||
use crate::error::Error;
|
||||
|
||||
use super::prepare::prepare_chat_completions_call;
|
||||
use super::prepare::{prepare_provider_request, resolve_request};
|
||||
use super::transformation::ChatCompletionsAuth;
|
||||
use super::types::ChatCompletionsRequest;
|
||||
use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
|
||||
|
||||
fn prepare_chat_completions_call(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
prepare_provider_request(resolve_request(request)?)
|
||||
}
|
||||
|
||||
fn request<'a>(
|
||||
model: &'a str,
|
||||
|
|
@ -29,7 +35,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 +202,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 +214,7 @@ fn rejects_an_unknown_provider() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
)),
|
||||
CoreError::InvalidProvider("openai".to_string())
|
||||
Error::InvalidProvider("openai".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -221,7 +227,7 @@ fn rejects_a_model_with_no_resolvable_provider() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
)),
|
||||
CoreError::InvalidProvider(_)
|
||||
Error::InvalidProvider(_)
|
||||
));
|
||||
}
|
||||
|
||||
|
|
@ -234,7 +240,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 +249,7 @@ fn rejects_an_empty_or_malformed_message_list() {
|
|||
json!("not a list"),
|
||||
json!({}),
|
||||
)),
|
||||
CoreError::InvalidRequest(_)
|
||||
Error::InvalidRequest(_)
|
||||
));
|
||||
}
|
||||
|
||||
|
|
@ -258,7 +264,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 +380,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 +733,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 +751,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 +769,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 +793,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 +803,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, .. }
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
@ -63,9 +62,8 @@ pub trait ChatCompletionsProviderConfig: Sync {
|
|||
false
|
||||
}
|
||||
|
||||
/// Provider parameter names (post-mapping) the Rust path knows how to place
|
||||
/// in the upstream body. Anything outside this set declines the request.
|
||||
fn supported_params(&self) -> &'static [&'static str];
|
||||
/// Supported OpenAI parameter names paired with their provider names.
|
||||
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)];
|
||||
|
||||
/// Parameters consumed as call configuration (credentials, endpoints)
|
||||
/// rather than placed in the body. Accepted, never serialized.
|
||||
|
|
@ -79,7 +77,7 @@ pub trait ChatCompletionsProviderConfig: Sync {
|
|||
optional_params: &Map<String, Value>,
|
||||
) -> Option<Unsupported> {
|
||||
unsupported_param(
|
||||
self.supported_params(),
|
||||
self.supported_openai_params(),
|
||||
self.config_params(),
|
||||
optional_params,
|
||||
)
|
||||
|
|
@ -91,17 +89,17 @@ 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(
|
||||
supported: &'static [&'static str],
|
||||
supported: &'static [(&'static str, &'static str)],
|
||||
config: &'static [&'static str],
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<Unsupported> {
|
||||
|
|
@ -116,7 +114,9 @@ pub fn unsupported_param(
|
|||
.keys()
|
||||
.any(|key| {
|
||||
key != STREAM_PARAM
|
||||
&& !supported.contains(&key.as_str())
|
||||
&& !supported
|
||||
.iter()
|
||||
.any(|(_, provider_name)| *provider_name == key)
|
||||
&& !config.contains(&key.as_str())
|
||||
})
|
||||
.then_some(Unsupported("unrecognized request parameter"))
|
||||
|
|
|
|||
|
|
@ -22,6 +22,17 @@ pub struct ChatCompletionsRequest<'a> {
|
|||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(super) struct ResolvedChatCompletionsRequest<'a> {
|
||||
pub(super) model: String,
|
||||
pub(super) config: &'static dyn ChatCompletionsProviderConfig,
|
||||
pub(super) messages: Vec<ChatMessage>,
|
||||
pub(super) optional_params: Map<String, Value>,
|
||||
pub(super) api_key: Option<&'a str>,
|
||||
pub(super) api_base: Option<&'a str>,
|
||||
pub(super) extra_headers: Option<Map<String, Value>>,
|
||||
pub(super) timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(super) struct ProviderChatCompletionsRequest {
|
||||
pub(super) model: String,
|
||||
pub(super) config: &'static dyn ChatCompletionsProviderConfig,
|
||||
|
|
|
|||
|
|
@ -30,6 +30,8 @@ pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
|
|||
/// Connect timeout for chat completions provider calls, in seconds.
|
||||
pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// `object` field every non-streaming chat completion response carries.
|
||||
pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion";
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,14 @@
|
|||
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};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn http_request(
|
||||
request: reqwest::RequestBuilder,
|
||||
) -> Result<reqwest::Response, reqwest::Error> {
|
||||
request.send().await
|
||||
}
|
||||
|
||||
/// Bound an upstream error body before it crosses a host boundary, so provider
|
||||
/// bodies stay data-minimized.
|
||||
|
|
@ -18,7 +25,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 +34,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 +88,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()
|
||||
)
|
||||
);
|
||||
|
|
|
|||
|
|
@ -13,4 +13,4 @@ pub mod responses;
|
|||
pub mod router;
|
||||
pub mod routing_utils;
|
||||
|
||||
pub use error::{CoreError, CoreResult};
|
||||
pub use error::Error;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
@ -11,6 +10,7 @@ pub(super) use crate::http_utils::{has_bearer_auth, has_header, truncate_error_b
|
|||
|
||||
const HEADER_CONTEXT: &str = "messages";
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(super) fn messages_provider_config(
|
||||
provider: &str,
|
||||
) -> Option<&'static dyn AnthropicMessagesProviderConfig> {
|
||||
|
|
@ -23,6 +23,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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,13 +1,17 @@
|
|||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::error::Error;
|
||||
use crate::http_utils::http_request;
|
||||
|
||||
use super::client::http_client;
|
||||
use super::common_utils::truncate_error_body;
|
||||
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
|
||||
use super::prepare::prepare_provider_request;
|
||||
use super::types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(super) async fn execute_messages_provider_call(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<AnthropicMessagesResponse> {
|
||||
request: MessagesRequest<'_>,
|
||||
) -> Result<AnthropicMessagesResponse, Error> {
|
||||
let request = prepare_provider_request(request)?;
|
||||
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);
|
||||
|
|
@ -16,35 +20,34 @@ pub(super) async fn execute_messages_provider_call(
|
|||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
let response = http_request(request_builder)
|
||||
.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> {
|
||||
request: MessagesRequest<'_>,
|
||||
) -> Result<reqwest::Response, Error> {
|
||||
let request = prepare_provider_request(request)?;
|
||||
if request.provider != ANTHROPIC_MESSAGES_PROVIDER {
|
||||
return Err(CoreError::InvalidRequest(
|
||||
return Err(Error::InvalidRequest(
|
||||
"streaming messages is not supported for this provider".to_string(),
|
||||
));
|
||||
}
|
||||
|
|
@ -57,17 +60,16 @@ pub(super) async fn execute_messages_provider_stream(
|
|||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
let response = http_request(request_builder)
|
||||
.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),
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,18 +15,16 @@ 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> {
|
||||
execute_messages_provider_call(prepare_messages_call(request)?).await
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
|
||||
execute_messages_provider_call(request).await
|
||||
}
|
||||
|
||||
pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult<reqwest::Response> {
|
||||
execute_messages_provider_stream(prepare_messages_call(request)?).await
|
||||
pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Response, Error> {
|
||||
execute_messages_provider_stream(request).await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
|
|||
|
|
@ -1,13 +1,14 @@
|
|||
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};
|
||||
use super::transformation::MessagesAuthStrategy;
|
||||
use super::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy};
|
||||
use super::types::{MessagesRequest, ProviderMessagesRequest};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub(super) fn prepare_messages_call(
|
||||
pub(super) fn prepare_provider_request(
|
||||
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 +19,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,16 +27,49 @@ 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)?;
|
||||
let headers =
|
||||
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
|
||||
|
||||
let typed_request = serde_json::from_value(request.body).map_err(|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| {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
))
|
||||
})?;
|
||||
|
||||
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
|
||||
|
||||
Ok(ProviderMessagesRequest {
|
||||
provider: provider.to_string(),
|
||||
model,
|
||||
config,
|
||||
url,
|
||||
body,
|
||||
upstream_headers: headers,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn validate_environment(
|
||||
config: &dyn AnthropicMessagesProviderConfig,
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
let mut headers = string_headers(extra_headers)?;
|
||||
|
||||
let auth_strategy = config.auth_strategy();
|
||||
let already_authorized = has_header(&headers, auth_strategy.header_name())
|
||||
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers));
|
||||
if !already_authorized {
|
||||
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
|
||||
let api_key = config.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match auth_strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
|
|
@ -51,24 +85,5 @@ 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}"))
|
||||
})?;
|
||||
let transformed = config.transform_request(typed_request)?;
|
||||
let body = serde_json::to_value(transformed).map_err(|err| {
|
||||
CoreError::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
))
|
||||
})?;
|
||||
|
||||
Ok(ProviderMessagesRequest {
|
||||
provider: provider.to_string(),
|
||||
model,
|
||||
config,
|
||||
url,
|
||||
body,
|
||||
upstream_headers: headers,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
Ok(headers)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -46,18 +45,20 @@ pub trait AnthropicMessagesProviderConfig: Sync {
|
|||
]
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_request(
|
||||
&self,
|
||||
request: AnthropicMessagesRequest,
|
||||
) -> CoreResult<AnthropicMessagesRequest> {
|
||||
) -> Result<AnthropicMessagesRequest, Error> {
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_response(
|
||||
&self,
|
||||
_model: &str,
|
||||
response: AnthropicMessagesResponse,
|
||||
) -> CoreResult<AnthropicMessagesResponse> {
|
||||
) -> Result<AnthropicMessagesResponse, Error> {
|
||||
Ok(response)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
@ -28,6 +27,7 @@ pub enum OcrResponseHandling {
|
|||
pub trait OcrProviderConfig: Sync {
|
||||
fn supported_ocr_params(&self) -> &'static [&'static str];
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn map_ocr_params(&self, non_default_params: &Map<String, Value>) -> Map<String, Value> {
|
||||
let mut mapped_params = Map::new();
|
||||
for (param, value) in non_default_params {
|
||||
|
|
@ -43,13 +43,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 +57,32 @@ 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>;
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn validate_environment(
|
||||
&self,
|
||||
headers: Vec<(String, String)>,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
let strategy = self.auth_strategy();
|
||||
if crate::http_utils::has_header(&headers, strategy.header_name()) {
|
||||
return Ok(headers);
|
||||
}
|
||||
let api_key = self.resolve_api_key(api_key, env_lookup)?;
|
||||
let auth_header = match strategy {
|
||||
OcrAuthStrategy::Bearer => ("Authorization".to_string(), format!("Bearer {api_key}")),
|
||||
OcrAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
Ok(std::iter::once(auth_header).chain(headers).collect())
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> OcrAuthStrategy {
|
||||
OcrAuthStrategy::Bearer
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
@ -27,7 +27,12 @@ use crate::chat_completions::response_utils::{finish_reason_for, unix_now, usage
|
|||
/// per-model gate inside `transform_request`, the function this route replaces.
|
||||
/// Forwarding it would send `top_k` to a model that removed sampling params and
|
||||
/// take a 400 after the call, where Python drops it and succeeds.
|
||||
const SUPPORTED_PARAMS: &[&str] = &["max_tokens", "temperature", "top_p", "stop_sequences"];
|
||||
const SUPPORTED_PARAMS: &[(&str, &str)] = &[
|
||||
("max_tokens", "max_tokens"),
|
||||
("temperature", "temperature"),
|
||||
("top_p", "top_p"),
|
||||
("stop", "stop_sequences"),
|
||||
];
|
||||
|
||||
pub struct AnthropicChatCompletionsConfig;
|
||||
|
||||
|
|
@ -74,7 +79,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 +89,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)?,
|
||||
|
|
@ -112,7 +117,8 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
|
|||
})
|
||||
}
|
||||
|
||||
fn supported_params(&self) -> &'static [&'static str] {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)] {
|
||||
SUPPORTED_PARAMS
|
||||
}
|
||||
|
||||
|
|
@ -121,7 +127,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
|
|||
messages: &[ChatMessage],
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<Unsupported> {
|
||||
unsupported_param(SUPPORTED_PARAMS, &[], optional_params)
|
||||
unsupported_param(self.supported_openai_params(), &[], optional_params)
|
||||
.or_else(|| messages.iter().find_map(unsupported_message))
|
||||
// Anthropic rejects a request whose first turn is not a user turn.
|
||||
// Python only repairs that under `litellm.modify_params`, which the
|
||||
|
|
@ -132,30 +138,33 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_request(
|
||||
&self,
|
||||
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),
|
||||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_response(
|
||||
&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 +172,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 +182,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 +190,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,
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
@ -47,12 +47,13 @@ pub fn complete_anthropic_url(
|
|||
}
|
||||
|
||||
impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
) -> Result<String, Error> {
|
||||
Ok(complete_anthropic_url(api_base, env_lookup))
|
||||
}
|
||||
|
||||
|
|
@ -60,7 +61,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 +122,7 @@ mod tests {
|
|||
);
|
||||
assert!(matches!(
|
||||
resolve_anthropic_api_key(None, &|_| None).expect_err("missing key"),
|
||||
CoreError::Auth(_)
|
||||
Error::Auth(_)
|
||||
));
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
@ -142,12 +142,13 @@ fn fold_system_role_messages(request: AnthropicMessagesRequest) -> AnthropicMess
|
|||
}
|
||||
|
||||
impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
) -> Result<String, Error> {
|
||||
complete_azure_anthropic_url(api_base, env_lookup)
|
||||
}
|
||||
|
||||
|
|
@ -155,7 +156,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 +175,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 +191,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
response: AnthropicMessagesResponse,
|
||||
) -> CoreResult<AnthropicMessagesResponse> {
|
||||
) -> Result<AnthropicMessagesResponse, Error> {
|
||||
self.anthropic.transform_response(model, response)
|
||||
}
|
||||
}
|
||||
|
|
@ -268,7 +269,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 +285,7 @@ mod tests {
|
|||
);
|
||||
assert!(matches!(
|
||||
resolve_azure_api_key(None, &|_| None).expect_err("missing key"),
|
||||
CoreError::Auth(_)
|
||||
Error::Auth(_)
|
||||
));
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
@ -14,7 +14,7 @@ const AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGE
|
|||
const AZURE_DOCUMENT_INTELLIGENCE_API_VERSION: &str = "2024-11-30";
|
||||
const AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI: i64 = 96;
|
||||
|
||||
const AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS: &[&str] = &["pages"];
|
||||
const AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS: &[&str] = &["pages", "features"];
|
||||
|
||||
pub struct AzureAiOcrConfig;
|
||||
pub struct AzureDocumentIntelligenceOcrConfig;
|
||||
|
|
@ -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,28 +176,68 @@ 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(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn feature_token_is_valid(token: &str) -> bool {
|
||||
let Some((first, rest)) = token.as_bytes().split_first() else {
|
||||
return false;
|
||||
};
|
||||
first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric)
|
||||
}
|
||||
|
||||
fn invalid_features_error(features: &Value) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"Invalid `features` for Azure Document Intelligence: {features:?}. Expected a list of feature names or a comma-separated string like 'keyValuePairs' or 'keyValuePairs,languages'."
|
||||
))
|
||||
}
|
||||
|
||||
fn normalize_features_param(features: &Value) -> Result<Option<String>, Error> {
|
||||
let normalized = match features {
|
||||
Value::String(value) => value
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.collect::<Vec<_>>()
|
||||
.join(","),
|
||||
Value::Array(values) if values.is_empty() => return Ok(None),
|
||||
Value::Array(values) => values
|
||||
.iter()
|
||||
.map(Value::as_str)
|
||||
.collect::<Option<Vec<_>>>()
|
||||
.ok_or_else(|| invalid_features_error(features))?
|
||||
.into_iter()
|
||||
.map(str::trim)
|
||||
.collect::<Vec<_>>()
|
||||
.join(","),
|
||||
_ => return Err(invalid_features_error(features)),
|
||||
};
|
||||
|
||||
if normalized.split(',').all(feature_token_is_valid) {
|
||||
Ok(Some(normalized))
|
||||
} else {
|
||||
Err(invalid_features_error(features))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn complete_document_intelligence_url(
|
||||
api_base: Option<&str>,
|
||||
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={}",
|
||||
|
|
@ -213,23 +253,30 @@ pub fn complete_document_intelligence_url(
|
|||
url.push_str(&normalized);
|
||||
}
|
||||
|
||||
if let Some(features) = optional_params.get("features")
|
||||
&& let Some(normalized) = normalize_features_param(features)?
|
||||
{
|
||||
url.push_str("&features=");
|
||||
url.push_str(&normalized);
|
||||
}
|
||||
|
||||
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 +285,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 +337,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 +345,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 +355,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 +363,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 +382,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 +406,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 +461,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 +469,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)
|
||||
}
|
||||
|
||||
|
|
@ -475,6 +522,103 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_normalizes_features() {
|
||||
let params = serde_json::Map::from_iter([(
|
||||
"features".to_string(),
|
||||
json!("keyValuePairs, languages"),
|
||||
)]);
|
||||
let url = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect("url builds");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&features=keyValuePairs,languages"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_combines_pages_and_feature_list() {
|
||||
let params = serde_json::Map::from_iter([
|
||||
("pages".to_string(), json!([0, 1, 2])),
|
||||
(
|
||||
"features".to_string(),
|
||||
json!([" keyValuePairs ", "languages"]),
|
||||
),
|
||||
]);
|
||||
let url = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect("url builds");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&pages=1,2,3&features=keyValuePairs,languages"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_omits_empty_feature_list() {
|
||||
let params = serde_json::Map::from_iter([("features".to_string(), json!([]))]);
|
||||
let url = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect("url builds");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_rejects_invalid_features() {
|
||||
for features in [
|
||||
json!("keyValuePairs&pages=9"),
|
||||
json!(""),
|
||||
json!(["keyValuePairs", 1]),
|
||||
json!({"feature": "keyValuePairs"}),
|
||||
] {
|
||||
let params = serde_json::Map::from_iter([("features".to_string(), features.clone())]);
|
||||
let error = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect_err("invalid features must fail");
|
||||
|
||||
assert!(
|
||||
matches!(error, Error::InvalidRequest(message) if message.contains("Invalid `features`")),
|
||||
"features={features:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_maps_features() {
|
||||
let params = Map::from_iter([
|
||||
("features".to_string(), json!(["keyValuePairs"])),
|
||||
("unsupported".to_string(), json!(true)),
|
||||
]);
|
||||
|
||||
assert_eq!(
|
||||
AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG.map_ocr_params(¶ms),
|
||||
Map::from_iter([("features".to_string(), json!(["keyValuePairs"]))])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_request_uses_base64_source_for_data_uri() {
|
||||
let body = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
}
|
||||
|
|
@ -46,16 +46,18 @@ fn optional_string<'a>(params: &'a Map<String, Value>, key: &str) -> Option<&'a
|
|||
}
|
||||
|
||||
impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn supported_transcription_params(&self) -> &'static [&'static str] {
|
||||
SUPPORTED_PARAMS
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_transcription_request(
|
||||
&self,
|
||||
_model: &str,
|
||||
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") {
|
||||
|
|
@ -83,18 +85,19 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_transcription_response(
|
||||
&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 +114,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 +136,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),
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)
|
||||
.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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
@ -23,11 +23,12 @@ use super::super::constants::{AWS_BEARER_TOKEN_BEDROCK, BEDROCK_RUNTIME_ENDPOINT
|
|||
/// `additionalModelRequestFields` for Anthropic base models and to
|
||||
/// `inferenceConfig` otherwise, and that branch reads the model catalog the
|
||||
/// core cannot see.
|
||||
const SUPPORTED_PARAMS: &[&str] = &["maxTokens", "temperature", "topP", "stopSequences"];
|
||||
|
||||
/// Params that belong in `inferenceConfig`, in the order Python's
|
||||
/// `AmazonConverseConfig` declares them, so bodies compare cleanly.
|
||||
const INFERENCE_CONFIG_PARAMS: &[&str] = SUPPORTED_PARAMS;
|
||||
const SUPPORTED_PARAMS: &[(&str, &str)] = &[
|
||||
("max_tokens", "maxTokens"),
|
||||
("temperature", "temperature"),
|
||||
("top_p", "topP"),
|
||||
("stop", "stopSequences"),
|
||||
];
|
||||
|
||||
const AWS_BEDROCK_RUNTIME_ENDPOINT: &str = "aws_bedrock_runtime_endpoint";
|
||||
|
||||
|
|
@ -66,7 +67,7 @@ fn converse_body(conversation: &Conversation, params: &Map<String, Value>) -> Va
|
|||
})
|
||||
.collect();
|
||||
|
||||
let inference_config = Map::from_iter(INFERENCE_CONFIG_PARAMS.iter().filter_map(|name| {
|
||||
let inference_config = Map::from_iter(SUPPORTED_PARAMS.iter().filter_map(|(_, name)| {
|
||||
params
|
||||
.get(*name)
|
||||
.map(|value| ((*name).to_string(), value.clone()))
|
||||
|
|
@ -110,7 +111,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 +138,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
|
||||
|
|
@ -162,7 +163,8 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
|
|||
&[("Content-Type", "application/json")]
|
||||
}
|
||||
|
||||
fn supported_params(&self) -> &'static [&'static str] {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn supported_openai_params(&self) -> &'static [(&'static str, &'static str)] {
|
||||
SUPPORTED_PARAMS
|
||||
}
|
||||
|
||||
|
|
@ -175,32 +177,36 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
|
|||
messages: &[ChatMessage],
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<Unsupported> {
|
||||
unsupported_param(SUPPORTED_PARAMS, CONFIG_PARAMS, optional_params)
|
||||
.or_else(|| messages.iter().find_map(unsupported_message))
|
||||
// Python's Converse translation drops blank text blocks instead of
|
||||
// substituting the placeholder the shared conversation builder
|
||||
// applies, so decline blank text rather than diverge.
|
||||
.or_else(|| {
|
||||
messages
|
||||
.iter()
|
||||
.any(has_blank_text)
|
||||
.then_some(Unsupported("blank message text"))
|
||||
})
|
||||
// Converse has no assistant prefill: Python inserts a continue turn
|
||||
// when a conversation opens or closes on an assistant message, and
|
||||
// only under `litellm.modify_params`, which the core cannot see.
|
||||
// Declining both ends also keeps the shared builder's final
|
||||
// assistant right-strip (an Anthropic rule) unreachable here.
|
||||
.or_else(|| {
|
||||
let conversation = build_conversation(messages);
|
||||
let ends_on_assistant = conversation
|
||||
.turns
|
||||
.last()
|
||||
.is_some_and(|turn| turn.role == TurnRole::Assistant);
|
||||
(!conversation.opens_on_user_turn() || ends_on_assistant).then_some(Unsupported(
|
||||
"conversation does not run user turn to user turn",
|
||||
))
|
||||
})
|
||||
unsupported_param(
|
||||
self.supported_openai_params(),
|
||||
CONFIG_PARAMS,
|
||||
optional_params,
|
||||
)
|
||||
.or_else(|| messages.iter().find_map(unsupported_message))
|
||||
// Python's Converse translation drops blank text blocks instead of
|
||||
// substituting the placeholder the shared conversation builder
|
||||
// applies, so decline blank text rather than diverge.
|
||||
.or_else(|| {
|
||||
messages
|
||||
.iter()
|
||||
.any(has_blank_text)
|
||||
.then_some(Unsupported("blank message text"))
|
||||
})
|
||||
// Converse has no assistant prefill: Python inserts a continue turn
|
||||
// when a conversation opens or closes on an assistant message, and
|
||||
// only under `litellm.modify_params`, which the core cannot see.
|
||||
// Declining both ends also keeps the shared builder's final
|
||||
// assistant right-strip (an Anthropic rule) unreachable here.
|
||||
.or_else(|| {
|
||||
let conversation = build_conversation(messages);
|
||||
let ends_on_assistant = conversation
|
||||
.turns
|
||||
.last()
|
||||
.is_some_and(|turn| turn.role == TurnRole::Assistant);
|
||||
(!conversation.opens_on_user_turn() || ends_on_assistant).then_some(Unsupported(
|
||||
"conversation does not run user turn to user turn",
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn transform_request(
|
||||
|
|
@ -208,7 +214,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 +224,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 +243,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 +253,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"),
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -70,18 +70,20 @@ pub struct MistralOcrConfig;
|
|||
pub const MISTRAL_OCR_CONFIG: MistralOcrConfig = MistralOcrConfig;
|
||||
|
||||
impl OcrProviderConfig for MistralOcrConfig {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn supported_ocr_params(&self) -> &'static [&'static str] {
|
||||
SUPPORTED_OCR_PARAMS
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_ocr_request(
|
||||
&self,
|
||||
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),
|
||||
});
|
||||
|
|
@ -100,14 +102,15 @@ impl OcrProviderConfig for MistralOcrConfig {
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn transform_ocr_response(
|
||||
&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),
|
||||
})?;
|
||||
|
|
@ -134,13 +137,14 @@ impl OcrProviderConfig for MistralOcrConfig {
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
_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,11 +152,12 @@ 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)
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub fn supported_ocr_params() -> &'static [&'static str] {
|
||||
MISTRAL_OCR_CONFIG.supported_ocr_params()
|
||||
}
|
||||
|
|
@ -161,15 +166,17 @@ pub fn map_ocr_params(non_default_params: &Map<String, Value>) -> Map<String, Va
|
|||
MISTRAL_OCR_CONFIG.map_ocr_params(non_default_params)
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
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> {
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub fn transform_ocr_response(model: &str, response_json: Value) -> Result<OcrResponseData, Error> {
|
||||
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
|
||||
}
|
||||
|
||||
|
|
@ -179,68 +186,180 @@ mod tests {
|
|||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn supported_params_match_python_mistral_ocr_config() {
|
||||
fn extract_header_is_a_supported_ocr_param() {
|
||||
assert!(supported_ocr_params().contains(&"extract_header"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extract_footer_is_a_supported_ocr_param() {
|
||||
assert!(supported_ocr_params().contains(&"extract_footer"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn existing_ocr_params_remain_supported() {
|
||||
for param in [
|
||||
"pages",
|
||||
"include_image_base64",
|
||||
"image_limit",
|
||||
"image_min_size",
|
||||
"bbox_annotation_format",
|
||||
"document_annotation_format",
|
||||
] {
|
||||
assert!(supported_ocr_params().contains(¶m));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn map_ocr_params_forwards_extract_header() {
|
||||
let params = json!({"extract_header": true});
|
||||
assert_eq!(
|
||||
supported_ocr_params(),
|
||||
&[
|
||||
"pages",
|
||||
"include_image_base64",
|
||||
"image_limit",
|
||||
"image_min_size",
|
||||
"bbox_annotation_format",
|
||||
"document_annotation_format",
|
||||
"document_annotation_prompt",
|
||||
"extract_header",
|
||||
"extract_footer",
|
||||
"table_format",
|
||||
"confidence_scores_granularity",
|
||||
"include_blocks",
|
||||
"id",
|
||||
]
|
||||
map_ocr_params(params.as_object().unwrap()),
|
||||
params.as_object().unwrap().clone()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn map_ocr_params_forwards_extract_footer() {
|
||||
let params = json!({"extract_footer": true});
|
||||
assert_eq!(
|
||||
map_ocr_params(params.as_object().unwrap()),
|
||||
params.as_object().unwrap().clone()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn map_ocr_params_forwards_extract_header_and_footer() {
|
||||
let params = json!({"extract_header": true, "extract_footer": false});
|
||||
assert_eq!(
|
||||
map_ocr_params(params.as_object().unwrap()),
|
||||
params.as_object().unwrap().clone()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn map_ocr_params_drops_unknown_params() {
|
||||
let params = json!({
|
||||
"extract_header": true,
|
||||
"unsupported_param": "value",
|
||||
"pages": [0, 1]
|
||||
});
|
||||
let params = json!({"extract_header": true, "unsupported_param": "value"});
|
||||
let mapped = map_ocr_params(params.as_object().unwrap());
|
||||
|
||||
assert_eq!(mapped.get("extract_header"), Some(&json!(true)));
|
||||
assert_eq!(mapped.get("pages"), Some(&json!([0, 1])));
|
||||
assert!(!mapped.contains_key("unsupported_param"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_ocr_request_builds_mistral_body() {
|
||||
fn new_ocr_params_are_supported() {
|
||||
for param in [
|
||||
"table_format",
|
||||
"confidence_scores_granularity",
|
||||
"document_annotation_prompt",
|
||||
"include_blocks",
|
||||
"id",
|
||||
] {
|
||||
assert!(supported_ocr_params().contains(¶m));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn map_ocr_params_forwards_new_ocr_params() {
|
||||
for (param, value) in [
|
||||
("table_format", json!("html")),
|
||||
("confidence_scores_granularity", json!("word")),
|
||||
(
|
||||
"document_annotation_prompt",
|
||||
json!("Extract all invoice line items"),
|
||||
),
|
||||
("include_blocks", json!(true)),
|
||||
("id", json!("req-123")),
|
||||
] {
|
||||
let params = json!({param: value});
|
||||
assert_eq!(
|
||||
map_ocr_params(params.as_object().unwrap()),
|
||||
params.as_object().unwrap().clone()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_ocr_request_includes_each_optional_param() {
|
||||
let document = json!({
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf"
|
||||
});
|
||||
for (param, value) in [
|
||||
("table_format", json!("html")),
|
||||
("confidence_scores_granularity", json!("word")),
|
||||
(
|
||||
"document_annotation_prompt",
|
||||
json!("Extract all invoice line items"),
|
||||
),
|
||||
("id", json!("req-123")),
|
||||
("extract_header", json!(true)),
|
||||
("include_blocks", json!(true)),
|
||||
("pages", json!([0, 1])),
|
||||
] {
|
||||
let result = transform_ocr_request(
|
||||
"mistral-ocr-latest",
|
||||
document.clone(),
|
||||
json!({param: value}).as_object().unwrap().clone(),
|
||||
)
|
||||
.expect("request should transform");
|
||||
assert_eq!(result.data.get(param), Some(&value));
|
||||
assert_eq!(result.data.get("model"), Some(&json!("mistral-ocr-latest")));
|
||||
assert_eq!(result.data.get("document"), Some(&document));
|
||||
assert_eq!(result.files, None);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_ocr_request_includes_multiple_new_params() {
|
||||
let document = json!({
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/doc.pdf"
|
||||
});
|
||||
let optional_params = json!({
|
||||
"include_image_base64": true,
|
||||
"table_format": "html"
|
||||
"table_format": "html",
|
||||
"confidence_scores_granularity": "page",
|
||||
"extract_header": true
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.clone();
|
||||
|
||||
let result = transform_ocr_request("mistral-ocr-latest", document.clone(), optional_params)
|
||||
let result = transform_ocr_request("mistral-ocr-latest", document, optional_params)
|
||||
.expect("request should transform");
|
||||
|
||||
assert_eq!(result.data.get("table_format"), Some(&json!("html")));
|
||||
assert_eq!(
|
||||
result.data,
|
||||
json!({
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": document,
|
||||
"include_image_base64": true,
|
||||
"table_format": "html"
|
||||
})
|
||||
result.data.get("confidence_scores_granularity"),
|
||||
Some(&json!("page"))
|
||||
);
|
||||
assert_eq!(result.files, None);
|
||||
assert_eq!(result.data.get("extract_header"), Some(&json!(true)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_ocr_response_preserves_blocks_and_confidence_scores() {
|
||||
let blocks = json!([{"type": "title", "content": "Invoice"}]);
|
||||
let confidence_scores = json!({"page": 0.98});
|
||||
let response = json!({
|
||||
"pages": [{"index": 0, "markdown": "# Invoice", "blocks": blocks, "confidence_scores": confidence_scores}],
|
||||
"model": "mistral-ocr-4-0",
|
||||
"usage_info": {"pages_processed": 1}
|
||||
});
|
||||
let result =
|
||||
transform_ocr_response("mistral-ocr-4-0", response).expect("response should transform");
|
||||
assert_eq!(result.pages[0].get("blocks"), Some(&blocks));
|
||||
assert_eq!(
|
||||
result.pages[0].get("confidence_scores"),
|
||||
Some(&confidence_scores)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_ocr_response_preserves_ocr4_page_fields() {
|
||||
let response = json!({
|
||||
"pages": [{"index": 0, "markdown": "table page", "tables": [{"rows": 2, "cols": 3}], "hyperlinks": ["https://example.com"], "header": "Acme Corp", "footer": "Page 1"}],
|
||||
"model": "mistral-ocr-4-0",
|
||||
"usage_info": {"pages_processed": 1}
|
||||
});
|
||||
let result = transform_ocr_response("mistral-ocr-4-0", response.clone())
|
||||
.expect("response should transform");
|
||||
assert_eq!(result.pages[0], response["pages"][0]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -250,7 +369,7 @@ mod tests {
|
|||
|
||||
assert_eq!(
|
||||
err,
|
||||
CoreError::InvalidType {
|
||||
Error::InvalidType {
|
||||
expected: "object",
|
||||
actual: "string",
|
||||
}
|
||||
|
|
@ -307,6 +426,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()));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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).";
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue