mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge remote-tracking branch 'origin/main' into litellm_bedrock_guardrail_attachment_scan
This commit is contained in:
commit
eeda4cb61d
527 changed files with 33182 additions and 6545 deletions
|
|
@ -17,3 +17,6 @@ rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
|||
|
||||
[target.aarch64-apple-darwin]
|
||||
rustflags = ["-C", "link-arg=-undefined", "-C", "link-arg=dynamic_lookup"]
|
||||
|
||||
[env]
|
||||
SQLX_OFFLINE = "true"
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ legacy_paths() {
|
|||
echo tests/unit/google_genai
|
||||
echo tests/unit/router_strategy
|
||||
echo tests/unit/router_utils
|
||||
echo tests/unit/proxy/common_utils/test_cache_aware_routing.py
|
||||
echo tests/unit/enterprise/enterprise_callbacks/send_emails
|
||||
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
|
||||
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py
|
||||
|
|
|
|||
7
.github/ci-coverage-allowlist.yml
vendored
7
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -111,3 +111,10 @@ dockerfiles:
|
|||
An example image under cookbook/ that is documentation rather than a shipped artifact
|
||||
paths:
|
||||
- cookbook/litellm-ollama-docker-image/Dockerfile
|
||||
- reason: >-
|
||||
The Rust gateway image compiles the whole workspace in release mode, which is too slow for
|
||||
a per-pull-request job while the gateway binary is still being assembled; the Rust lint,
|
||||
clippy, and compile jobs already cover the code it packages. Revisit when the gateway is
|
||||
published
|
||||
paths:
|
||||
- litellm-rust/crates/gateway/Dockerfile
|
||||
|
|
|
|||
2
.github/scripts/assert_ci_coverage.py
vendored
2
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -516,7 +516,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
|
|||
str(path.relative_to(repo_root))
|
||||
for folders in groups.values()
|
||||
for folder in folders
|
||||
for path in (integration_root / folder).glob("test_*.py")
|
||||
for path in (integration_root / folder).rglob("test_*.py")
|
||||
)
|
||||
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
|
||||
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()
|
||||
|
|
|
|||
2
.github/scripts/verify_linux_native_wheel.py
vendored
2
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -214,7 +214,7 @@ def main(
|
|||
native_module: Final = load_native_module(native_path)
|
||||
native_module_loads: Final = native_module is not None
|
||||
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
|
||||
native_size_limit: Final = 40_000_000
|
||||
native_size_limit: Final = 45_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),
|
||||
|
|
|
|||
6
Makefile
6
Makefile
|
|
@ -4,7 +4,7 @@
|
|||
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
|
||||
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
|
||||
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
|
||||
test-rust-extension \
|
||||
test-rust-extension rust-sqlx-prepare \
|
||||
info lint lint-inner lint-dev lint-checks format \
|
||||
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
|
||||
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
|
||||
|
|
@ -56,6 +56,7 @@ help:
|
|||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
@echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container"
|
||||
@echo ""
|
||||
@echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide"
|
||||
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
|
||||
|
|
@ -306,6 +307,9 @@ test-rust-extension:
|
|||
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
|
||||
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust
|
||||
|
||||
rust-sqlx-prepare:
|
||||
cd litellm-rust && cargo run -p litellm-db-testing --bin sqlx-prepare
|
||||
|
||||
test: install-test-deps
|
||||
$(UV_RUN) pytest tests/
|
||||
|
||||
|
|
|
|||
37
cookbook/litellm_proxy_server/mcp/README.md
Normal file
37
cookbook/litellm_proxy_server/mcp/README.md
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
# Publish MCP servers in the AI Hub
|
||||
|
||||
Set `litellm_settings.public_mcp_servers` to the concrete IDs of the servers you want listed in the public AI Hub. Pin `server_id` in each configuration entry so the publication list stays stable across deployments
|
||||
|
||||
```yaml
|
||||
mcp_servers:
|
||||
documentation:
|
||||
server_id: documentation-mcp
|
||||
url: https://mcp.example.com/mcp
|
||||
transport: http
|
||||
available_on_public_internet: true
|
||||
|
||||
litellm_settings:
|
||||
public_mcp_hub_strict_whitelist: true
|
||||
public_mcp_servers:
|
||||
- documentation-mcp
|
||||
```
|
||||
|
||||
Use `documentation-mcp`, the `server_id`, in the publication list. The configuration key `documentation`, display names, and aliases are not publication IDs. Database-created servers use the ID returned by `/v1/mcp/server`
|
||||
|
||||
The dashboard's **AI Hub > MCP Hub > Manage MCP Hub Visibility** dialog edits this same list. Its YAML example includes the selected server IDs. With database-backed configuration (`store_model_in_db: true`), a value declared in YAML is owned by that file: edit the file and reload, or remove that key from YAML to let the dashboard manage it in the database. File-backed deployments can save the list directly to their configuration file
|
||||
|
||||
To remove all explicit entries, save an empty selection in the dialog or configure:
|
||||
|
||||
```yaml
|
||||
litellm_settings:
|
||||
public_mcp_hub_strict_whitelist: true
|
||||
public_mcp_servers: []
|
||||
```
|
||||
|
||||
## Hub listing and network access
|
||||
|
||||
The **Hub listing** column in AI Hub identifies servers that appear in `/public/mcp_hub`. The dashboard derives this status from the current registry and publication settings. Setting `mcp_info.is_public` on a server does not publish it; that response field is derived metadata. `mcp_info.is_public_explicit` identifies registered servers included in the explicit publication list
|
||||
|
||||
Gateway cards and server details show **All Networks** when `available_on_public_internet` is enabled or the server is explicitly published in `public_mcp_servers`. They show **Internal Only** when both are false. The per-server flag defaults to `true`; explicit publication overrides a disabled flag for compatibility. Older proxies that omit the metadata needed to determine access show **Unknown**. These labels describe allowed client IPs; authentication and tool permissions still apply
|
||||
|
||||
The default `public_mcp_hub_strict_whitelist: true` lists only registered servers in `public_mcp_servers`. Legacy mode (`false`) additionally lists registered servers with `available_on_public_internet: true`. In legacy mode, clearing the explicit publication list leaves these automatically listed servers visible. Enable strict mode when the publication list should fully determine hub visibility
|
||||
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
22
litellm-rust/.agents/skills/rust-tracing/SKILL.md
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
---
|
||||
name: rust-tracing
|
||||
description: Add or change Rust diagnostic tracing in litellm-rust, including route spans, subscriber layers, and Python logger delivery
|
||||
---
|
||||
|
||||
# Rust tracing
|
||||
|
||||
Use upstream `tracing` throughout Rust, including `#[tracing::instrument]`, events, and span propagation. Centralize collection and delivery infrastructure in `crates/tracing`. Direct upstream imports still reach our configured subscriber; re-exporting macros does not control delivery. Do not introduce Rust `log` or `pyo3-log` for this path
|
||||
|
||||
`litellm-tracing` owns shared subscriber layers, span field collection, and diagnostic processing. Keep adapters composable as `tracing_subscriber::Layer`s, with `Logger` providing host setup. Runtime-specific delivery belongs in the host bridge. The Python bridge delivers directly to the existing Python SDK logger, preserving its handlers, filtering, redaction, and request correlation. Keep Python dependencies out of `crates/tracing`
|
||||
|
||||
Hosts configure subscribers. Keep Python execution scoped to its captured dispatch rather than installing a process-wide subscriber. Propagate both span context and dispatch across spawned work and returned streams
|
||||
|
||||
In core, instrument execution shared by native calls and hosted machines. Use consistent route, model, provider, streaming, and outcome fields. Put status recording at shared provider boundaries instead of scattering basic logging through handlers. Keep upstream HTTP status separate from route success
|
||||
|
||||
Use `skip_all` and explicitly selected fields. Basic tracing excludes bodies, credentials, headers, and raw error strings. Avoid automatic `ret` or `err` capture of sensitive values. Keep payload diagnostics separate and subject to existing redaction
|
||||
|
||||
A returned stream retains its route span until exhaustion, error, or drop, with exactly one terminal outcome. Builder construction does not start a trace. Never hold a span entry guard across an await. Diagnostic tracing remains separate from lifecycle callbacks and `CustomLogger` dispatch
|
||||
|
||||
Use `litellm_tracing::sink_layer` to compose a sink with other subscriber layers. It inherits span fields into events and emits span-close summaries with elapsed time. Test observable records, concurrent isolation, dynamic filtering, sensitive-field exclusion, and stream cancellation when changing this behavior
|
||||
|
||||
Consult the [tracing API](https://docs.rs/tracing/latest/tracing/) and [subscriber layers](https://docs.rs/tracing-subscriber/latest/tracing_subscriber/layer/index.html) for implementation details
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
# Rust workspace rules
|
||||
|
||||
For diagnostic tracing changes, follow [.agents/skills/rust-tracing/SKILL.md](.agents/skills/rust-tracing/SKILL.md)
|
||||
|
||||
## Test placement
|
||||
|
||||
- Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;`
|
||||
|
|
@ -16,7 +18,9 @@ Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new
|
|||
## Error definitions
|
||||
|
||||
- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs`
|
||||
- Put message templates in the variant's `#[error(...)]` declaration. Callers pass only the small typed arguments needed to fill them, never `Error::Variant(format!(...))` or a preformatted message. Keep the smallest set of neutral variants that callers need to distinguish; different wording or providers do not justify new variants
|
||||
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message
|
||||
- Keep shared error enums minimal and provider-neutral. Provider names, credential types, configuration fields, and setup guidance belong in caller-supplied data, not dedicated variants or hardcoded shared messages. Reuse a variant for the same failure mode across providers, such as `MissingApiBase { provider: "Azure", guidance: "..." }`. An exact parity message does not justify a provider-specific variant when caller-supplied context can preserve it
|
||||
- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string
|
||||
- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return
|
||||
- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it
|
||||
|
|
|
|||
1403
litellm-rust/Cargo.lock
generated
1403
litellm-rust/Cargo.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -13,11 +13,15 @@ litellm-config = { path = "crates/config" }
|
|||
litellm-router = { path = "crates/router" }
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
litellm-gateway-inference = { path = "crates/gateway-inference" }
|
||||
litellm-gateway-auth = { path = "crates/gateway-auth" }
|
||||
litellm-gateway-management = { path = "crates/gateway-management" }
|
||||
litellm-gateway-ui = { path = "crates/gateway-ui" }
|
||||
litellm-coroutine = { path = "crates/coroutine" }
|
||||
litellm-host = { path = "crates/host" }
|
||||
litellm-host-http = { path = "crates/host-http" }
|
||||
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
|
||||
litellm-framing = { path = "crates/framer" }
|
||||
litellm-auth = { path = "crates/auth" }
|
||||
|
|
@ -36,6 +40,8 @@ litellm-http = { path = "crates/http" }
|
|||
litellm-llms = { path = "crates/llms" }
|
||||
litellm-types = { path = "crates/types" }
|
||||
litellm-core-utils = { path = "crates/core-utils" }
|
||||
litellm-db = { path = "crates/db" }
|
||||
litellm-db-testing = { path = "crates/db-testing" }
|
||||
litellm-cache = { path = "crates/cache" }
|
||||
litellm-cache-azure-blob = { path = "crates/cache-azure-blob" }
|
||||
litellm-cache-memory = { path = "crates/cache-memory" }
|
||||
|
|
@ -56,6 +62,8 @@ litellm-python-compat = { path = "crates/python-compat" }
|
|||
|
||||
tracing = "0.1"
|
||||
axum = { version = "0.8.9", default-features = false, features = ["http1", "tokio", "multipart"] }
|
||||
axum-login = "0.18.0"
|
||||
tower-sessions = { version = "0.14.0", features = ["memory-store"] }
|
||||
bytes = "1"
|
||||
http = "1"
|
||||
google-cloud-auth = { version = "1.16.0", default-features = false }
|
||||
|
|
@ -80,6 +88,7 @@ serde = { version = "1.0", features = ["derive"] }
|
|||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
serde_with = { version = "=3.16.1", default-features = false, features = ["std", "macros"] }
|
||||
sha2 = "0.10"
|
||||
sqlx = { version = "0.9.0", default-features = false, features = ["json", "macros", "postgres", "runtime-tokio", "chrono", "tls-rustls-ring-native-roots"] }
|
||||
subtle = "2"
|
||||
thiserror = "2.0"
|
||||
tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] }
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
# The Tokio runtime is reached only through `host-python/src/execution.rs`, whose fork gate
|
||||
# The Tokio runtime is reached only through `host-python/src/runtime.rs`, whose fork gate
|
||||
# must see every entry. Going around it makes a fork-after-use hang instead of raising.
|
||||
disallowed-methods = [
|
||||
{ path = "pyo3_async_runtimes::tokio::get_runtime", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
|
|
@ -12,6 +12,13 @@ disallowed-methods = [
|
|||
{ path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" },
|
||||
{ path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" },
|
||||
{ path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" },
|
||||
{ path = "sqlx::query", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_with", reason = "use sqlx::query! or query_file! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_as_with", reason = "use sqlx::query_as! or query_file_as! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::query_scalar_with", reason = "use sqlx::query_scalar! so the SQL is checked against the migrated schema" },
|
||||
{ path = "sqlx::raw_sql", reason = "raw_sql is unchecked; use the checked query macros" },
|
||||
]
|
||||
|
||||
# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS,
|
||||
|
|
|
|||
|
|
@ -133,7 +133,7 @@ impl NativeAzureTokenAcquirer {
|
|||
let token = credential
|
||||
.get_token(&[scope.as_str()], None)
|
||||
.await
|
||||
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?;
|
||||
.map_err(|error| Error::CredentialAcquisition(error.to_string().into()))?;
|
||||
let expires_on = u64::try_from(token.expires_on.unix_timestamp())
|
||||
.ok()
|
||||
.map(|seconds| UNIX_EPOCH + Duration::from_secs(seconds));
|
||||
|
|
@ -250,7 +250,12 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
|
|||
let Some(authority) = authority else {
|
||||
return Ok(());
|
||||
};
|
||||
let url = url::Url::parse(authority.value()).map_err(|_| Error::InvalidAzureAuthority)?;
|
||||
let url = url::Url::parse(authority.value()).map_err(|_| {
|
||||
Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into(),
|
||||
)
|
||||
})?;
|
||||
if url.scheme() != "https"
|
||||
|| url.host_str().is_none()
|
||||
|| !url.username().is_empty()
|
||||
|
|
@ -259,7 +264,10 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> {
|
|||
|| url.fragment().is_some()
|
||||
|| !matches!(url.path(), "" | "/")
|
||||
{
|
||||
return Err(Error::InvalidAzureAuthority);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -368,7 +376,9 @@ fn trusted_source(sources: &[InputSource]) -> InputSource {
|
|||
}
|
||||
|
||||
fn mixed_sources<T>() -> Result<T, Error> {
|
||||
Err(Error::MixedAzureCredentialSources)
|
||||
Err(Error::InvalidConfiguration(
|
||||
"request-controlled Azure auth inputs cannot be combined with host credentials".into(),
|
||||
))
|
||||
}
|
||||
|
||||
fn build_credential(
|
||||
|
|
@ -433,7 +443,12 @@ fn build_credential(
|
|||
NativeAzureRequest::DeveloperTools { .. } => DeveloperToolsCredential::new(None)
|
||||
.map(|credential| credential as Arc<dyn TokenCredential>),
|
||||
}
|
||||
.map_err(|error| Error::AzureCredentialInitialization(error.to_string()))
|
||||
.map_err(|error| {
|
||||
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
|
||||
"Azure credential initialization",
|
||||
error,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn client_options(
|
||||
|
|
@ -638,7 +653,7 @@ mod tests {
|
|||
assert_eq!(transport.requests.lock().unwrap().len(), 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn request_authority_requires_request_owned_client_secret_identity() {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
|
|
@ -647,10 +662,13 @@ mod tests {
|
|||
))
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(
|
||||
assert_eq!(
|
||||
error,
|
||||
litellm_auth_types::Error::MixedAzureCredentialSources
|
||||
));
|
||||
litellm_auth_types::Error::InvalidConfiguration(
|
||||
"request-controlled Azure auth inputs cannot be combined with host credentials"
|
||||
.into()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -665,24 +683,24 @@ mod tests {
|
|||
assert_eq!(request.credential_source(), InputSource::Request);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn authority_is_restricted_to_an_https_origin() {
|
||||
for authority in [
|
||||
"http://login.example",
|
||||
"https://user@login.example",
|
||||
"https://login.example/tenant",
|
||||
"https://login.example?target=other",
|
||||
] {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
InputSource::Deployment,
|
||||
authority,
|
||||
))
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
error,
|
||||
litellm_auth_types::Error::InvalidAzureAuthority
|
||||
));
|
||||
}
|
||||
#[rstest::rstest]
|
||||
#[case::http("http://login.example")]
|
||||
#[case::userinfo("https://user@login.example")]
|
||||
#[case::path("https://login.example/tenant")]
|
||||
#[case::query("https://login.example?target=other")]
|
||||
fn authority_is_restricted_to_an_https_origin(#[case] authority: &str) {
|
||||
let error = ValidatedAzureRequest::new(sourced_client_secret(
|
||||
InputSource::Deployment,
|
||||
InputSource::Deployment,
|
||||
authority,
|
||||
))
|
||||
.unwrap_err();
|
||||
assert_eq!(
|
||||
error,
|
||||
litellm_auth_types::Error::InvalidConfiguration(
|
||||
"Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
.into()
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -91,7 +91,9 @@ impl AzureAuthService {
|
|||
AzureCredentialPlan::Caller(caller) => {
|
||||
let credential = caller.acquire().await?;
|
||||
if credential.secret().expose().is_empty() {
|
||||
return Err(Error::EmptyAzureToken);
|
||||
return Err(Error::EmptyCallerCredential(
|
||||
"Azure AD token provider returned an empty token",
|
||||
));
|
||||
}
|
||||
Ok(Some(Sourced::new(credential, InputSource::Deployment)))
|
||||
}
|
||||
|
|
@ -104,7 +106,11 @@ impl AzureAuthService {
|
|||
} => {
|
||||
let assertion = resolve_reference(inputs, env_lookup, reference.value())
|
||||
.await?
|
||||
.ok_or(Error::UnresolvedOidcReference)?;
|
||||
.ok_or_else(|| {
|
||||
Error::CredentialAcquisition(
|
||||
"Azure OIDC reference did not resolve to a value".into(),
|
||||
)
|
||||
})?;
|
||||
let request = ValidatedAzureRequest::new(NativeAzureRequest::ClientAssertion {
|
||||
tenant_id,
|
||||
client_id,
|
||||
|
|
@ -167,7 +173,7 @@ pub(crate) fn select_auth_plan(
|
|||
.map(|selector| Sourced::new(selector, value.source()))
|
||||
})
|
||||
.transpose()
|
||||
.map_err(|_| Error::InvalidAzureSelector)?;
|
||||
.map_err(|_| Error::InvalidConfiguration("invalid Azure credential selector".into()))?;
|
||||
let federated_token_file = configured_string(
|
||||
&inputs.federated_token_file,
|
||||
AZURE_FEDERATED_TOKEN_FILE_ENV,
|
||||
|
|
@ -257,7 +263,9 @@ fn select_native_plan(
|
|||
let selection_source = selected.source();
|
||||
|
||||
match selected.into_value() {
|
||||
AzureCredentialType::ClientSecretCredential => Err(Error::MissingClientSecretFields),
|
||||
AzureCredentialType::ClientSecretCredential => Err(Error::InvalidConfiguration(
|
||||
"ClientSecretCredential requires tenant_id, client_id, and client_secret".into(),
|
||||
)),
|
||||
AzureCredentialType::WorkloadIdentityCredential => {
|
||||
Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new(
|
||||
workload_request(tenant_id, client_id, federated_token_file, scope, authority)?,
|
||||
|
|
@ -341,9 +349,17 @@ fn workload_request(
|
|||
authority: Option<Sourced<String>>,
|
||||
) -> Result<NativeAzureRequest, Error> {
|
||||
Ok(NativeAzureRequest::WorkloadIdentity {
|
||||
tenant_id: tenant_id.ok_or(Error::MissingWorkloadTenant)?,
|
||||
client_id: client_id.ok_or(Error::MissingWorkloadClient)?,
|
||||
token_file_path: token_file_path.ok_or(Error::MissingWorkloadTokenFile)?,
|
||||
tenant_id: tenant_id.ok_or_else(|| {
|
||||
Error::InvalidConfiguration("WorkloadIdentityCredential requires tenant_id".into())
|
||||
})?,
|
||||
client_id: client_id.ok_or_else(|| {
|
||||
Error::InvalidConfiguration("WorkloadIdentityCredential requires client_id".into())
|
||||
})?,
|
||||
token_file_path: token_file_path.ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"WorkloadIdentityCredential requires azure_federated_token_file".into(),
|
||||
)
|
||||
})?,
|
||||
scope,
|
||||
authority,
|
||||
})
|
||||
|
|
@ -394,10 +410,11 @@ async fn resolve_reference(
|
|||
.map_or(CredentialLookup::Missing, CredentialLookup::Found),
|
||||
CredentialRef::None => return Ok(None),
|
||||
CredentialRef::File(_) | CredentialRef::Request(_) | CredentialRef::Host(_) => {
|
||||
let resolver = inputs
|
||||
.credential_resolver
|
||||
.as_ref()
|
||||
.ok_or(Error::MissingHostResolver)?;
|
||||
let resolver = inputs.credential_resolver.as_ref().ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"credential reference requires a host credential resolver".into(),
|
||||
)
|
||||
})?;
|
||||
resolver.resolve(reference).await?
|
||||
}
|
||||
};
|
||||
|
|
@ -415,7 +432,9 @@ fn oidc_reference(
|
|||
};
|
||||
let value = token.value().expose();
|
||||
if token.source() == InputSource::Request && value.starts_with("oidc/") {
|
||||
return Err(Error::RequestAzureCredentialReference);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"request-controlled Azure credential references are not allowed".into(),
|
||||
));
|
||||
}
|
||||
if let Some(name) = value.strip_prefix("oidc/env/") {
|
||||
return non_empty_reference(name, "OIDC environment reference")
|
||||
|
|
@ -437,14 +456,20 @@ fn oidc_reference(
|
|||
)));
|
||||
}
|
||||
if value.starts_with("oidc/") {
|
||||
return Err(Error::UnsupportedOidcReference);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"unsupported OIDC reference".into(),
|
||||
));
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn non_empty_reference(value: &str, kind: &str) -> Result<String, Error> {
|
||||
if value.is_empty() {
|
||||
return Err(Error::EmptyReference(kind.to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::Empty {
|
||||
subject: kind.into(),
|
||||
},
|
||||
));
|
||||
}
|
||||
Ok(value.to_string())
|
||||
}
|
||||
|
|
@ -493,7 +518,7 @@ mod tests {
|
|||
expires_on: None,
|
||||
})
|
||||
} else {
|
||||
Err(Error::AzureTokenAcquisition(format!("{kind} failed")))
|
||||
Err(Error::CredentialAcquisition(kind.into()))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
|
@ -602,7 +627,7 @@ mod tests {
|
|||
assert!(error.to_string().contains("unsupported OIDC reference"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn request_oidc_reference_is_rejected_before_lookup() {
|
||||
let params = json!({
|
||||
"azure_ad_token": "oidc/env/ASSERTION",
|
||||
|
|
@ -624,7 +649,12 @@ mod tests {
|
|||
})
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, Error::RequestAzureCredentialReference));
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::InvalidConfiguration(
|
||||
"request-controlled Azure credential references are not allowed".into()
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -723,6 +753,7 @@ mod tests {
|
|||
assert_eq!(credential.value().secret().expose(), "caller-token");
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn empty_caller_token_is_rejected() {
|
||||
let error = AzureAuthService::default()
|
||||
|
|
@ -730,6 +761,9 @@ mod tests {
|
|||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(error, Error::EmptyAzureToken));
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::EmptyCallerCredential("Azure AD token provider returned an empty token")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -117,7 +117,12 @@ fn string_config(
|
|||
None => Ok(ConfigValue::Absent),
|
||||
Some(Value::Null) => Ok(ConfigValue::ExplicitNone(source)),
|
||||
Some(Value::String(value)) => Ok(ConfigValue::Value(Sourced::new(value.clone(), source))),
|
||||
Some(_) => Err(Error::InvalidFieldType(name.to_string())),
|
||||
Some(_) => Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: name.into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -19,3 +19,6 @@ tokio.workspace = true
|
|||
gcp_auth = "0.12.7"
|
||||
google-cloud-auth = { workspace = true, optional = true }
|
||||
http = { workspace = true, optional = true }
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -299,7 +299,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> {
|
|||
.map(str::to_string)
|
||||
});
|
||||
if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) {
|
||||
return Err(Error::RequestVertexTokenEndpoint);
|
||||
return Err(Error::InvalidConfiguration("request-controlled Vertex credentials must use the canonical Google OAuth token endpoint".into()));
|
||||
}
|
||||
Ok(configured)
|
||||
}
|
||||
|
|
@ -376,10 +376,20 @@ fn optional_credentials(
|
|||
.map(SecretValue::new)
|
||||
.map(|value| Sourced::new(value, source))
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidFieldType(format!("{}: {error}", names[0])));
|
||||
.map_err(|error| {
|
||||
Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed(
|
||||
"credential serialization",
|
||||
error,
|
||||
))
|
||||
});
|
||||
}
|
||||
Some(_) => {
|
||||
return Err(Error::InvalidFieldType(names[0].to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: names[0].into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -397,7 +407,12 @@ fn optional_string(params: &Map<String, Value>, names: &[&str]) -> Result<Option
|
|||
Some(Value::String(value)) if value.trim().is_empty() => continue,
|
||||
Some(Value::String(value)) => return Ok(Some(value.clone())),
|
||||
Some(_) => {
|
||||
return Err(Error::InvalidFieldType(names[0].to_string()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
litellm_auth_types::ErrorDetail::InvalidType {
|
||||
field: names[0].into(),
|
||||
expected: "a string or null",
|
||||
},
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -411,7 +426,10 @@ fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option<String>, name: &str) -> Opt
|
|||
}
|
||||
|
||||
fn auth_acquisition_error(error: gcp_auth::Error) -> Error {
|
||||
Error::VertexTokenAcquisition(error.to_string())
|
||||
Error::CredentialAcquisition(litellm_auth_types::ErrorDetail::failed(
|
||||
"Vertex AI credentials",
|
||||
error,
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -612,20 +630,15 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_credentials_require_canonical_token_endpoint() {
|
||||
assert!(
|
||||
validate_request_credentials(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#)
|
||||
.is_ok()
|
||||
);
|
||||
assert!(matches!(
|
||||
validate_request_credentials(r#"{"token_uri":"http://127.0.0.1/token"}"#),
|
||||
Err(Error::RequestVertexTokenEndpoint)
|
||||
));
|
||||
assert!(matches!(
|
||||
validate_request_credentials("{}"),
|
||||
Err(Error::RequestVertexTokenEndpoint)
|
||||
));
|
||||
#[rstest::rstest]
|
||||
#[case::canonical_endpoint(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#, true)]
|
||||
#[case::noncanonical_endpoint(r#"{"token_uri":"http://127.0.0.1/token"}"#, false)]
|
||||
#[case::missing_endpoint("{}", false)]
|
||||
fn request_credentials_require_canonical_token_endpoint(
|
||||
#[case] credentials: &str,
|
||||
#[case] accepted: bool,
|
||||
) {
|
||||
assert_eq!(validate_request_credentials(credentials).is_ok(), accepted);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
|
|||
|
|
@ -12,4 +12,5 @@ thiserror.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -86,7 +86,9 @@ impl CredentialPlan {
|
|||
Self::Caller(caller) => {
|
||||
let credential = caller.acquire().await?;
|
||||
if credential.secret().expose().is_empty() {
|
||||
return Err(Error::EmptyCallerCredential);
|
||||
return Err(Error::EmptyCallerCredential(
|
||||
"credential caller returned an empty credential",
|
||||
));
|
||||
}
|
||||
Ok(CredentialPlanResolution::Resolved(credential))
|
||||
}
|
||||
|
|
@ -147,10 +149,15 @@ mod tests {
|
|||
|
||||
impl CredentialResolver for FailingResolver {
|
||||
fn resolve<'a>(&'a self, _reference: &'a CredentialRef) -> CredentialLookupFuture<'a> {
|
||||
Box::pin(async { Err(Error::UnresolvedOidcReference) })
|
||||
Box::pin(async {
|
||||
Err(Error::CredentialAcquisition(
|
||||
"host credential lookup failed".into(),
|
||||
))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn acquisition_failure_is_terminal() {
|
||||
let resolver = CredentialResolverHandle::new(Arc::new(FailingResolver));
|
||||
|
|
@ -161,6 +168,9 @@ mod tests {
|
|||
.await
|
||||
.expect_err("acquisition errors cannot become fallback");
|
||||
|
||||
assert_eq!(error, Error::UnresolvedOidcReference);
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::CredentialAcquisition("host credential lookup failed".into())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,84 +2,16 @@ use thiserror::Error as ThisError;
|
|||
|
||||
#[derive(Clone, Debug, ThisError, PartialEq, Eq)]
|
||||
pub enum Error {
|
||||
#[error("invalid authentication configuration: credential header already exists")]
|
||||
ExistingCredentialHeader,
|
||||
#[error(
|
||||
"invalid authentication configuration: credential plan is not allowed by the provider auth policy"
|
||||
)]
|
||||
DisallowedCredentialPlan,
|
||||
#[error("invalid authentication configuration: credential cannot be empty")]
|
||||
EmptyCredential,
|
||||
#[error("invalid authentication configuration: invalid Azure credential selector")]
|
||||
InvalidAzureSelector,
|
||||
#[error(
|
||||
"invalid authentication configuration: ClientSecretCredential requires tenant_id, client_id, and client_secret"
|
||||
)]
|
||||
MissingClientSecretFields,
|
||||
#[error("invalid authentication configuration: WorkloadIdentityCredential requires tenant_id")]
|
||||
MissingWorkloadTenant,
|
||||
#[error("invalid authentication configuration: WorkloadIdentityCredential requires client_id")]
|
||||
MissingWorkloadClient,
|
||||
#[error(
|
||||
"invalid authentication configuration: WorkloadIdentityCredential requires azure_federated_token_file"
|
||||
)]
|
||||
MissingWorkloadTokenFile,
|
||||
#[error(
|
||||
"invalid authentication configuration: credential reference requires a host credential resolver"
|
||||
)]
|
||||
MissingHostResolver,
|
||||
#[error(
|
||||
"invalid authentication configuration: caller credential plan requires provider-specific inputs"
|
||||
)]
|
||||
MissingCallerInputs,
|
||||
#[error("invalid authentication configuration: credential header {0} already exists")]
|
||||
DuplicateHeader(&'static str),
|
||||
#[error("invalid authentication configuration: {0} must be a string or null")]
|
||||
InvalidFieldType(String),
|
||||
#[error("invalid authentication configuration: unsupported OIDC reference")]
|
||||
UnsupportedOidcReference,
|
||||
#[error("invalid authentication configuration: {0} cannot be empty")]
|
||||
EmptyReference(String),
|
||||
#[error("invalid authentication configuration: Azure credential initialization failed: {0}")]
|
||||
AzureCredentialInitialization(String),
|
||||
#[error(
|
||||
"invalid authentication configuration: Azure authority must be an HTTPS origin without credentials, query, or fragment"
|
||||
)]
|
||||
InvalidAzureAuthority,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Azure auth inputs cannot be combined with host credentials"
|
||||
)]
|
||||
MixedAzureCredentialSources,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Azure credential references are not allowed"
|
||||
)]
|
||||
RequestAzureCredentialReference,
|
||||
#[error(
|
||||
"invalid authentication configuration: host credentials cannot be sent to a request-controlled Azure endpoint"
|
||||
)]
|
||||
RequestAzureCredentialDestination,
|
||||
#[error(
|
||||
"invalid authentication configuration: credentials cannot be sent to a request-controlled Vertex AI endpoint"
|
||||
)]
|
||||
RequestVertexCredentialDestination,
|
||||
#[error(
|
||||
"invalid authentication configuration: request-controlled Vertex credentials must use the canonical Google OAuth token endpoint"
|
||||
)]
|
||||
RequestVertexTokenEndpoint,
|
||||
#[error("invalid authentication configuration: {0}")]
|
||||
InvalidConfiguration(#[source] ErrorDetail),
|
||||
#[error("credential acquisition failed: {0}")]
|
||||
AzureTokenAcquisition(String),
|
||||
#[error("credential acquisition failed: Vertex AI credentials: {0}")]
|
||||
VertexTokenAcquisition(String),
|
||||
CredentialAcquisition(#[source] ErrorDetail),
|
||||
#[error("credential caller failed: {0}")]
|
||||
EmptyCallerCredential(&'static str),
|
||||
#[error("{0}")]
|
||||
ProviderAuthentication(String),
|
||||
#[error("credential acquisition failed: {}", .0.iter().map(ToString::to_string).collect::<Vec<_>>().join("; "))]
|
||||
CredentialChain(Vec<Error>),
|
||||
#[error("credential caller failed: credential caller returned an empty credential")]
|
||||
EmptyCallerCredential,
|
||||
#[error("credential caller failed: Azure AD token provider returned an empty token")]
|
||||
EmptyAzureToken,
|
||||
#[error("credential acquisition failed: Azure OIDC reference did not resolve to a value")]
|
||||
UnresolvedOidcReference,
|
||||
#[error(
|
||||
"Missing {provider} API Key - Set `api_key` or the {environment_variable} environment variable"
|
||||
)]
|
||||
|
|
@ -87,34 +19,87 @@ pub enum Error {
|
|||
provider: &'static str,
|
||||
environment_variable: &'static str,
|
||||
},
|
||||
#[error(
|
||||
"Missing {provider} API Base - Set {environment_variable} environment variable or pass api_base parameter"
|
||||
)]
|
||||
#[error("Missing {provider} API Base - {guidance}")]
|
||||
MissingApiBase {
|
||||
provider: &'static str,
|
||||
environment_variable: &'static str,
|
||||
guidance: &'static str,
|
||||
},
|
||||
#[error(
|
||||
"Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
|
||||
)]
|
||||
MissingAzureApiBase,
|
||||
#[error("invalid authentication header")]
|
||||
InvalidHeader,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::Error;
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum ErrorDetail {
|
||||
#[error("{0}")]
|
||||
Message(String),
|
||||
#[error("{field} must be {expected}")]
|
||||
InvalidType {
|
||||
field: String,
|
||||
expected: &'static str,
|
||||
},
|
||||
#[error("{subject} cannot be empty")]
|
||||
Empty { subject: String },
|
||||
#[error("credential header {0} already exists")]
|
||||
DuplicateHeader(&'static str),
|
||||
#[error("{operation} failed: {source}")]
|
||||
Failed {
|
||||
operation: &'static str,
|
||||
#[source]
|
||||
source: ErrorSource,
|
||||
},
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_api_key_names_provider_and_environment_variable() {
|
||||
assert_eq!(
|
||||
Error::MissingApiKey {
|
||||
provider: "Anthropic",
|
||||
environment_variable: "ANTHROPIC_API_KEY",
|
||||
}
|
||||
.to_string(),
|
||||
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
|
||||
);
|
||||
impl ErrorDetail {
|
||||
pub fn failed(
|
||||
operation: &'static str,
|
||||
source: impl std::error::Error + Send + Sync + 'static,
|
||||
) -> Self {
|
||||
Self::Failed {
|
||||
operation,
|
||||
source: ErrorSource::new(source),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for ErrorDetail {
|
||||
fn from(message: String) -> Self {
|
||||
Self::Message(message)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for ErrorDetail {
|
||||
fn from(message: &str) -> Self {
|
||||
Self::Message(message.into())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ErrorSource(std::sync::Arc<dyn std::error::Error + Send + Sync>);
|
||||
|
||||
impl std::ops::Deref for ErrorSource {
|
||||
type Target = dyn std::error::Error + Send + Sync;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
self.0.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ErrorSource {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
std::fmt::Display::fmt(&self.0, formatter)
|
||||
}
|
||||
}
|
||||
|
||||
impl ErrorSource {
|
||||
pub fn new(error: impl std::error::Error + Send + Sync + 'static) -> Self {
|
||||
Self(std::sync::Arc::new(error))
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for ErrorSource {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
std::sync::Arc::ptr_eq(&self.0, &other.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl Eq for ErrorSource {}
|
||||
|
|
|
|||
|
|
@ -21,13 +21,17 @@ pub fn apply_credential(
|
|||
placement: CredentialPlacement,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
if credential.trim().is_empty() {
|
||||
return Err(Error::EmptyCredential);
|
||||
return Err(Error::InvalidConfiguration(
|
||||
"credential cannot be empty".into(),
|
||||
));
|
||||
}
|
||||
if headers
|
||||
.iter()
|
||||
.any(|(name, _)| name.eq_ignore_ascii_case(placement.header_name()))
|
||||
{
|
||||
return Err(Error::DuplicateHeader(placement.header_name()));
|
||||
return Err(Error::InvalidConfiguration(
|
||||
crate::ErrorDetail::DuplicateHeader(placement.header_name()),
|
||||
));
|
||||
}
|
||||
let value = match placement {
|
||||
CredentialPlacement::Bearer => format!("Bearer {credential}"),
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ pub use credential::{
|
|||
CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan,
|
||||
CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle,
|
||||
};
|
||||
pub use error::Error;
|
||||
pub use error::{Error, ErrorDetail, ErrorSource};
|
||||
pub use http::CredentialPlacement;
|
||||
pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy};
|
||||
pub use secret::SecretValue;
|
||||
|
|
|
|||
|
|
@ -47,14 +47,20 @@ impl ProviderAuthPolicy {
|
|||
if self.has_existing_credential(&headers) {
|
||||
return match self.existing_header_behavior {
|
||||
ExistingHeaderBehavior::Preserve => Ok(headers),
|
||||
ExistingHeaderBehavior::Reject => Err(Error::ExistingCredentialHeader),
|
||||
ExistingHeaderBehavior::Reject => Err(Error::InvalidConfiguration(
|
||||
"credential header already exists".into(),
|
||||
)),
|
||||
};
|
||||
}
|
||||
let rule = self
|
||||
.rules
|
||||
.iter()
|
||||
.find(|rule| rule.kind == kind)
|
||||
.ok_or(Error::DisallowedCredentialPlan)?;
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidConfiguration(
|
||||
"credential plan is not allowed by the provider auth policy".into(),
|
||||
)
|
||||
})?;
|
||||
apply_credential(headers, credential.secret().expose(), rule.placement)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
74
litellm-rust/crates/auth-types/tests/error.rs
Normal file
74
litellm-rust/crates/auth-types/tests/error.rs
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
use litellm_auth_types::Error;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key(
|
||||
Error::MissingApiKey { provider: "Example", environment_variable: "EXAMPLE_API_KEY" },
|
||||
"Missing Example API Key - Set `api_key` or the EXAMPLE_API_KEY environment variable"
|
||||
)]
|
||||
#[case::another_api_key(
|
||||
Error::MissingApiKey { provider: "Custom", environment_variable: "CUSTOM_KEY" },
|
||||
"Missing Custom API Key - Set `api_key` or the CUSTOM_KEY environment variable"
|
||||
)]
|
||||
#[case::api_base(
|
||||
Error::MissingApiBase { provider: "Example", guidance: "Pass api_base" },
|
||||
"Missing Example API Base - Pass api_base"
|
||||
)]
|
||||
#[case::another_api_base(
|
||||
Error::MissingApiBase { provider: "Custom", guidance: "Set CUSTOM_ENDPOINT" },
|
||||
"Missing Custom API Base - Set CUSTOM_ENDPOINT"
|
||||
)]
|
||||
#[case::configuration(
|
||||
Error::InvalidConfiguration("credential selector is invalid".into()),
|
||||
"invalid authentication configuration: credential selector is invalid"
|
||||
)]
|
||||
#[case::acquisition(
|
||||
Error::CredentialAcquisition("token expired".into()),
|
||||
"credential acquisition failed: token expired"
|
||||
)]
|
||||
#[case::caller(
|
||||
Error::EmptyCallerCredential("empty token"),
|
||||
"credential caller failed: empty token"
|
||||
)]
|
||||
#[case::provider(
|
||||
Error::ProviderAuthentication("provider rejected credentials".into()),
|
||||
"provider rejected credentials"
|
||||
)]
|
||||
#[case::chain(
|
||||
Error::CredentialChain(vec![
|
||||
Error::CredentialAcquisition("token expired".into()),
|
||||
Error::EmptyCallerCredential("empty token"),
|
||||
]),
|
||||
"credential acquisition failed: credential acquisition failed: token expired; credential caller failed: empty token"
|
||||
)]
|
||||
fn display_preserves_failure_phase_and_caller_context(
|
||||
#[case] error: Error,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
assert_eq!(error.to_string(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::configuration(true)]
|
||||
#[case::acquisition(false)]
|
||||
fn contextual_failures_keep_the_original_source(#[case] configuration: bool) {
|
||||
use litellm_auth_types::ErrorDetail;
|
||||
|
||||
let detail = ErrorDetail::failed(
|
||||
"test credential",
|
||||
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
|
||||
);
|
||||
let error = if configuration {
|
||||
Error::InvalidConfiguration(detail)
|
||||
} else {
|
||||
Error::CredentialAcquisition(detail)
|
||||
};
|
||||
let source = std::iter::successors(Some(&error as &dyn std::error::Error), |error| {
|
||||
error.source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<std::io::Error>())
|
||||
.expect("the original credential error remains available");
|
||||
assert_eq!(source.kind(), std::io::ErrorKind::PermissionDenied);
|
||||
assert!(error.to_string().contains("test credential failed:"));
|
||||
assert!(error.to_string().ends_with(&source.to_string()));
|
||||
}
|
||||
|
|
@ -9,7 +9,7 @@ repository.workspace = true
|
|||
litellm-cache.workspace = true
|
||||
py_literal = "0.4.0"
|
||||
rand.workspace = true
|
||||
rusqlite = { version = "0.40", features = ["bundled"] }
|
||||
rusqlite = { version = "0.39", features = ["bundled"] }
|
||||
serde-pickle = "1.2"
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
- This crate is the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on)
|
||||
- Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here
|
||||
- SDK request policy (credential inheritance, the budget and retry-count limits) is the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call
|
||||
- The driver in `litellm-host-python`, the routes and core see one `PythonCallHooks`; they never learn which Python objects consume a call
|
||||
- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json`
|
||||
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
|
||||
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`
|
||||
|
|
|
|||
|
|
@ -2,12 +2,12 @@
|
|||
//! raises is answered with the same `Logging` calls, in the same order, as the Python
|
||||
//! `@client` path makes them.
|
||||
|
||||
use litellm_host_python::PythonOwned;
|
||||
|
||||
use litellm_host::event::{
|
||||
FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest, epoch_seconds,
|
||||
};
|
||||
use litellm_host_python::{
|
||||
LifecycleEvent, LifecycleStep, PythonLifecycle, from_py, missing_state, to_py,
|
||||
};
|
||||
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, from_py, missing_state, to_py};
|
||||
use pyo3::{
|
||||
exceptions::{PyBaseException, PyException},
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
|
|
@ -48,11 +48,10 @@ struct DeliveredStream {
|
|||
first_chunk: Option<Py<PyAny>>,
|
||||
}
|
||||
|
||||
enum Pending {
|
||||
DeploymentPreCall,
|
||||
DeploymentPostCall,
|
||||
DeploymentFailure,
|
||||
AsyncFailure,
|
||||
struct LoggedRequest {
|
||||
body: Py<PyDict>,
|
||||
headers: Py<PyDict>,
|
||||
context: RequestContext,
|
||||
}
|
||||
|
||||
pub struct LegacyLogging {
|
||||
|
|
@ -63,13 +62,10 @@ pub struct LegacyLogging {
|
|||
end: Option<Py<PyAny>>,
|
||||
response: Option<Py<PyAny>>,
|
||||
error: Option<Py<PyBaseException>>,
|
||||
body: Option<Py<PyDict>>,
|
||||
headers: Option<Py<PyDict>>,
|
||||
context: Option<RequestContext>,
|
||||
request: Option<LoggedRequest>,
|
||||
stream: Option<DeliveredStream>,
|
||||
asynchronous: bool,
|
||||
internal: bool,
|
||||
pending: Option<Pending>,
|
||||
}
|
||||
|
||||
fn datetime(py: Python<'_>, epoch_seconds: f64) -> PyResult<Py<PyAny>> {
|
||||
|
|
@ -95,13 +91,10 @@ impl LegacyLogging {
|
|||
end: None,
|
||||
response: None,
|
||||
error: None,
|
||||
body: None,
|
||||
headers: None,
|
||||
context: None,
|
||||
request: None,
|
||||
stream: None,
|
||||
asynchronous,
|
||||
internal: false,
|
||||
pending: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -120,14 +113,14 @@ impl LegacyLogging {
|
|||
/// The keyword view the rest of the call reads: a copy, so the deployment hook's own
|
||||
/// dict is left as the hook returned it, carrying the logger as `@client` injects it.
|
||||
/// The driver's preflight rewrites this same dict before the host projects from it.
|
||||
fn prepare(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
|
||||
fn prepare(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, Py<PyDict>>> {
|
||||
let prepared = self.call.kwargs().bind(py).copy()?;
|
||||
prepared.set_item("litellm_logging_obj", self.logger()?.object(py))?;
|
||||
self.call.set_kwargs(prepared.unbind());
|
||||
Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py)))
|
||||
Ok(HookStep::Ready(self.call.kwargs().clone_ref(py)))
|
||||
}
|
||||
|
||||
fn finalize(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
|
||||
fn finalize(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, Py<PyAny>>> {
|
||||
finalize(
|
||||
py,
|
||||
&self.response,
|
||||
|
|
@ -138,7 +131,7 @@ impl LegacyLogging {
|
|||
)?;
|
||||
self.response
|
||||
.as_ref()
|
||||
.map(|response| LifecycleStep::Response(response.clone_ref(py)))
|
||||
.map(|response| HookStep::Ready(response.clone_ref(py)))
|
||||
.ok_or_else(missing_state)
|
||||
}
|
||||
|
||||
|
|
@ -195,7 +188,10 @@ impl LegacyLogging {
|
|||
logger.object(py),
|
||||
billing.url_route,
|
||||
billing.endpoint_type,
|
||||
&self.body,
|
||||
&self
|
||||
.request
|
||||
.as_ref()
|
||||
.map(|request| request.body.clone_ref(py)),
|
||||
&stream.chunks,
|
||||
&self.start,
|
||||
&self.end,
|
||||
|
|
@ -214,11 +210,11 @@ impl LegacyLogging {
|
|||
/// A failure after the stream reached the caller bills the delivered chunks as
|
||||
/// partial usage. The sync path has no loop to schedule that on, so it falls back to
|
||||
/// the plain failure handler.
|
||||
fn stream_failure(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
|
||||
fn stream_failure(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, ()>> {
|
||||
let (Some(logger), Some(error), Some(stream), Some(billing)) =
|
||||
(&self.logger, &self.error, &self.stream, self.surface.stream)
|
||||
else {
|
||||
return Ok(LifecycleStep::Done);
|
||||
return Ok(HookStep::Ready(()));
|
||||
};
|
||||
if !self.asynchronous {
|
||||
return self.dispatch_failure(py);
|
||||
|
|
@ -228,30 +224,33 @@ impl LegacyLogging {
|
|||
(
|
||||
logger.object(py),
|
||||
billing.endpoint_type,
|
||||
&self.body,
|
||||
&self
|
||||
.request
|
||||
.as_ref()
|
||||
.map(|request| request.body.clone_ref(py)),
|
||||
&stream.chunks,
|
||||
error,
|
||||
),
|
||||
);
|
||||
match scheduled {
|
||||
Ok(awaitable) => {
|
||||
self.pending = Some(Pending::AsyncFailure);
|
||||
Ok(LifecycleStep::Await(awaitable.unbind()))
|
||||
}
|
||||
Ok(awaitable) => Ok(HookStep::Await(
|
||||
awaitable.unbind(),
|
||||
Self::resume_async_failure,
|
||||
)),
|
||||
Err(failure) if is_cancellation(py, &failure) => Err(failure),
|
||||
Err(_) => Ok(LifecycleStep::Done),
|
||||
Err(_) => Ok(HookStep::Ready(())),
|
||||
}
|
||||
}
|
||||
|
||||
/// The sync failure handler, then the async one for async calls. Ordinary handler
|
||||
/// errors never replace the selected failure or suppress the other family; a
|
||||
/// cancellation does end the call.
|
||||
fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
|
||||
fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult<HookStep<Self, ()>> {
|
||||
let (Some(logger), Some(error)) = (&self.logger, &self.error) else {
|
||||
return Ok(LifecycleStep::Done);
|
||||
return Ok(HookStep::Ready(()));
|
||||
};
|
||||
if self.asynchronous && self.internal {
|
||||
return Ok(LifecycleStep::Done);
|
||||
return Ok(HookStep::Ready(()));
|
||||
}
|
||||
if let Err(failure) = logger.failure(py, error, &self.start, &self.end, false)
|
||||
&& is_cancellation(py, &failure)
|
||||
|
|
@ -259,27 +258,64 @@ impl LegacyLogging {
|
|||
return Err(failure);
|
||||
}
|
||||
if !self.asynchronous {
|
||||
return Ok(LifecycleStep::Done);
|
||||
return Ok(HookStep::Ready(()));
|
||||
}
|
||||
match logger.failure(py, error, &self.start, &self.end, true) {
|
||||
Ok(Some(awaitable)) => {
|
||||
self.pending = Some(Pending::AsyncFailure);
|
||||
Ok(LifecycleStep::Await(awaitable))
|
||||
}
|
||||
Ok(None) => Ok(LifecycleStep::Done),
|
||||
Ok(Some(awaitable)) => Ok(HookStep::Await(awaitable, Self::resume_async_failure)),
|
||||
Ok(None) => Ok(HookStep::Ready(())),
|
||||
Err(failure) if is_cancellation(py, &failure) => Err(failure),
|
||||
Err(_) => Ok(LifecycleStep::Done),
|
||||
Err(_) => Ok(HookStep::Ready(())),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PythonLifecycle for LegacyLogging {
|
||||
fn begin(
|
||||
impl LegacyLogging {
|
||||
fn resume_begin(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: PyResult<Py<PyAny>>,
|
||||
) -> PyResult<HookStep<Self, Py<PyDict>>> {
|
||||
self.call
|
||||
.set_kwargs(result?.into_bound(py).cast_into::<PyDict>()?.unbind());
|
||||
self.prepare(py)
|
||||
}
|
||||
|
||||
fn resume_after_success(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: PyResult<Py<PyAny>>,
|
||||
) -> PyResult<HookStep<Self, Py<PyAny>>> {
|
||||
self.response = Some(result?);
|
||||
self.finalize(py)
|
||||
}
|
||||
|
||||
fn resume_deployment_failure(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
_: PyResult<Py<PyAny>>,
|
||||
) -> PyResult<HookStep<Self, ()>> {
|
||||
self.dispatch_failure(py)
|
||||
}
|
||||
|
||||
fn resume_async_failure(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: PyResult<Py<PyAny>>,
|
||||
) -> PyResult<HookStep<Self, ()>> {
|
||||
match result {
|
||||
Err(error) if is_cancellation(py, &error) => Err(error),
|
||||
_ => Ok(HookStep::Ready(())),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PythonCallHooks for LegacyLogging {
|
||||
fn prepare_arguments(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
arguments: Py<PyDict>,
|
||||
started_at: f64,
|
||||
) -> PyResult<LifecycleStep> {
|
||||
) -> PyResult<HookStep<Self, Py<PyDict>>> {
|
||||
self.call.set_kwargs(arguments);
|
||||
self.start = datetime(py, started_at)?;
|
||||
self.internal = is_internal_call(py)?;
|
||||
|
|
@ -294,22 +330,20 @@ impl PythonLifecycle for LegacyLogging {
|
|||
self.logger = Some(result.logger()?);
|
||||
self.call.set_kwargs(result.kwargs()?);
|
||||
if self.runs_deployment_hooks() {
|
||||
self.pending = Some(Pending::DeploymentPreCall);
|
||||
return Ok(LifecycleStep::Await(DeploymentHooks::before_call(
|
||||
py,
|
||||
self.call.kwargs(),
|
||||
self.surface.call_type,
|
||||
)?));
|
||||
return Ok(HookStep::Await(
|
||||
DeploymentHooks::before_call(py, self.call.kwargs(), self.surface.call_type)?,
|
||||
Self::resume_begin,
|
||||
));
|
||||
}
|
||||
self.prepare(py)
|
||||
}
|
||||
|
||||
fn before_send(
|
||||
fn before_provider_request(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
wire: Box<WireRequest>,
|
||||
context: &RequestContext,
|
||||
) -> PyResult<LifecycleStep> {
|
||||
) -> PyResult<HookStep<Self, Box<WireRequest>>> {
|
||||
let logger = self.logger()?;
|
||||
logger.update_from_kwargs(py, self.call.kwargs(), &wire, context)?;
|
||||
let body = to_py(py, &wire.body)?
|
||||
|
|
@ -326,9 +360,11 @@ impl PythonLifecycle for LegacyLogging {
|
|||
for (name, value) in &wire.headers {
|
||||
headers.set_item(name, value)?;
|
||||
}
|
||||
self.body = Some(body.clone().unbind());
|
||||
self.headers = Some(headers.clone().unbind());
|
||||
self.context = Some(context.clone());
|
||||
self.request = Some(LoggedRequest {
|
||||
body: body.clone().unbind(),
|
||||
headers: headers.clone().unbind(),
|
||||
context: context.clone(),
|
||||
});
|
||||
self.logger()?.pre_call(
|
||||
py,
|
||||
self.surface.input_description,
|
||||
|
|
@ -341,61 +377,63 @@ impl PythonLifecycle for LegacyLogging {
|
|||
.iter()
|
||||
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
Ok(LifecycleStep::Wire(Box::new(WireRequest {
|
||||
Ok(HookStep::Ready(Box::new(WireRequest {
|
||||
body: from_py(&body)?,
|
||||
headers,
|
||||
..*wire
|
||||
})))
|
||||
}
|
||||
|
||||
fn after_success(
|
||||
fn transform_response(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
response: Py<PyAny>,
|
||||
timing: Timing,
|
||||
) -> PyResult<LifecycleStep> {
|
||||
) -> PyResult<HookStep<Self, Py<PyAny>>> {
|
||||
self.end = Some(datetime(py, timing.end_time)?);
|
||||
self.response = Some(response);
|
||||
if self.runs_deployment_hooks() {
|
||||
self.pending = Some(Pending::DeploymentPostCall);
|
||||
return Ok(LifecycleStep::Await(DeploymentHooks::after_success(
|
||||
py,
|
||||
self.call.kwargs(),
|
||||
&self.response,
|
||||
self.surface.call_type,
|
||||
)?));
|
||||
return Ok(HookStep::Await(
|
||||
DeploymentHooks::after_success(
|
||||
py,
|
||||
self.call.kwargs(),
|
||||
&self.response,
|
||||
self.surface.call_type,
|
||||
)?,
|
||||
Self::resume_after_success,
|
||||
));
|
||||
}
|
||||
self.finalize(py)
|
||||
}
|
||||
|
||||
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep> {
|
||||
fn on_event(&mut self, py: Python<'_>, event: HookEvent<'_>) -> PyResult<HookStep<Self, ()>> {
|
||||
match event {
|
||||
LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done),
|
||||
LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
HookEvent::Started { .. } => Ok(HookStep::Ready(())),
|
||||
HookEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
let api_key = self
|
||||
.context
|
||||
.request
|
||||
.as_ref()
|
||||
.and_then(|context| context.api_key.as_ref())
|
||||
.and_then(|request| request.context.api_key.as_ref())
|
||||
.map(|api_key| api_key.expose());
|
||||
self.logger()?.post_call(
|
||||
py,
|
||||
&raw.body,
|
||||
api_key,
|
||||
self.body.as_ref(),
|
||||
self.headers.as_ref(),
|
||||
self.request.as_ref().map(|request| &request.body),
|
||||
self.request.as_ref().map(|request| &request.headers),
|
||||
)?;
|
||||
Ok(LifecycleStep::Done)
|
||||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
LifecycleEvent::Succeeded { timing, response } => {
|
||||
HookEvent::Succeeded { timing, response } => {
|
||||
self.end = Some(datetime(py, timing.end_time)?);
|
||||
self.response = Some(response.clone_ref(py));
|
||||
match &self.stream {
|
||||
Some(stream) => self.stream_success(py, stream)?,
|
||||
None => self.dispatch_success(py)?,
|
||||
}
|
||||
Ok(LifecycleStep::Done)
|
||||
Ok(HookStep::Ready(()))
|
||||
}
|
||||
LifecycleEvent::Failed {
|
||||
HookEvent::Failed {
|
||||
timing,
|
||||
origin,
|
||||
error,
|
||||
|
|
@ -410,20 +448,22 @@ impl PythonLifecycle for LegacyLogging {
|
|||
&& self.runs_deployment_hooks()
|
||||
{
|
||||
let error = self.error.as_ref().ok_or_else(missing_state)?;
|
||||
self.pending = Some(Pending::DeploymentFailure);
|
||||
return Ok(LifecycleStep::Await(DeploymentHooks::after_failure(
|
||||
py,
|
||||
self.call.kwargs(),
|
||||
error,
|
||||
self.surface.call_type,
|
||||
)?));
|
||||
return Ok(HookStep::Await(
|
||||
DeploymentHooks::after_failure(
|
||||
py,
|
||||
self.call.kwargs(),
|
||||
error,
|
||||
self.surface.call_type,
|
||||
)?,
|
||||
Self::resume_deployment_failure,
|
||||
));
|
||||
}
|
||||
self.dispatch_failure(py)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn opened(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
fn on_stream_open(&mut self, py: Python<'_>) -> PyResult<()> {
|
||||
if self.surface.stream.is_none() {
|
||||
return Err(missing_state());
|
||||
}
|
||||
|
|
@ -435,45 +475,25 @@ impl PythonLifecycle for LegacyLogging {
|
|||
Ok(())
|
||||
}
|
||||
|
||||
fn delivered(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
fn on_stream_chunk(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
|
||||
let stream = self.stream.as_mut().ok_or_else(missing_state)?;
|
||||
if stream.first_chunk.is_none() {
|
||||
stream.first_chunk = Some(datetime(py, epoch_seconds())?);
|
||||
}
|
||||
stream.chunks.bind(py).append(chunk)
|
||||
}
|
||||
}
|
||||
|
||||
fn resume(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<LifecycleStep> {
|
||||
match self.pending.take().ok_or_else(missing_state)? {
|
||||
Pending::DeploymentPreCall => {
|
||||
self.call
|
||||
.set_kwargs(result?.into_bound(py).cast_into::<PyDict>()?.unbind());
|
||||
self.prepare(py)
|
||||
}
|
||||
Pending::DeploymentPostCall => {
|
||||
self.response = Some(result?);
|
||||
self.finalize(py)
|
||||
}
|
||||
Pending::DeploymentFailure => self.dispatch_failure(py),
|
||||
Pending::AsyncFailure => match result {
|
||||
Err(failure) if is_cancellation(py, &failure) => Err(failure),
|
||||
_ => Ok(LifecycleStep::Done),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
impl PythonOwned for LegacyLogging {
|
||||
fn close(&mut self, py: Python<'_>) {
|
||||
if let Some(logger) = self.logger.take()
|
||||
&& let Err(error) = logger.restore_context(py)
|
||||
{
|
||||
error.write_unraisable(py, None);
|
||||
}
|
||||
self.body = None;
|
||||
self.headers = None;
|
||||
self.context = None;
|
||||
self.request = None;
|
||||
self.stream = None;
|
||||
}
|
||||
|
||||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
self.call.traverse(visit)?;
|
||||
if let Some(logger) = &self.logger {
|
||||
|
|
@ -487,8 +507,11 @@ impl PythonLifecycle for LegacyLogging {
|
|||
visit.call(&stream.chunks)?;
|
||||
visit.call(&stream.first_chunk)?;
|
||||
}
|
||||
visit.call(&self.body)?;
|
||||
visit.call(&self.headers)
|
||||
if let Some(request) = &self.request {
|
||||
visit.call(&request.body)?;
|
||||
visit.call(&request.headers)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -497,7 +520,7 @@ mod deployment_hooks_tests {
|
|||
use std::ffi::CStr;
|
||||
|
||||
use litellm_host::event::{FailureOrigin, Timing};
|
||||
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
|
||||
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks};
|
||||
use pyo3::exceptions::asyncio::CancelledError;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
|
|
@ -506,6 +529,18 @@ mod deployment_hooks_tests {
|
|||
use super::LegacyLogging;
|
||||
use crate::test_support::{legacy_call, local, namespace, run};
|
||||
|
||||
fn resume<T>(
|
||||
logging: &mut LegacyLogging,
|
||||
step: HookStep<LegacyLogging, T>,
|
||||
py: Python<'_>,
|
||||
value: PyResult<Py<PyAny>>,
|
||||
) -> PyResult<HookStep<LegacyLogging, T>> {
|
||||
let HookStep::Await(_, continuation) = step else {
|
||||
panic!("expected suspension")
|
||||
};
|
||||
continuation(logging, py, value)
|
||||
}
|
||||
|
||||
const CALL: &CStr = c"
|
||||
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
|
||||
kwargs = {'logger': logger, 'document': document}
|
||||
|
|
@ -520,25 +555,28 @@ kwargs = {'logger': logger, 'document': document}
|
|||
py: Python<'py>,
|
||||
locals: &Bound<'py, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> (LegacyLogging, LifecycleStep) {
|
||||
) -> (LegacyLogging, HookStep<LegacyLogging, Py<PyDict>>) {
|
||||
let mut logging = legacy_call(py, locals, asynchronous);
|
||||
let kwargs = local(locals, "kwargs")
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap()
|
||||
.unbind();
|
||||
let step = logging.begin(py, kwargs, 0.0).unwrap();
|
||||
let step = logging.prepare_arguments(py, kwargs, 0.0).unwrap();
|
||||
(logging, step)
|
||||
}
|
||||
|
||||
fn arguments<'py>(py: Python<'py>, step: LifecycleStep) -> Bound<'py, PyDict> {
|
||||
let LifecycleStep::Arguments(arguments) = step else {
|
||||
fn arguments<'py>(
|
||||
py: Python<'py>,
|
||||
step: HookStep<LegacyLogging, Py<PyDict>>,
|
||||
) -> Bound<'py, PyDict> {
|
||||
let HookStep::Ready(arguments) = step else {
|
||||
panic!("expected the prepared arguments");
|
||||
};
|
||||
arguments.into_bound(py)
|
||||
}
|
||||
|
||||
fn awaits_deployment_hook(step: &LifecycleStep) -> bool {
|
||||
matches!(step, LifecycleStep::Await(_))
|
||||
fn awaits_deployment_hook<T>(step: &HookStep<LegacyLogging, T>) -> bool {
|
||||
matches!(step, HookStep::Await(_, _))
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -559,7 +597,7 @@ kwargs = {'logger': logger, 'document': document}
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn kwargs_returned_by_the_pre_call_hook_are_what_the_call_prepares() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -574,9 +612,13 @@ replaced_kwargs = {'logger': logger, 'document': replacement, 'pages': [0]}
|
|||
);
|
||||
let (mut logging, step) = begin(py, &locals, true);
|
||||
assert!(awaits_deployment_hook(&step));
|
||||
let step = logging
|
||||
.resume(py, Ok(local(&locals, "replaced_kwargs").unbind()))
|
||||
.unwrap();
|
||||
let step = resume(
|
||||
&mut logging,
|
||||
step,
|
||||
py,
|
||||
Ok(local(&locals, "replaced_kwargs").unbind()),
|
||||
)
|
||||
.unwrap();
|
||||
locals.set_item("prepared", arguments(py, step)).unwrap();
|
||||
run(
|
||||
py,
|
||||
|
|
@ -610,7 +652,9 @@ kwargs = {'logger': logger, 'vendor_extension': opaque}
|
|||
);
|
||||
let (mut logging, step) = begin(py, &locals, asynchronous);
|
||||
let step = match step {
|
||||
LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(),
|
||||
HookStep::Await(hook_result, resume) => {
|
||||
resume(&mut logging, py, Ok(hook_result)).unwrap()
|
||||
}
|
||||
step => step,
|
||||
};
|
||||
locals.set_item("prepared", arguments(py, step)).unwrap();
|
||||
|
|
@ -626,7 +670,7 @@ assert hooked == ([opaque] if asynchronous else []), hooked
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn response_returned_by_the_post_call_hook_is_finalized_and_returned() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -639,18 +683,26 @@ replacement = object()
|
|||
logger.hooks = {'pre': lambda kwargs: kwargs}
|
||||
",
|
||||
);
|
||||
let (mut logging, _) = begin(py, &locals, true);
|
||||
logging
|
||||
.resume(py, Ok(local(&locals, "kwargs").unbind()))
|
||||
.unwrap();
|
||||
let (mut logging, step) = begin(py, &locals, true);
|
||||
resume(
|
||||
&mut logging,
|
||||
step,
|
||||
py,
|
||||
Ok(local(&locals, "kwargs").unbind()),
|
||||
)
|
||||
.unwrap();
|
||||
let step = logging
|
||||
.after_success(py, local(&locals, "response").unbind(), TIMING)
|
||||
.transform_response(py, local(&locals, "response").unbind(), TIMING)
|
||||
.unwrap();
|
||||
assert!(awaits_deployment_hook(&step));
|
||||
let step = logging
|
||||
.resume(py, Ok(local(&locals, "replacement").unbind()))
|
||||
.unwrap();
|
||||
let LifecycleStep::Response(returned) = step else {
|
||||
let step = resume(
|
||||
&mut logging,
|
||||
step,
|
||||
py,
|
||||
Ok(local(&locals, "replacement").unbind()),
|
||||
)
|
||||
.unwrap();
|
||||
let HookStep::Ready(returned) = step else {
|
||||
panic!("expected the finalized response");
|
||||
};
|
||||
assert!(returned.bind(py).is(local(&locals, "replacement")));
|
||||
|
|
@ -672,18 +724,28 @@ assert finalized is replacement
|
|||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = namespace(py, c"kwargs = {'logger': logger}\nresponse = object()");
|
||||
let (mut logging, _) = begin(py, &locals, true);
|
||||
if post_call {
|
||||
logging
|
||||
.resume(py, Ok(local(&locals, "kwargs").unbind()))
|
||||
.unwrap();
|
||||
logging
|
||||
.after_success(py, local(&locals, "response").unbind(), TIMING)
|
||||
.unwrap();
|
||||
}
|
||||
let (mut logging, step) = begin(py, &locals, true);
|
||||
let cancellation = CancelledError::new_err("cancelled");
|
||||
let cancelled = cancellation.value(py).clone();
|
||||
let error = logging.resume(py, Err(cancellation)).err().unwrap();
|
||||
let error = if post_call {
|
||||
resume(
|
||||
&mut logging,
|
||||
step,
|
||||
py,
|
||||
Ok(local(&locals, "kwargs").unbind()),
|
||||
)
|
||||
.unwrap();
|
||||
let step = logging
|
||||
.transform_response(py, local(&locals, "response").unbind(), TIMING)
|
||||
.unwrap();
|
||||
resume(&mut logging, step, py, Err(cancellation))
|
||||
.err()
|
||||
.unwrap()
|
||||
} else {
|
||||
resume(&mut logging, step, py, Err(cancellation))
|
||||
.err()
|
||||
.unwrap()
|
||||
};
|
||||
assert!(error.value(py).is(&cancelled));
|
||||
let names: Vec<String> = local(&locals, "logger")
|
||||
.call_method0("names")
|
||||
|
|
@ -704,17 +766,21 @@ assert finalized is replacement
|
|||
py,
|
||||
c"kwargs = {'logger': logger}\nfailure = ValueError('provider')",
|
||||
);
|
||||
let (mut logging, _) = begin(py, &locals, true);
|
||||
logging
|
||||
.resume(py, Ok(local(&locals, "kwargs").unbind()))
|
||||
.unwrap();
|
||||
let (mut logging, step) = begin(py, &locals, true);
|
||||
resume(
|
||||
&mut logging,
|
||||
step,
|
||||
py,
|
||||
Ok(local(&locals, "kwargs").unbind()),
|
||||
)
|
||||
.unwrap();
|
||||
let failure = PyErr::from_value(local(&locals, "failure"));
|
||||
let failed = LifecycleEvent::Failed {
|
||||
let failed = HookEvent::Failed {
|
||||
timing: TIMING,
|
||||
origin: FailureOrigin::Call,
|
||||
error: &failure,
|
||||
};
|
||||
let step = logging.emit(py, failed).unwrap();
|
||||
let step = logging.on_event(py, failed).unwrap();
|
||||
assert!(awaits_deployment_hook(&step));
|
||||
let hook_result = if cancelled {
|
||||
Err(CancelledError::new_err("cancelled"))
|
||||
|
|
@ -722,8 +788,8 @@ assert finalized is replacement
|
|||
Ok(py.None())
|
||||
};
|
||||
assert!(matches!(
|
||||
logging.resume(py, hook_result).unwrap(),
|
||||
LifecycleStep::Await(_)
|
||||
resume(&mut logging, step, py, hook_result).unwrap(),
|
||||
HookStep::Await(_, _)
|
||||
));
|
||||
run(
|
||||
py,
|
||||
|
|
@ -743,7 +809,7 @@ mod payload_tests {
|
|||
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py};
|
||||
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned, to_py};
|
||||
use proptest::prelude::*;
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -788,11 +854,11 @@ check = lambda: None
|
|||
json!({"type": "document_url", "document_url": source})
|
||||
}
|
||||
|
||||
fn before_send(script: &CStr, body: Value) -> WireRequest {
|
||||
fn before_provider_request(script: &CStr, body: Value) -> WireRequest {
|
||||
before_send_with_secrets(script, json!({}), body, &[])
|
||||
}
|
||||
|
||||
/// Runs `before_send` over `body` for a route whose parameters are `optional_params`, with
|
||||
/// Runs `before_provider_request` over `body` for a route whose parameters are `optional_params`, with
|
||||
/// the Python objects `script` binds, then delivers the provider's raw response the way the
|
||||
/// driver does and runs the script's `check()`.
|
||||
fn before_send_with_secrets(
|
||||
|
|
@ -837,30 +903,35 @@ check = lambda: None
|
|||
};
|
||||
let (_, step) = send_and_receive(py, &mut logging, wire, &context);
|
||||
run(py, &locals, c"check()");
|
||||
let LifecycleStep::Wire(wire) = step else {
|
||||
panic!("before_send did not hand back the wire request");
|
||||
let HookStep::Ready(wire) = step else {
|
||||
panic!("before_provider_request did not hand back the wire request");
|
||||
};
|
||||
*wire
|
||||
})
|
||||
}
|
||||
|
||||
/// `before_send` over `wire`, then the provider's raw response the way the driver
|
||||
/// `before_provider_request` over `wire`, then the provider's raw response the way the driver
|
||||
/// delivers it, so `pre_call` and `post_call` have both seen the retained payload.
|
||||
fn send_and_receive<'a>(
|
||||
py: Python<'_>,
|
||||
logging: &'a mut LegacyLogging,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
) -> (&'a mut LegacyLogging, LifecycleStep) {
|
||||
let step = logging.before_send(py, Box::new(wire), context).unwrap();
|
||||
) -> (
|
||||
&'a mut LegacyLogging,
|
||||
HookStep<LegacyLogging, Box<WireRequest>>,
|
||||
) {
|
||||
let step = logging
|
||||
.before_provider_request(py, Box::new(wire), context)
|
||||
.unwrap();
|
||||
let raw = MachineEvent::ResponseReceived {
|
||||
raw: RawResponse {
|
||||
body: "raw response".into(),
|
||||
},
|
||||
};
|
||||
assert!(matches!(
|
||||
logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(),
|
||||
LifecycleStep::Done
|
||||
logging.on_event(py, HookEvent::Machine(&raw)).unwrap(),
|
||||
HookStep::Ready(())
|
||||
));
|
||||
(logging, step)
|
||||
}
|
||||
|
|
@ -904,7 +975,7 @@ check = lambda: None
|
|||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn a_cycle_through_the_retained_headers_is_collected() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -940,7 +1011,7 @@ assert reference() is None
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn close_releases_the_retained_headers() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -998,13 +1069,13 @@ def check():
|
|||
")]
|
||||
fn passthrough_keys_reach_pre_call_as_the_callers_own_objects(#[case] script: &CStr) {
|
||||
let body = json!({"model": "model", "document": document(DOCUMENT), "pages": [0]});
|
||||
let wire = before_send(script, body.clone());
|
||||
let wire = before_provider_request(script, body.clone());
|
||||
assert_eq!(wire.body, body);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn pre_call_edit_of_a_passthrough_object_reaches_the_caller_and_the_wire() {
|
||||
let wire = before_send(
|
||||
let wire = before_provider_request(
|
||||
c"
|
||||
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
|
||||
kwargs = {'document': document}
|
||||
|
|
@ -1018,9 +1089,9 @@ def check():
|
|||
assert_eq!(wire.body["document"], document(EDITED));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn a_body_key_the_route_rewrote_is_not_the_callers_object() {
|
||||
let wire = before_send(
|
||||
let wire = before_provider_request(
|
||||
c"
|
||||
document = {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'}
|
||||
kwargs = {'document': document}
|
||||
|
|
@ -1040,10 +1111,10 @@ def check():
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn a_caller_value_with_no_json_form_is_left_out_of_realiasing() {
|
||||
let body = json!({"pages": [0]});
|
||||
let wire = before_send(
|
||||
let wire = before_provider_request(
|
||||
c"
|
||||
opaque = object()
|
||||
kwargs = {'pages': opaque}
|
||||
|
|
@ -1072,14 +1143,14 @@ def on_pre_call(args):
|
|||
)]
|
||||
fn rebinding_the_payload_envelope_does_not_reach_the_wire(#[case] script: &CStr) {
|
||||
let body = json!({"document": document(DOCUMENT)});
|
||||
let wire = before_send(script, body.clone());
|
||||
let wire = before_provider_request(script, body.clone());
|
||||
assert_eq!(wire.body, body);
|
||||
assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn pre_call_header_edit_reaches_the_wire() {
|
||||
let wire = before_send(
|
||||
let wire = before_provider_request(
|
||||
c"
|
||||
def on_pre_call(args):
|
||||
args['headers']['x-callback'] = 'edited'
|
||||
|
|
@ -1095,7 +1166,7 @@ def on_pre_call(args):
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn pre_call_receives_the_wire_request_and_the_logger_its_redacted_request() {
|
||||
let body = json!({"model": "model", "document": document(DOCUMENT)});
|
||||
before_send_with_secrets(
|
||||
|
|
@ -1167,13 +1238,13 @@ def on_pre_call(args):
|
|||
)]
|
||||
fn pre_call_body_edits_reach_the_wire(#[case] script: &CStr, #[case] expected: Value) {
|
||||
let body = json!({"document": document(DOCUMENT)});
|
||||
let wire = before_send(script, body);
|
||||
let wire = before_provider_request(script, body);
|
||||
assert_eq!(wire.body, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn retained_headers_edited_after_rebinding_reach_the_wire() {
|
||||
let wire = before_send(
|
||||
let wire = before_provider_request(
|
||||
c"
|
||||
def on_pre_call(args):
|
||||
retained = args['headers']
|
||||
|
|
@ -1191,9 +1262,9 @@ def on_pre_call(args):
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn post_call_receives_the_raw_response_the_route_key_and_the_body_and_headers_pre_call_saw() {
|
||||
before_send(
|
||||
before_provider_request(
|
||||
c"
|
||||
def check():
|
||||
original_response, api_key, additional_args = logger.post
|
||||
|
|
@ -1210,9 +1281,9 @@ def check():
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn every_request_runs_the_full_pre_call_and_post_call() {
|
||||
let wire = before_send(
|
||||
let wire = before_provider_request(
|
||||
c"
|
||||
def on_pre_call(args):
|
||||
args['complete_input_dict']['include_image_base64'] = True
|
||||
|
|
@ -1343,7 +1414,7 @@ def check():
|
|||
/// For any body, any caller keywords and any callback edit: every keyword the route
|
||||
/// sends unchanged reaches `pre_call` as the caller's own object, and the provider is
|
||||
/// sent exactly what the model says, so a callback that edits nothing changes nothing.
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it(
|
||||
fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5),
|
||||
edit in edit(),
|
||||
|
|
@ -1389,7 +1460,7 @@ mod terminal_tests {
|
|||
use std::ffi::CStr;
|
||||
|
||||
use litellm_host::event::{FailureOrigin, Timing};
|
||||
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
|
||||
use litellm_host_python::{HookEvent, HookStep, PythonCallHooks, PythonOwned};
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::exceptions::asyncio::CancelledError;
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -1416,12 +1487,12 @@ mod terminal_tests {
|
|||
py: Python<'_>,
|
||||
locals: &Bound<'_, PyDict>,
|
||||
logging: &mut LegacyLogging,
|
||||
) -> LifecycleStep {
|
||||
) -> HookStep<LegacyLogging, ()> {
|
||||
let response = local(locals, "response").unbind();
|
||||
logging
|
||||
.emit(
|
||||
.on_event(
|
||||
py,
|
||||
LifecycleEvent::Succeeded {
|
||||
HookEvent::Succeeded {
|
||||
timing: TIMING,
|
||||
response: &response,
|
||||
},
|
||||
|
|
@ -1433,12 +1504,12 @@ mod terminal_tests {
|
|||
py: Python<'_>,
|
||||
locals: &Bound<'_, PyDict>,
|
||||
logging: &mut LegacyLogging,
|
||||
) -> LifecycleStep {
|
||||
) -> HookStep<LegacyLogging, ()> {
|
||||
let failure = PyErr::from_value(local(locals, "failure"));
|
||||
logging
|
||||
.emit(
|
||||
.on_event(
|
||||
py,
|
||||
LifecycleEvent::Failed {
|
||||
HookEvent::Failed {
|
||||
timing: TIMING,
|
||||
origin: FailureOrigin::Host,
|
||||
error: &failure,
|
||||
|
|
@ -1468,7 +1539,7 @@ mod terminal_tests {
|
|||
let mut logging = logged(py, &locals, asynchronous);
|
||||
assert!(matches!(
|
||||
succeed(py, &locals, &mut logging),
|
||||
LifecycleStep::Done
|
||||
HookStep::Ready(())
|
||||
));
|
||||
let names: Vec<String> = local(&locals, "logger")
|
||||
.call_method0("names")
|
||||
|
|
@ -1503,7 +1574,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy
|
|||
};
|
||||
assert!(matches!(
|
||||
fail(py, &locals, &mut logging),
|
||||
LifecycleStep::Done
|
||||
HookStep::Ready(())
|
||||
));
|
||||
let names: Vec<String> = local(&locals, "logger")
|
||||
.call_method0("names")
|
||||
|
|
@ -1514,7 +1585,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn internal_async_calls_skip_the_async_success_fan_out() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -1532,7 +1603,7 @@ assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_asy
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn a_failing_success_callback_is_reported_without_replacing_the_response() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -1552,7 +1623,7 @@ logger = FailingLogger()
|
|||
let mut logging = logged(py, &locals, true);
|
||||
assert!(matches!(
|
||||
succeed(py, &locals, &mut logging),
|
||||
LifecycleStep::Done
|
||||
HookStep::Ready(())
|
||||
));
|
||||
assert!(
|
||||
logging
|
||||
|
|
@ -1581,10 +1652,7 @@ logger = FailingLogger()
|
|||
let mut logging = logged(py, &locals, asynchronous);
|
||||
let step = fail(py, &locals, &mut logging);
|
||||
let awaits_async_handler = expected.contains(&"async_failure_handler");
|
||||
assert_eq!(
|
||||
matches!(step, LifecycleStep::Await(_)),
|
||||
awaits_async_handler
|
||||
);
|
||||
assert_eq!(matches!(step, HookStep::Await(_, _)), awaits_async_handler);
|
||||
let names: Vec<String> = local(&locals, "logger")
|
||||
.call_method0("names")
|
||||
.unwrap()
|
||||
|
|
@ -1599,7 +1667,7 @@ logger = FailingLogger()
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn a_failing_sync_failure_callback_keeps_the_error_and_still_runs_the_async_family() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -1619,7 +1687,7 @@ logger = FailingLogger()
|
|||
let mut logging = logged(py, &locals, true);
|
||||
assert!(matches!(
|
||||
fail(py, &locals, &mut logging),
|
||||
LifecycleStep::Await(_)
|
||||
HookStep::Await(_, _)
|
||||
));
|
||||
assert!(
|
||||
logging
|
||||
|
|
@ -1649,15 +1717,17 @@ logger = FailingLogger()
|
|||
Python::attach(|py| {
|
||||
let locals = namespace(py, c"failure = ValueError('provider')");
|
||||
let mut logging = logged(py, &locals, true);
|
||||
fail(py, &locals, &mut logging);
|
||||
let HookStep::Await(_, resume) = fail(py, &locals, &mut logging) else {
|
||||
panic!("expected async failure handler")
|
||||
};
|
||||
let result = match error {
|
||||
None => Ok(py.None()),
|
||||
Some(false) => Err(PyRuntimeError::new_err("handler failed")),
|
||||
Some(true) => Err(CancelledError::new_err("cancelled")),
|
||||
};
|
||||
let expected = result.as_ref().err().map(|error| error.value(py).clone());
|
||||
match logging.resume(py, result) {
|
||||
Ok(step) => assert!(done && matches!(step, LifecycleStep::Done)),
|
||||
match resume(&mut logging, py, result) {
|
||||
Ok(step) => assert!(done && matches!(step, HookStep::Ready(()))),
|
||||
Err(propagated) => {
|
||||
assert!(!done);
|
||||
assert!(propagated.value(py).is(expected.unwrap()));
|
||||
|
|
@ -1666,7 +1736,7 @@ logger = FailingLogger()
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn closing_restores_the_correlation_context_once() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@
|
|||
//! lifetime. No other callback host has that obligation, which is why nothing outside
|
||||
//! this crate holds them.
|
||||
|
||||
use litellm_host::{machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call};
|
||||
use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{Preflight, PythonBinding, PythonHostCalls, lookup, run_call};
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
|
|
@ -71,21 +71,22 @@ pub fn run_legacy_call<H, M>(
|
|||
py: Python<'_>,
|
||||
surface: LegacySurface,
|
||||
call: PublicCall,
|
||||
machine: M,
|
||||
start: impl FnOnce(<H::Protocol as Protocol>::Request) -> M + Send + Sync + 'static,
|
||||
host: H,
|
||||
preflight: Preflight,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
H: ProtocolHost + 'static,
|
||||
M: Machine<Protocol = H::Protocol, Complete = <H::Protocol as Protocol>::Response> + 'static,
|
||||
H: PythonBinding + PythonHostCalls<H::Protocol> + 'static,
|
||||
M: Machine<Protocol = H::Protocol> + 'static,
|
||||
M::Complete: Into<HostedCompletion<<H::Protocol as Protocol>::Response>>,
|
||||
{
|
||||
let arguments = call.kwargs.clone_ref(py);
|
||||
run_call(
|
||||
py,
|
||||
machine,
|
||||
start,
|
||||
host,
|
||||
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
|
||||
LegacyLogging::new(py, surface, call, asynchronous),
|
||||
preflight,
|
||||
arguments,
|
||||
asynchronous,
|
||||
|
|
@ -110,7 +111,7 @@ mod tests {
|
|||
(call, locals)
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn capture_copies_the_keyword_dict_without_copying_its_values() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
//! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the
|
||||
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
|
||||
//! proxy release. All of it sits behind one
|
||||
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
|
||||
//! [`PythonCallHooks`](litellm_host_python::PythonCallHooks), so the driver, the routes and
|
||||
//! core never learn which Python object is on the other end. The SDK's own request policy
|
||||
//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this
|
||||
//! crate's.
|
||||
|
|
|
|||
61
litellm-rust/crates/config/src/includes.rs
Normal file
61
litellm-rust/crates/config/src/includes.rs
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
use std::{
|
||||
collections::{BTreeSet, VecDeque},
|
||||
path::{Path, PathBuf},
|
||||
};
|
||||
|
||||
use serde::Deserialize;
|
||||
use serde_yaml_ng::{Mapping, Value};
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Includes {
|
||||
#[serde(default)]
|
||||
include: Vec<String>,
|
||||
}
|
||||
|
||||
fn read(path: &Path) -> Result<Mapping, Error> {
|
||||
Ok(serde_yaml_ng::from_str(&std::fs::read_to_string(path)?)?)
|
||||
}
|
||||
|
||||
fn entries(config: &Mapping, path: &Path) -> Result<Vec<(String, PathBuf)>, Error> {
|
||||
let includes: Includes = serde_yaml_ng::from_value(Value::Mapping(config.clone()))?;
|
||||
Ok(includes
|
||||
.include
|
||||
.into_iter()
|
||||
.map(|entry| (entry, path.to_owned()))
|
||||
.collect())
|
||||
}
|
||||
|
||||
pub(super) fn load(path: &Path) -> Result<Value, Error> {
|
||||
let root = path.canonicalize()?;
|
||||
let mut merged = read(&root)?;
|
||||
let mut pending: VecDeque<_> = entries(&merged, &root)?.into();
|
||||
let mut loaded = BTreeSet::from([root.clone()]);
|
||||
merged.remove(Value::String("include".into()));
|
||||
while let Some((entry, declaring)) = pending.pop_front() {
|
||||
let declared = declaring.parent().unwrap_or(Path::new(".")).join(&entry);
|
||||
let fallback = root.parent().unwrap_or(Path::new(".")).join(&entry);
|
||||
let location = if declared.exists() {
|
||||
declared
|
||||
} else {
|
||||
fallback
|
||||
}
|
||||
.canonicalize()?;
|
||||
if !loaded.insert(location.clone()) {
|
||||
continue;
|
||||
}
|
||||
let mut included = read(&location)?;
|
||||
pending.extend(entries(&included, &location)?);
|
||||
included.remove(Value::String("include".into()));
|
||||
for (key, value) in included {
|
||||
match (merged.get_mut(&key), value) {
|
||||
(Some(Value::Sequence(base)), Value::Sequence(extra)) => base.extend(extra),
|
||||
(_, value) => {
|
||||
merged.insert(key, value);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(Value::Mapping(merged))
|
||||
}
|
||||
|
|
@ -1,40 +1,82 @@
|
|||
mod error;
|
||||
mod includes;
|
||||
mod mcp;
|
||||
mod model;
|
||||
mod settings;
|
||||
mod value;
|
||||
|
||||
use std::path::Path;
|
||||
use std::{fmt, path::Path};
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
pub use error::Error;
|
||||
pub use mcp::{McpAuth, McpServer, McpTransport};
|
||||
pub use model::{LiteLlmParams, Model};
|
||||
pub use settings::{GeneralSettings, LiteLlmSettings, RouterSettings};
|
||||
pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct Config {
|
||||
pub model_list: Box<[Model]>,
|
||||
#[serde(default)]
|
||||
pub general_settings: GeneralSettings,
|
||||
pub router_settings: RouterSettings,
|
||||
pub litellm_settings: LiteLlmSettings,
|
||||
pub environment_variables: Object,
|
||||
pub callback_settings: Object,
|
||||
pub assistant_settings: Object,
|
||||
pub default_vertex_config: Object,
|
||||
pub mcp_servers: std::collections::BTreeMap<String, McpServer>,
|
||||
pub credential_list: Box<[Object]>,
|
||||
pub guardrails: Box<[Object]>,
|
||||
pub prompts: Box<[Object]>,
|
||||
pub sandbox_tools: Box<[Object]>,
|
||||
pub search_tools: Box<[Object]>,
|
||||
pub files_settings: Box<[Object]>,
|
||||
pub finetune_settings: Box<[Object]>,
|
||||
pub mcp_tools: Box<[Object]>,
|
||||
pub vector_store_registry: Box<[Object]>,
|
||||
pub worker_registry: Box<[Object]>,
|
||||
pub agents: Box<[Object]>,
|
||||
pub agent_list: Box<[Object]>,
|
||||
pub policies: Object,
|
||||
pub policy_attachments: Box<[Object]>,
|
||||
pub include: Box<[String]>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct GeneralSettings {
|
||||
pub master_key: Option<SecretValue>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct Model {
|
||||
pub model_name: String,
|
||||
pub litellm_params: LiteLlmParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct LiteLlmParams {
|
||||
pub model: String,
|
||||
pub api_key: Option<SecretValue>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
impl fmt::Debug for Config {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Config")
|
||||
.field("model_list", &self.model_list)
|
||||
.field("general_settings", &self.general_settings)
|
||||
.field("router_settings", &self.router_settings)
|
||||
.field("litellm_settings", &self.litellm_settings)
|
||||
.field("environment_variables", &self.environment_variables)
|
||||
.field("callback_settings", &self.callback_settings)
|
||||
.field("assistant_settings", &self.assistant_settings)
|
||||
.field("default_vertex_config", &self.default_vertex_config)
|
||||
.field("mcp_servers", &self.mcp_servers)
|
||||
.field("credential_list", &self.credential_list)
|
||||
.field("guardrails", &self.guardrails)
|
||||
.field("prompts", &self.prompts)
|
||||
.field("sandbox_tools", &self.sandbox_tools)
|
||||
.field("search_tools", &self.search_tools)
|
||||
.field("files_settings", &self.files_settings)
|
||||
.field("finetune_settings", &self.finetune_settings)
|
||||
.field("mcp_tools", &self.mcp_tools)
|
||||
.field("vector_store_registry", &self.vector_store_registry)
|
||||
.field("worker_registry", &self.worker_registry)
|
||||
.field("agents", &self.agents)
|
||||
.field("agent_list", &self.agent_list)
|
||||
.field("policies", &self.policies)
|
||||
.field("policy_attachments", &self.policy_attachments)
|
||||
.field("include", &self.include)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl Config {
|
||||
|
|
@ -43,6 +85,6 @@ impl Config {
|
|||
}
|
||||
|
||||
pub fn load(path: impl AsRef<Path>) -> Result<Self, Error> {
|
||||
Self::from_yaml(&std::fs::read_to_string(path)?)
|
||||
Ok(serde_yaml_ng::from_value(includes::load(path.as_ref())?)?)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
97
litellm-rust/crates/config/src/mcp.rs
Normal file
97
litellm-rust/crates/config/src/mcp.rs
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
use std::{collections::BTreeMap, fmt};
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::Object;
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct McpServer {
|
||||
pub server_id: Option<String>,
|
||||
pub alias: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub mcp_info: Object,
|
||||
pub transport: McpTransport,
|
||||
pub url: Option<SecretValue>,
|
||||
pub command: Option<String>,
|
||||
pub args: Box<[String]>,
|
||||
pub env: BTreeMap<String, SecretValue>,
|
||||
pub auth_type: Option<McpAuth>,
|
||||
#[serde(alias = "auth_value")]
|
||||
pub authentication_token: Option<SecretValue>,
|
||||
pub static_headers: BTreeMap<String, SecretValue>,
|
||||
pub upstream_token_header: Option<String>,
|
||||
pub allowed_tools: Option<Box<[String]>>,
|
||||
pub timeout: Option<f64>,
|
||||
pub max_concurrent_requests: Option<usize>,
|
||||
#[serde(flatten)]
|
||||
pub unsupported: Object,
|
||||
}
|
||||
|
||||
impl fmt::Debug for McpServer {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("McpServer")
|
||||
.field("transport", &self.transport)
|
||||
.field("auth_type", &self.auth_type)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("max_concurrent_requests", &self.max_concurrent_requests)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpTransport {
|
||||
#[default]
|
||||
Http,
|
||||
Sse,
|
||||
Stdio,
|
||||
}
|
||||
|
||||
impl McpTransport {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Http => "http",
|
||||
Self::Sse => "sse",
|
||||
Self::Stdio => "stdio",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpAuth {
|
||||
None,
|
||||
ApiKey,
|
||||
BearerToken,
|
||||
Basic,
|
||||
Authorization,
|
||||
Token,
|
||||
Oauth2,
|
||||
AwsSigv4,
|
||||
Oauth2TokenExchange,
|
||||
Oauth2IdJag,
|
||||
TruePassthrough,
|
||||
OauthDelegate,
|
||||
}
|
||||
|
||||
impl McpAuth {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::None => "none",
|
||||
Self::ApiKey => "api_key",
|
||||
Self::BearerToken => "bearer_token",
|
||||
Self::Basic => "basic",
|
||||
Self::Authorization => "authorization",
|
||||
Self::Token => "token",
|
||||
Self::Oauth2 => "oauth2",
|
||||
Self::AwsSigv4 => "aws_sigv4",
|
||||
Self::Oauth2TokenExchange => "oauth2_token_exchange",
|
||||
Self::Oauth2IdJag => "oauth2_id_jag",
|
||||
Self::TruePassthrough => "true_passthrough",
|
||||
Self::OauthDelegate => "oauth_delegate",
|
||||
}
|
||||
}
|
||||
}
|
||||
95
litellm-rust/crates/config/src/model.rs
Normal file
95
litellm-rust/crates/config/src/model.rs
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
use std::fmt;
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{AdditionalFields, Flag, NumberOrString, Object};
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
pub struct Model {
|
||||
pub model_name: String,
|
||||
pub litellm_params: LiteLlmParams,
|
||||
#[serde(default)]
|
||||
pub model_info: Object,
|
||||
pub blocked: Option<bool>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for Model {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Model")
|
||||
.field("model_name", &self.model_name)
|
||||
.field("litellm_params", &self.litellm_params)
|
||||
.field("model_info", &self.model_info)
|
||||
.field("blocked", &self.blocked)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
pub struct LiteLlmParams {
|
||||
pub model: String,
|
||||
pub api_key: Option<SecretValue>,
|
||||
pub api_base: Option<String>,
|
||||
pub api_version: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub timeout: Option<NumberOrString>,
|
||||
pub stream_timeout: Option<NumberOrString>,
|
||||
pub max_retries: Option<NumberOrString>,
|
||||
pub tpm: Option<NumberOrString>,
|
||||
pub rpm: Option<NumberOrString>,
|
||||
pub itpm: Option<NumberOrString>,
|
||||
pub otpm: Option<NumberOrString>,
|
||||
pub max_parallel_requests: Option<u64>,
|
||||
pub organization: Option<serde_yaml_ng::Value>,
|
||||
pub drop_params: Option<Flag>,
|
||||
pub tags: Option<Box<[String]>>,
|
||||
pub tag_regex: Option<Box<[String]>>,
|
||||
pub max_budget: Option<f64>,
|
||||
pub budget_duration: Option<String>,
|
||||
pub default_api_key_tpm_limit: Option<u64>,
|
||||
pub default_api_key_rpm_limit: Option<u64>,
|
||||
pub use_in_pass_through: Option<bool>,
|
||||
pub use_chat_completions_api: Option<bool>,
|
||||
pub litellm_credential_name: Option<String>,
|
||||
pub provider_affinity_header: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LiteLlmParams {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("LiteLlmParams")
|
||||
.field("model", &self.model)
|
||||
.field("api_key", &self.api_key)
|
||||
.field("api_base", &self.api_base)
|
||||
.field("api_version", &self.api_version)
|
||||
.field("custom_llm_provider", &self.custom_llm_provider)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("stream_timeout", &self.stream_timeout)
|
||||
.field("max_retries", &self.max_retries)
|
||||
.field("tpm", &self.tpm)
|
||||
.field("rpm", &self.rpm)
|
||||
.field("itpm", &self.itpm)
|
||||
.field("otpm", &self.otpm)
|
||||
.field("max_parallel_requests", &self.max_parallel_requests)
|
||||
.field("organization", &self.organization)
|
||||
.field("drop_params", &self.drop_params)
|
||||
.field("tags", &self.tags)
|
||||
.field("tag_regex", &self.tag_regex)
|
||||
.field("max_budget", &self.max_budget)
|
||||
.field("budget_duration", &self.budget_duration)
|
||||
.field("default_api_key_tpm_limit", &self.default_api_key_tpm_limit)
|
||||
.field("default_api_key_rpm_limit", &self.default_api_key_rpm_limit)
|
||||
.field("use_in_pass_through", &self.use_in_pass_through)
|
||||
.field("use_chat_completions_api", &self.use_chat_completions_api)
|
||||
.field("litellm_credential_name", &self.litellm_credential_name)
|
||||
.field("provider_affinity_header", &self.provider_affinity_header)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
221
litellm-rust/crates/config/src/settings.rs
Normal file
221
litellm-rust/crates/config/src/settings.rs
Normal file
|
|
@ -0,0 +1,221 @@
|
|||
use std::fmt;
|
||||
|
||||
use litellm_auth_types::SecretValue;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
|
||||
|
||||
#[derive(Clone, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct GeneralSettings {
|
||||
pub completion_model: Option<String>,
|
||||
pub max_in_flight_requests_per_worker: Option<u64>,
|
||||
pub max_queued_requests_per_worker: Option<u64>,
|
||||
pub admission_queue_timeout_seconds: f64,
|
||||
pub master_key: Option<SecretValue>,
|
||||
pub database_url: Option<SecretValue>,
|
||||
pub database_connection_pool_limit: Option<u64>,
|
||||
pub database_connection_timeout: Option<f64>,
|
||||
pub database_connect_timeout: Option<f64>,
|
||||
pub database_socket_timeout: Option<f64>,
|
||||
pub database_max_idle_connection_lifetime: Option<f64>,
|
||||
pub max_parallel_requests: Option<u64>,
|
||||
pub global_max_parallel_requests: Option<u64>,
|
||||
pub max_request_size_mb: Option<u64>,
|
||||
pub max_response_size_mb: Option<u64>,
|
||||
pub proxy_config_reload_interval_seconds: u64,
|
||||
pub background_health_checks: Option<bool>,
|
||||
pub health_check_interval: u64,
|
||||
pub health_check_concurrency: Option<u64>,
|
||||
pub store_model_in_db: Option<bool>,
|
||||
pub forward_client_headers_to_llm_api: Option<bool>,
|
||||
pub cancel_on_disconnect: Option<bool>,
|
||||
pub infer_model_from_keys: Option<bool>,
|
||||
pub enable_public_model_hub: bool,
|
||||
pub dangerously_permit_weak_or_unset_master_key: Option<bool>,
|
||||
pub plugins: Option<Box<[Object]>>,
|
||||
pub coordination_redis: Option<Object>,
|
||||
pub mcp_allowed_hosts: Option<Box<[String]>>,
|
||||
pub mcp_allowed_origins: Box<[String]>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl Default for GeneralSettings {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
completion_model: None,
|
||||
max_in_flight_requests_per_worker: None,
|
||||
max_queued_requests_per_worker: None,
|
||||
admission_queue_timeout_seconds: 1.0,
|
||||
master_key: None,
|
||||
database_url: None,
|
||||
database_connection_pool_limit: Some(10),
|
||||
database_connection_timeout: Some(60.0),
|
||||
database_connect_timeout: None,
|
||||
database_socket_timeout: None,
|
||||
database_max_idle_connection_lifetime: Some(60.0),
|
||||
max_parallel_requests: None,
|
||||
global_max_parallel_requests: None,
|
||||
max_request_size_mb: None,
|
||||
max_response_size_mb: None,
|
||||
proxy_config_reload_interval_seconds: 30,
|
||||
background_health_checks: None,
|
||||
health_check_interval: 300,
|
||||
health_check_concurrency: None,
|
||||
store_model_in_db: None,
|
||||
forward_client_headers_to_llm_api: None,
|
||||
cancel_on_disconnect: None,
|
||||
infer_model_from_keys: None,
|
||||
enable_public_model_hub: false,
|
||||
dangerously_permit_weak_or_unset_master_key: None,
|
||||
plugins: None,
|
||||
coordination_redis: None,
|
||||
mcp_allowed_hosts: None,
|
||||
mcp_allowed_origins: Box::default(),
|
||||
additional_fields: AdditionalFields::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for GeneralSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("GeneralSettings")
|
||||
.field("completion_model", &self.completion_model)
|
||||
.field(
|
||||
"max_in_flight_requests_per_worker",
|
||||
&self.max_in_flight_requests_per_worker,
|
||||
)
|
||||
.field(
|
||||
"max_queued_requests_per_worker",
|
||||
&self.max_queued_requests_per_worker,
|
||||
)
|
||||
.field(
|
||||
"admission_queue_timeout_seconds",
|
||||
&self.admission_queue_timeout_seconds,
|
||||
)
|
||||
.field("master_key", &self.master_key)
|
||||
.field("database_url", &self.database_url)
|
||||
.field("store_model_in_db", &self.store_model_in_db)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct RouterSettings {
|
||||
pub routing_strategy: Option<String>,
|
||||
pub routing_strategy_args: Option<Object>,
|
||||
pub routing_groups: Option<Box<[Object]>>,
|
||||
pub retry_policy: Option<Object>,
|
||||
pub model_group_retry_policy: Option<Object>,
|
||||
pub model_group_affinity_config: Option<Object>,
|
||||
pub allowed_fails: Option<u64>,
|
||||
pub cooldown_time: Option<f64>,
|
||||
pub num_retries: Option<u64>,
|
||||
pub timeout: Option<f64>,
|
||||
pub max_retries: Option<u64>,
|
||||
pub retry_after: Option<f64>,
|
||||
pub fallbacks: Option<Box<[Object]>>,
|
||||
pub context_window_fallbacks: Option<Box<[Object]>>,
|
||||
pub model_group_alias: Option<Object>,
|
||||
pub enable_tag_filtering: Option<bool>,
|
||||
pub weights: Option<Object>,
|
||||
pub tag_routing_prefix: Option<String>,
|
||||
pub optional_pre_call_checks: Option<Box<[String]>>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for RouterSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("RouterSettings")
|
||||
.field("routing_strategy", &self.routing_strategy)
|
||||
.field("routing_strategy_args", &self.routing_strategy_args)
|
||||
.field("routing_groups", &self.routing_groups)
|
||||
.field("retry_policy", &self.retry_policy)
|
||||
.field("model_group_retry_policy", &self.model_group_retry_policy)
|
||||
.field(
|
||||
"model_group_affinity_config",
|
||||
&self.model_group_affinity_config,
|
||||
)
|
||||
.field("allowed_fails", &self.allowed_fails)
|
||||
.field("cooldown_time", &self.cooldown_time)
|
||||
.field("num_retries", &self.num_retries)
|
||||
.field("timeout", &self.timeout)
|
||||
.field("max_retries", &self.max_retries)
|
||||
.field("retry_after", &self.retry_after)
|
||||
.field("fallbacks", &self.fallbacks)
|
||||
.field("context_window_fallbacks", &self.context_window_fallbacks)
|
||||
.field("model_group_alias", &self.model_group_alias)
|
||||
.field("enable_tag_filtering", &self.enable_tag_filtering)
|
||||
.field("weights", &self.weights)
|
||||
.field("tag_routing_prefix", &self.tag_routing_prefix)
|
||||
.field("optional_pre_call_checks", &self.optional_pre_call_checks)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(default)]
|
||||
pub struct LiteLlmSettings {
|
||||
pub ssl_verify: Option<Flag>,
|
||||
pub ssl_certificate: Option<String>,
|
||||
pub ssl_security_level: Option<String>,
|
||||
pub ssl_ecdh_curve: Option<String>,
|
||||
pub force_ipv4: Option<bool>,
|
||||
pub http2: Option<bool>,
|
||||
pub aiohttp_trust_env: Option<bool>,
|
||||
pub disable_aiohttp_trust_env: Option<bool>,
|
||||
pub disable_aiohttp_transport: Option<bool>,
|
||||
pub drop_params: Option<Flag>,
|
||||
pub request_timeout: Option<NumberOrString>,
|
||||
pub num_retries: Option<u64>,
|
||||
pub cache: Option<bool>,
|
||||
pub cache_params: Option<Object>,
|
||||
pub callbacks: Option<OneOrMany<Value>>,
|
||||
pub success_callback: Option<OneOrMany<Value>>,
|
||||
pub failure_callback: Option<OneOrMany<Value>>,
|
||||
pub json_logs: Option<bool>,
|
||||
pub set_verbose: Option<bool>,
|
||||
#[serde(flatten)]
|
||||
pub additional_fields: AdditionalFields,
|
||||
}
|
||||
|
||||
impl fmt::Debug for LiteLlmSettings {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("LiteLlmSettings")
|
||||
.field("drop_params", &self.drop_params)
|
||||
.field("request_timeout", &self.request_timeout)
|
||||
.field("num_retries", &self.num_retries)
|
||||
.field("cache", &self.cache)
|
||||
.field("cache_params", &self.cache_params)
|
||||
.field(
|
||||
"callbacks",
|
||||
&self.callbacks.as_ref().map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field(
|
||||
"success_callback",
|
||||
&self
|
||||
.success_callback
|
||||
.as_ref()
|
||||
.map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field(
|
||||
"failure_callback",
|
||||
&self
|
||||
.failure_callback
|
||||
.as_ref()
|
||||
.map(|callbacks| callbacks.len()),
|
||||
)
|
||||
.field("json_logs", &self.json_logs)
|
||||
.field("set_verbose", &self.set_verbose)
|
||||
.field("additional_fields", &self.additional_fields.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
78
litellm-rust/crates/config/src/value.rs
Normal file
78
litellm-rust/crates/config/src/value.rs
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
use std::{collections::BTreeMap, fmt, ops::Deref};
|
||||
|
||||
use serde::Deserialize;
|
||||
|
||||
pub type Value = serde_yaml_ng::Value;
|
||||
pub type AdditionalFields = BTreeMap<String, Value>;
|
||||
|
||||
#[derive(Clone, Default, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct Object(BTreeMap<String, Value>);
|
||||
|
||||
impl Object {
|
||||
pub fn get(&self, key: &str) -> Option<&Value> {
|
||||
self.0.get(key)
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.0.is_empty()
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize {
|
||||
self.0.len()
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for Object {
|
||||
type Target = BTreeMap<String, Value>;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for Object {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("Object")
|
||||
.field("keys", &self.0.keys())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum NumberOrString {
|
||||
Number(f64),
|
||||
String(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Eq)]
|
||||
#[serde(untagged)]
|
||||
pub enum Flag {
|
||||
Boolean(bool),
|
||||
String(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum OneOrMany<T> {
|
||||
Many(Box<[T]>),
|
||||
One(T),
|
||||
}
|
||||
|
||||
impl<T> OneOrMany<T> {
|
||||
pub fn len(&self) -> usize {
|
||||
match self {
|
||||
Self::Many(values) => values.len(),
|
||||
Self::One(_) => 1,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
match self {
|
||||
Self::Many(values) => values.is_empty(),
|
||||
Self::One(_) => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_config::{Config, Error};
|
||||
use litellm_config::{Config, Error, Flag, NumberOrString};
|
||||
use rstest::{fixture, rstest};
|
||||
use tempfile::TempDir;
|
||||
|
||||
|
|
@ -73,14 +73,9 @@ fn config_debug_redacts_api_keys() {
|
|||
|
||||
#[rstest]
|
||||
#[case::malformed_yaml("model_list: [")]
|
||||
#[case::missing_model_list("{}")]
|
||||
#[case::missing_params("model_list: [{model_name: assistant}]")]
|
||||
#[case::missing_model("model_list: [{model_name: assistant, litellm_params: {api_key: key}}]")]
|
||||
#[case::unsupported_settings("model_list: []\ngeneral_settings: {unknown: true}")]
|
||||
#[case::misspelled_param(
|
||||
"model_list: [{model_name: assistant, litellm_params: {model: test, api_bsae: url}}]"
|
||||
)]
|
||||
fn rejects_malformed_incomplete_and_unsupported_config(#[case] yaml: &str) {
|
||||
fn rejects_malformed_and_incomplete_config(#[case] yaml: &str) {
|
||||
assert!(matches!(Config::from_yaml(yaml), Err(Error::Parse(_))));
|
||||
}
|
||||
|
||||
|
|
@ -117,3 +112,264 @@ fn missing_general_settings_has_no_master_key() {
|
|||
let config = Config::from_yaml("model_list: []").unwrap();
|
||||
assert!(config.general_settings.master_key.is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn empty_config_matches_python_defaults() {
|
||||
let config = Config::from_yaml("{}").unwrap();
|
||||
assert!(config.model_list.is_empty());
|
||||
assert_eq!(config.general_settings.admission_queue_timeout_seconds, 1.0);
|
||||
assert_eq!(
|
||||
config.general_settings.database_connection_pool_limit,
|
||||
Some(10)
|
||||
);
|
||||
assert_eq!(
|
||||
config.general_settings.proxy_config_reload_interval_seconds,
|
||||
30
|
||||
);
|
||||
assert_eq!(config.general_settings.health_check_interval, 300);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn parses_typed_settings_and_preserves_extension_fields() {
|
||||
let config = Config::from_yaml(
|
||||
r#"
|
||||
model_list:
|
||||
- model_name: assistant
|
||||
litellm_params:
|
||||
model: vertex_ai/test-model
|
||||
timeout: os.environ/REQUEST_TIMEOUT
|
||||
tpm: os.environ/TPM_LIMIT
|
||||
rpm: 5
|
||||
drop_params: "true"
|
||||
vertex_project: test-project
|
||||
model_info:
|
||||
mode: chat
|
||||
access_groups: [internal]
|
||||
general_settings:
|
||||
master_key: secret-master-key
|
||||
store_model_in_db: true
|
||||
custom_auth: auth.py
|
||||
router_settings:
|
||||
routing_strategy: simple-shuffle
|
||||
allowed_fails: 2
|
||||
redis_host: cache.internal
|
||||
litellm_settings:
|
||||
drop_params: true
|
||||
cache: true
|
||||
custom_callback_name: audit
|
||||
future_section:
|
||||
enabled: true
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let model = &config.model_list[0];
|
||||
assert_eq!(
|
||||
model.litellm_params.timeout,
|
||||
Some(NumberOrString::String(
|
||||
"os.environ/REQUEST_TIMEOUT".to_string()
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
model.litellm_params.tpm,
|
||||
Some(NumberOrString::String("os.environ/TPM_LIMIT".to_string()))
|
||||
);
|
||||
assert_eq!(model.litellm_params.rpm, Some(NumberOrString::Number(5.0)));
|
||||
assert_eq!(
|
||||
model.litellm_params.drop_params,
|
||||
Some(Flag::String("true".to_string()))
|
||||
);
|
||||
assert!(
|
||||
model
|
||||
.litellm_params
|
||||
.additional_fields
|
||||
.contains_key("vertex_project")
|
||||
);
|
||||
assert!(model.additional_fields.contains_key("access_groups"));
|
||||
assert_eq!(config.router_settings.allowed_fails, Some(2));
|
||||
assert!(
|
||||
config
|
||||
.router_settings
|
||||
.additional_fields
|
||||
.contains_key("redis_host")
|
||||
);
|
||||
assert!(
|
||||
config
|
||||
.litellm_settings
|
||||
.additional_fields
|
||||
.contains_key("custom_callback_name")
|
||||
);
|
||||
assert!(config.additional_fields.contains_key("future_section"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn parses_python_config_sections() {
|
||||
let config = Config::from_yaml(
|
||||
r#"
|
||||
environment_variables:
|
||||
REDIS_PORT: 6379
|
||||
callback_settings:
|
||||
otel:
|
||||
message_logging: false
|
||||
assistant_settings:
|
||||
custom_llm_provider: openai
|
||||
credential_list:
|
||||
- credential_name: bedrock
|
||||
credential_values:
|
||||
aws_region_name: us-east-1
|
||||
guardrails:
|
||||
- guardrail_name: pii
|
||||
litellm_params:
|
||||
guardrail: presidio
|
||||
prompts:
|
||||
- prompt_id: support
|
||||
sandbox_tools:
|
||||
- sandbox_tool_name: e2b
|
||||
search_tools:
|
||||
- search_tool_name: web
|
||||
files_settings:
|
||||
- custom_llm_provider: openai
|
||||
finetune_settings:
|
||||
- custom_llm_provider: openai
|
||||
mcp_tools:
|
||||
- name: lookup
|
||||
mcp_servers:
|
||||
docs:
|
||||
url: https://example.test/mcp
|
||||
vector_store_registry:
|
||||
- vector_store_name: docs
|
||||
worker_registry:
|
||||
- worker_id: regional
|
||||
agents:
|
||||
- agent_name: reviewer
|
||||
agent_list:
|
||||
- agent_name: legacy-reviewer
|
||||
policies:
|
||||
safe:
|
||||
guardrails:
|
||||
add: [pii]
|
||||
policy_attachments:
|
||||
- policy_id: safe
|
||||
include:
|
||||
- models.yaml
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(config.environment_variables.len(), 1);
|
||||
assert_eq!(config.credential_list.len(), 1);
|
||||
assert_eq!(config.guardrails.len(), 1);
|
||||
assert_eq!(config.prompts.len(), 1);
|
||||
assert_eq!(config.sandbox_tools.len(), 1);
|
||||
assert_eq!(config.search_tools.len(), 1);
|
||||
assert_eq!(config.files_settings.len(), 1);
|
||||
assert_eq!(config.finetune_settings.len(), 1);
|
||||
assert_eq!(config.mcp_tools.len(), 1);
|
||||
assert_eq!(config.mcp_servers.len(), 1);
|
||||
assert_eq!(config.vector_store_registry.len(), 1);
|
||||
assert_eq!(config.worker_registry.len(), 1);
|
||||
assert_eq!(config.agents.len(), 1);
|
||||
assert_eq!(config.agent_list.len(), 1);
|
||||
assert!(config.policies.contains_key("safe"));
|
||||
assert_eq!(config.policy_attachments.len(), 1);
|
||||
assert_eq!(&*config.include, &["models.yaml"]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::policy_pipeline("../../../litellm/proxy/example_config_yaml/test_pipeline_config.yaml")]
|
||||
#[case::gateway("../../../tests/e2e/gateway/stage_mirror_ci_config.yml")]
|
||||
fn parses_representative_python_configs(#[case] relative_path: &str) {
|
||||
let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join(relative_path);
|
||||
let config = Config::load(path).unwrap();
|
||||
assert!(!config.model_list.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::one("callbacks: custom_callbacks.logger", 1)]
|
||||
#[case::many("callbacks: [prometheus, otel]", 2)]
|
||||
fn accepts_python_callback_shorthand(#[case] setting: &str, #[case] expected_len: usize) {
|
||||
let config = Config::from_yaml(&format!("litellm_settings:\n {setting}")).unwrap();
|
||||
assert_eq!(
|
||||
config.litellm_settings.callbacks.as_ref().unwrap().len(),
|
||||
expected_len
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key("api_key", "provider-secret")]
|
||||
#[case::provider_extension("aws_secret_access_key", "aws-secret")]
|
||||
#[case::general_extension("custom_auth_secret", "auth-secret")]
|
||||
#[case::router_extension("redis_password", "redis-secret")]
|
||||
#[case::litellm_extension("callback_token", "callback-secret")]
|
||||
#[case::root_extension("private_token", "root-secret")]
|
||||
fn debug_output_does_not_expose_config_values(#[case] field: &str, #[case] secret: &str) {
|
||||
let yaml = match field {
|
||||
"api_key" => format!(
|
||||
"model_list: [{{model_name: assistant, litellm_params: {{model: test, api_key: {secret}}}}}]"
|
||||
),
|
||||
"aws_secret_access_key" => format!(
|
||||
"model_list: [{{model_name: assistant, litellm_params: {{model: test, aws_secret_access_key: {secret}}}}}]"
|
||||
),
|
||||
"custom_auth_secret" => format!("general_settings: {{{field}: {secret}}}"),
|
||||
"redis_password" => format!("router_settings: {{{field}: {secret}}}"),
|
||||
"callback_token" => format!("litellm_settings: {{{field}: {secret}}}"),
|
||||
"private_token" => format!("{field}: {secret}"),
|
||||
_ => unreachable!(),
|
||||
};
|
||||
let config = Config::from_yaml(&yaml).unwrap();
|
||||
assert!(!format!("{config:?}").contains(secret));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn resolves_nested_includes_once_in_breadth_first_order() {
|
||||
let directory = TempDir::new().unwrap();
|
||||
let root = directory.path();
|
||||
std::fs::create_dir(root.join("nested")).unwrap();
|
||||
std::fs::write(root.join("config.yaml"), "include: [nested/first.yaml, second.yaml]\nmodel_list: [{model_name: root, litellm_params: {model: root}}]\n").unwrap();
|
||||
std::fs::write(root.join("nested/first.yaml"), "include: [third.yaml]\nmodel_list: [{model_name: first, litellm_params: {model: first}}]\n").unwrap();
|
||||
std::fs::write(root.join("second.yaml"), "general_settings: {master_key: second}\nmodel_list: [{model_name: second, litellm_params: {model: second}}]\n").unwrap();
|
||||
std::fs::write(root.join("nested/third.yaml"), "include: [../config.yaml]\ngeneral_settings: {master_key: third}\nmodel_list: [{model_name: third, litellm_params: {model: third}}]\n").unwrap();
|
||||
|
||||
let config = Config::load(root.join("config.yaml")).unwrap();
|
||||
assert_eq!(
|
||||
config
|
||||
.model_list
|
||||
.iter()
|
||||
.map(|model| model.model_name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["root", "first", "second", "third"]
|
||||
);
|
||||
assert_eq!(
|
||||
config.general_settings.master_key.unwrap().expose(),
|
||||
"third"
|
||||
);
|
||||
assert!(config.include.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn mcp_config_redacts_nested_credentials_and_preserves_policy_for_validation() {
|
||||
let config = Config::from_yaml("mcp_servers:\n docs:\n url: https://example.test/private-secret/mcp\n authentication_token: upstream-secret\n static_headers: {x-token: header-secret}\n env: {TOKEN: env-secret}\n args: [argument-secret]\n client_secret: oauth-secret\n allowed_tools: [search]\n").unwrap();
|
||||
let server = &config.mcp_servers["docs"];
|
||||
assert_eq!(server.allowed_tools.as_deref().unwrap(), ["search"]);
|
||||
assert!(server.unsupported.contains_key("client_secret"));
|
||||
let debug = format!("{config:?}");
|
||||
for secret in [
|
||||
"private-secret",
|
||||
"upstream-secret",
|
||||
"header-secret",
|
||||
"env-secret",
|
||||
"oauth-secret",
|
||||
"argument-secret",
|
||||
] {
|
||||
assert!(!debug.contains(secret));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::transport("transport: invalid")]
|
||||
#[case::auth("auth_type: invalid")]
|
||||
#[case::concurrency("max_concurrent_requests: -1")]
|
||||
#[case::headers("static_headers: {x-token: [not, a, string]}")]
|
||||
fn rejects_invalid_typed_mcp_settings(#[case] setting: &str) {
|
||||
assert!(Config::from_yaml(&format!("mcp_servers:\n docs:\n {setting}\n")).is_err());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -109,6 +109,7 @@ impl IntoIterator for CallArguments {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
|
@ -140,22 +141,31 @@ mod tests {
|
|||
assert_eq!(serde_json::to_value(arguments).unwrap(), original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty() {
|
||||
for value in [json!(false), json!(0), json!([]), json!("")] {
|
||||
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
|
||||
assert_eq!(
|
||||
compose_body(&arguments, &json!({}), &[]),
|
||||
Err(crate::params::Error::ExtraBody)
|
||||
);
|
||||
}
|
||||
let arguments = serde_json::from_value(json!({"extra_body":null})).unwrap();
|
||||
#[rstest]
|
||||
#[case::boolean(json!(false))]
|
||||
#[case::number(json!(0))]
|
||||
#[case::array(json!([]))]
|
||||
#[case::string(json!(""))]
|
||||
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty(
|
||||
#[case] value: serde_json::Value,
|
||||
) {
|
||||
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
|
||||
assert_eq!(
|
||||
compose_body(&arguments, &json!({}), &[]).unwrap(),
|
||||
json!({})
|
||||
compose_body(&arguments, &json!({}), &[]),
|
||||
Err(crate::params::Error::ExtraBody)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::null(json!(null), json!({}))]
|
||||
fn null_extra_body_is_coerced_to_empty_object(
|
||||
#[case] value: serde_json::Value,
|
||||
#[case] expected: serde_json::Value,
|
||||
) {
|
||||
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
|
||||
assert_eq!(compose_body(&arguments, &json!({}), &[]).unwrap(), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typed_views_preserve_missing_and_explicit_null_in_the_source() {
|
||||
#[derive(Deserialize)]
|
||||
|
|
|
|||
|
|
@ -68,27 +68,25 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str {
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rstest::rstest;
|
||||
|
||||
#[test]
|
||||
fn maps_every_reason_the_route_can_observe() {
|
||||
assert_eq!(finish_reason_for("end_turn"), "stop");
|
||||
assert_eq!(finish_reason_for("stop_sequence"), "stop");
|
||||
assert_eq!(finish_reason_for("max_tokens"), "length");
|
||||
assert_eq!(finish_reason_for("refusal"), "content_filter");
|
||||
assert_eq!(finish_reason_for("guardrail_intervened"), "content_filter");
|
||||
// Converse emits these two, and folding them into `stop` would report a
|
||||
// filtered completion as a normal one.
|
||||
assert_eq!(finish_reason_for("content_filtered"), "content_filter");
|
||||
assert_eq!(finish_reason_for("content_filter"), "content_filter");
|
||||
#[rstest]
|
||||
#[case::end_turn("end_turn", "stop")]
|
||||
#[case::stop_sequence("stop_sequence", "stop")]
|
||||
#[case::max_tokens("max_tokens", "length")]
|
||||
#[case::refusal("refusal", "content_filter")]
|
||||
#[case::guardrail_intervened("guardrail_intervened", "content_filter")]
|
||||
#[case::content_filtered("content_filtered", "content_filter")]
|
||||
#[case::content_filter("content_filter", "content_filter")]
|
||||
fn maps_every_reason_the_route_can_observe(#[case] reason: &str, #[case] expected: &str) {
|
||||
assert_eq!(finish_reason_for(reason), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn defaults_an_unmapped_reason_to_stop_like_python() {
|
||||
// Python warns and falls back to `stop` for a reason its own map does
|
||||
// not carry, so only a reason absent from `_FINISH_REASON_MAP` belongs
|
||||
// here.
|
||||
assert_eq!(finish_reason_for("something_new"), "stop");
|
||||
assert_eq!(finish_reason_for(""), "stop");
|
||||
#[rstest]
|
||||
#[case::unknown("something_new")]
|
||||
#[case::empty("")]
|
||||
fn defaults_an_unmapped_reason_to_stop_like_python(#[case] reason: &str) {
|
||||
assert_eq!(finish_reason_for(reason), "stop");
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -1,9 +1,26 @@
|
|||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct CustomLlmProvider<'a> {
|
||||
pub model: &'a str,
|
||||
pub custom_llm_provider: &'a str,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub enum LlmProviders {
|
||||
Anthropic,
|
||||
AwsTextract,
|
||||
AzureAi,
|
||||
Bedrock,
|
||||
Cohere,
|
||||
Mistral,
|
||||
Openai,
|
||||
OpenaiLike,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
pub fn get_custom_llm_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
|
|
|
|||
|
|
@ -129,6 +129,7 @@ fn integral_float(value: f64) -> Option<i64> {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::json;
|
||||
use serde_with::serde_as;
|
||||
|
|
@ -144,20 +145,20 @@ mod tests {
|
|||
float: Option<f64>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens() {
|
||||
for (input, expected) in [
|
||||
(" True ", Some(true)),
|
||||
("\u{1c}TRUE\u{1f}", Some(true)),
|
||||
("\u{a0}False\u{2003}", Some(false)),
|
||||
("true\u{200b}", None),
|
||||
("yes", None),
|
||||
("1", None),
|
||||
("", None),
|
||||
("unknown", None),
|
||||
] {
|
||||
assert_eq!(parse_str_bool(input), expected, "{input:?}");
|
||||
}
|
||||
#[rstest]
|
||||
#[case::trimmed_true(" True ", Some(true))]
|
||||
#[case::control_whitespace_true("\u{1c}TRUE\u{1f}", Some(true))]
|
||||
#[case::unicode_whitespace_false("\u{a0}False\u{2003}", Some(false))]
|
||||
#[case::zero_width_space("true\u{200b}", None)]
|
||||
#[case::yes("yes", None)]
|
||||
#[case::one("1", None)]
|
||||
#[case::empty("", None)]
|
||||
#[case::unknown("unknown", None)]
|
||||
fn boolean_tokens_follow_python_string_trimming_without_redis_tokens(
|
||||
#[case] input: &str,
|
||||
#[case] expected: Option<bool>,
|
||||
) {
|
||||
assert_eq!(parse_str_bool(input), expected, "{input:?}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -37,6 +37,22 @@ impl Lookup for ProcessEnvironment {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn resolve_non_empty(
|
||||
value: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
names: &[&str],
|
||||
) -> Option<String> {
|
||||
value
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| {
|
||||
names
|
||||
.iter()
|
||||
.find_map(|name| env_lookup(name).filter(|value| !value.trim().is_empty()))
|
||||
})
|
||||
}
|
||||
|
||||
pub trait Layer: Default {
|
||||
fn or(self, lower: Self) -> Self;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -72,20 +72,18 @@ impl ApiUrl<Complete> {
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rstest::rstest;
|
||||
|
||||
#[test]
|
||||
fn completion_appends_only_the_missing_path_suffix() {
|
||||
for (base, expected) in [
|
||||
("https://example.test", "https://example.test/v1/ocr"),
|
||||
("https://example.test/v1", "https://example.test/v1/ocr"),
|
||||
("https://example.test/v1/ocr", "https://example.test/v1/ocr"),
|
||||
] {
|
||||
let actual = ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.expect("url builds");
|
||||
assert_eq!(actual, expected);
|
||||
}
|
||||
#[rstest]
|
||||
#[case::root("https://example.test", "https://example.test/v1/ocr")]
|
||||
#[case::version_prefix("https://example.test/v1", "https://example.test/v1/ocr")]
|
||||
#[case::complete("https://example.test/v1/ocr", "https://example.test/v1/ocr")]
|
||||
fn completion_appends_only_the_missing_path_suffix(#[case] base: &str, #[case] expected: &str) {
|
||||
let actual = ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.expect("url builds");
|
||||
assert_eq!(actual, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
33
litellm-rust/crates/core-utils/tests/settings.rs
Normal file
33
litellm-rust/crates/core-utils/tests/settings.rs
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
use litellm_core_utils::settings::resolve_non_empty;
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::explicit_wins(Some(" explicit "), &["FIRST", "SECOND"], Some("explicit"))]
|
||||
#[case::absent_falls_back(None, &["FIRST", "SECOND"], Some(" first "))]
|
||||
#[case::blank_falls_back(Some(" \t "), &["BLANK", "SECOND"], Some("second"))]
|
||||
#[case::skips_missing_and_blank(None, &["MISSING", "BLANK", "SECOND"], Some("second"))]
|
||||
#[case::environment_order(None, &["SECOND", "FIRST"], Some("second"))]
|
||||
#[case::missing(None, &["MISSING", "BLANK"], None)]
|
||||
#[case::no_environment(None, &[], None)]
|
||||
fn resolves_explicit_value_then_first_nonblank_environment_value(
|
||||
#[case] value: Option<&str>,
|
||||
#[case] names: &[&str],
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
let env = |name: &str| match name {
|
||||
"FIRST" => Some(" first ".to_string()),
|
||||
"SECOND" => Some("second".to_string()),
|
||||
"BLANK" => Some(" \t ".to_string()),
|
||||
_ => None,
|
||||
};
|
||||
assert_eq!(resolve_non_empty(value, &env, names).as_deref(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn explicit_value_does_not_read_the_environment() {
|
||||
let env = |_: &str| panic!("an explicit value must short-circuit environment lookup");
|
||||
assert_eq!(
|
||||
resolve_non_empty(Some("key"), &env, &["KEY"]).as_deref(),
|
||||
Some("key")
|
||||
);
|
||||
}
|
||||
|
|
@ -1,9 +1,17 @@
|
|||
litellm-core is the LiteLLM SDK in Rust. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route. `messages::messages()` returns `MessagesResponse::Message` for a completed response or `MessagesResponse::Stream { headers, chunks }` when the request sets `stream: true`. The chunks are Anthropic SSE bytes in a `Stream<Item = Result<Bytes, Error>>`. Dropping the stream cancels the call. The Python bridge drives `messages::route::messages_machine()` instead, because Python has to answer the call's operations on its own thread; the gateway and the Rust SDK call the plain entrypoint
|
||||
litellm-core owns route orchestration. Messages and HTTP Responses return `litellm_host::call::CallOutput`, containing either a completed response or a stream head and chunks. OCR and currently non-streaming Chat Completions return their completed response directly
|
||||
|
||||
A route module has the same five pieces, in the order Python runs them. `types.rs` holds the call, the provider request, and the response. `prepare.rs` resolves the provider and credentials and shapes the request (Python's `validate_environment`, `get_complete_url`, `transform_request`). `handler.rs` resolves auth, offers the wire request to `litellm_host::hooks::RouteHooks::before_send`, sends it, reports the raw response through `emit`, and normalizes the response or stream (`pre_call`, `post`, `post_call`, `transform_response`). `mod.rs` exposes the entrypoint that runs prepare then handler with no hooks (`()`). `route.rs`, where a host needs it, wraps the same two calls in a `CallMachine` whose `HostChannel` is the hooks, and pumps a stream through `open` and `deliver`. A handler takes `&impl RouteHooks<Error>` and never a `HostChannel` directly, so it runs without a coroutine. Keep provider transport and transformation details out of the machine driver
|
||||
Hosts assemble route objects from shared `CoreResources`, HTTP settings, and secret sources. Each route owns its provider client and authentication dependencies. Gateway routes live for the gateway lifetime; Python assembles routes per call from its settings snapshot
|
||||
|
||||
Chat Completions, Messages, and OCR execute through their route objects. Calls pass `RouteHooks` directly; use `&()` when no hooks are needed. Construction does no work; preparation and lifecycle observation begin when the future is polled. Handlers accept `RouteHooks`, never a concrete `ChannelHooks`. Native observers receive start and terminal events through the shared call runner; a stream retains its lifecycle until exhaustion, error, or drop. A host channel has no native observer because its driver owns terminal dispatch
|
||||
|
||||
`route.rs` declares the concrete `Protocol` and implements a route method that accepts a typed request and constructs a `litellm_host::call::HostedMachine` with `hosted_call`. The shared call plumbing owns stream opening, delivery, backpressure, and detachment. Request decoding belongs to the boundary before the machine starts. Route closures only supply execution dependencies and route-specific host capabilities such as an OCR token provider. Use `run_hosted` for a native host so detachment is reported as cancellation. Python uses its own shared driver and preserves caller-task callback execution
|
||||
|
||||
Responses WebSocket sessions remain separate from the HTTP call driver because a connection can accept multiple requests while receiving events
|
||||
|
||||
## Crate layering
|
||||
|
||||
For Messages, Responses, Chat Completions, OCR, and other API formats, `core/src/<format>/` owns orchestration. Shared API data contracts belong in `litellm-types`, adapter contracts and shared transformation machinery in `llms/src/base_llm/<format>/`, and provider policy in `llms/src/<provider>/<format>/`. A repeated format directory name does not imply interchangeable responsibilities. Select concrete adapters here, then invoke their contracts instead of applying one provider's policy to every call. Route types describe call envelopes and execution state, not duplicate public payload schemas
|
||||
|
||||
Each crate mirrors one top-level Python package, so a Rust path reads as its Python path with the crate name in place of the package directory. Dependencies only point down:
|
||||
|
||||
- `litellm-types` mirrors `litellm/types/`: pure serde data, no I/O
|
||||
|
|
|
|||
|
|
@ -18,12 +18,11 @@ litellm-auth-aws.workspace = true
|
|||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
tracing.workspace = true
|
||||
moka.workspace = true
|
||||
mime_guess = "2.0.5"
|
||||
rand.workspace = true
|
||||
reqwest.workspace = true
|
||||
rustls.workspace = true
|
||||
rustls-native-certs.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json = { workspace = true, features = ["preserve_order"] }
|
||||
strum.workspace = true
|
||||
|
|
|
|||
|
|
@ -15,9 +15,9 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderAudioTranscriptionRequest,
|
||||
) -> Result<Value, Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let env_lookup = |key: &str| request.secrets.get(key);
|
||||
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
|
||||
let response = crate::outbound::outbound_request(
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
authenticated,
|
||||
request.url.clone(),
|
||||
&request.body,
|
||||
|
|
@ -26,12 +26,12 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
.timeout
|
||||
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
|
||||
),
|
||||
)?
|
||||
.send(http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
let status = response.status();
|
||||
let text = response.text().await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
|
|
@ -42,8 +42,12 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
body: truncate_error_body(&text),
|
||||
}));
|
||||
}
|
||||
let response_json = serde_json::from_str(&text)
|
||||
.map_err(|error| Error::InvalidResponse(format!("invalid audio response JSON: {error}")))?;
|
||||
let response_json = serde_json::from_str(&text).map_err(|error| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"audio response JSON",
|
||||
error,
|
||||
))
|
||||
})?;
|
||||
Ok(request
|
||||
.config
|
||||
.transform_audio_transcription_response(&request.model, response_json)?
|
||||
|
|
|
|||
|
|
@ -3,18 +3,52 @@ pub use crate::error::RouteError as Error;
|
|||
mod handler;
|
||||
mod prepare;
|
||||
pub use handler::execute_audio_transcription_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
pub use prepare::prepare_audio_transcription_provider_call;
|
||||
use serde_json::Value;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::audio_transcription::types::AudioTranscriptionRequest;
|
||||
|
||||
pub async fn audio_transcription(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
) -> Result<Value, Error> {
|
||||
let request = prepare_audio_transcription_provider_call(request)?;
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
execute_audio_transcription_provider_call(&http, &resources.auth, request).await
|
||||
#[derive(Clone)]
|
||||
pub struct AudioTranscriptionRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl AudioTranscriptionRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "audio_transcription",
|
||||
model = request.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub async fn execute(&self, request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let request =
|
||||
prepare_audio_transcription_provider_call(request, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.model, &request.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<Value, Error>> = Box::pin(
|
||||
execute_audio_transcription_provider_call(&self.http, &self.auth, request),
|
||||
);
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,52 +1,55 @@
|
|||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_http::request::string_headers;
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::{
|
||||
base_llm::{
|
||||
audio_transcription::transformation::BaseAudioTranscriptionConfig,
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
auth::ValidatedEnvironment,
|
||||
},
|
||||
bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
use super::Error;
|
||||
use crate::audio_transcription::types::{
|
||||
AudioTranscriptionRequest, ProviderAudioTranscriptionRequest,
|
||||
};
|
||||
use crate::provider::{LlmProviders, resolve_llm_provider};
|
||||
|
||||
fn provider_config(provider: &str) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
if provider == "bedrock" {
|
||||
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);
|
||||
fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
match provider {
|
||||
LlmProviders::Bedrock => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
let _ = provider;
|
||||
None
|
||||
}
|
||||
|
||||
pub fn prepare_audio_transcription_provider_call(
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub async fn prepare_audio_transcription_provider_call(
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderAudioTranscriptionRequest, Error> {
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
.or_else(|| {
|
||||
request
|
||||
.custom_llm_provider
|
||||
.map(|provider| CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for audio transcription request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let provider_info = resolve_llm_provider(
|
||||
request.model,
|
||||
request.custom_llm_provider,
|
||||
"audio transcription",
|
||||
)?;
|
||||
let model = provider_info.model.to_string();
|
||||
let config = provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let config = provider_config(provider_info.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?;
|
||||
let snapshot = secrets.resolve(&config.secret_names()).await?;
|
||||
let env_lookup = |key: &str| snapshot.get(key);
|
||||
let forwarded = string_headers("audio transcription", request.extra_headers)?;
|
||||
let validated =
|
||||
config.validate_environment(forwarded, &model, &request.optional_params, &env_lookup)?;
|
||||
let environment = ValidatedEnvironment {
|
||||
headers: with_default_headers(validated.headers, &[("Content-Type", "application/json")]),
|
||||
headers: with_default_headers(validated.headers, config.default_headers()),
|
||||
auth: validated.auth,
|
||||
};
|
||||
let url = config.get_complete_url(
|
||||
|
|
@ -60,11 +63,12 @@ pub fn prepare_audio_transcription_provider_call(
|
|||
config.transform_audio_transcription_request(&model, request.audio, filtered_params)?;
|
||||
Ok(ProviderAudioTranscriptionRequest {
|
||||
model,
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
|
||||
config,
|
||||
url,
|
||||
body: transformed.body,
|
||||
environment,
|
||||
secrets: snapshot,
|
||||
timeout: request.timeout,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_secrets::source::Secrets;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_llms::base_llm::{
|
||||
|
|
@ -24,6 +25,7 @@ pub struct ProviderAudioTranscriptionRequest {
|
|||
pub url: String,
|
||||
pub body: Value,
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub secrets: Secrets,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -3,18 +3,43 @@ use litellm_llms::{
|
|||
anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
|
||||
base_llm::chat::transformation::BaseConfig,
|
||||
bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG,
|
||||
openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
const HEADER_CONTEXT: &str = "chat completions";
|
||||
|
||||
pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'static dyn BaseConfig> {
|
||||
pub(super) enum ChatProvider {
|
||||
Anthropic,
|
||||
Bedrock,
|
||||
OpenaiLike,
|
||||
}
|
||||
|
||||
impl ChatProvider {
|
||||
pub(super) fn config(self) -> &'static dyn BaseConfig {
|
||||
match self {
|
||||
Self::Anthropic => &ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
|
||||
Self::Bedrock => &BEDROCK_CHAT_COMPLETIONS_CONFIG,
|
||||
Self::OpenaiLike => &OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option<ChatProvider> {
|
||||
match provider {
|
||||
"anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG),
|
||||
"bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG),
|
||||
_ => None,
|
||||
LlmProviders::Anthropic => Some(ChatProvider::Anthropic),
|
||||
LlmProviders::Bedrock => Some(ChatProvider::Bedrock),
|
||||
LlmProviders::OpenaiLike => Some(ChatProvider::OpenaiLike),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::AzureAi
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ pub(super) async fn execute(
|
|||
body,
|
||||
optional_params,
|
||||
environment,
|
||||
secrets,
|
||||
timeout,
|
||||
api_key,
|
||||
} = request;
|
||||
|
|
@ -43,9 +44,9 @@ pub(super) async fn execute(
|
|||
secret_fields: Vec::new(),
|
||||
api_key,
|
||||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let authenticated = resolve_auth(auth, environment, &|key| secrets.get(key)).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
|
|
@ -64,7 +65,7 @@ pub(super) async fn execute(
|
|||
timeout,
|
||||
)?;
|
||||
|
||||
let response = outbound.send(http).await.map_err(|err| {
|
||||
let response = crate::outbound::send(outbound, http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
|
|
@ -87,13 +88,17 @@ pub(super) async fn execute(
|
|||
}));
|
||||
}
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
.on_event(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
|
||||
let body: Value = serde_json::from_str(&text).map_err(|err| {
|
||||
Error::InvalidResponse(format!("invalid chat completions response JSON: {err}"))
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_response(&model, ProviderChatResponseData { body })
|
||||
|
|
@ -114,7 +119,7 @@ pub(super) fn as_response_error(err: Error) -> Error {
|
|||
match err {
|
||||
already @ (Error::InvalidResponse(_)
|
||||
| Error::Transport(litellm_http::transport::Error::Http { .. })) => already,
|
||||
other => Error::InvalidResponse(other.to_string()),
|
||||
other => Error::InvalidResponse(other.to_string().into()),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -164,7 +169,7 @@ mod tests {
|
|||
}
|
||||
|
||||
impl RouteHooks<Error> for RecordingHooks {
|
||||
async fn before_send(
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
|
|
@ -183,7 +188,7 @@ mod tests {
|
|||
})
|
||||
}
|
||||
|
||||
async fn emit(&self, event: MachineEvent) -> Result<(), Error> {
|
||||
async fn on_event(&self, event: MachineEvent) -> Result<(), Error> {
|
||||
let MachineEvent::ResponseReceived { raw } = event;
|
||||
self.raw.lock().unwrap().push(raw.body);
|
||||
Ok(())
|
||||
|
|
@ -203,6 +208,7 @@ mod tests {
|
|||
timeout: None,
|
||||
})
|
||||
.unwrap(),
|
||||
std::sync::Arc::new(|_: &str| None),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
|
@ -235,7 +241,7 @@ mod tests {
|
|||
assert_eq!(request.headers["x-host"], "seen");
|
||||
assert_eq!(request.headers["x-api-key"], "sk-test");
|
||||
let [context] = <[RequestContext; 1]>::try_from(hooks.contexts.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
.unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len()));
|
||||
assert_eq!(
|
||||
(context.model.as_str(), context.custom_llm_provider.as_str()),
|
||||
("claude-sonnet-4-5", "anthropic")
|
||||
|
|
@ -270,12 +276,12 @@ mod tests {
|
|||
assert!(hooks.raw.into_inner().unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
|
||||
for original in [
|
||||
Error::MissingField("usage"),
|
||||
Error::Unsupported("non-text response content block"),
|
||||
Error::InvalidRequest("whatever".to_string()),
|
||||
Error::InvalidRequest("whatever".to_string().into()),
|
||||
Error::Auth(litellm_auth::Error::InvalidHeader),
|
||||
] {
|
||||
let label = format!("{original:?}");
|
||||
|
|
|
|||
|
|
@ -1,56 +1,72 @@
|
|||
//! The `/chat/completions` call, the Rust equivalent of Python's
|
||||
//! `litellm.completion()`.
|
||||
//!
|
||||
//! [`chat_completions`] is the top-level entrypoint: give it a model, the
|
||||
//! OpenAI-shaped message list, the provider-mapped optional params, and
|
||||
//! credentials, and it resolves the provider, translates the conversation,
|
||||
//! calls the provider, and returns a typed OpenAI-shaped response.
|
||||
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
pub use crate::error::RouteError as Error;
|
||||
mod common_utils;
|
||||
pub(crate) mod handler;
|
||||
mod prepare;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use prepare::{parse_messages, prepare_provider_request, resolve_provider_config, resolve_request};
|
||||
use serde_json::{Map, Value};
|
||||
use prepare::{prepare_provider_request, resolve_request};
|
||||
|
||||
use crate::chat_completions::types::ChatCompletionsRequest;
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub async fn chat_completions(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let request = prepare_provider_request(resolve_request(request)?)?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
#[derive(Clone)]
|
||||
pub struct ChatCompletionsRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
/// touching the network.
|
||||
///
|
||||
/// A host that keeps the Python implementation asks this first so it can emit
|
||||
/// its pre-call logging exactly once, on whichever path is about to run.
|
||||
/// Returns the decline reason, or `None` when the request is accepted.
|
||||
pub fn chat_completions_decline_reason(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
messages: Value,
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Option<&'static str> {
|
||||
let Ok(resolved) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
return Some("provider is not on the rust chat completions path");
|
||||
};
|
||||
let config = resolved.config;
|
||||
let Ok(messages) = parse_messages(messages) else {
|
||||
return Some("unreadable message list");
|
||||
};
|
||||
if messages.is_empty() {
|
||||
return Some("empty message list");
|
||||
impl ChatCompletionsRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "chat_completions",
|
||||
model = %request.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let resolved = resolve_request(request)?;
|
||||
let snapshot = self
|
||||
.secrets
|
||||
.resolve(&resolved.config.secret_names())
|
||||
.await?;
|
||||
let prepared = prepare_provider_request(resolved, snapshot)?;
|
||||
crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<
|
||||
'_,
|
||||
Result<ChatCompletionsResponse, Error>,
|
||||
> = Box::pin(handler::execute(&self.http, &self.auth, prepared, hooks));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
config
|
||||
.unsupported_reason(&messages, optional_params)
|
||||
.map(|reason| reason.0)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::base_llm::{
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
chat::transformation::BaseConfig,
|
||||
};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::base_llm::{auth::ValidatedEnvironment, chat::transformation::BaseConfig};
|
||||
use litellm_secrets::source::Secrets;
|
||||
use litellm_types::llms::openai::ChatMessage;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::{chat_completions_provider_config, string_headers},
|
||||
common_utils::{chat_completions_provider, string_headers},
|
||||
};
|
||||
use crate::chat_completions::types::{
|
||||
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
|
||||
};
|
||||
use crate::provider::resolve_llm_provider;
|
||||
|
||||
pub(super) struct ResolvedProvider {
|
||||
pub(super) model: String,
|
||||
|
|
@ -25,30 +25,24 @@ pub(super) fn resolve_provider_config<'a>(
|
|||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for chat completions request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
let provider_info = resolve_llm_provider(model, custom_llm_provider, "chat completions")?;
|
||||
let config = chat_completions_provider(provider_info.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(provider_info.provider).to_string()))?
|
||||
.config();
|
||||
Ok(ResolvedProvider {
|
||||
model: provider_info.model.to_string(),
|
||||
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
|
||||
custom_llm_provider: <&str>::from(provider_info.provider).to_string(),
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
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}")))
|
||||
serde_json::from_value(messages).map_err(|err| {
|
||||
Error::InvalidRequest(litellm_llms::ErrorDetail::invalid(
|
||||
"chat completions messages",
|
||||
err,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn resolve_request(
|
||||
|
|
@ -62,7 +56,7 @@ pub(super) fn resolve_request(
|
|||
let messages = parse_messages(request.messages)?;
|
||||
if messages.is_empty() {
|
||||
return Err(Error::InvalidRequest(
|
||||
"chat completions requires at least one message".to_string(),
|
||||
"chat completions requires at least one message".into(),
|
||||
));
|
||||
}
|
||||
if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) {
|
||||
|
|
@ -85,8 +79,9 @@ fn validate_environment(
|
|||
request: &ResolvedChatCompletionsRequest<'_>,
|
||||
model: &str,
|
||||
config: &dyn BaseConfig,
|
||||
secrets: &dyn Lookup,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
let forwarded = string_headers(request.extra_headers.clone())?;
|
||||
let validated = config.validate_environment(
|
||||
forwarded,
|
||||
|
|
@ -101,13 +96,16 @@ fn validate_environment(
|
|||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
secrets: Secrets,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
let environment = validate_environment(&request, &request.model, request.config)?;
|
||||
let environment =
|
||||
validate_environment(&request, &request.model, request.config, secrets.as_ref())?;
|
||||
let model = request.model;
|
||||
let config = request.config;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
let url = config.get_complete_url(
|
||||
request.api_base,
|
||||
&model,
|
||||
|
|
@ -125,6 +123,7 @@ pub(super) fn prepare_provider_request(
|
|||
body: transformed.body,
|
||||
optional_params: request.optional_params,
|
||||
environment,
|
||||
secrets,
|
||||
timeout: request.timeout,
|
||||
api_key: request.api_key.map(|key| SecretValue::new(key.to_string())),
|
||||
})
|
||||
|
|
@ -145,7 +144,10 @@ mod tests {
|
|||
fn prepare_chat_completions_call(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
prepare_provider_request(resolve_request(request)?)
|
||||
prepare_provider_request(
|
||||
resolve_request(request)?,
|
||||
std::sync::Arc::new(|_: &str| None),
|
||||
)
|
||||
}
|
||||
|
||||
/// The headers as they go on the wire, credential applied.
|
||||
|
|
@ -183,13 +185,10 @@ mod tests {
|
|||
}
|
||||
}
|
||||
|
||||
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
|
||||
/// carry resolved credentials), so unwrap the failure case by hand.
|
||||
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
|
||||
match prepare_chat_completions_call(request) {
|
||||
Err(error) => error,
|
||||
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
|
||||
}
|
||||
fn preparation_error(request: ChatCompletionsRequest<'_>) -> Error {
|
||||
prepare_chat_completions_call(request)
|
||||
.err()
|
||||
.expect("request preparation should fail")
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -335,24 +334,19 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_an_unsupported_request_before_resolving_credentials() {
|
||||
let mut call = request(
|
||||
"claude-sonnet-4-5",
|
||||
Some("anthropic"),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
);
|
||||
call.api_key = None;
|
||||
// No api_key is set and no env is consulted: the gate must run first, so the
|
||||
// error is the decline rather than a missing-credential error.
|
||||
assert_eq!(decline(call), Error::Unsupported("streaming"));
|
||||
#[rstest::rstest]
|
||||
fn rejects_empty_messages_before_resolving_credentials() {
|
||||
let call = ChatCompletionsRequest {
|
||||
api_key: None,
|
||||
..request("claude-sonnet-4-5", Some("anthropic"), json!([]), json!({}))
|
||||
};
|
||||
assert!(matches!(preparation_error(call), Error::InvalidRequest(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_an_unknown_provider() {
|
||||
assert_eq!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"openai/gpt-4o",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -365,7 +359,7 @@ mod tests {
|
|||
#[test]
|
||||
fn rejects_a_model_with_no_resolvable_provider() {
|
||||
assert!(matches!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -378,16 +372,20 @@ mod tests {
|
|||
#[test]
|
||||
fn rejects_an_empty_or_malformed_message_list() {
|
||||
assert_eq!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([]),
|
||||
json!({}),
|
||||
)),
|
||||
Error::InvalidRequest("chat completions requires at least one message".to_string())
|
||||
Error::InvalidRequest(
|
||||
"chat completions requires at least one message"
|
||||
.to_string()
|
||||
.into()
|
||||
)
|
||||
);
|
||||
assert!(matches!(
|
||||
decline(request(
|
||||
preparation_error(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("not a list"),
|
||||
|
|
@ -407,7 +405,7 @@ mod tests {
|
|||
);
|
||||
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
|
||||
assert_eq!(
|
||||
decline(call),
|
||||
preparation_error(call),
|
||||
Error::Headers(litellm_http::request::HeaderError {
|
||||
context: "chat completions",
|
||||
name: "x-trace".to_string(),
|
||||
|
|
@ -505,18 +503,17 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::authorization("Authorization")]
|
||||
#[case::amz_date("x-amz-date")]
|
||||
#[case::security_token("x-amz-security-token")]
|
||||
#[case::date("Date")]
|
||||
#[tokio::test]
|
||||
async fn a_forwarded_header_the_signer_computes_declines_to_python() {
|
||||
// Reattaching the caller's copy next to the computed one puts the name on
|
||||
// the wire twice and Bedrock rejects the pair, so a request carrying one
|
||||
// has to go to Python instead of being signed here.
|
||||
for forwarded in [
|
||||
"Authorization",
|
||||
"x-amz-date",
|
||||
"x-amz-security-token",
|
||||
"Date",
|
||||
] {
|
||||
let mut call = request(
|
||||
async fn rejects_a_forwarded_header_the_signer_computes(#[case] forwarded: &str) {
|
||||
let call = ChatCompletionsRequest {
|
||||
api_key: None,
|
||||
extra_headers: Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])),
|
||||
..request(
|
||||
"bedrock/us-east-1/anthropic.claude-v2",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -525,29 +522,27 @@ mod tests {
|
|||
"aws_access_key_id": "AKIDEXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
|
||||
}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment,
|
||||
&|_| None,
|
||||
)
|
||||
.await
|
||||
.expect("resolves");
|
||||
let error = crate::chat_completions::handler::outbound_request(
|
||||
authenticated,
|
||||
prepared.url,
|
||||
&prepared.body,
|
||||
prepared.timeout,
|
||||
)
|
||||
.expect_err("{forwarded} should decline instead of being signed");
|
||||
assert!(
|
||||
matches!(error, Error::Unsupported(_)),
|
||||
"{forwarded} declined as {error:?}, which the host would not fall back on"
|
||||
);
|
||||
}
|
||||
};
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let authenticated = resolve_auth(
|
||||
&litellm_auth::AuthServices::default(),
|
||||
prepared.environment,
|
||||
&|_| None,
|
||||
)
|
||||
.await
|
||||
.expect("resolves");
|
||||
let error = crate::chat_completions::handler::outbound_request(
|
||||
authenticated,
|
||||
prepared.url,
|
||||
&prepared.body,
|
||||
prepared.timeout,
|
||||
)
|
||||
.expect_err("conflicting signing headers must fail");
|
||||
assert!(
|
||||
matches!(error, Error::Unsupported(_)),
|
||||
"{forwarded} returned {error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -640,112 +635,4 @@ mod tests {
|
|||
"prepare did not carry the bearer token"
|
||||
);
|
||||
}
|
||||
|
||||
fn decline_reason(
|
||||
model: &str,
|
||||
provider: Option<&str>,
|
||||
messages: Value,
|
||||
params: Value,
|
||||
) -> Option<&'static str> {
|
||||
let params = match params {
|
||||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
};
|
||||
crate::chat_completions::chat_completions_decline_reason(model, provider, messages, ¶ms)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_accepts_what_prepare_accepts() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
),
|
||||
Some("streaming")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"openai/gpt-4o",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("nope"),
|
||||
json!({})
|
||||
),
|
||||
Some("unreadable message list")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
|
||||
Some("empty message list")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
|
||||
// A gate that accepts what prepare then declines would make the host emit
|
||||
// its pre-call logging on a path that falls back, so pin the agreement.
|
||||
for (messages, params) in [
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 8}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
|
||||
json!({"temperature": 0.1}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
|
||||
json!({}),
|
||||
),
|
||||
] {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params.clone()
|
||||
),
|
||||
None,
|
||||
"gate declined {messages}"
|
||||
);
|
||||
prepare_chat_completions_call(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params,
|
||||
))
|
||||
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
44
litellm-rust/crates/core/src/chat_completions/route.rs
Normal file
44
litellm-rust/crates/core/src/chat_completions/route.rs
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
use std::convert::Infallible;
|
||||
|
||||
use litellm_host::{
|
||||
call::{CallOutput, HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
|
||||
use super::{
|
||||
ChatCompletionsRoute, Error,
|
||||
types::{ChatCompletionsCall, ChatCompletionsRequest},
|
||||
};
|
||||
|
||||
pub struct ChatCompletions;
|
||||
|
||||
impl Protocol for ChatCompletions {
|
||||
type Response = ChatCompletionsResponse;
|
||||
type Error = Error;
|
||||
type Request = ChatCompletionsCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Infallible;
|
||||
type StreamHead = Infallible;
|
||||
}
|
||||
|
||||
impl ChatCompletionsRoute {
|
||||
pub fn machine(self, call: ChatCompletionsCall) -> HostedMachine<ChatCompletions> {
|
||||
hosted_call(
|
||||
call,
|
||||
move |call: ChatCompletionsCall, _, hooks| async move {
|
||||
let request = ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers,
|
||||
timeout: call.timeout,
|
||||
};
|
||||
self.run(request, &hooks).await.map(CallOutput::Complete)
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_secrets::source::Secrets;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_auth::SecretValue;
|
||||
|
|
@ -22,6 +23,32 @@ pub struct ChatCompletionsRequest<'a> {
|
|||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub struct ChatCompletionsCall {
|
||||
pub model: String,
|
||||
pub messages: Value,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
impl From<ChatCompletionsRequest<'_>> for ChatCompletionsCall {
|
||||
fn from(request: ChatCompletionsRequest<'_>) -> Self {
|
||||
Self {
|
||||
model: request.model.into(),
|
||||
messages: request.messages,
|
||||
optional_params: request.optional_params,
|
||||
api_key: request.api_key.map(str::to_owned),
|
||||
api_base: request.api_base.map(str::to_owned),
|
||||
custom_llm_provider: request.custom_llm_provider.map(str::to_owned),
|
||||
extra_headers: request.extra_headers,
|
||||
timeout: request.timeout,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ResolvedChatCompletionsRequest<'a> {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
|
|
@ -46,6 +73,7 @@ pub struct ProviderChatCompletionsRequest {
|
|||
/// The forwarded and default headers plus how the call authenticates; the credential
|
||||
/// itself is applied when the request is sent.
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub secrets: Secrets,
|
||||
pub timeout: Option<Duration>,
|
||||
pub api_key: Option<SecretValue>,
|
||||
}
|
||||
|
|
|
|||
324
litellm-rust/crates/core/src/diagnostic.rs
Normal file
324
litellm-rust/crates/core/src/diagnostic.rs
Normal file
|
|
@ -0,0 +1,324 @@
|
|||
use std::{
|
||||
future::Future,
|
||||
pin::Pin,
|
||||
task::{Context, Poll},
|
||||
};
|
||||
|
||||
use futures_util::{Stream, stream::BoxStream};
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_tracing::Logger;
|
||||
use tracing::Span;
|
||||
|
||||
struct Completion {
|
||||
span: Span,
|
||||
outcome: &'static str,
|
||||
}
|
||||
|
||||
impl Completion {
|
||||
fn new(name: &str) -> Self {
|
||||
let current = Span::current();
|
||||
Self {
|
||||
span: if current
|
||||
.metadata()
|
||||
.is_some_and(|metadata| metadata.name() == name)
|
||||
{
|
||||
current
|
||||
} else {
|
||||
Span::none()
|
||||
},
|
||||
outcome: "cancelled",
|
||||
}
|
||||
}
|
||||
|
||||
fn finish(mut self, outcome: &'static str) {
|
||||
self.outcome = outcome;
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Completion {
|
||||
fn drop(&mut self) {
|
||||
self.span.record("outcome", self.outcome);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn provider(model: &str, provider: &str) {
|
||||
let span = Span::current();
|
||||
span.record("resolved_model", model);
|
||||
span.record("provider", provider);
|
||||
}
|
||||
|
||||
pub(crate) async fn unary<R, E>(execute: impl Future<Output = Result<R, E>>) -> Result<R, E> {
|
||||
operation("litellm.route", execute).await
|
||||
}
|
||||
|
||||
pub(crate) async fn operation<R, E>(
|
||||
name: &str,
|
||||
execute: impl Future<Output = Result<R, E>>,
|
||||
) -> Result<R, E> {
|
||||
let completion = Completion::new(name);
|
||||
let result = execute.await;
|
||||
completion.finish(if result.is_ok() { "success" } else { "failure" });
|
||||
result
|
||||
}
|
||||
|
||||
pub(crate) async fn call<R, H, C, E>(
|
||||
execute: impl Future<Output = Result<CallOutput<R, H, C, E>, E>>,
|
||||
) -> Result<CallOutput<R, H, C, E>, E>
|
||||
where
|
||||
C: Send + 'static,
|
||||
E: Send + 'static,
|
||||
{
|
||||
let completion = Completion::new("litellm.route");
|
||||
match execute.await {
|
||||
Err(error) => {
|
||||
completion.finish("failure");
|
||||
Err(error)
|
||||
}
|
||||
Ok(CallOutput::Complete(response)) => {
|
||||
completion.span.record("stream", false);
|
||||
completion.finish("success");
|
||||
Ok(CallOutput::Complete(response))
|
||||
}
|
||||
Ok(CallOutput::Stream { head, chunks }) => {
|
||||
completion.span.record("stream", true);
|
||||
Ok(CallOutput::Stream {
|
||||
head,
|
||||
chunks: Box::pin(TracedStream {
|
||||
state: Some(StreamState {
|
||||
chunks,
|
||||
completion,
|
||||
logger: Logger::current(),
|
||||
}),
|
||||
}),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct StreamState<C, E> {
|
||||
chunks: BoxStream<'static, Result<C, E>>,
|
||||
completion: Completion,
|
||||
logger: Logger,
|
||||
}
|
||||
|
||||
impl<C, E> StreamState<C, E> {
|
||||
fn close(self, outcome: &'static str) {
|
||||
self.logger
|
||||
.scope(|| self.completion.span.in_scope(|| drop(self.chunks)));
|
||||
self.completion.finish(outcome);
|
||||
}
|
||||
}
|
||||
|
||||
struct TracedStream<C, E> {
|
||||
state: Option<StreamState<C, E>>,
|
||||
}
|
||||
|
||||
impl<C, E> Stream for TracedStream<C, E> {
|
||||
type Item = Result<C, E>;
|
||||
|
||||
fn poll_next(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
let Some(state) = self.state.as_mut() else {
|
||||
return Poll::Ready(None);
|
||||
};
|
||||
let next = state.logger.scope(|| {
|
||||
state
|
||||
.completion
|
||||
.span
|
||||
.in_scope(|| state.chunks.as_mut().poll_next(context))
|
||||
});
|
||||
let outcome = match &next {
|
||||
Poll::Ready(None) => "success",
|
||||
Poll::Ready(Some(Err(_))) => "failure",
|
||||
_ => return next,
|
||||
};
|
||||
if let Some(state) = self.state.take() {
|
||||
state.close(outcome);
|
||||
}
|
||||
next
|
||||
}
|
||||
}
|
||||
|
||||
impl<C, E> Drop for TracedStream<C, E> {
|
||||
fn drop(&mut self) {
|
||||
if let Some(state) = self.state.take() {
|
||||
state.close("cancelled");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{sync::mpsc, task::Context};
|
||||
|
||||
use futures_util::{StreamExt, task::noop_waker_ref};
|
||||
use litellm_tracing::{Metadata, Record, Sink};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
struct Capture(mpsc::Sender<Value>);
|
||||
|
||||
impl Sink for Capture {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
|
||||
*metadata.level() <= tracing::Level::INFO
|
||||
}
|
||||
fn emit(&self, record: &Record) {
|
||||
self.0.send(Value::Object(record.fields.clone())).unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn logger() -> (Logger, mpsc::Receiver<Value>) {
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
(Logger::new(Capture(sender)), receiver)
|
||||
}
|
||||
|
||||
struct Chunks(std::vec::IntoIter<Result<u8, &'static str>>);
|
||||
|
||||
impl Stream for Chunks {
|
||||
type Item = Result<u8, &'static str>;
|
||||
fn poll_next(mut self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Option<Self::Item>> {
|
||||
tracing::info!(event = "poll");
|
||||
Poll::Ready(self.0.next())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for Chunks {
|
||||
fn drop(&mut self) {
|
||||
tracing::info!(event = "drop");
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.route",
|
||||
skip_all,
|
||||
fields(route = "fixture", stream, outcome)
|
||||
)]
|
||||
async fn streamed() -> Result<CallOutput<(), (), u8, &'static str>, &'static str> {
|
||||
call(async {
|
||||
Ok(CallOutput::Stream {
|
||||
head: (),
|
||||
chunks: Box::pin(Chunks(vec![Ok(1), Err("broken"), Ok(2)].into_iter())),
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn stream_errors_finish_once_and_poll_and_drop_use_the_captured_context(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let CallOutput::Stream { mut chunks, .. } = logger.instrument(streamed()).await.unwrap()
|
||||
else {
|
||||
panic!()
|
||||
};
|
||||
assert!(records.try_recv().is_err());
|
||||
tokio::spawn(async move {
|
||||
assert_eq!(chunks.next().await, Some(Ok(1)));
|
||||
assert_eq!(chunks.next().await, Some(Err("broken")));
|
||||
assert_eq!(chunks.next().await, None);
|
||||
let emitted = records.try_iter().collect::<Vec<_>>();
|
||||
assert_eq!(emitted.len(), 4);
|
||||
assert!(emitted.iter().all(|record| record["route"] == "fixture"));
|
||||
assert_eq!(emitted[0]["event"], "poll");
|
||||
assert_eq!(emitted[1]["event"], "poll");
|
||||
assert_eq!(emitted[2]["event"], "drop");
|
||||
assert_eq!(emitted[3]["outcome"], "failure");
|
||||
drop(chunks);
|
||||
assert!(records.try_recv().is_err());
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(route = "waiting", outcome))]
|
||||
async fn waiting(streaming: bool) {
|
||||
if streaming {
|
||||
let _: Result<CallOutput<(), (), u8, ()>, ()> = call(std::future::pending()).await;
|
||||
} else {
|
||||
let _: Result<(), ()> = unary(std::future::pending()).await;
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unary(false)]
|
||||
#[case::streaming(true)]
|
||||
fn cancellation_before_headers_closes_the_span(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
#[case] streaming: bool,
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let mut future = Box::pin(logger.instrument(waiting(streaming)));
|
||||
assert!(
|
||||
future
|
||||
.as_mut()
|
||||
.poll(&mut Context::from_waker(noop_waker_ref()))
|
||||
.is_pending()
|
||||
);
|
||||
assert!(records.try_recv().is_err());
|
||||
drop(future);
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["outcome"], "cancelled");
|
||||
assert_eq!(summary["route"], "waiting");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn dropped_stream_teardown_uses_its_original_logger(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
let output = logger.instrument(streamed()).await.unwrap();
|
||||
Logger::default().scope(|| drop(output));
|
||||
let emitted = records.try_iter().collect::<Vec<_>>();
|
||||
assert_eq!(emitted.len(), 2);
|
||||
assert_eq!(
|
||||
emitted[0],
|
||||
json!({"route":"fixture", "stream":true, "event":"drop"})
|
||||
);
|
||||
assert_eq!(emitted[1]["outcome"], "cancelled");
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "disabled", level = "debug", skip_all, fields(outcome))]
|
||||
async fn disabled_child() {
|
||||
let _: Result<(), ()> = operation("disabled", async { Err(()) }).await;
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_filtered_operation_does_not_overwrite_its_parent_outcome(
|
||||
logger: (Logger, mpsc::Receiver<Value>),
|
||||
) {
|
||||
let (logger, records) = logger;
|
||||
logger
|
||||
.instrument(async {
|
||||
let parent = tracing::info_span!("parent", outcome = "original");
|
||||
tracing::Instrument::instrument(disabled_child(), parent).await;
|
||||
})
|
||||
.await;
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["span_name"], "parent");
|
||||
assert_eq!(summary["outcome"], "original");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(stream = true, outcome))]
|
||||
async fn completed() -> Result<CallOutput<(), (), u8, ()>, ()> {
|
||||
call(async { Ok(CallOutput::Complete(())) }).await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn streaming_mode_reflects_the_returned_output(logger: (Logger, mpsc::Receiver<Value>)) {
|
||||
let (logger, records) = logger;
|
||||
logger.instrument(completed()).await.unwrap();
|
||||
let summary = records.try_recv().unwrap();
|
||||
assert_eq!(summary["stream"], false);
|
||||
assert_eq!(summary["outcome"], "success");
|
||||
assert!(records.try_recv().is_err());
|
||||
}
|
||||
}
|
||||
|
|
@ -22,9 +22,9 @@ pub enum RouteError {
|
|||
#[error("invalid provider: {0}")]
|
||||
InvalidProvider(String),
|
||||
#[error("invalid request: {0}")]
|
||||
InvalidRequest(String),
|
||||
InvalidRequest(#[source] litellm_llms::ErrorDetail),
|
||||
#[error("invalid response: {0}")]
|
||||
InvalidResponse(String),
|
||||
InvalidResponse(#[source] litellm_llms::ErrorDetail),
|
||||
#[error("unsupported by the rust path: {0}")]
|
||||
Unsupported(&'static str),
|
||||
#[error(transparent)]
|
||||
|
|
@ -37,34 +37,23 @@ pub enum RouteError {
|
|||
Http(#[from] litellm_http::Error),
|
||||
#[error(transparent)]
|
||||
Secret(#[from] SecretError),
|
||||
#[error("post-call hook failed: {0}")]
|
||||
PostCallHook(#[source] Arc<RouteError>),
|
||||
}
|
||||
|
||||
/// Whether the provider had already been called when the route failed. Before the send, a
|
||||
/// host may retry on another path; after it, the provider has done the work and billed for it.
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Phase {
|
||||
BeforeSend,
|
||||
AfterSend,
|
||||
impl From<litellm_host::machine::MachineFault> for RouteError {
|
||||
fn from(fault: litellm_host::machine::MachineFault) -> Self {
|
||||
use litellm_host::machine::MachineFault;
|
||||
Self::InvalidRequest(match fault {
|
||||
MachineFault::Abandoned => "host driver was abandoned".into(),
|
||||
MachineFault::Protocol(message) => format!("host {message}").into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl RouteError {
|
||||
pub fn phase(&self) -> Phase {
|
||||
match self {
|
||||
Self::InvalidResponse(_)
|
||||
| Self::Transport(TransportError::Http { .. } | TransportError::Network(_)) => {
|
||||
Phase::AfterSend
|
||||
}
|
||||
Self::Transport(TransportError::Connect(_))
|
||||
| Self::InvalidType { .. }
|
||||
| Self::MissingField(_)
|
||||
| Self::InvalidProvider(_)
|
||||
| Self::InvalidRequest(_)
|
||||
| Self::Unsupported(_)
|
||||
| Self::Auth(_)
|
||||
| Self::Headers(_)
|
||||
| Self::Http(_)
|
||||
| Self::Secret(_) => Phase::BeforeSend,
|
||||
}
|
||||
pub(crate) fn post_call(error: Self) -> Self {
|
||||
Self::PostCallHook(Arc::new(error))
|
||||
}
|
||||
|
||||
/// The caller's request is what is wrong, as opposed to the environment, the wire, or
|
||||
|
|
@ -78,9 +67,11 @@ impl RouteError {
|
|||
| Self::Unsupported(_)
|
||||
| Self::Headers(_) => true,
|
||||
Self::Auth(error) => !matches!(error, litellm_auth::Error::MissingApiKey { .. }),
|
||||
Self::InvalidResponse(_) | Self::Transport(_) | Self::Http(_) | Self::Secret(_) => {
|
||||
false
|
||||
}
|
||||
Self::InvalidResponse(_)
|
||||
| Self::Transport(_)
|
||||
| Self::Http(_)
|
||||
| Self::Secret(_)
|
||||
| Self::PostCallHook(_) => false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -124,31 +115,9 @@ impl Eq for SecretError {}
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{Phase, RouteError};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
|
||||
#[test]
|
||||
fn only_a_provider_answer_or_a_lost_connection_counts_as_after_send() {
|
||||
let after = [
|
||||
RouteError::InvalidResponse("bad json".into()),
|
||||
RouteError::Transport(TransportError::Http {
|
||||
status: 500,
|
||||
body: "boom".into(),
|
||||
}),
|
||||
RouteError::Transport(TransportError::Network("reset".into())),
|
||||
];
|
||||
for error in after {
|
||||
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
|
||||
}
|
||||
let before = [
|
||||
RouteError::Transport(TransportError::Connect("refused".into())),
|
||||
RouteError::Unsupported("streaming"),
|
||||
RouteError::Auth(litellm_auth::Error::InvalidHeader),
|
||||
];
|
||||
for error in before {
|
||||
assert_eq!(error.phase(), Phase::BeforeSend, "{error:?}");
|
||||
}
|
||||
}
|
||||
use super::RouteError;
|
||||
use litellm_llms::{Error as LlmError, ErrorDetail};
|
||||
use rstest::rstest;
|
||||
|
||||
#[test]
|
||||
fn a_missing_api_key_is_the_environment_not_the_request() {
|
||||
|
|
@ -163,4 +132,29 @@ mod tests {
|
|||
assert!(RouteError::InvalidRequest("top_k".into()).is_request());
|
||||
assert!(!RouteError::InvalidResponse("bad json".into()).is_request());
|
||||
}
|
||||
#[rstest]
|
||||
#[case::request(true)]
|
||||
#[case::response(false)]
|
||||
fn contextual_errors_preserve_sources_and_route_classification(#[case] request: bool) {
|
||||
let source = serde_json::from_str::<serde_json::Value>("{").unwrap_err();
|
||||
let source_message = source.to_string();
|
||||
let detail = ErrorDetail::invalid("test payload", source);
|
||||
let error = RouteError::from(if request {
|
||||
LlmError::InvalidRequest(detail)
|
||||
} else {
|
||||
LlmError::InvalidResponse(detail)
|
||||
});
|
||||
assert_eq!(error.is_request(), request);
|
||||
let category = if request { "request" } else { "response" };
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
format!("invalid {category}: invalid test payload: {source_message}")
|
||||
);
|
||||
let source = std::iter::successors(Some(&error as &dyn std::error::Error), |error| {
|
||||
error.source()
|
||||
})
|
||||
.find_map(|error| error.downcast_ref::<serde_json::Error>())
|
||||
.expect("the original JSON error remains available");
|
||||
assert_eq!(source.to_string(), source_message);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
mod diagnostic;
|
||||
|
||||
pub mod audio_transcription;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
|
|
@ -5,7 +7,8 @@ pub mod error;
|
|||
pub mod messages;
|
||||
pub mod ocr;
|
||||
mod outbound;
|
||||
mod provider;
|
||||
pub mod resources;
|
||||
pub mod responses;
|
||||
|
||||
pub use error::{Phase, RouteError};
|
||||
pub use error::RouteError;
|
||||
|
|
|
|||
7
litellm-rust/crates/core/src/messages/AGENTS.md
Normal file
7
litellm-rust/crates/core/src/messages/AGENTS.md
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
This directory owns provider-independent Messages call orchestration: the entrypoint, call envelopes, provider selection, credential resolution, transport coordination, hooks, and stream lifecycle. Shared API data contracts belong in `litellm-types::messages`, adapter contracts and execution inputs in `llms/src/base_llm/messages`, and provider implementations in `llms/src/<provider>/messages`
|
||||
|
||||
Select concrete provider adapters and invoke their contracts. Delegate authentication policy, beta selection, payload rewriting, and response interpretation to those adapters. Keep provider policy out of request preparation and transport handlers. Calling a concrete provider helper for every provider is still a policy dependency
|
||||
|
||||
Route types such as `MessagesCall`, prepared requests, and response wrappers containing live streams describe execution. Reuse the shared Messages payload types inside them instead of defining another request or response schema here
|
||||
|
||||
Preserve the order of validation, normalization, caller-requested parameter removal, and provider transformation when that order affects observable behavior. Test provider dispatch, auth precedence, header handling, transformations, and responses through behavior, not source structure
|
||||
|
|
@ -2,19 +2,18 @@ use litellm_http::request::string_headers as shared_string_headers;
|
|||
pub(super) use litellm_http::request::truncate_error_body;
|
||||
use litellm_llms::{
|
||||
anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::messages::transformation::BaseAnthropicMessagesConfig,
|
||||
bedrock::messages::invoke_transformations::anthropic_claude3_transformation::BEDROCK_ANTHROPIC_MESSAGES_CONFIG,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
use super::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
const HEADER_CONTEXT: &str = "messages";
|
||||
|
||||
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum MessagesProvider {
|
||||
Anthropic,
|
||||
AzureAi,
|
||||
|
|
@ -23,7 +22,12 @@ pub(crate) enum MessagesProvider {
|
|||
|
||||
impl MessagesProvider {
|
||||
pub(crate) fn as_str(self) -> &'static str {
|
||||
self.into()
|
||||
match self {
|
||||
Self::Anthropic => LlmProviders::Anthropic,
|
||||
Self::AzureAi => LlmProviders::AzureAi,
|
||||
Self::Bedrock => LlmProviders::Bedrock,
|
||||
}
|
||||
.into()
|
||||
}
|
||||
|
||||
pub(crate) fn config(self) -> &'static dyn BaseAnthropicMessagesConfig {
|
||||
|
|
@ -35,6 +39,21 @@ impl MessagesProvider {
|
|||
}
|
||||
}
|
||||
|
||||
pub(crate) fn messages_provider(provider: LlmProviders) -> Option<MessagesProvider> {
|
||||
match provider {
|
||||
LlmProviders::Anthropic => Some(MessagesProvider::Anthropic),
|
||||
LlmProviders::AzureAi => Some(MessagesProvider::AzureAi),
|
||||
LlmProviders::Bedrock => Some(MessagesProvider::Bedrock),
|
||||
LlmProviders::AwsTextract
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Mistral
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn string_headers(
|
||||
extra_headers: Option<Map<String, Value>>,
|
||||
) -> Result<Vec<(String, String)>, Error> {
|
||||
|
|
@ -47,8 +66,9 @@ mod tests {
|
|||
|
||||
use rstest::rstest;
|
||||
|
||||
use super::{MessagesProvider, string_headers, truncate_error_body};
|
||||
use super::{MessagesProvider, messages_provider, string_headers, truncate_error_body};
|
||||
use crate::messages::Error;
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic", MessagesProvider::Anthropic)]
|
||||
|
|
@ -58,13 +78,16 @@ mod tests {
|
|||
#[case] name: &str,
|
||||
#[case] provider: MessagesProvider,
|
||||
) {
|
||||
assert_eq!(name.parse::<MessagesProvider>(), Ok(provider));
|
||||
assert_eq!(
|
||||
messages_provider(name.parse::<LlmProviders>().unwrap()),
|
||||
Some(provider)
|
||||
);
|
||||
assert_eq!(provider.as_str(), name);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_without_a_messages_config_is_rejected() {
|
||||
assert!("openai".parse::<MessagesProvider>().is_err());
|
||||
assert_eq!(messages_provider(LlmProviders::Openai), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -9,13 +9,13 @@ use litellm_host::{
|
|||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_llms::base_llm::{
|
||||
anthropic_messages::{
|
||||
auth::{Authenticated, resolve_auth},
|
||||
messages::{
|
||||
streaming::{ByteStream, StreamDecoder, encode_anthropic_sse},
|
||||
transformation::BaseAnthropicMessagesConfig,
|
||||
},
|
||||
auth::{Authenticated, resolve_auth},
|
||||
};
|
||||
use litellm_tracing::{ByteChunk, debug};
|
||||
use litellm_tracing::ByteChunk;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
|
|
@ -48,7 +48,7 @@ pub(super) async fn execute(
|
|||
};
|
||||
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
|
||||
let wire = hooks
|
||||
.before_send(
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url,
|
||||
headers: authenticated.headers,
|
||||
|
|
@ -58,7 +58,7 @@ pub(super) async fn execute(
|
|||
)
|
||||
.await?;
|
||||
let provider_name = provider.as_str();
|
||||
debug!(provider = provider_name, stream, body = %wire.body, "provider request");
|
||||
log_request_body(provider_name, stream, &wire.body);
|
||||
let response = send(
|
||||
http,
|
||||
Authenticated {
|
||||
|
|
@ -70,11 +70,6 @@ pub(super) async fn execute(
|
|||
timeout,
|
||||
)
|
||||
.await?;
|
||||
debug!(
|
||||
provider = provider_name,
|
||||
status = response.status().as_u16(),
|
||||
"provider response headers"
|
||||
);
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
|
|
@ -87,19 +82,21 @@ pub(super) async fn execute(
|
|||
));
|
||||
}
|
||||
let text = response.text().await.map_err(network)?;
|
||||
debug!(body = text.as_str(), "provider response body");
|
||||
log_response_body(&text);
|
||||
hooks
|
||||
.emit(MachineEvent::ResponseReceived {
|
||||
.on_event(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: text.clone() },
|
||||
})
|
||||
.await?;
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
decode_response(config, &body.model, &text)
|
||||
.map(|message| MessagesResponse::Message(Box::new(message)))
|
||||
.map(|message| MessagesResponse::Complete(Box::new(message)))
|
||||
}
|
||||
|
||||
fn serialize_failure(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
Error::InvalidRequest(litellm_llms::ErrorDetail::failed(
|
||||
"Anthropic messages request serialization",
|
||||
err,
|
||||
))
|
||||
}
|
||||
|
||||
|
|
@ -120,14 +117,14 @@ async fn send(
|
|||
body,
|
||||
Some(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
|
||||
)?;
|
||||
request.send(http).await.map_err(network)
|
||||
crate::outbound::send(request, http).await.map_err(network)
|
||||
}
|
||||
|
||||
async fn provider_error(response: reqwest::Response) -> Error {
|
||||
let status = response.status().as_u16();
|
||||
match response.text().await {
|
||||
Ok(text) => {
|
||||
litellm_tracing::debug!(status, body = text.as_str(), "provider error body");
|
||||
log_error_body(status, &text);
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
|
|
@ -142,8 +139,12 @@ fn decode_response(
|
|||
model: &str,
|
||||
text: &str,
|
||||
) -> Result<AnthropicMessagesResponse, Error> {
|
||||
let response = serde_json::from_str(text)
|
||||
.map_err(|err| Error::InvalidResponse(format!("invalid messages response JSON: {err}")))?;
|
||||
let response = serde_json::from_str(text).map_err(|err| {
|
||||
Error::InvalidResponse(litellm_llms::ErrorDetail::invalid(
|
||||
"messages response JSON",
|
||||
err,
|
||||
))
|
||||
})?;
|
||||
config
|
||||
.transform_anthropic_messages_response(model, response)
|
||||
.map_err(Error::from)
|
||||
|
|
@ -170,7 +171,10 @@ fn streaming_response(
|
|||
.boxed(),
|
||||
Some(decode) => decoded_chunks(response, decode, provider),
|
||||
};
|
||||
MessagesResponse::Stream { headers, chunks }
|
||||
MessagesResponse::Stream {
|
||||
head: super::route::MessagesStreamHead { headers },
|
||||
chunks,
|
||||
}
|
||||
}
|
||||
|
||||
fn decoded_chunks(
|
||||
|
|
@ -194,14 +198,26 @@ fn decoded_chunks(
|
|||
.boxed()
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &Bytes) {
|
||||
fn log_request_body(provider: &str, stream: bool, body: &serde_json::Value) {
|
||||
tracing::debug!(provider, stream, body = %body, "provider request");
|
||||
}
|
||||
|
||||
fn log_response_body(body: &str) {
|
||||
tracing::debug!(body, "provider response body");
|
||||
}
|
||||
|
||||
fn log_error_body(status: u16, body: &str) {
|
||||
tracing::debug!(status, body, "provider error body");
|
||||
}
|
||||
|
||||
fn log_chunk(provider: &str, stage: &str, data: &bytes::Bytes) {
|
||||
let chunk = ByteChunk::new(data);
|
||||
debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
|
||||
tracing::debug!(provider, stage, encoding = chunk.encoding(), chunk = %chunk, "stream chunk");
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::base_llm::anthropic_messages::streaming::anthropic_sse_event_stream;
|
||||
use litellm_llms::base_llm::messages::streaming::anthropic_sse_event_stream;
|
||||
use rstest::rstest;
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any};
|
||||
|
||||
|
|
@ -212,6 +228,7 @@ mod tests {
|
|||
"data: {\"type\":\"ping\"}\n\n",
|
||||
Some("event: ping\ndata: {\"type\":\"ping\"}\n\n")
|
||||
)]
|
||||
#[rstest::rstest]
|
||||
#[case::invalid_event("data: invalid\n\ndata: {\"type\":\"ping\"}\n\n", None)]
|
||||
#[tokio::test]
|
||||
async fn decoded_streams_encode_events_and_stop_at_the_first_error(
|
||||
|
|
|
|||
|
|
@ -1,27 +1,64 @@
|
|||
//! The Anthropic Messages call, the Rust equivalent of Python's `litellm.messages()`.
|
||||
//!
|
||||
//! [`messages`] prepares the provider request and sends it in process. [`route`] runs the
|
||||
//! same two steps as a machine for a host that answers the call's operations itself.
|
||||
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
mod types;
|
||||
|
||||
use litellm_http::{ClientVariant, HttpClientConfig};
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub use crate::error::RouteError as Error;
|
||||
pub use types::{MessagesCall, MessagesResponse, MessagesShaping, messages_body};
|
||||
|
||||
pub async fn messages(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: &dyn SecretSource,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let request = prepare::prepare(call, secrets).await?;
|
||||
handler::execute(&http, &resources.auth, request, &()).await
|
||||
#[derive(Clone)]
|
||||
pub struct MessagesRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl MessagesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "messages",
|
||||
model = %call.body.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = call.body.params.stream == Some(true),
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
call: MessagesCall,
|
||||
hooks: &impl litellm_host::hooks::RouteHooks<Error>,
|
||||
) -> Result<MessagesResponse, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.body.model, request.provider.as_str());
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<MessagesResponse, Error>> =
|
||||
Box::pin(handler::execute(&self.http, &self.auth, request, hooks));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,25 +3,21 @@ use std::time::Duration;
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::{
|
||||
dot_notation_indexing::delete_nested_value,
|
||||
get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider},
|
||||
get_provider_specific_headers::get_provider_specific_headers,
|
||||
settings::Lookup,
|
||||
get_provider_specific_headers::get_provider_specific_headers, settings::Lookup,
|
||||
};
|
||||
use litellm_llms::{
|
||||
anthropic::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::{
|
||||
anthropic_messages::transformation::MessagesTransformContext,
|
||||
auth::{ValidatedEnvironment, with_default_headers},
|
||||
},
|
||||
use litellm_http::request::with_default_headers;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::ValidatedEnvironment, messages::context::MessagesTransformContext,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
|
||||
|
||||
use super::{
|
||||
Error, MessagesCall,
|
||||
common_utils::{MessagesProvider, string_headers},
|
||||
common_utils::{MessagesProvider, messages_provider, string_headers},
|
||||
types::invalid_request,
|
||||
};
|
||||
use crate::provider::resolve_llm_provider;
|
||||
|
||||
struct ResolvedProvider {
|
||||
model: String,
|
||||
|
|
@ -38,6 +34,7 @@ pub(super) struct ProviderMessagesRequest {
|
|||
pub(super) api_key: Option<SecretValue>,
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) async fn prepare(
|
||||
call: MessagesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
|
|
@ -53,26 +50,11 @@ fn resolve_provider(
|
|||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<ResolvedProvider, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(
|
||||
"unable to resolve custom_llm_provider for messages request".to_string(),
|
||||
)
|
||||
})?;
|
||||
let provider = provider
|
||||
.parse()
|
||||
.map_err(|_| Error::InvalidProvider(provider.to_string()))?;
|
||||
let resolved = resolve_llm_provider(model, custom_llm_provider, "messages")?;
|
||||
let provider = messages_provider(resolved.provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(<&str>::from(resolved.provider).to_string()))?;
|
||||
Ok(ResolvedProvider {
|
||||
model: model.to_string(),
|
||||
model: resolved.model.to_string(),
|
||||
provider,
|
||||
})
|
||||
}
|
||||
|
|
@ -96,7 +78,7 @@ fn prepare_provider_request(
|
|||
let config = provider.config();
|
||||
let env_lookup = |key: &str| secrets.get(key);
|
||||
|
||||
let sanitized = shape_anthropic_messages_request(
|
||||
let sanitized = config.shape_request(
|
||||
AnthropicMessagesRequest { model, ..body },
|
||||
shaping.reasoning_auto_summary,
|
||||
)?;
|
||||
|
|
@ -468,7 +450,11 @@ mod tests {
|
|||
shaping,
|
||||
),
|
||||
Err(Error::InvalidRequest(
|
||||
"metadata.user_id must be a string, got 123".to_string()
|
||||
litellm_llms::ErrorDetail::InvalidValue {
|
||||
field: "metadata.user_id",
|
||||
expected: "a string",
|
||||
actual: json!(123),
|
||||
}
|
||||
))
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,26 +1,15 @@
|
|||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_host::{
|
||||
host::{Demand, Host},
|
||||
machine::{CallMachine, HostChannel, MachineFault},
|
||||
call::{HostedCompletion, HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_http::{Client, ClientVariant, HttpClientConfig};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
|
||||
use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare};
|
||||
use super::{Error, MessagesCall};
|
||||
|
||||
pub enum MessagesOutput {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
/// Every chunk already reached the host through `Deliver`.
|
||||
Streamed,
|
||||
}
|
||||
pub type MessagesOutput = HostedCompletion<Box<AnthropicMessagesResponse>>;
|
||||
|
||||
/// The upstream response as the caller sees it at stream hand-off, before any chunk.
|
||||
pub struct MessagesStreamHead {
|
||||
|
|
@ -30,91 +19,20 @@ pub struct MessagesStreamHead {
|
|||
pub struct Messages;
|
||||
|
||||
impl Protocol for Messages {
|
||||
type Response = MessagesOutput;
|
||||
type Response = Box<AnthropicMessagesResponse>;
|
||||
type Error = Error;
|
||||
type Projection = MessagesCall;
|
||||
type Op = Infallible;
|
||||
type Request = MessagesCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = MessagesStreamHead;
|
||||
}
|
||||
|
||||
impl From<MachineFault> for Error {
|
||||
fn from(fault: MachineFault) -> Self {
|
||||
Self::InvalidRequest(match fault {
|
||||
MachineFault::Abandoned => "messages host driver was abandoned".into(),
|
||||
MachineFault::Protocol(message) => format!("messages {message}"),
|
||||
pub type MessagesMachine = HostedMachine<Messages>;
|
||||
|
||||
impl super::MessagesRoute {
|
||||
pub fn machine(self, request: super::MessagesCall) -> MessagesMachine {
|
||||
hosted_call(request, move |call, _, hooks| async move {
|
||||
self.run(call, &hooks).await
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub type MessagesHost = HostChannel<Messages>;
|
||||
pub type MessagesMachine = CallMachine<Messages>;
|
||||
|
||||
/// The in-process host for a request already in hand. It answers projection once and
|
||||
/// observes nothing.
|
||||
pub struct LocalMessagesHost {
|
||||
call: Mutex<Option<MessagesCall>>,
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
pub fn new(call: MessagesCall) -> Self {
|
||||
Self {
|
||||
call: Mutex::new(Some(call)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for LocalMessagesHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn messages_machine(
|
||||
resources: &crate::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesMachine, litellm_http::Error> {
|
||||
let http = resources.pool.client(config, ClientVariant::Provider)?;
|
||||
let auth = resources.auth.clone();
|
||||
Ok(CallMachine::new(move |host| {
|
||||
Box::pin(drive(host, http, auth, secrets))
|
||||
}))
|
||||
}
|
||||
|
||||
/// The call as its host sees it: projection first, then the same prepare and execute as
|
||||
/// [`super::messages`], with each chunk of a stream handed over as it arrives.
|
||||
async fn drive(
|
||||
host: MessagesHost,
|
||||
http: Client,
|
||||
auth: Arc<litellm_auth::AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let call = host.project().await?;
|
||||
let request = prepare(call, secrets.as_ref()).await?;
|
||||
match execute(&http, &auth, request, &host).await? {
|
||||
MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)),
|
||||
MessagesResponse::Stream {
|
||||
headers,
|
||||
mut chunks,
|
||||
} => {
|
||||
if host.open(MessagesStreamHead { headers }).await? == Demand::Detached {
|
||||
return Ok(MessagesOutput::Streamed);
|
||||
}
|
||||
while let Some(chunk) = chunks.try_next().await? {
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(MessagesOutput::Streamed)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream::BoxStream;
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
|
|
@ -30,16 +30,11 @@ pub fn messages_body(body: Map<String, Value>) -> Result<AnthropicMessagesReques
|
|||
}
|
||||
|
||||
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into())
|
||||
}
|
||||
|
||||
pub enum MessagesResponse {
|
||||
Message(Box<AnthropicMessagesResponse>),
|
||||
Stream {
|
||||
headers: Vec<(String, String)>,
|
||||
chunks: BoxStream<'static, Result<Bytes, Error>>,
|
||||
},
|
||||
}
|
||||
pub type MessagesResponse =
|
||||
CallOutput<Box<AnthropicMessagesResponse>, super::route::MessagesStreamHead, Bytes, Error>;
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct MessagesShaping {
|
||||
|
|
@ -55,7 +50,7 @@ pub struct MessagesShaping {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use litellm_llms::anthropic::common_utils::SupportedEffortTiers;
|
||||
use litellm_llms::base_llm::messages::context::SupportedEffortTiers;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
|
|
|
|||
|
|
@ -1,15 +1,74 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_host::hooks::RouteHooks;
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
|
||||
};
|
||||
|
||||
use crate::ocr::{
|
||||
route::{LocalOcrHost, ocr_machine},
|
||||
types::LiteLLMOcrRequest,
|
||||
use super::{
|
||||
handler::perform_ocr_request,
|
||||
types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest},
|
||||
};
|
||||
|
||||
pub async fn perform(
|
||||
client: &OcrClient,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_host::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await
|
||||
#[derive(Clone)]
|
||||
pub struct OcrRoute {
|
||||
client: OcrClient,
|
||||
}
|
||||
|
||||
impl OcrRoute {
|
||||
pub fn new(client: OcrClient) -> Self {
|
||||
Self { client }
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
request: LiteLLMOcrRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_host::lifecycle::observe_unary(hooks.observer(), self.run(request, hooks)).await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "ocr",
|
||||
model = %request.model,
|
||||
resolved_model = %request.model,
|
||||
provider = <&str>::from(request.config.provider()),
|
||||
stream = false,
|
||||
outcome
|
||||
))]
|
||||
pub(super) async fn run(
|
||||
&self,
|
||||
request: LiteLLMOcrRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
let caller_document = matches!(&request.document, OcrDocumentInput::Document(_));
|
||||
let prepared = prepare_request_document(request).await?;
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<LiteLLMOcrResponse, Error>> =
|
||||
Box::pin(perform_ocr_request(
|
||||
&self.client,
|
||||
prepared,
|
||||
hooks,
|
||||
caller_document,
|
||||
));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
async fn prepare_request_document(
|
||||
request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
) -> Result<ResolvedOcrRequest, Error> {
|
||||
if let OcrDocumentInput::Document(_) = &request.document {
|
||||
return request.map_document(super::document::prepare_document);
|
||||
}
|
||||
let logger = litellm_tracing::Logger::current();
|
||||
let span = tracing::Span::current();
|
||||
tokio::task::spawn_blocking(move || {
|
||||
logger.scope(|| span.in_scope(|| request.map_document(super::document::prepare_document)))
|
||||
})
|
||||
.await
|
||||
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
|
||||
}
|
||||
|
|
|
|||
|
|
@ -110,6 +110,7 @@ mod tests {
|
|||
use std::collections::BTreeMap as Map;
|
||||
|
||||
use litellm_llms::base_llm::ocr::document::InlineDocument;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
@ -135,24 +136,21 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn file_name_mime_mapping_matches_python() {
|
||||
for (name, expected) in [
|
||||
("document.pdf", "application/pdf"),
|
||||
("image.png", "image/png"),
|
||||
("photo.jpg", "image/jpeg"),
|
||||
("photo.jpeg", "image/jpeg"),
|
||||
("animation.gif", "image/gif"),
|
||||
("image.webp", "image/webp"),
|
||||
("scan.tiff", "image/tiff"),
|
||||
("scan.tif", "image/tiff"),
|
||||
("bitmap.bmp", "image/bmp"),
|
||||
("DOCUMENT.PDF", "application/pdf"),
|
||||
("IMAGE.PNG", "image/png"),
|
||||
("file.unknown-extension", "application/octet-stream"),
|
||||
] {
|
||||
assert_eq!(mime_type_for_name(name), expected);
|
||||
}
|
||||
#[rstest]
|
||||
#[case::pdf("document.pdf", "application/pdf")]
|
||||
#[case::png("image.png", "image/png")]
|
||||
#[case::jpg("photo.jpg", "image/jpeg")]
|
||||
#[case::jpeg("photo.jpeg", "image/jpeg")]
|
||||
#[case::gif("animation.gif", "image/gif")]
|
||||
#[case::webp("image.webp", "image/webp")]
|
||||
#[case::tiff("scan.tiff", "image/tiff")]
|
||||
#[case::tif("scan.tif", "image/tiff")]
|
||||
#[case::bmp("bitmap.bmp", "image/bmp")]
|
||||
#[case::uppercase_pdf("DOCUMENT.PDF", "application/pdf")]
|
||||
#[case::uppercase_png("IMAGE.PNG", "image/png")]
|
||||
#[case::unknown("file.unknown-extension", "application/octet-stream")]
|
||||
fn file_name_mime_mapping_matches_python(#[case] name: &str, #[case] expected: &str) {
|
||||
assert_eq!(mime_type_for_name(name), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
use futures_util::future::BoxFuture;
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::{CallHooks, OcrClient},
|
||||
|
|
@ -8,16 +10,13 @@ use litellm_llms::base_llm::ocr::{
|
|||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind,
|
||||
route::OcrHost,
|
||||
};
|
||||
use super::{arguments::is_secret_param, prepare::prepare_request, provider_config::OcrConfigKind};
|
||||
use crate::ocr::types::ResolvedOcrRequest;
|
||||
|
||||
pub(crate) async fn perform_ocr_request(
|
||||
client: &OcrClient,
|
||||
request: ResolvedOcrRequest,
|
||||
host: &OcrHost,
|
||||
host: &impl RouteHooks<Error>,
|
||||
caller_document: bool,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
request.response_format()?;
|
||||
|
|
@ -28,53 +27,48 @@ pub(crate) async fn perform_ocr_request(
|
|||
.await
|
||||
.map_err(|error| Error::Secret(std::sync::Arc::new(error)))?;
|
||||
let request = prepare_request(request, caller_document, client, secrets);
|
||||
let hooks = OcrCallHooks::new(host.clone(), &request, config);
|
||||
let hooks = OcrCallHooks::new(host, &request, config);
|
||||
config.ocr(client, &request, &hooks).await
|
||||
}
|
||||
|
||||
/// Lets provider code reach the host mid-call, filling in the request context only the
|
||||
/// route knows.
|
||||
pub(crate) struct OcrCallHooks {
|
||||
host: OcrHost,
|
||||
model: String,
|
||||
custom_llm_provider: &'static str,
|
||||
optional_params: Value,
|
||||
secret_fields: Vec<String>,
|
||||
api_key: Option<SecretValue>,
|
||||
struct OcrCallHooks<'a, H> {
|
||||
hooks: &'a H,
|
||||
context: RequestContext,
|
||||
}
|
||||
|
||||
impl OcrCallHooks {
|
||||
pub(crate) fn new(host: OcrHost, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self {
|
||||
impl<'a, H> OcrCallHooks<'a, H> {
|
||||
fn new(hooks: &'a H, request: &PreparedOcrRequest, config: OcrConfigKind) -> Self {
|
||||
Self {
|
||||
host,
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: config.provider().into(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
secret_fields: request
|
||||
.optional_params
|
||||
.keys()
|
||||
.filter(|name| is_secret_param(name))
|
||||
.cloned()
|
||||
.collect(),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
hooks,
|
||||
context: RequestContext {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: <&str>::from(config.provider()).to_owned(),
|
||||
optional_params: Value::Object(request.optional_params.clone().into()),
|
||||
secret_fields: request
|
||||
.optional_params
|
||||
.keys()
|
||||
.filter(|name| is_secret_param(name))
|
||||
.cloned()
|
||||
.collect(),
|
||||
api_key: request.connection.api_key.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl CallHooks<Error> for OcrCallHooks {
|
||||
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
let context = RequestContext {
|
||||
model: self.model.clone(),
|
||||
custom_llm_provider: self.custom_llm_provider.into(),
|
||||
optional_params: self.optional_params.clone(),
|
||||
secret_fields: self.secret_fields.clone(),
|
||||
api_key: self.api_key.clone(),
|
||||
};
|
||||
Box::pin(self.host.before_send(wire, context))
|
||||
impl<H: RouteHooks<Error>> CallHooks<Error> for OcrCallHooks<'_, H> {
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
Box::pin(
|
||||
self.hooks
|
||||
.before_provider_request(wire, self.context.clone()),
|
||||
)
|
||||
}
|
||||
|
||||
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
|
||||
Box::pin(self.host.emit(MachineEvent::ResponseReceived {
|
||||
Box::pin(self.hooks.on_event(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse {
|
||||
body: String::from_utf8_lossy(body).into_owned(),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod arguments;
|
||||
pub mod client;
|
||||
mod client;
|
||||
pub use client::OcrRoute;
|
||||
pub mod document;
|
||||
pub(crate) mod handler;
|
||||
pub(crate) mod prepare;
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ use litellm_llms::base_llm::ocr::{
|
|||
};
|
||||
use litellm_secrets::source::Secrets;
|
||||
|
||||
use super::provider_config::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest};
|
||||
use crate::provider::LlmProviders;
|
||||
|
||||
pub(crate) fn prepare_request(
|
||||
request: ResolvedOcrRequest,
|
||||
|
|
@ -16,15 +16,19 @@ pub(crate) fn prepare_request(
|
|||
) -> PreparedOcrRequest {
|
||||
let credentials = request.credentials.clone();
|
||||
let (preferred_api_key_env, api_base_env) = match request.config.provider() {
|
||||
OcrProvider::Mistral => (
|
||||
LlmProviders::Mistral => (
|
||||
Some("MISTRAL_AZURE_API_KEY"),
|
||||
Some("MISTRAL_AZURE_API_BASE"),
|
||||
),
|
||||
OcrProvider::AzureAi => (None, Some("AZURE_AI_API_BASE")),
|
||||
OcrProvider::AwsTextract
|
||||
| OcrProvider::Cohere
|
||||
| OcrProvider::Reducto
|
||||
| OcrProvider::VertexAi => (None, None),
|
||||
LlmProviders::AzureAi => (None, Some("AZURE_AI_API_BASE")),
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::AwsTextract
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Cohere
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike
|
||||
| LlmProviders::Reducto
|
||||
| LlmProviders::VertexAi => (None, None),
|
||||
};
|
||||
let secret = |name: &str| secrets.truthy(name);
|
||||
let dynamic_api_key = credentials.dynamic_api_key.or_else(|| {
|
||||
|
|
@ -100,7 +104,10 @@ mod tests {
|
|||
struct NoHooks;
|
||||
|
||||
impl CallHooks<Error> for NoHooks {
|
||||
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
) -> BoxFuture<'_, Result<WireRequest, Error>> {
|
||||
Box::pin(async move { Ok(wire) })
|
||||
}
|
||||
|
||||
|
|
@ -144,6 +151,7 @@ mod tests {
|
|||
json!({"type": "image_url", "image_url": url})
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn cohere_body_keeps_native_document_fields_and_untyped_overrides() {
|
||||
let request = request(
|
||||
|
|
@ -176,6 +184,7 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn explicit_null_options_use_defaults_before_http() {
|
||||
let request = request(
|
||||
|
|
@ -199,6 +208,7 @@ mod tests {
|
|||
assert!(body.get("req_format").is_none());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn direct_and_vertex_mistral_build_the_same_request_and_share_normalization() {
|
||||
let options = json!({
|
||||
|
|
@ -276,7 +286,7 @@ mod tests {
|
|||
pages: Option<Vec<i64>>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest::rstest]
|
||||
fn parsed_provider_params_separates_known_and_extra_params() {
|
||||
let arguments: CallArguments = serde_json::from_value(json!({
|
||||
"pages": [0, 2],
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::provider::LlmProviders;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use litellm_llms::{
|
||||
aws_textract::ocr::{
|
||||
|
|
@ -24,7 +25,6 @@ use litellm_llms::{
|
|||
deepseek_transformation::VertexAIDeepSeekOCRConfig, transformation::VertexAiOcrConfig,
|
||||
},
|
||||
};
|
||||
use strum::{EnumString, IntoStaticStr};
|
||||
|
||||
macro_rules! with_config {
|
||||
($kind:expr, $config:ident => $body:expr) => {
|
||||
|
|
@ -93,16 +93,16 @@ pub(crate) enum OcrConfigKind {
|
|||
}
|
||||
|
||||
impl OcrConfigKind {
|
||||
pub(crate) const fn provider(self) -> OcrProvider {
|
||||
pub(crate) const fn provider(self) -> LlmProviders {
|
||||
match self {
|
||||
Self::AwsTextract | Self::AwsTextractAnalyze => OcrProvider::AwsTextract,
|
||||
Self::Cohere => OcrProvider::Cohere,
|
||||
Self::Mistral => OcrProvider::Mistral,
|
||||
Self::AwsTextract | Self::AwsTextractAnalyze => LlmProviders::AwsTextract,
|
||||
Self::Cohere => LlmProviders::Cohere,
|
||||
Self::Mistral => LlmProviders::Mistral,
|
||||
Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => {
|
||||
OcrProvider::AzureAi
|
||||
LlmProviders::AzureAi
|
||||
}
|
||||
Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto,
|
||||
Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi,
|
||||
Self::ReductoLegacy | Self::ReductoV3 => LlmProviders::Reducto,
|
||||
Self::VertexAi | Self::VertexDeepSeek => LlmProviders::VertexAi,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -187,17 +187,6 @@ pub fn passthrough_response(
|
|||
.map(Some)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub(crate) enum OcrProvider {
|
||||
AwsTextract,
|
||||
Cohere,
|
||||
Mistral,
|
||||
AzureAi,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_provider_config(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
|
|
@ -205,37 +194,45 @@ pub(crate) fn resolve_provider_config(
|
|||
let provider =
|
||||
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: OcrProvider::Mistral.into(),
|
||||
custom_llm_provider: LlmProviders::Mistral.into(),
|
||||
});
|
||||
let ocr_provider = provider
|
||||
let llm_provider = provider
|
||||
.custom_llm_provider
|
||||
.parse::<OcrProvider>()
|
||||
.parse::<LlmProviders>()
|
||||
.map_err(|_| Error::InvalidProvider(provider.custom_llm_provider.to_string()))?;
|
||||
let config = match ocr_provider {
|
||||
OcrProvider::AwsTextract => match TextractOperation::from_model(provider.model)? {
|
||||
let config = match llm_provider {
|
||||
LlmProviders::AwsTextract => match TextractOperation::from_model(provider.model)? {
|
||||
TextractOperation::DetectDocumentText => OcrConfigKind::AwsTextract,
|
||||
TextractOperation::AnalyzeDocument => OcrConfigKind::AwsTextractAnalyze,
|
||||
},
|
||||
OcrProvider::Cohere => OcrConfigKind::Cohere,
|
||||
OcrProvider::Mistral => OcrConfigKind::Mistral,
|
||||
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
LlmProviders::Cohere => OcrConfigKind::Cohere,
|
||||
LlmProviders::Mistral => OcrConfigKind::Mistral,
|
||||
LlmProviders::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
OcrConfigKind::AzureDocumentIntelligence
|
||||
}
|
||||
OcrProvider::AzureAi
|
||||
LlmProviders::AzureAi
|
||||
if provider.model.to_ascii_lowercase().contains("cohere")
|
||||
&& provider.model.to_ascii_lowercase().contains("parse") =>
|
||||
{
|
||||
OcrConfigKind::AzureCohere
|
||||
}
|
||||
OcrProvider::AzureAi => OcrConfigKind::AzureAi,
|
||||
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
LlmProviders::AzureAi => OcrConfigKind::AzureAi,
|
||||
LlmProviders::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
OcrConfigKind::ReductoLegacy
|
||||
}
|
||||
OcrProvider::Reducto => OcrConfigKind::ReductoV3,
|
||||
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
LlmProviders::Reducto => OcrConfigKind::ReductoV3,
|
||||
LlmProviders::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
OcrConfigKind::VertexDeepSeek
|
||||
}
|
||||
OcrProvider::VertexAi => OcrConfigKind::VertexAi,
|
||||
LlmProviders::VertexAi => OcrConfigKind::VertexAi,
|
||||
LlmProviders::Anthropic
|
||||
| LlmProviders::Bedrock
|
||||
| LlmProviders::Openai
|
||||
| LlmProviders::OpenaiLike => {
|
||||
return Err(Error::InvalidProvider(
|
||||
provider.custom_llm_provider.to_string(),
|
||||
));
|
||||
}
|
||||
};
|
||||
Ok((provider.model.to_string(), config))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,25 +1,20 @@
|
|||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use litellm_auth::ResolvedCredential;
|
||||
use litellm_host::{
|
||||
event::{CallEvent, RequestContext, WireRequest},
|
||||
host::Reply,
|
||||
machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol},
|
||||
call::{CallOutput, HostedMachine, hosted_call},
|
||||
machine::{HostTokenProvider, TokenProtocol},
|
||||
protocol::Protocol,
|
||||
protocol::Reply,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse};
|
||||
|
||||
use super::handler::perform_ocr_request;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput};
|
||||
|
||||
pub enum OcrOp {
|
||||
AcquireAzureAdToken(Reply<ResolvedCredential>),
|
||||
}
|
||||
|
||||
/// The caller's request as the host projects it.
|
||||
pub struct OcrProjection {
|
||||
pub struct OcrCall {
|
||||
pub request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
/// The caller passed its own Azure AD token provider, which the host keeps.
|
||||
pub caller_token: bool,
|
||||
|
|
@ -30,8 +25,8 @@ pub struct Ocr;
|
|||
impl Protocol for Ocr {
|
||||
type Response = LiteLLMOcrResponse;
|
||||
type Error = Error;
|
||||
type Projection = OcrProjection;
|
||||
type Op = OcrOp;
|
||||
type Request = OcrCall;
|
||||
type HostCall = OcrOp;
|
||||
type Chunk = std::convert::Infallible;
|
||||
type StreamHead = std::convert::Infallible;
|
||||
}
|
||||
|
|
@ -42,122 +37,22 @@ impl TokenProtocol for Ocr {
|
|||
}
|
||||
}
|
||||
|
||||
pub type OcrHost = HostChannel<Ocr>;
|
||||
pub type OcrMachine = CallMachine<Ocr>;
|
||||
pub type OcrMachine = HostedMachine<Ocr>;
|
||||
|
||||
/// The OCR call as a machine: projection and token acquisition are host operations;
|
||||
/// everything else runs in Rust.
|
||||
pub fn ocr_machine(client: OcrClient) -> OcrMachine {
|
||||
CallMachine::new(move |host| Box::pin(execute(client, host)))
|
||||
}
|
||||
|
||||
async fn execute(client: OcrClient, host: OcrHost) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let OcrProjection {
|
||||
request,
|
||||
caller_token,
|
||||
} = host.project().await?;
|
||||
let request = LiteLLMOcrRequest {
|
||||
azure_ad_token_provider: caller_token
|
||||
.then(|| HostTokenProvider::handle(host.clone()))
|
||||
.or(request.azure_ad_token_provider),
|
||||
..request
|
||||
};
|
||||
let caller_document = matches!(request.document, OcrDocumentInput::Document(_));
|
||||
let request = prepare_request_document(request).await?;
|
||||
perform_ocr_request(&client, request, &host, caller_document).await
|
||||
}
|
||||
|
||||
async fn prepare_request_document(
|
||||
request: LiteLLMOcrRequest<OcrDocumentInput>,
|
||||
) -> Result<ResolvedOcrRequest, Error> {
|
||||
if let OcrDocumentInput::Document(_) = &request.document {
|
||||
return request.map_document(super::document::prepare_document);
|
||||
}
|
||||
tokio::task::spawn_blocking(move || request.map_document(super::document::prepare_document))
|
||||
.await
|
||||
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
|
||||
}
|
||||
|
||||
type BeforeSend =
|
||||
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
|
||||
|
||||
/// The in-process host for a request that is already in hand: the request answers
|
||||
/// projection, and the optional observer sees and may rewrite the wire request.
|
||||
pub struct LocalOcrHost {
|
||||
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
|
||||
before_send: Option<BeforeSend>,
|
||||
observer: Option<Observer>,
|
||||
}
|
||||
|
||||
impl LocalOcrHost {
|
||||
pub fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
|
||||
Self {
|
||||
request: Mutex::new(Some(request)),
|
||||
before_send: None,
|
||||
observer: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_before_send(
|
||||
self,
|
||||
before_send: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
) -> Self {
|
||||
Self {
|
||||
before_send: Some(Box::new(before_send)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self {
|
||||
Self {
|
||||
observer: Some(Box::new(observer)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_host::host::Host<Ocr> for LocalOcrHost {
|
||||
async fn project(&self) -> Result<OcrProjection, Error> {
|
||||
self.request
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.map(|request| OcrProjection {
|
||||
request,
|
||||
caller_token: false,
|
||||
})
|
||||
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into()))
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
|
||||
match op {
|
||||
OcrOp::AcquireAzureAdToken(_) => {
|
||||
Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition(
|
||||
"OCR host has no Azure AD token provider".into(),
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
match &self.before_send {
|
||||
Some(before_send) => before_send(wire, context),
|
||||
None => Ok(wire),
|
||||
}
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
if let Some(observer) = &self.observer {
|
||||
observer(event);
|
||||
}
|
||||
Ok(())
|
||||
impl crate::ocr::OcrRoute {
|
||||
pub fn machine(self, request: OcrCall) -> OcrMachine {
|
||||
hosted_call(
|
||||
request,
|
||||
move |projection: OcrCall, services, hooks| async move {
|
||||
let request = LiteLLMOcrRequest {
|
||||
azure_ad_token_provider: projection
|
||||
.caller_token
|
||||
.then(|| HostTokenProvider::handle(services))
|
||||
.or(projection.request.azure_ad_token_provider),
|
||||
..projection.request
|
||||
};
|
||||
self.run(request, &hooks).await.map(CallOutput::Complete)
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,6 +4,21 @@ use litellm_http::outbound::OutboundRequest;
|
|||
use litellm_llms::base_llm::auth::Authenticated;
|
||||
use serde_json::Value;
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.provider.send",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(status)
|
||||
)]
|
||||
pub(crate) async fn send(
|
||||
request: OutboundRequest,
|
||||
client: &litellm_http::Client,
|
||||
) -> Result<reqwest::Response, reqwest::Error> {
|
||||
request.send(client).await.inspect(|response| {
|
||||
tracing::Span::current().record("status", response.status().as_u16());
|
||||
})
|
||||
}
|
||||
|
||||
/// Header credentials are already in `headers`; SigV4 is applied here, over the
|
||||
/// bytes that are sent.
|
||||
pub(crate) fn outbound_request(
|
||||
|
|
|
|||
36
litellm-rust/crates/core/src/provider.rs
Normal file
36
litellm-rust/crates/core/src/provider.rs
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
pub use litellm_core_utils::get_llm_provider_logic::LlmProviders;
|
||||
use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
use crate::error::RouteError as Error;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) struct ResolvedProvider<'a> {
|
||||
pub(crate) model: &'a str,
|
||||
pub(crate) provider: LlmProviders,
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_llm_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
route: &'static str,
|
||||
) -> Result<ResolvedProvider<'a>, Error> {
|
||||
let CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
} = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
Error::InvalidProvider(format!(
|
||||
"unable to resolve custom_llm_provider for {route} request"
|
||||
))
|
||||
})?;
|
||||
let provider = custom_llm_provider
|
||||
.parse()
|
||||
.map_err(|_| Error::InvalidProvider(custom_llm_provider.to_string()))?;
|
||||
Ok(ResolvedProvider { model, provider })
|
||||
}
|
||||
|
|
@ -1,9 +1,7 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_http::{HttpClientConfig, HttpClientPool, media::UrlPolicy};
|
||||
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_http::HttpClientPool;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CoreResources {
|
||||
|
|
@ -18,21 +16,4 @@ impl CoreResources {
|
|||
auth: Arc::new(AuthServices::default()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn ocr_client(
|
||||
&self,
|
||||
config: &HttpClientConfig,
|
||||
url_policy: UrlPolicy,
|
||||
settings: OcrSettings,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<OcrClient, litellm_http::Error> {
|
||||
OcrClient::new(
|
||||
&self.pool,
|
||||
config,
|
||||
url_policy,
|
||||
self.auth.clone(),
|
||||
settings,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
91
litellm-rust/crates/core/src/responses/handler.rs
Normal file
91
litellm-rust/crates/core/src/responses/handler.rs
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use futures_util::StreamExt;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_llms::base_llm::auth::{Authenticated, resolve_auth};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesOutput, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub(super) async fn execute(
|
||||
http: &litellm_http::Client,
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderResponsesRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let authenticated = resolve_auth(auth, request.environment, &|_| None).await?;
|
||||
let wire = hooks
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url: request.url,
|
||||
headers: authenticated.headers,
|
||||
body: request.body,
|
||||
},
|
||||
request.context,
|
||||
)
|
||||
.await?;
|
||||
let stream = match wire.body.get("stream") {
|
||||
None => false,
|
||||
Some(serde_json::Value::Bool(value)) => *value,
|
||||
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
|
||||
};
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(network)?;
|
||||
let status = response.status().as_u16();
|
||||
if !response.status().is_success() {
|
||||
let body = response.text().await.map_err(network)?;
|
||||
return Err(litellm_http::transport::Error::Http {
|
||||
status,
|
||||
body: litellm_http::request::truncate_error_body(&body),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
if stream {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_owned())))
|
||||
.collect();
|
||||
let chunks = response
|
||||
.bytes_stream()
|
||||
.map(|chunk| chunk.map_err(network))
|
||||
.boxed();
|
||||
return Ok(ResponsesOutput::Stream {
|
||||
head: ResponsesStreamHead { headers },
|
||||
chunks,
|
||||
});
|
||||
}
|
||||
let body = response.text().await.map_err(network)?;
|
||||
hooks
|
||||
.on_event(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: body.clone() },
|
||||
})
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
let value = serde_json::from_str(&body)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into()))?;
|
||||
request
|
||||
.config
|
||||
.transform_response_api_response(value)
|
||||
.map(ResponsesOutput::Complete)
|
||||
.map_err(Error::from)
|
||||
}
|
||||
|
||||
fn network(error: reqwest::Error) -> Error {
|
||||
litellm_http::transport::Error::Network(error.to_string()).into()
|
||||
}
|
||||
|
|
@ -1,2 +1,69 @@
|
|||
pub use crate::error::RouteError as Error;
|
||||
pub mod websocket;
|
||||
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::hooks::RouteHooks;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use types::{ResponsesCall, ResponsesOutput};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await
|
||||
}
|
||||
|
||||
#[tracing::instrument(name = "litellm.route", skip_all, fields(
|
||||
route = "responses",
|
||||
model = %call.model,
|
||||
provider,
|
||||
resolved_model,
|
||||
stream = call.optional_params.get("stream").and_then(serde_json::Value::as_bool).unwrap_or(false),
|
||||
outcome
|
||||
))]
|
||||
async fn run(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(
|
||||
&request.context.model,
|
||||
&request.context.custom_llm_provider,
|
||||
);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
|
||||
Box::pin(handler::execute(&self.http, &self.auth, request, hooks));
|
||||
execute.await
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
58
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
58
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
use litellm_host::event::RequestContext;
|
||||
use litellm_llms::{
|
||||
base_llm::responses::transformation::BaseResponsesApiConfig,
|
||||
openai::responses::transformation::OpenAiResponsesApiConfig,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesCall},
|
||||
};
|
||||
|
||||
#[tracing::instrument(name = "litellm.prepare", level = "debug", skip_all)]
|
||||
pub(super) async fn prepare(
|
||||
call: ResponsesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderResponsesRequest, Error> {
|
||||
let provider = call.custom_llm_provider.as_deref().unwrap_or("openai");
|
||||
if provider != "openai" {
|
||||
return Err(Error::Unsupported("native HTTP responses provider"));
|
||||
}
|
||||
let model = call.model.strip_prefix("openai/").unwrap_or(&call.model);
|
||||
if model.is_empty() || model.contains('/') {
|
||||
return Err(Error::InvalidProvider(call.model));
|
||||
}
|
||||
let config: &'static dyn BaseResponsesApiConfig = &OpenAiResponsesApiConfig;
|
||||
let snapshot = secrets
|
||||
.resolve(config.secret_names(call.api_key.as_deref(), call.api_base.as_deref()))
|
||||
.await?;
|
||||
let lookup = |name: &str| snapshot.get(name);
|
||||
let environment = config.validate_environment(
|
||||
litellm_http::request::string_headers("responses", call.extra_headers)?,
|
||||
call.api_key.as_deref(),
|
||||
&lookup,
|
||||
)?;
|
||||
let context = RequestContext {
|
||||
model: model.into(),
|
||||
custom_llm_provider: provider.into(),
|
||||
optional_params: Value::Object(call.optional_params.clone()),
|
||||
secret_fields: Vec::new(),
|
||||
api_key: match &environment.auth {
|
||||
litellm_llms::base_llm::auth::AuthScheme::Credential { secret, .. } => {
|
||||
Some(secret.clone())
|
||||
}
|
||||
_ => None,
|
||||
},
|
||||
};
|
||||
let body = config.transform_responses_api_request(model, call.input, call.optional_params)?;
|
||||
Ok(ProviderResponsesRequest {
|
||||
url: config.get_complete_url(call.api_base.as_deref(), &lookup),
|
||||
config,
|
||||
environment,
|
||||
body,
|
||||
context,
|
||||
timeout: call.timeout,
|
||||
})
|
||||
}
|
||||
32
litellm-rust/crates/core/src/responses/route.rs
Normal file
32
litellm-rust/crates/core/src/responses/route.rs
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::{
|
||||
call::{HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
|
||||
use super::{
|
||||
Error, ResponsesRoute,
|
||||
types::{ResponsesCall, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub struct Responses;
|
||||
|
||||
impl Protocol for Responses {
|
||||
type Response = ResponsesApiResponse;
|
||||
type Error = Error;
|
||||
type Request = ResponsesCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = ResponsesStreamHead;
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn machine(self, call: ResponsesCall) -> HostedMachine<Responses> {
|
||||
hosted_call(call, move |call, _, hooks| async move {
|
||||
self.run(call, &hooks).await
|
||||
})
|
||||
}
|
||||
}
|
||||
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::ValidatedEnvironment, responses::transformation::BaseResponsesApiConfig,
|
||||
};
|
||||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
|
||||
pub struct ResponsesCall {
|
||||
pub model: String,
|
||||
pub input: Value,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub struct ResponsesStreamHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub type ResponsesOutput = CallOutput<ResponsesApiResponse, ResponsesStreamHead, Bytes, Error>;
|
||||
|
||||
pub(super) struct ProviderResponsesRequest {
|
||||
pub config: &'static dyn BaseResponsesApiConfig,
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
pub context: litellm_host::event::RequestContext,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
|
@ -1,23 +1,13 @@
|
|||
use std::{
|
||||
collections::HashMap,
|
||||
io,
|
||||
sync::{Arc, OnceLock},
|
||||
time::Duration,
|
||||
};
|
||||
use std::{collections::HashMap, sync::Arc, time::Duration};
|
||||
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use litellm_http::websocket::{UpstreamWebSocket, connect_upstream};
|
||||
use litellm_types::responses::streaming_websocket::ResponsesWsEventType;
|
||||
use rustls::{ClientConfig, RootCertStore};
|
||||
use tokio::{net::TcpStream, sync::Mutex};
|
||||
use tokio_tungstenite::{
|
||||
Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config,
|
||||
tungstenite::{
|
||||
Message,
|
||||
client::IntoClientRequest,
|
||||
error::TlsError,
|
||||
handshake::client::Response,
|
||||
http::{HeaderName, HeaderValue},
|
||||
},
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_tungstenite::tungstenite::{
|
||||
Message,
|
||||
client::IntoClientRequest,
|
||||
http::{HeaderName, HeaderValue},
|
||||
};
|
||||
|
||||
use super::Error;
|
||||
|
|
@ -33,139 +23,127 @@ pub fn is_terminal_event(event_type: &ResponsesWsEventType) -> bool {
|
|||
)
|
||||
}
|
||||
|
||||
pub type ResponsesUpstreamWs = WebSocketStream<MaybeTlsStream<TcpStream>>;
|
||||
|
||||
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
|
||||
|
||||
fn build_tls_config() -> Result<ClientConfig, Box<tokio_tungstenite::tungstenite::Error>> {
|
||||
let native = rustls_native_certs::load_native_certs();
|
||||
let mut store = RootCertStore::empty();
|
||||
let (added, _ignored) = store.add_parsable_certificates(native.certs);
|
||||
if added == 0 {
|
||||
return Err(Box::new(tokio_tungstenite::tungstenite::Error::Io(
|
||||
io::Error::other(format!(
|
||||
"no usable native root certificates: {:?}",
|
||||
native.errors
|
||||
)),
|
||||
)));
|
||||
}
|
||||
ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider()))
|
||||
.with_safe_default_protocol_versions()
|
||||
.map(|builder| builder.with_root_certificates(store).with_no_client_auth())
|
||||
.map_err(|error| {
|
||||
Box::new(tokio_tungstenite::tungstenite::Error::Tls(
|
||||
TlsError::Rustls(error),
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn tls_config() -> Result<Arc<ClientConfig>, Box<tokio_tungstenite::tungstenite::Error>> {
|
||||
if let Some(config) = TLS_CONFIG.get() {
|
||||
return Ok(Arc::clone(config));
|
||||
}
|
||||
let built = Arc::new(build_tls_config()?);
|
||||
Ok(Arc::clone(TLS_CONFIG.get_or_init(|| built)))
|
||||
}
|
||||
|
||||
pub async fn connect_upstream<R>(
|
||||
request: R,
|
||||
) -> Result<(ResponsesUpstreamWs, Response), Box<tokio_tungstenite::tungstenite::Error>>
|
||||
where
|
||||
R: IntoClientRequest + Unpin,
|
||||
{
|
||||
let request = request.into_client_request().map_err(Box::new)?;
|
||||
let connector = match request.uri().scheme_str() {
|
||||
Some("wss") => Some(Connector::Rustls(tls_config()?)),
|
||||
_ => None,
|
||||
};
|
||||
connect_async_tls_with_config(request, None, false, connector)
|
||||
.await
|
||||
.map_err(Box::new)
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesWebSocketConnection {
|
||||
socket: Arc<Mutex<Option<ResponsesUpstreamWs>>>,
|
||||
socket: Arc<Mutex<Option<UpstreamWebSocket>>>,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketConnection {
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.connect_url",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn connect_url(
|
||||
url: &str,
|
||||
headers: &HashMap<String, String>,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<Self, Error> {
|
||||
let mut request = url.into_client_request().map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
for (name, value) in headers {
|
||||
let header_name = name
|
||||
.parse::<HeaderName>()
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
|
||||
let header_value = HeaderValue::from_str(value)
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
|
||||
request.headers_mut().insert(header_name, header_value);
|
||||
}
|
||||
let connect = connect_upstream(request);
|
||||
let result = match timeout {
|
||||
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket connection timed out".into(),
|
||||
))
|
||||
})?,
|
||||
None => connect.await,
|
||||
};
|
||||
let (socket, _) = result.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => {
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
})
|
||||
}
|
||||
other => Error::Transport(litellm_http::transport::Error::Network(other.to_string())),
|
||||
})?;
|
||||
Ok(Self {
|
||||
socket: Arc::new(Mutex::new(Some(socket))),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn send_text(&self, text: String) -> Result<(), Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket is closed".into(),
|
||||
)));
|
||||
};
|
||||
socket.send(Message::Text(text)).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Ok(None);
|
||||
};
|
||||
match socket.next().await {
|
||||
Some(Ok(Message::Text(text))) => Ok(Some(text)),
|
||||
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string())),
|
||||
Some(Ok(Message::Close(_))) | None => Ok(None),
|
||||
Some(Ok(_)) => Ok(None),
|
||||
Some(Err(error)) => Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
error.to_string(),
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn close(&self) -> Result<(), Error> {
|
||||
let mut socket = self.socket.lock().await;
|
||||
if let Some(socket) = socket.as_mut() {
|
||||
socket.close(None).await.map_err(|error| {
|
||||
crate::diagnostic::operation("litellm.websocket.connect_url", async {
|
||||
let mut request = url.into_client_request().map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
}
|
||||
*socket = None;
|
||||
Ok(())
|
||||
for (name, value) in headers {
|
||||
let header_name = name
|
||||
.parse::<HeaderName>()
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
let header_value = HeaderValue::from_str(value)
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string().into()))?;
|
||||
request.headers_mut().insert(header_name, header_value);
|
||||
}
|
||||
let connect = connect_upstream(request);
|
||||
let result = match timeout {
|
||||
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket connection timed out".into(),
|
||||
))
|
||||
})?,
|
||||
None => connect.await,
|
||||
};
|
||||
let (socket, _) = result.map_err(|error| match *error {
|
||||
tokio_tungstenite::tungstenite::Error::Http(response) => {
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: response.status().as_u16(),
|
||||
body: String::new(),
|
||||
})
|
||||
}
|
||||
other => {
|
||||
Error::Transport(litellm_http::transport::Error::Network(other.to_string()))
|
||||
}
|
||||
})?;
|
||||
Ok(Self {
|
||||
socket: Arc::new(Mutex::new(Some(socket))),
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.send_text",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn send_text(&self, text: String) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.send_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
"Responses WebSocket is closed".into(),
|
||||
)));
|
||||
};
|
||||
socket.send(Message::Text(text)).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.recv_text",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.recv_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Ok(None);
|
||||
};
|
||||
match socket.next().await {
|
||||
Some(Ok(Message::Text(text))) => Ok(Some(text)),
|
||||
Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec())
|
||||
.map(Some)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into())),
|
||||
Some(Ok(Message::Close(_))) | None => Ok(None),
|
||||
Some(Ok(_)) => Ok(None),
|
||||
Some(Err(error)) => Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
error.to_string(),
|
||||
))),
|
||||
}
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "litellm.websocket.close",
|
||||
level = "debug",
|
||||
skip_all,
|
||||
fields(outcome)
|
||||
)]
|
||||
pub async fn close(&self) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.close", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
if let Some(socket) = socket.as_mut() {
|
||||
socket.close(None).await.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
}
|
||||
*socket = None;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,6 +1,4 @@
|
|||
use litellm_core::audio_transcription::{
|
||||
Error, audio_transcription, types::AudioTranscriptionRequest,
|
||||
};
|
||||
use litellm_core::audio_transcription::{Error, types::AudioTranscriptionRequest};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
|
@ -11,7 +9,7 @@ use support::*;
|
|||
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
|
||||
|
||||
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
audio_transcription(&support::resources(), &http_config(), request).await
|
||||
audio_transcription_route().execute(request).await
|
||||
}
|
||||
|
||||
fn transcript_response(text: &str) -> ResponseTemplate {
|
||||
|
|
@ -252,3 +250,31 @@ async fn an_unreadable_success_body_is_an_invalid_response(
|
|||
|
||||
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn transcription_records_route_and_resolved_provider(
|
||||
request: AudioTranscriptionRequest<'static>,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
let model = request.model;
|
||||
traces
|
||||
.logger()
|
||||
.instrument(transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "audio_transcription");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(summaries[0]["resolved_model"], model);
|
||||
assert_eq!(summaries[0]["provider"], "bedrock");
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
assert!(!format!("{:?}", traces.records()).contains("secret-key"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,6 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::chat_completions::{
|
||||
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
|
||||
};
|
||||
use litellm_core::chat_completions::{Error, types::ChatCompletionsRequest};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use rstest::{fixture, rstest};
|
||||
|
|
@ -15,7 +13,7 @@ use support::*;
|
|||
const ANTHROPIC_MESSAGE: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
|
||||
|
||||
async fn complete(request: ChatCompletionsRequest<'_>) -> Result<ChatCompletionsResponse, Error> {
|
||||
chat_completions(&support::resources(), &http_config(), request).await
|
||||
chat_completions_route().execute(request, &()).await
|
||||
}
|
||||
|
||||
fn object(value: Value) -> Map<String, Value> {
|
||||
|
|
@ -156,8 +154,6 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq
|
|||
assert_eq!(response.usage.total_tokens, 15);
|
||||
}
|
||||
|
||||
/// The provider already answered and billed these, so the host must not retry them on
|
||||
/// its own path: they surface as `InvalidResponse`, never as a pre-send decline.
|
||||
#[rstest]
|
||||
#[case::missing_usage(
|
||||
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#
|
||||
|
|
@ -209,10 +205,9 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
|
|||
);
|
||||
}
|
||||
|
||||
/// Nothing was sent, so nothing was billed and the host can still serve the request.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
||||
async fn a_connection_that_is_never_established_returns_a_connect_error(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
|
|
@ -230,9 +225,7 @@ async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
|||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
async fn a_timeout_after_sending_returns_a_network_error(request: ChatCompletionsRequest<'static>) {
|
||||
let upstream =
|
||||
upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await;
|
||||
let base = upstream.uri();
|
||||
|
|
@ -252,74 +245,141 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::accepted("anthropic/claude-sonnet-4-5", None, hi(), json!({"max_tokens": 16}), None)]
|
||||
#[case::accepted_bedrock("bedrock/anthropic.claude-sonnet-4-5", None, hi(), json!({}), None)]
|
||||
#[case::unknown_provider(
|
||||
"gpt-4o",
|
||||
Some("openai"),
|
||||
hi(),
|
||||
json!({}),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
)]
|
||||
#[case::unreadable_messages(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("hi"),
|
||||
json!({}),
|
||||
Some("unreadable message list")
|
||||
)]
|
||||
#[case::empty_messages("anthropic/claude-sonnet-4-5", None, json!([]), json!({}), Some("empty message list"))]
|
||||
#[case::streaming(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"stream": true}),
|
||||
Some("streaming")
|
||||
)]
|
||||
#[case::unrecognized_param(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
hi(),
|
||||
json!({"not_a_param": 1}),
|
||||
Some("unrecognized request parameter")
|
||||
)]
|
||||
#[case::opens_on_assistant_turn(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "assistant", "content": "hi"}]),
|
||||
json!({}),
|
||||
Some("conversation does not open on a user turn")
|
||||
)]
|
||||
fn decline_reason_names_why_the_core_would_not_serve_the_request(
|
||||
#[case] model: &str,
|
||||
#[case] provider: Option<&str>,
|
||||
#[case] messages: Value,
|
||||
#[case] params: Value,
|
||||
#[case] reason: Option<&str>,
|
||||
#[case::direct(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
#[case] hosted: bool,
|
||||
) {
|
||||
assert_eq!(
|
||||
chat_completions_decline_reason(model, provider, messages, &object(params)),
|
||||
reason
|
||||
use litellm_core::chat_completions::route::ChatCompletions;
|
||||
use litellm_host::{call::HostedCompletion, event::CallEvent};
|
||||
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
let host = RecordingCall::<ChatCompletions>::new(
|
||||
ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
}
|
||||
.into(),
|
||||
);
|
||||
let response = if hosted {
|
||||
let result = litellm_host::in_process::run_hosted(
|
||||
chat_completions_route().machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let HostedCompletion::Complete(response) = result else {
|
||||
panic!("expected a complete response")
|
||||
};
|
||||
response
|
||||
} else {
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
chat_completions_route()
|
||||
.execute(
|
||||
ChatCompletionsRequest {
|
||||
model: &call.model,
|
||||
messages: call.messages,
|
||||
optional_params: call.optional_params,
|
||||
api_key: call.api_key.as_deref(),
|
||||
api_base: call.api_base.as_deref(),
|
||||
custom_llm_provider: call.custom_llm_provider.as_deref(),
|
||||
extra_headers: call.extra_headers,
|
||||
timeout: call.timeout,
|
||||
},
|
||||
&host,
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
};
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello")
|
||||
);
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("x-hook"),
|
||||
Some("called")
|
||||
);
|
||||
let events = host.events.0.lock().unwrap();
|
||||
assert!(matches!(
|
||||
&events[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Machine(_),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
}
|
||||
|
||||
/// A request the decline check accepts must not be declined by the call itself.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_declined_request_fails_the_call_before_sending(
|
||||
async fn a_post_call_hook_failure_never_looks_safe_to_retry(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RequestContext, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
struct FailingHook;
|
||||
impl RouteHooks<Error> for FailingHook {
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
_: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
Ok(wire)
|
||||
}
|
||||
async fn on_event(&self, _: MachineEvent) -> Result<(), Error> {
|
||||
Err(Error::InvalidRequest("callback rejected".into()))
|
||||
}
|
||||
}
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
let error = chat_completions_route()
|
||||
.execute(
|
||||
ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
},
|
||||
&FailingHook,
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
let Error::PostCallHook(source) = error else {
|
||||
panic!("expected retained callback error")
|
||||
};
|
||||
assert_eq!(*source, Error::InvalidRequest("callback rejected".into()));
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn completed_chat_records_route_and_resolved_provider(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
optional_params: object(json!({"stream": true})),
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
.await
|
||||
.expect_err("streaming is declined");
|
||||
|
||||
assert_eq!(error, Error::Unsupported("streaming"));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
let model = request.model;
|
||||
traces
|
||||
.logger()
|
||||
.instrument(complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
}))
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "chat_completions");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(summaries[0]["provider"], "anthropic");
|
||||
assert_eq!(
|
||||
summaries[0]["resolved_model"],
|
||||
only_request(&upstream).await.json()["model"]
|
||||
);
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,18 +1,15 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
use std::sync::Mutex;
|
||||
|
||||
use litellm_core::messages::route::Messages;
|
||||
use litellm_host::{
|
||||
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
|
||||
host::Host,
|
||||
};
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
|
||||
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
|
||||
/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps
|
||||
/// Projects like `LocalMessagesHost`, answers `before_provider_request` through `rewrite`, and keeps
|
||||
/// every event the driver emits.
|
||||
struct RecordingHost {
|
||||
call: LocalMessagesHost,
|
||||
|
|
@ -50,19 +47,32 @@ impl RecordingHost {
|
|||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for RecordingHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.project().await
|
||||
impl RecordingHost {
|
||||
pub fn request(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.request()
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, ()> {
|
||||
litellm_host::in_process::Host {
|
||||
services: &(),
|
||||
hooks: self,
|
||||
stream: &(),
|
||||
observer: Some(self),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
impl litellm_host::lifecycle::CallObserver for RecordingHost {
|
||||
fn observe(&self, event: litellm_host::event::CallEvent) {
|
||||
self.events.lock().unwrap().push(event.clone());
|
||||
}
|
||||
}
|
||||
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
|
||||
for RecordingHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
context: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
self.optional_params
|
||||
.lock()
|
||||
|
|
@ -70,15 +80,24 @@ impl Host<Messages> for RecordingHost {
|
|||
.push(context.optional_params.clone());
|
||||
(self.rewrite)(wire)
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
self.events.lock().unwrap().push(event.clone());
|
||||
async fn on_event(
|
||||
&self,
|
||||
event: litellm_host::event::MachineEvent,
|
||||
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::event::CallEvent::Machine(event),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
litellm_host::in_process::run_hosted(
|
||||
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
|
|
@ -129,8 +148,7 @@ async fn a_before_send_failure_never_sends(call: MessagesCall) {
|
|||
|
||||
let error = run_through(&host)
|
||||
.await
|
||||
.err()
|
||||
.expect("the host failure fails the call");
|
||||
.expect_err("the host failure fails the call");
|
||||
|
||||
assert_eq!(error, Error::InvalidRequest("vetoed by the host".into()));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
|
@ -146,7 +164,7 @@ async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall)
|
|||
|
||||
let output = run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
assert!(matches!(output, MessagesOutput::Complete(_)));
|
||||
let [emitted] = <[String; 1]>::try_from(host.raw_responses())
|
||||
.unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len()));
|
||||
assert_eq!(serde_json::from_str::<Value>(&emitted).unwrap(), raw);
|
||||
|
|
@ -198,6 +216,6 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
|
|||
run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
.unwrap_or_else(|seen| panic!("before_provider_request runs once, saw {}", seen.len()));
|
||||
assert_eq!(optional_params, json!({"max_tokens": 16}));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
use std::{
|
||||
sync::{Arc, Mutex},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use litellm_core::messages::{
|
||||
Error, MessagesCall, MessagesShaping,
|
||||
route::{LocalMessagesHost, MessagesMachine, MessagesOutput, messages_machine},
|
||||
route::{Messages, MessagesMachine, MessagesOutput},
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
|
@ -95,16 +98,16 @@ fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Ma
|
|||
)
|
||||
}
|
||||
|
||||
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
messages_machine(&support::resources(), &http_config(), secrets)
|
||||
.expect("default HTTP settings build a client")
|
||||
fn machine(secrets: Arc<dyn SecretSource>) -> impl FnOnce(MessagesCall) -> MessagesMachine {
|
||||
move |request| messages_route(secrets).machine(request)
|
||||
}
|
||||
|
||||
async fn run_with(
|
||||
secrets: Arc<RecordingSecrets>,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
|
||||
let host = LocalMessagesHost::new(call);
|
||||
litellm_host::in_process::run_hosted(machine(secrets)(host.request()?), host.runtime()).await
|
||||
}
|
||||
|
||||
/// Runs the route with a secret source that knows nothing, so no environment leaks in.
|
||||
|
|
@ -114,7 +117,67 @@ async fn run(call: MessagesCall) -> Result<MessagesOutput, Error> {
|
|||
|
||||
async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse {
|
||||
match run(call).await.expect("messages call succeeds") {
|
||||
MessagesOutput::Message(message) => *message,
|
||||
MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"),
|
||||
MessagesOutput::Complete(message) => *message,
|
||||
MessagesOutput::StreamEnded | MessagesOutput::Detached => {
|
||||
panic!("a non-streaming call returned a stream")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct LocalMessagesHost {
|
||||
call: Mutex<Option<MessagesCall>>,
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
fn new(call: MessagesCall) -> Self {
|
||||
Self {
|
||||
call: Mutex::new(Some(call)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalMessagesHost {
|
||||
pub fn request(&self) -> Result<MessagesCall, Error> {
|
||||
self.call
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
|
||||
}
|
||||
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, ()> {
|
||||
litellm_host::in_process::Host {
|
||||
services: &(),
|
||||
hooks: self,
|
||||
stream: &(),
|
||||
observer: Some(self),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_host::lifecycle::CallObserver for LocalMessagesHost {
|
||||
fn observe(&self, _: litellm_host::event::CallEvent) {}
|
||||
}
|
||||
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
|
||||
for LocalMessagesHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: litellm_host::event::WireRequest,
|
||||
_: litellm_host::event::RequestContext,
|
||||
) -> Result<
|
||||
litellm_host::event::WireRequest,
|
||||
<Messages as litellm_host::protocol::Protocol>::Error,
|
||||
> {
|
||||
Ok(wire)
|
||||
}
|
||||
async fn on_event(
|
||||
&self,
|
||||
event: litellm_host::event::MachineEvent,
|
||||
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::event::CallEvent::Machine(event),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_llms::anthropic::common_utils::{AnthropicModelCapabilities, SupportedEffortTiers};
|
||||
use litellm_llms::base_llm::messages::context::{MessagesModelCapabilities, SupportedEffortTiers};
|
||||
use litellm_types::llms::anthropic::{AnthropicBeta, BetaSet};
|
||||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use rstest::rstest;
|
||||
|
|
@ -87,8 +87,7 @@ async fn a_call_without_credentials_fails_before_sending(
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("a call without credentials fails");
|
||||
.expect_err("a call without credentials fails");
|
||||
|
||||
assert!(
|
||||
matches!(
|
||||
|
|
@ -158,8 +157,7 @@ async fn unsupported_providers_are_rejected_before_sending(
|
|||
..with_model(call, model)
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("unsupported provider errors");
|
||||
.expect_err("unsupported provider errors");
|
||||
|
||||
assert_eq!(error, Error::InvalidProvider(reported.into()));
|
||||
}
|
||||
|
|
@ -194,12 +192,18 @@ async fn caller_headers_and_provider_scoped_headers_are_forwarded(call: Messages
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::azure("azure_ai", json!({"type": "ephemeral", "ttl": "1h", "future": "kept"}))]
|
||||
#[case::anthropic("anthropic", json!({"type": "ephemeral", "ttl": "1h", "scope": "global", "future": "kept"}))]
|
||||
#[tokio::test]
|
||||
async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCall) {
|
||||
async fn cache_scope_removal_is_selected_by_the_provider(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk-azure".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
body: body(json!({
|
||||
|
|
@ -210,7 +214,7 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa
|
|||
"content": [{
|
||||
"type": "text",
|
||||
"text": "hi",
|
||||
"cache_control": {"type": "ephemeral", "scope": "global"}
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h", "scope": "global", "future": "kept"}
|
||||
}]
|
||||
}]
|
||||
})),
|
||||
|
|
@ -220,7 +224,7 @@ async fn azure_strips_the_cache_control_scope_anthropic_rejects(call: MessagesCa
|
|||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.json()["messages"][0]["content"][0]["cache_control"],
|
||||
json!({"type": "ephemeral"})
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -278,9 +282,9 @@ async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
|
|||
#[case] features: &[AnthropicBeta],
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let capabilities = AnthropicModelCapabilities {
|
||||
let capabilities = MessagesModelCapabilities {
|
||||
supports_speed: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
..MessagesModelCapabilities::default()
|
||||
};
|
||||
|
||||
run_message(with_fields(
|
||||
|
|
@ -358,20 +362,20 @@ async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[cas
|
|||
);
|
||||
}
|
||||
|
||||
fn sampling_removed() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
fn sampling_removed() -> MessagesModelCapabilities {
|
||||
MessagesModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..AnthropicModelCapabilities::default()
|
||||
..MessagesModelCapabilities::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::sampling_params(sampling_removed(), json!({"temperature": 0.2, "top_p": 0.9, "top_k": 5}), &["temperature", "top_p", "top_k"], "temperature=0.2")]
|
||||
#[case::speed(AnthropicModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")]
|
||||
#[case::speed(MessagesModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")]
|
||||
#[tokio::test]
|
||||
async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_it(
|
||||
call: MessagesCall,
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] capabilities: MessagesModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] dropped: &[&str],
|
||||
#[case] rejected_as: &str,
|
||||
|
|
@ -399,10 +403,9 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
|
|||
|
||||
let error = run(shaped(false))
|
||||
.await
|
||||
.err()
|
||||
.expect("an unsupported param is rejected without drop_params");
|
||||
.expect_err("an unsupported param is rejected without drop_params");
|
||||
assert!(
|
||||
matches!(&error, Error::InvalidRequest(message) if message.contains(rejected_as)),
|
||||
matches!(&error, Error::InvalidRequest(message) if message.to_string().contains(rejected_as)),
|
||||
"{error:?}"
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
|
@ -431,10 +434,10 @@ async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
|
|||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
capabilities: MessagesModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
..MessagesModelCapabilities::default()
|
||||
},
|
||||
reasoning_auto_summary: true,
|
||||
..MessagesShaping::default()
|
||||
|
|
@ -450,41 +453,41 @@ async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
|
|||
|
||||
#[rstest]
|
||||
#[case::reasoning_effort_on_an_adaptive_model(
|
||||
AnthropicModelCapabilities {
|
||||
MessagesModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_output_config: true,
|
||||
effort_tiers: SupportedEffortTiers { high: true, ..SupportedEffortTiers::default() },
|
||||
..AnthropicModelCapabilities::default()
|
||||
..MessagesModelCapabilities::default()
|
||||
},
|
||||
json!({"reasoning_effort": "high"}),
|
||||
json!({"thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}})
|
||||
)]
|
||||
#[case::reasoning_effort_on_a_legacy_model_caps_the_budget_below_max_tokens(
|
||||
AnthropicModelCapabilities {
|
||||
MessagesModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
..MessagesModelCapabilities::default()
|
||||
},
|
||||
json!({"reasoning_effort": "high"}),
|
||||
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
|
||||
)]
|
||||
#[case::adaptive_payload_on_a_legacy_model_becomes_a_capped_budget(
|
||||
AnthropicModelCapabilities {
|
||||
MessagesModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
..MessagesModelCapabilities::default()
|
||||
},
|
||||
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}, "temperature": 0}),
|
||||
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
|
||||
)]
|
||||
#[case::adaptive_payload_on_a_model_without_reasoning_is_dropped(
|
||||
AnthropicModelCapabilities::default(),
|
||||
MessagesModelCapabilities::default(),
|
||||
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}),
|
||||
json!({})
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn reasoning_is_translated_by_the_model_capabilities(
|
||||
call: MessagesCall,
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] capabilities: MessagesModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
|
|
@ -556,12 +559,15 @@ async fn replayed_history_is_cleaned_before_sending(
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn metadata_is_reduced_to_the_user_id(call: MessagesCall) {
|
||||
async fn metadata_is_reduced_to_the_user_id(call: MessagesCall, #[case] provider: &str) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
|
|
@ -592,23 +598,43 @@ async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fie
|
|||
fields,
|
||||
))
|
||||
.await
|
||||
.err()
|
||||
.expect("the request is rejected");
|
||||
.expect_err("the request is rejected");
|
||||
|
||||
assert!(error.is_request(), "{error:?}");
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::azure("azure_ai", &[], json!([
|
||||
{"type": "text", "text": "top level"},
|
||||
{"type": "text", "text": "from a message"}
|
||||
]), json!([{"role": "user", "content": "hi"}]))]
|
||||
#[case::anthropic("anthropic", &[], json!("top level"), json!([
|
||||
{"role": "system", "content": "from a message"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]))]
|
||||
#[case::azure_folds_after_caller_drops("azure_ai", &["system"], json!([
|
||||
{"type": "text", "text": "from a message"}
|
||||
]), json!([{"role": "user", "content": "hi"}]))]
|
||||
#[tokio::test]
|
||||
async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesCall) {
|
||||
async fn system_message_folding_is_selected_by_the_provider(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
#[case] drop_params: &[&str],
|
||||
#[case] expected_system: Value,
|
||||
#[case] expected_messages: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk-azure".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
additional_drop_params: drop_params.iter().map(ToString::to_string).collect(),
|
||||
..call.shaping
|
||||
},
|
||||
..call
|
||||
},
|
||||
json!({
|
||||
|
|
@ -622,14 +648,8 @@ async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesC
|
|||
.await;
|
||||
|
||||
let sent = only_request(&upstream).await.json();
|
||||
assert_eq!(
|
||||
sent["system"],
|
||||
json!([
|
||||
{"type": "text", "text": "top level"},
|
||||
{"type": "text", "text": "from a message"}
|
||||
])
|
||||
);
|
||||
assert_eq!(sent["messages"], json!([{"role": "user", "content": "hi"}]));
|
||||
assert_eq!(sent["system"], expected_system);
|
||||
assert_eq!(sent["messages"], expected_messages);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -656,3 +676,37 @@ async fn the_provider_prefix_is_stripped_exactly_once(
|
|||
|
||||
assert_eq!(only_request(&upstream).await.json()["model"], sent_model);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure("azure_ai")]
|
||||
#[case::bedrock("bedrock")]
|
||||
#[tokio::test]
|
||||
async fn provider_validation_runs_before_caller_parameter_removal(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let result = run(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk-test".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
additional_drop_params: vec!["metadata".into()],
|
||||
..call.shaping
|
||||
},
|
||||
..call
|
||||
},
|
||||
json!({"metadata": {"user_id": 7}}),
|
||||
))
|
||||
.await;
|
||||
let error = result.expect_err("metadata is validated before removal");
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("metadata.user_id must be a string"),
|
||||
"{error}"
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,12 +1,62 @@
|
|||
use litellm_core::{
|
||||
Phase,
|
||||
messages::{MessagesResponse, messages, messages_body},
|
||||
};
|
||||
use litellm_core::messages::{MessagesResponse, messages_body};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::without_hooks(false)]
|
||||
#[case::with_hooks(true)]
|
||||
#[tokio::test]
|
||||
async fn calls_defer_execution_until_polled(call: MessagesCall, #[case] with_hooks: bool) {
|
||||
use futures_util::future::BoxFuture;
|
||||
|
||||
use litellm_host::event::CallEvent;
|
||||
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let secrets = Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "test-key")]));
|
||||
let route = messages_route(secrets.clone());
|
||||
let host = RecordingCall::<Messages>::new(MessagesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
});
|
||||
let request = host.request().unwrap();
|
||||
let future: BoxFuture<'_, Result<MessagesResponse, Error>> = if with_hooks {
|
||||
Box::pin(route.execute(request, &host))
|
||||
} else {
|
||||
Box::pin(route.execute(request, &()))
|
||||
};
|
||||
|
||||
assert!(secrets.requested().is_empty());
|
||||
assert!(host.events.0.lock().unwrap().is_empty());
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
||||
let MessagesResponse::Complete(response) = future.await.unwrap() else {
|
||||
panic!("expected a completed message");
|
||||
};
|
||||
assert_eq!(
|
||||
response.content,
|
||||
message_body()["content"].as_array().unwrap().as_slice()
|
||||
);
|
||||
assert!(secrets.requested().contains(&"ANTHROPIC_API_KEY".into()));
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.header("x-api-key"), Some("test-key"));
|
||||
assert_eq!(sent.header("x-hook"), with_hooks.then_some("called"));
|
||||
let events = host.events.0.lock().unwrap();
|
||||
if with_hooks {
|
||||
assert!(matches!(
|
||||
&events[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Machine(_),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
} else {
|
||||
assert!(events.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
|
|
@ -75,8 +125,7 @@ async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) {
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
let Error::Transport(TransportError::Http { status, body }) = error else {
|
||||
panic!("{error:?}");
|
||||
|
|
@ -97,8 +146,7 @@ async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
|
|
@ -126,8 +174,7 @@ async fn an_upstream_error_keeps_its_status_and_body(call: MessagesCall, #[case]
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
|
|
@ -154,10 +201,9 @@ async fn an_unreadable_success_body_is_an_invalid_response(
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("an unreadable body fails");
|
||||
.expect_err("an unreadable body fails");
|
||||
|
||||
assert_eq!(error.phase(), Phase::AfterSend, "{error:?}");
|
||||
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -172,8 +218,7 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) {
|
|||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("the call times out");
|
||||
.expect_err("the call times out");
|
||||
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
}
|
||||
|
|
@ -188,20 +233,24 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
|
|||
..HttpSettings::default()
|
||||
};
|
||||
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&Resolution::from(&settings).config,
|
||||
&RecordingSecrets::empty(),
|
||||
let resources = support::resources();
|
||||
let response = litellm_core::messages::MessagesRoute::new(
|
||||
provider_http(&resources, &Resolution::from(&settings).config),
|
||||
resources.auth,
|
||||
no_secrets(),
|
||||
)
|
||||
.execute(
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(base),
|
||||
..call
|
||||
},
|
||||
&(),
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
let MessagesResponse::Message(message) = response else {
|
||||
let MessagesResponse::Complete(message) = response else {
|
||||
panic!("a non-streaming request returns a message");
|
||||
};
|
||||
assert_eq!(message.id, "msg_1");
|
||||
|
|
@ -217,7 +266,38 @@ fn a_body_that_does_not_parse_is_an_invalid_request(#[case] raw: Value) {
|
|||
let error = messages_body(object(raw)).expect_err("the body is rejected");
|
||||
|
||||
assert!(
|
||||
matches!(&error, Error::InvalidRequest(message) if message.starts_with("invalid Anthropic messages request: ")),
|
||||
matches!(&error, Error::InvalidRequest(message) if message.to_string().starts_with("invalid Anthropic messages request: ")),
|
||||
"{error:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn message_route_summary_excludes_payload_diagnostics(
|
||||
call: MessagesCall,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let model = call.body.model.clone();
|
||||
traces
|
||||
.logger()
|
||||
.instrument(run_message(MessagesCall {
|
||||
api_key: Some("private-key-sentinel".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
}))
|
||||
.await;
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "messages");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(
|
||||
summaries[0]["resolved_model"],
|
||||
only_request(&upstream).await.json()["model"]
|
||||
);
|
||||
assert_eq!(summaries[0]["provider"], "anthropic");
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
assert!(summaries[0].get("body").is_none());
|
||||
assert!(!format!("{:?}", traces.records()).contains("private-key-sentinel"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ async fn the_credential_and_base_come_from_the_secret_source(
|
|||
.await
|
||||
.expect("messages call succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
assert!(matches!(output, MessagesOutput::Complete(_)));
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.url.path(), path);
|
||||
assert_eq!(request.header("x-api-key"), Some("sk-from-manager"));
|
||||
|
|
@ -90,8 +90,7 @@ async fn a_secret_manager_failure_fails_the_call_before_sending(call: MessagesCa
|
|||
},
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("a secret manager failure fails the call");
|
||||
.expect_err("a secret manager failure fails the call");
|
||||
|
||||
assert!(
|
||||
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
|
||||
|
|
@ -193,8 +192,13 @@ async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall)
|
|||
},
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("azure needs a base");
|
||||
.expect_err("azure needs a base");
|
||||
|
||||
assert_eq!(error, Error::Auth(litellm_auth::Error::MissingAzureApiBase));
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Auth(litellm_auth::Error::MissingApiBase {
|
||||
provider: "Azure",
|
||||
guidance: "Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
|
||||
})
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,15 +1,12 @@
|
|||
use std::{
|
||||
convert::Infallible,
|
||||
sync::{Mutex, mpsc},
|
||||
};
|
||||
use std::sync::{Mutex, mpsc};
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt};
|
||||
use litellm_core::messages::{
|
||||
MessagesResponse, messages,
|
||||
MessagesResponse,
|
||||
route::{Messages, MessagesStreamHead},
|
||||
};
|
||||
use litellm_host::host::{Demand, Host};
|
||||
use litellm_host::protocol::Demand;
|
||||
use litellm_tracing::{Logger, Metadata, Record, Sink};
|
||||
use rstest::rstest;
|
||||
use tokio::{
|
||||
|
|
@ -73,23 +70,55 @@ impl RecordingStreamHost {
|
|||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for RecordingStreamHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.project().await
|
||||
impl RecordingStreamHost {
|
||||
pub fn request(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.request()
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, Self> {
|
||||
litellm_host::in_process::Host {
|
||||
services: &(),
|
||||
hooks: self,
|
||||
stream: self,
|
||||
observer: Some(self),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn open(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
|
||||
impl litellm_host::in_process::StreamConsumer<Messages> for RecordingStreamHost {
|
||||
async fn open_stream(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
|
||||
Ok(self.record(Seen::Open(head.headers)))
|
||||
}
|
||||
|
||||
async fn deliver(&self, chunk: Bytes) -> Result<Demand, Error> {
|
||||
async fn send_chunk(&self, chunk: Bytes) -> Result<Demand, Error> {
|
||||
Ok(self.record(Seen::Deliver(chunk)))
|
||||
}
|
||||
}
|
||||
impl litellm_host::lifecycle::CallObserver for RecordingStreamHost {
|
||||
fn observe(&self, _: litellm_host::event::CallEvent) {}
|
||||
}
|
||||
impl litellm_host::hooks::RouteHooks<<Messages as litellm_host::protocol::Protocol>::Error>
|
||||
for RecordingStreamHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: litellm_host::event::WireRequest,
|
||||
_: litellm_host::event::RequestContext,
|
||||
) -> Result<
|
||||
litellm_host::event::WireRequest,
|
||||
<Messages as litellm_host::protocol::Protocol>::Error,
|
||||
> {
|
||||
Ok(wire)
|
||||
}
|
||||
async fn on_event(
|
||||
&self,
|
||||
event: litellm_host::event::MachineEvent,
|
||||
) -> Result<(), <Messages as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::event::CallEvent::Machine(event),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn streaming(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
MessagesCall {
|
||||
|
|
@ -107,7 +136,11 @@ fn sse_response() -> ResponseTemplate {
|
|||
}
|
||||
|
||||
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
litellm_host::in_process::run_hosted(
|
||||
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -118,7 +151,7 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me
|
|||
|
||||
let outcome = stream_through(&host).await.expect("streamed call succeeds");
|
||||
|
||||
assert!(matches!(outcome, MessagesOutput::Streamed));
|
||||
assert!(matches!(outcome, MessagesOutput::StreamEnded));
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
let [Seen::Open(headers), chunks @ ..] = seen.as_slice() else {
|
||||
panic!("the stream opens before any chunk is delivered");
|
||||
|
|
@ -186,7 +219,7 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det
|
|||
.await
|
||||
.expect("a detached stream still completes");
|
||||
|
||||
assert!(matches!(outcome, MessagesOutput::Streamed));
|
||||
assert!(matches!(outcome, MessagesOutput::Detached));
|
||||
assert_eq!(host.seen.into_inner().unwrap().len(), detach_after);
|
||||
}
|
||||
|
||||
|
|
@ -196,6 +229,7 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det
|
|||
status_response(429, json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}})),
|
||||
r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"#
|
||||
)]
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_error_fails_the_call_without_opening_the_stream(
|
||||
call: MessagesCall,
|
||||
|
|
@ -207,8 +241,7 @@ async fn an_upstream_error_fails_the_call_without_opening_the_stream(
|
|||
|
||||
let error = stream_through(&host)
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
.expect_err("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
|
|
@ -280,8 +313,7 @@ async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
|
|||
let error = tokio::time::timeout(Duration::from_secs(5), stream_through(&host))
|
||||
.await
|
||||
.expect("the stalled stream gives up within the timeout")
|
||||
.err()
|
||||
.expect("a stalled body fails the call");
|
||||
.expect_err("a stalled body fails the call");
|
||||
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
|
|
@ -305,23 +337,22 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
|
|||
#[case] provider: &str,
|
||||
) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
..streaming(call, upstream.uri())
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let response = messages_route(no_secrets())
|
||||
.execute(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
..streaming(call, upstream.uri())
|
||||
},
|
||||
&(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { headers, chunks } = response else {
|
||||
let MessagesResponse::Stream { head, chunks } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
};
|
||||
for (name, value) in UPSTREAM_HEADERS {
|
||||
assert!(headers.contains(&(name.into(), value.into())));
|
||||
assert!(head.headers.contains(&(name.into(), value.into())));
|
||||
}
|
||||
let delivered = chunks.try_collect::<Vec<_>>().await.unwrap().concat();
|
||||
assert_eq!(delivered, SSE_BODY.as_bytes());
|
||||
|
|
@ -332,15 +363,11 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
|
|||
#[tokio::test]
|
||||
async fn the_sdk_returns_http_errors_before_opening_a_stream(call: MessagesCall) {
|
||||
let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await;
|
||||
let error = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
streaming(call, upstream.uri()),
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream failure is returned by messages()");
|
||||
let error = messages_route(no_secrets())
|
||||
.execute(streaming(call, upstream.uri()), &())
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream failure is returned by messages()");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
|
|
@ -362,14 +389,12 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
|
|||
let (base, connection) = stalling_upstream().await;
|
||||
let response = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
messages_route(no_secrets()).execute(
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_secs(30)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
&(),
|
||||
),
|
||||
)
|
||||
.await
|
||||
|
|
@ -399,17 +424,16 @@ async fn dropping_the_sdk_stream_closes_the_unfinished_upstream(
|
|||
#[tokio::test]
|
||||
async fn the_sdk_yields_a_body_error_once_after_delivered_chunks(call: MessagesCall) {
|
||||
let (base, connection) = stalling_upstream().await;
|
||||
let response = messages(
|
||||
&support::resources(),
|
||||
&http_config(),
|
||||
&RecordingSecrets::empty(),
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_millis(300)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let response = messages_route(no_secrets())
|
||||
.execute(
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_millis(300)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
&(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let MessagesResponse::Stream { mut chunks, .. } = response else {
|
||||
panic!("a streaming request returns a stream");
|
||||
|
|
@ -445,7 +469,7 @@ async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) {
|
|||
|
||||
let outcome = stream_through(&host).await.expect("azure streams");
|
||||
|
||||
assert!(matches!(outcome, MessagesOutput::Streamed));
|
||||
assert!(matches!(outcome, MessagesOutput::StreamEnded));
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
let delivered: Vec<u8> = seen
|
||||
.iter()
|
||||
|
|
|
|||
|
|
@ -224,7 +224,7 @@ async fn credential_precedence(
|
|||
numbered_token,
|
||||
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::MissingApiBase {
|
||||
provider: "Azure AI",
|
||||
environment_variable: "AZURE_AI_API_BASE",
|
||||
guidance: "Set AZURE_AI_API_BASE environment variable or pass api_base parameter",
|
||||
})),
|
||||
0
|
||||
)]
|
||||
|
|
@ -232,7 +232,7 @@ async fn credential_precedence(
|
|||
true,
|
||||
json!({"azure_ad_token": "oidc/assertion", "client_id": "client", "tenant_id": "tenant"}),
|
||||
numbered_token,
|
||||
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::UnsupportedOidcReference)),
|
||||
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::InvalidConfiguration(detail)) if detail.to_string() == "unsupported OIDC reference"),
|
||||
0
|
||||
)]
|
||||
#[case::empty_provider_token_ignores_static_token(
|
||||
|
|
|
|||
|
|
@ -183,16 +183,16 @@ async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() {
|
|||
"analyzeResult": {"pages": [{"pageNumber": 1, "width": 8.5, "height": 11, "unit": "inch"}]}
|
||||
}))])
|
||||
.await;
|
||||
let client = ocr_client().with_settings(OcrSettings {
|
||||
let route = ocr_route_with(OcrSettings {
|
||||
document_intelligence_api_version: "2099-01-01".into(),
|
||||
document_intelligence_dpi: 72,
|
||||
..OcrSettings::default()
|
||||
});
|
||||
|
||||
let result =
|
||||
litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({})))
|
||||
.await
|
||||
.unwrap();
|
||||
let result = route
|
||||
.execute(read_request(&upstream.uri(), json!({})), &())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream)
|
||||
|
|
@ -347,14 +347,14 @@ async fn the_polling_deadline_bounds_the_retry_delay() {
|
|||
],
|
||||
)
|
||||
.await;
|
||||
let client = ocr_client().with_settings(OcrSettings {
|
||||
let route = ocr_route_with(OcrSettings {
|
||||
poll_timeout: Duration::from_millis(100),
|
||||
..OcrSettings::default()
|
||||
});
|
||||
|
||||
let error = tokio::time::timeout(
|
||||
Duration::from_secs(1),
|
||||
litellm_core::ocr::client::perform(&client, read_request(&upstream.uri(), json!({}))),
|
||||
route.execute(read_request(&upstream.uri(), json!({})), &()),
|
||||
)
|
||||
.await
|
||||
.expect("the deadline cuts the retry delay short")
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ impl Route {
|
|||
}
|
||||
}
|
||||
|
||||
/// What the host does to the wire request in `before_send`.
|
||||
/// What the host does to the wire request in `before_provider_request`.
|
||||
#[derive(Clone, Copy, Debug)]
|
||||
enum Guardrail {
|
||||
Detached,
|
||||
|
|
@ -53,7 +53,7 @@ enum Guardrail {
|
|||
}
|
||||
|
||||
impl Guardrail {
|
||||
fn before_send(self, wire: WireRequest) -> WireRequest {
|
||||
fn before_provider_request(self, wire: WireRequest) -> WireRequest {
|
||||
let Value::Object(fields) = wire.body else {
|
||||
return wire;
|
||||
};
|
||||
|
|
@ -96,8 +96,8 @@ async fn provider_document(route: Route, guardrail: Guardrail) -> Value {
|
|||
json!({"type": document_type, document_type: format!("{}/scan.png", documents.uri())}),
|
||||
route.options(),
|
||||
);
|
||||
let host =
|
||||
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(guardrail.before_send(wire)));
|
||||
let host = LocalOcrHost::new(request)
|
||||
.with_before_send(move |wire, _| Ok(guardrail.before_provider_request(wire)));
|
||||
|
||||
perform_with(host).await.unwrap();
|
||||
|
||||
|
|
@ -139,6 +139,7 @@ async fn a_document_replaced_by_the_host_reaches_the_provider(
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn an_empty_byte_document_fails_before_sending() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
|
|
@ -156,6 +157,7 @@ async fn an_empty_byte_document_fails_before_sending() {
|
|||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn a_missing_path_document_fails_before_sending() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
|
|
@ -180,3 +182,51 @@ async fn a_missing_path_document_fails_before_sending() {
|
|||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::blocked(false)]
|
||||
#[case::allowed(true)]
|
||||
#[tokio::test]
|
||||
async fn configured_client_preserves_document_url_policy(#[case] allowed: bool) {
|
||||
let documents = document_server().await;
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
let document_url = format!("{}/scan.png", documents.uri());
|
||||
let authority = documents.address().to_string();
|
||||
let route = build_ocr_route(
|
||||
&resources(),
|
||||
&http_config(),
|
||||
litellm_http::media::UrlPolicy {
|
||||
validate: true,
|
||||
allowed_hosts: allowed.then_some(authority).into_iter().collect(),
|
||||
},
|
||||
Default::default(),
|
||||
no_secrets(),
|
||||
);
|
||||
let host = LocalOcrHost::new(ocr_request_with_document(
|
||||
"azure_ai/model",
|
||||
&upstream.uri(),
|
||||
json!({"type": "document_url", "document_url": document_url}),
|
||||
json!({}),
|
||||
));
|
||||
let result = litellm_host::in_process::run_hosted(
|
||||
route.machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await;
|
||||
|
||||
if !allowed {
|
||||
assert!(matches!(result, Err(Error::BlockedDocumentUrl)));
|
||||
assert!(received(&documents).await.is_empty());
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
return;
|
||||
}
|
||||
result.unwrap();
|
||||
assert_eq!(only_request(&documents).await.url.path(), "/scan.png");
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.json()["document"]["document_url"],
|
||||
format!(
|
||||
"data:image/png;base64,{}",
|
||||
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
|
||||
)
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,13 +1,10 @@
|
|||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use litellm_core::ocr::{
|
||||
route::{Ocr, OcrOp, OcrProjection, ocr_machine},
|
||||
route::{Ocr, OcrCall, OcrOp},
|
||||
types::OcrDocumentInput,
|
||||
};
|
||||
use litellm_host::{
|
||||
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
|
||||
host::Host,
|
||||
};
|
||||
use litellm_host::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
|
@ -18,6 +15,7 @@ pub(crate) fn event_name(event: &CallEvent) -> &'static str {
|
|||
CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response",
|
||||
CallEvent::Succeeded { .. } => "success",
|
||||
CallEvent::Failed { .. } => "failure",
|
||||
CallEvent::Cancelled { .. } => "cancelled",
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -29,7 +27,10 @@ fn recording_host(
|
|||
let before_send_events = events.clone();
|
||||
LocalOcrHost::new(request)
|
||||
.with_before_send(move |wire, _| {
|
||||
before_send_events.lock().unwrap().push("before_send");
|
||||
before_send_events
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push("before_provider_request");
|
||||
match block {
|
||||
true => Err(Error::InvalidRequest("blocked".into())),
|
||||
false => Ok(wire),
|
||||
|
|
@ -38,6 +39,7 @@ fn recording_host(
|
|||
.with_observer(move |event| events.lock().unwrap().push(event_name(event)))
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn hooks_run_in_order_and_one_success_is_emitted() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
|
|
@ -53,11 +55,12 @@ async fn hooks_run_in_order_and_one_success_is_emitted() {
|
|||
|
||||
assert_eq!(
|
||||
*events.lock().unwrap(),
|
||||
["started", "before_send", "response", "success"]
|
||||
["started", "before_provider_request", "response", "success"]
|
||||
);
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
|
|
@ -77,11 +80,12 @@ async fn a_blocking_before_send_prevents_the_call_and_emits_one_failure() {
|
|||
);
|
||||
assert_eq!(
|
||||
*events.lock().unwrap(),
|
||||
["started", "before_send", "failure"]
|
||||
["started", "before_provider_request", "failure"]
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_failure_emits_one_terminal_failure() {
|
||||
let upstream = upstream([status_response(500, json!({"error": "failed"}))]).await;
|
||||
|
|
@ -97,11 +101,12 @@ async fn an_upstream_failure_emits_one_terminal_failure() {
|
|||
assert!(result.is_err());
|
||||
assert_eq!(
|
||||
*events.lock().unwrap(),
|
||||
["started", "before_send", "failure"]
|
||||
["started", "before_provider_request", "failure"]
|
||||
);
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn an_invalid_provider_response_is_observed_before_normalization_fails() {
|
||||
let upstream = upstream([json_response(json!({"pages": "invalid"}))]).await;
|
||||
|
|
@ -120,6 +125,7 @@ async fn an_invalid_provider_response_is_observed_before_normalization_fails() {
|
|||
assert_eq!(*observed.lock().unwrap(), [r#"{"pages":"invalid"}"#]);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn headers_returned_by_before_send_are_sent() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
|
|
@ -147,9 +153,10 @@ async fn before_send_context(request: LiteLLMOcrRequest) -> (WireRequest, Reques
|
|||
});
|
||||
perform_with(host).await.unwrap();
|
||||
let context = observed.lock().unwrap().take();
|
||||
context.expect("before_send ran")
|
||||
context.expect("before_provider_request ran")
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn before_send_sees_the_route_its_params_and_the_body() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
|
|
@ -187,22 +194,31 @@ async fn before_send_names_the_secret_params(#[case] options: Value, #[case] sec
|
|||
assert_eq!(context.secret_fields, secrets);
|
||||
}
|
||||
|
||||
/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_send`.
|
||||
/// Hands the route a caller-owned Azure token and rewrites the bearer in `before_provider_request`.
|
||||
struct CallerTokenHost {
|
||||
request: Mutex<Option<LiteLLMOcrRequest>>,
|
||||
trace: Mutex<Vec<String>>,
|
||||
}
|
||||
|
||||
impl Host<Ocr> for CallerTokenHost {
|
||||
async fn project(&self) -> Result<OcrProjection, Error> {
|
||||
impl CallerTokenHost {
|
||||
pub fn request(&self) -> Result<OcrCall, Error> {
|
||||
self.trace.lock().unwrap().push("project".into());
|
||||
Ok(OcrProjection {
|
||||
Ok(OcrCall {
|
||||
request: self.request.lock().unwrap().take().unwrap(),
|
||||
caller_token: true,
|
||||
})
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
|
||||
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, Self, Self, ()> {
|
||||
litellm_host::in_process::Host {
|
||||
services: self,
|
||||
hooks: self,
|
||||
stream: &(),
|
||||
observer: Some(self),
|
||||
}
|
||||
}
|
||||
}
|
||||
impl litellm_host::services::HostCallHandler<Ocr> for CallerTokenHost {
|
||||
async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> {
|
||||
match op {
|
||||
OcrOp::AcquireAzureAdToken(reply) => {
|
||||
self.trace.lock().unwrap().push("token".into());
|
||||
|
|
@ -213,11 +229,18 @@ impl Host<Ocr> for CallerTokenHost {
|
|||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
impl litellm_host::lifecycle::CallObserver for CallerTokenHost {
|
||||
fn observe(&self, _: litellm_host::event::CallEvent) {}
|
||||
}
|
||||
impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::Error>
|
||||
for CallerTokenHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
_: &RequestContext,
|
||||
_: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization");
|
||||
let authorization = wire
|
||||
|
|
@ -229,7 +252,7 @@ impl Host<Ocr> for CallerTokenHost {
|
|||
self.trace
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(format!("before_send:{authorization}"));
|
||||
.push(format!("before_provider_request:{authorization}"));
|
||||
let headers = wire
|
||||
.headers
|
||||
.into_iter()
|
||||
|
|
@ -240,8 +263,19 @@ impl Host<Ocr> for CallerTokenHost {
|
|||
.collect();
|
||||
Ok(WireRequest { headers, ..wire })
|
||||
}
|
||||
async fn on_event(
|
||||
&self,
|
||||
event: litellm_host::event::MachineEvent,
|
||||
) -> Result<(), <Ocr as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::event::CallEvent::Machine(event),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_replace_it() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
|
|
@ -254,16 +288,116 @@ async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_
|
|||
trace: Mutex::new(Vec::new()),
|
||||
};
|
||||
|
||||
litellm_host::run::run(ocr_machine(ocr_client()), &host)
|
||||
.await
|
||||
.unwrap();
|
||||
litellm_host::in_process::run_hosted(
|
||||
ocr_route().machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
*host.trace.lock().unwrap(),
|
||||
["project", "token", "before_send:Bearer caller-token"]
|
||||
[
|
||||
"project",
|
||||
"token",
|
||||
"before_provider_request:Bearer caller-token"
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header_values("authorization"),
|
||||
["Bearer edited"]
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn direct_execution_uses_hooks_without_a_machine() {
|
||||
use litellm_host::{hooks::RouteHooks, lifecycle::CallObserver};
|
||||
|
||||
struct Hooks(Arc<super::support::CallEvents>);
|
||||
|
||||
impl RouteHooks<Error> for Hooks {
|
||||
fn observer(&self) -> Option<Arc<dyn CallObserver>> {
|
||||
Some(self.0.clone())
|
||||
}
|
||||
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
_: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
Ok(WireRequest {
|
||||
headers: wire
|
||||
.headers
|
||||
.into_iter()
|
||||
.chain([("x-direct-hook".into(), "called".into())])
|
||||
.collect(),
|
||||
..wire
|
||||
})
|
||||
}
|
||||
|
||||
async fn on_event(&self, event: MachineEvent) -> Result<(), Error> {
|
||||
self.0.observe(CallEvent::Machine(event));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
let upstream = upstream([json_response(
|
||||
json!({"pages":[{"index":0,"markdown":"direct"}]}),
|
||||
)])
|
||||
.await;
|
||||
let events = Arc::new(super::support::CallEvents::default());
|
||||
let route = ocr_route();
|
||||
let hooks = Hooks(events.clone());
|
||||
let builder = route.execute(
|
||||
ocr_request("mistral/model", &upstream.uri(), json!({})),
|
||||
&hooks,
|
||||
);
|
||||
assert!(events.0.lock().unwrap().is_empty());
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
let result = builder.await.unwrap();
|
||||
assert_eq!(result.pages[0].markdown, "direct");
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("x-direct-hook"),
|
||||
Some("called")
|
||||
);
|
||||
assert!(matches!(
|
||||
&events.0.lock().unwrap()[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Machine(_),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::native(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn ocr_records_one_route_summary_across_both_execution_paths(
|
||||
traces: TraceCapture,
|
||||
#[case] hosted: bool,
|
||||
) {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
let request = ocr_request("mistral/model", &upstream.uri(), json!({}));
|
||||
let model = request.model.clone();
|
||||
traces
|
||||
.logger()
|
||||
.instrument(async {
|
||||
if hosted {
|
||||
perform_with(LocalOcrHost::new(request)).await
|
||||
} else {
|
||||
perform(request).await
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["route"], "ocr");
|
||||
assert_eq!(summaries[0]["model"], model);
|
||||
assert_eq!(summaries[0]["provider"], "mistral");
|
||||
assert_eq!(summaries[0]["outcome"], "success");
|
||||
assert_eq!(summaries[0]["stream"], false);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_host::protocol::HookRequest;
|
||||
use std::{
|
||||
sync::{
|
||||
Arc,
|
||||
|
|
@ -7,13 +8,15 @@ use std::{
|
|||
};
|
||||
|
||||
use litellm_core::ocr::{
|
||||
route::{OcrMachine, OcrOp, OcrProjection},
|
||||
route::{OcrCall, OcrMachine, OcrOp},
|
||||
types::OcrDocumentInput,
|
||||
};
|
||||
use litellm_host::{
|
||||
event::{CallEvent, WireRequest},
|
||||
host::{Host, HostOp},
|
||||
hooks::RouteHooks,
|
||||
machine::{HostFailure, Machine, MachineStep},
|
||||
protocol::Suspension,
|
||||
services::HostCallHandler,
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::transformation::OcrTransportConfig;
|
||||
use rstest::rstest;
|
||||
|
|
@ -21,7 +24,7 @@ use tokio::{io::AsyncReadExt, net::TcpListener, sync::Notify};
|
|||
|
||||
use super::{lifecycle::event_name, *};
|
||||
|
||||
/// Drives the machine by hand, answering every op through `host` except `before_send`,
|
||||
/// Drives the machine by hand, answering every op through `host` except `before_provider_request`,
|
||||
/// which `intercept` answers so a test can fail or cancel exactly there.
|
||||
async fn drive_until(
|
||||
host: &LocalOcrHost,
|
||||
|
|
@ -31,43 +34,39 @@ async fn drive_until(
|
|||
Vec<&'static str>,
|
||||
OcrMachine,
|
||||
) {
|
||||
let mut machine = ocr_machine(ocr_client());
|
||||
let mut machine = ocr_route().machine(host.request().unwrap());
|
||||
let mut ops = Vec::new();
|
||||
let outcome = loop {
|
||||
let op = match machine.resume().await {
|
||||
Ok(MachineStep::Host(op)) => op,
|
||||
Ok(MachineStep::Complete(response)) => break Ok(response),
|
||||
Ok(MachineStep::Suspended(op)) => op,
|
||||
Ok(MachineStep::Complete(response)) => break Ok(completed(response)),
|
||||
Err(error) => break Err(error),
|
||||
};
|
||||
let answer = match op {
|
||||
HostOp::Project(reply) => {
|
||||
ops.push("Project");
|
||||
host.project()
|
||||
.await
|
||||
.map(|projection| reply.send(projection))
|
||||
.map_err(HostFailure::Error)
|
||||
}
|
||||
HostOp::Custom(op) => {
|
||||
Suspension::Stream(stream) => match stream {
|
||||
litellm_host::protocol::StreamDelivery::Open(head, _) => match head {},
|
||||
litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {},
|
||||
},
|
||||
Suspension::HostCall(op) => {
|
||||
ops.push(match op {
|
||||
OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken",
|
||||
});
|
||||
host.custom_op(op).await.map_err(HostFailure::Error)
|
||||
host.handle_host_call(op).await.map_err(HostFailure::Error)
|
||||
}
|
||||
HostOp::BeforeSend { wire, reply, .. } => {
|
||||
Suspension::Hook(HookRequest::BeforeProviderRequest { wire, reply, .. }) => {
|
||||
ops.push("BeforeSend");
|
||||
intercept(*wire).map(|wire| reply.send(wire))
|
||||
}
|
||||
HostOp::Emit(event, reply) => {
|
||||
let event = CallEvent::Machine(event);
|
||||
ops.push(event_name(&event));
|
||||
host.emit(&event)
|
||||
Suspension::Hook(HookRequest::Event(event, reply)) => {
|
||||
ops.push(event_name(&CallEvent::Machine(event.clone())));
|
||||
host.on_event(event)
|
||||
.await
|
||||
.map(|()| reply.send(()))
|
||||
.map_err(HostFailure::Error)
|
||||
}
|
||||
};
|
||||
if let Err(failure) = answer {
|
||||
break machine.interrupt(failure).await;
|
||||
break machine.interrupt(failure).await.map(completed);
|
||||
}
|
||||
};
|
||||
(outcome, ops, machine)
|
||||
|
|
@ -81,10 +80,13 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto
|
|||
_ = stop.notified() => break,
|
||||
step = machine.resume() => {
|
||||
match step.unwrap() {
|
||||
MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()),
|
||||
MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(),
|
||||
MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire),
|
||||
MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()),
|
||||
MachineStep::Suspended(Suspension::HostCall(op)) => host.handle_host_call(op).await.unwrap(),
|
||||
MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest { wire, reply, .. })) => reply.send(*wire),
|
||||
MachineStep::Suspended(Suspension::Hook(HookRequest::Event(_, reply))) => reply.send(()),
|
||||
MachineStep::Suspended(Suspension::Stream(stream)) => match stream {
|
||||
litellm_host::protocol::StreamDelivery::Open(head, _) => match head {},
|
||||
litellm_host::protocol::StreamDelivery::Chunk(chunk, _) => match chunk {},
|
||||
},
|
||||
MachineStep::Complete(_) => panic!("the stalled call completed"),
|
||||
}
|
||||
}
|
||||
|
|
@ -95,6 +97,7 @@ async fn drive_until_notified(machine: &mut OcrMachine, host: &LocalOcrHost, sto
|
|||
.expect("the call reached the stall point");
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn a_hand_driven_machine_performs_the_same_call() {
|
||||
let upstream = upstream([json_response(json!({
|
||||
|
|
@ -107,13 +110,14 @@ async fn a_hand_driven_machine_performs_the_same_call() {
|
|||
|
||||
assert_eq!(outcome.unwrap().pages[0].markdown, "native");
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
|
||||
assert_eq!(ops, ["BeforeSend", "response"]);
|
||||
assert!(matches!(
|
||||
machine.resume().await,
|
||||
Err(Error::InvalidRequest(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn a_path_document_is_read_by_core_without_a_host_operation() {
|
||||
let upstream = upstream([json_response(json!({
|
||||
|
|
@ -135,7 +139,7 @@ async fn a_path_document_is_read_by_core_without_a_host_operation() {
|
|||
std::fs::remove_dir_all(&dir).unwrap();
|
||||
|
||||
assert_eq!(response.unwrap().pages[0].markdown, "path");
|
||||
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
|
||||
assert_eq!(ops, ["BeforeSend", "response"]);
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.json()["document"]["image_url"],
|
||||
"data:image/png;base64,YWJj"
|
||||
|
|
@ -143,7 +147,7 @@ async fn a_path_document_is_read_by_core_without_a_host_operation() {
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::failed(HostFailure::Error(Error::InvalidRequest("before_send failed".into())), "before_send failed")]
|
||||
#[case::failed(HostFailure::Error(Error::InvalidRequest("before_provider_request failed".into())), "before_provider_request failed")]
|
||||
#[case::cancelled(HostFailure::Cancelled(Error::InvalidRequest("cancelled".into())), "cancelled")]
|
||||
#[tokio::test]
|
||||
async fn a_before_send_failure_ends_the_call_without_reaching_transport(
|
||||
|
|
@ -159,7 +163,7 @@ async fn a_before_send_failure_ends_the_call_without_reaching_transport(
|
|||
.lock()
|
||||
.unwrap()
|
||||
.take()
|
||||
.expect("before_send is asked once"))
|
||||
.expect("before_provider_request is asked once"))
|
||||
})
|
||||
.await;
|
||||
|
||||
|
|
@ -167,27 +171,35 @@ async fn a_before_send_failure_ends_the_call_without_reaching_transport(
|
|||
matches!(&outcome, Err(Error::InvalidRequest(actual)) if actual == message),
|
||||
"{outcome:?}"
|
||||
);
|
||||
assert_eq!(ops, ["Project", "BeforeSend"]);
|
||||
assert_eq!(ops, ["BeforeSend"]);
|
||||
assert!(machine.resume().await.is_err());
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn resuming_before_answering_keeps_the_pending_operation() {
|
||||
let request = ocr_request("mistral/model", UNREACHABLE_BASE, json!({}));
|
||||
let mut machine = ocr_machine(ocr_client());
|
||||
let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else {
|
||||
panic!("expected the projection op first");
|
||||
};
|
||||
|
||||
assert!(machine.resume().await.is_err());
|
||||
reply.send(OcrProjection {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
let request = ocr_request("mistral/model", &upstream.uri(), json!({}));
|
||||
let mut machine = ocr_route().machine(OcrCall {
|
||||
request,
|
||||
caller_token: false,
|
||||
});
|
||||
let Ok(MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest {
|
||||
wire,
|
||||
reply,
|
||||
..
|
||||
}))) = machine.resume().await
|
||||
else {
|
||||
panic!("expected the provider request hook");
|
||||
};
|
||||
assert!(machine.resume().await.is_err());
|
||||
reply.send(*wire);
|
||||
assert!(matches!(
|
||||
machine.resume().await,
|
||||
Ok(MachineStep::Host(HostOp::BeforeSend { .. }))
|
||||
Ok(MachineStep::Suspended(Suspension::Hook(
|
||||
HookRequest::Event(_, _)
|
||||
)))
|
||||
));
|
||||
}
|
||||
|
||||
|
|
@ -215,6 +227,7 @@ impl litellm_auth::TokenProvider for PendingToken {
|
|||
}
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn interrupt_drops_provider_captures_before_returning() {
|
||||
let entered = Arc::new(Notify::new());
|
||||
|
|
@ -231,7 +244,7 @@ async fn interrupt_drops_provider_captures_before_returning() {
|
|||
},
|
||||
)));
|
||||
let host = LocalOcrHost::new(request);
|
||||
let mut machine = ocr_machine(ocr_client());
|
||||
let mut machine = ocr_route().machine(host.request().unwrap());
|
||||
|
||||
drive_until_notified(&mut machine, &host, &entered).await;
|
||||
assert!(!dropped.load(Ordering::SeqCst));
|
||||
|
|
@ -248,6 +261,7 @@ async fn interrupt_drops_provider_captures_before_returning() {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[tokio::test]
|
||||
async fn interrupting_an_in_flight_provider_request_closes_its_connection() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
|
|
@ -266,7 +280,7 @@ async fn interrupting_an_in_flight_provider_request_closes_its_connection() {
|
|||
while socket.read(&mut buffer).await.unwrap() != 0 {}
|
||||
});
|
||||
let host = LocalOcrHost::new(ocr_request("mistral/model", &base, json!({})));
|
||||
let mut machine = ocr_machine(ocr_client());
|
||||
let mut machine = ocr_route().machine(host.request().unwrap());
|
||||
|
||||
drive_until_notified(&mut machine, &host, &received).await;
|
||||
let cancelled = Error::InvalidRequest("cancelled".into());
|
||||
|
|
|
|||
|
|
@ -1,16 +1,18 @@
|
|||
use litellm_core::ocr::{
|
||||
OcrRoute,
|
||||
document::prepare_document,
|
||||
route::{LocalOcrHost, ocr_machine},
|
||||
types::LiteLLMOcrRequest,
|
||||
route::{Ocr, OcrCall, OcrOp},
|
||||
types::{LiteLLMOcrRequest, OcrDocumentInput},
|
||||
wire::{OcrWireRequest, decode_request},
|
||||
};
|
||||
use litellm_http::Client;
|
||||
use litellm_host::event::{CallEvent, RequestContext, WireRequest};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::OcrClient,
|
||||
settings::OcrSettings,
|
||||
transformation::{LiteLLMOcrResponse, OcrDocument},
|
||||
};
|
||||
use serde_json::{Map, Value, json};
|
||||
use std::sync::Mutex;
|
||||
use wiremock::{MockServer, ResponseTemplate};
|
||||
|
||||
#[path = "../support/mod.rs"]
|
||||
|
|
@ -37,16 +39,31 @@ fn object(value: Value) -> Map<String, Value> {
|
|||
map
|
||||
}
|
||||
|
||||
fn ocr_client() -> OcrClient {
|
||||
OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test())
|
||||
fn ocr_route() -> OcrRoute {
|
||||
ocr_route_with(OcrSettings::default())
|
||||
}
|
||||
|
||||
fn ocr_route_with(settings: OcrSettings) -> OcrRoute {
|
||||
build_ocr_route(
|
||||
&resources(),
|
||||
&http_config(),
|
||||
litellm_http::media::UrlPolicy {
|
||||
validate: false,
|
||||
allowed_hosts: Vec::new(),
|
||||
},
|
||||
settings,
|
||||
no_secrets(),
|
||||
)
|
||||
}
|
||||
|
||||
async fn perform(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_core::ocr::client::perform(&ocr_client(), request).await
|
||||
ocr_route().execute(request, &()).await
|
||||
}
|
||||
|
||||
async fn perform_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
|
||||
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
|
||||
litellm_host::in_process::run_hosted(ocr_route().machine(host.request()?), host.runtime())
|
||||
.await
|
||||
.map(completed)
|
||||
}
|
||||
|
||||
fn wire(model: &str, base: &str, document: Value, options: Value) -> OcrWireRequest {
|
||||
|
|
@ -120,3 +137,117 @@ fn accepted(server: &MockServer, body: Value) -> ResponseTemplate {
|
|||
.insert_header("Operation-Location", format!("{}/operation", server.uri()))
|
||||
.set_body_json(body)
|
||||
}
|
||||
|
||||
fn completed(
|
||||
result: litellm_host::call::HostedCompletion<LiteLLMOcrResponse>,
|
||||
) -> LiteLLMOcrResponse {
|
||||
match result {
|
||||
litellm_host::call::HostedCompletion::Complete(response) => response,
|
||||
other => panic!("unexpected OCR completion: {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
type BeforeSend =
|
||||
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
|
||||
|
||||
struct LocalOcrHost {
|
||||
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
|
||||
before_provider_request: Option<BeforeSend>,
|
||||
observer: Option<Observer>,
|
||||
}
|
||||
|
||||
impl LocalOcrHost {
|
||||
fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
|
||||
Self {
|
||||
request: Mutex::new(Some(request)),
|
||||
before_provider_request: None,
|
||||
observer: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn with_before_send(
|
||||
self,
|
||||
before_provider_request: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
) -> Self {
|
||||
Self {
|
||||
before_provider_request: Some(Box::new(before_provider_request)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
|
||||
fn with_observer(self, observer: impl Fn(&CallEvent) + Send + Sync + 'static) -> Self {
|
||||
Self {
|
||||
observer: Some(Box::new(observer)),
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOcrHost {
|
||||
pub fn request(&self) -> Result<OcrCall, Error> {
|
||||
self.request
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.take()
|
||||
.map(|request| OcrCall {
|
||||
request,
|
||||
caller_token: false,
|
||||
})
|
||||
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into()))
|
||||
}
|
||||
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, Self, Self, ()> {
|
||||
litellm_host::in_process::Host {
|
||||
services: self,
|
||||
hooks: self,
|
||||
stream: &(),
|
||||
observer: Some(self),
|
||||
}
|
||||
}
|
||||
}
|
||||
impl litellm_host::services::HostCallHandler<Ocr> for LocalOcrHost {
|
||||
async fn handle_host_call(&self, op: OcrOp) -> Result<(), Error> {
|
||||
match op {
|
||||
OcrOp::AcquireAzureAdToken(_) => {
|
||||
Err(Error::Auth(litellm_auth::Error::CredentialAcquisition(
|
||||
"OCR host has no Azure AD token provider".into(),
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_host::lifecycle::CallObserver for LocalOcrHost {
|
||||
fn observe(&self, event: litellm_host::event::CallEvent) {
|
||||
if let Some(observer) = &self.observer {
|
||||
observer(&event);
|
||||
}
|
||||
}
|
||||
}
|
||||
impl litellm_host::hooks::RouteHooks<<Ocr as litellm_host::protocol::Protocol>::Error>
|
||||
for LocalOcrHost
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
match &self.before_provider_request {
|
||||
Some(before_provider_request) => before_provider_request(wire, &context),
|
||||
None => Ok(wire),
|
||||
}
|
||||
}
|
||||
async fn on_event(
|
||||
&self,
|
||||
event: litellm_host::event::MachineEvent,
|
||||
) -> Result<(), <Ocr as litellm_host::protocol::Protocol>::Error> {
|
||||
litellm_host::lifecycle::CallObserver::observe(
|
||||
self,
|
||||
litellm_host::event::CallEvent::Machine(event),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,11 +1,8 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_http::{HttpSettings, Resolution, media::UrlPolicy};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
use litellm_llms::{
|
||||
base_llm::ocr::{
|
||||
settings::OcrSettings,
|
||||
transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES},
|
||||
},
|
||||
base_llm::ocr::transformation::{BaseOcrConfig, OCR_RESPONSE_MAX_BYTES},
|
||||
mistral::ocr::transformation::MistralOcrConfig,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
|
@ -156,7 +153,13 @@ async fn missing_credentials_come_from_the_injected_secret_source(
|
|||
.copied()
|
||||
.chain([("MISTRAL_AZURE_API_BASE", base.as_str())]),
|
||||
));
|
||||
let client = ocr_client().with_secrets(source.clone());
|
||||
let route = build_ocr_route(
|
||||
&resources(),
|
||||
&http_config(),
|
||||
Default::default(),
|
||||
Default::default(),
|
||||
source.clone(),
|
||||
);
|
||||
let request = decode_request(OcrWireRequest {
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
|
|
@ -169,9 +172,7 @@ async fn missing_credentials_come_from_the_injected_secret_source(
|
|||
})
|
||||
.unwrap();
|
||||
|
||||
litellm_core::ocr::client::perform(&client, request)
|
||||
.await
|
||||
.unwrap();
|
||||
route.execute(request, &()).await.unwrap();
|
||||
|
||||
assert_eq!(source.requested(), MistralOcrConfig.secret_names());
|
||||
assert_eq!(
|
||||
|
|
@ -188,25 +189,21 @@ async fn the_client_uses_the_injected_http_pool_configuration() {
|
|||
user_agent: Some("host-owned/1".into()),
|
||||
..HttpSettings::default()
|
||||
};
|
||||
let client = resources()
|
||||
.ocr_client(
|
||||
&Resolution::from(&settings).config,
|
||||
UrlPolicy::default(),
|
||||
OcrSettings::default(),
|
||||
Arc::new(
|
||||
litellm_secrets::source::EnvironmentSecrets::python_compatible(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
),
|
||||
),
|
||||
)
|
||||
.unwrap();
|
||||
let route = build_ocr_route(
|
||||
&resources(),
|
||||
&Resolution::from(&settings).config,
|
||||
Default::default(),
|
||||
Default::default(),
|
||||
no_secrets(),
|
||||
);
|
||||
|
||||
litellm_core::ocr::client::perform(
|
||||
&client,
|
||||
ocr_request("mistral/model", &upstream.uri(), json!({})),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
route
|
||||
.execute(
|
||||
ocr_request("mistral/model", &upstream.uri(), json!({})),
|
||||
&(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("user-agent"),
|
||||
|
|
|
|||
|
|
@ -45,18 +45,19 @@ async fn mistral_is_served_at_the_resolved_project_and_location() {
|
|||
#[tokio::test]
|
||||
async fn configured_project_and_location_apply_when_the_call_sets_neither() {
|
||||
let upstream = upstream([pages_response()]).await;
|
||||
let client = ocr_client().with_settings(OcrSettings {
|
||||
let route = ocr_route_with(OcrSettings {
|
||||
vertex_project: Some("configured-project".into()),
|
||||
vertex_location: Some("europe-west4".into()),
|
||||
..OcrSettings::default()
|
||||
});
|
||||
|
||||
litellm_core::ocr::client::perform(
|
||||
&client,
|
||||
ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
route
|
||||
.execute(
|
||||
ocr_request("vertex_ai/mistral-ocr-maas", &upstream.uri(), json!({})),
|
||||
&(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.url.path(),
|
||||
|
|
|
|||
|
|
@ -10,10 +10,7 @@ use litellm_auth_gcp::{
|
|||
CredentialSource, VertexAuth, VertexAuthFuture, VertexProviderLoader, VertexTokenSource,
|
||||
};
|
||||
use litellm_core::{
|
||||
ocr::{
|
||||
client::perform,
|
||||
wire::{OcrWireRequest, decode_request},
|
||||
},
|
||||
ocr::wire::{OcrWireRequest, decode_request},
|
||||
resources::CoreResources,
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
|
|
@ -105,17 +102,16 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets(
|
|||
..HttpSettings::default()
|
||||
})
|
||||
.config;
|
||||
let client = owner
|
||||
.ocr_client(
|
||||
&http,
|
||||
Default::default(),
|
||||
OcrSettings {
|
||||
vertex_location: Some(location.into()),
|
||||
..OcrSettings::default()
|
||||
},
|
||||
Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])),
|
||||
)
|
||||
.unwrap();
|
||||
let route = support::build_ocr_route(
|
||||
owner,
|
||||
&http,
|
||||
Default::default(),
|
||||
OcrSettings {
|
||||
vertex_location: Some(location.into()),
|
||||
..OcrSettings::default()
|
||||
},
|
||||
Arc::new(RecordingSecrets::new([("VERTEXAI_CREDENTIALS", identity)])),
|
||||
);
|
||||
let request = decode_request(OcrWireRequest {
|
||||
model: "vertex_ai/mistral-ocr-maas".into(),
|
||||
document: json!({"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}),
|
||||
|
|
@ -127,7 +123,7 @@ async fn auth_survives_per_call_clients_without_freezing_settings_or_secrets(
|
|||
input_sources: Default::default(),
|
||||
timeout_seconds: Some(5.0),
|
||||
}).unwrap();
|
||||
let result = perform(&client, request).await.unwrap();
|
||||
let result = route.execute(request, &()).await.unwrap();
|
||||
assert!(!result.pages.is_empty());
|
||||
}
|
||||
let requests = upstream.received_requests().await.unwrap();
|
||||
|
|
|
|||
442
litellm-rust/crates/core/tests/responses.rs
Normal file
442
litellm-rust/crates/core/tests/responses.rs
Normal file
|
|
@ -0,0 +1,442 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_core::responses::{
|
||||
route::Responses,
|
||||
types::{ResponsesCall, ResponsesOutput},
|
||||
};
|
||||
use litellm_host::{call::HostedCompletion, event::CallEvent};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
use wiremock::ResponseTemplate;
|
||||
|
||||
mod support;
|
||||
use support::*;
|
||||
|
||||
#[fixture]
|
||||
fn call() -> ResponsesCall {
|
||||
ResponsesCall {
|
||||
model: "openai/test-model".into(),
|
||||
input: json!("hello"),
|
||||
optional_params: Default::default(),
|
||||
api_key: Some("test-key".into()),
|
||||
api_base: None,
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::direct(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn http_responses_share_execution_and_hooks(call: ResponsesCall, #[case] hosted: bool) {
|
||||
let body = json!({"id": "response-1", "model": "test-model", "output": [{"type":"message", "content":[]}], "usage":{"total_tokens":7}, "provider_extra": true});
|
||||
let upstream = upstream([json_response(body.clone())]).await;
|
||||
let host = RecordingCall::<Responses>::new(ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
});
|
||||
let response = if hosted {
|
||||
let HostedCompletion::Complete(response) = litellm_host::in_process::run_hosted(
|
||||
responses_route(no_secrets()).machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap() else {
|
||||
panic!()
|
||||
};
|
||||
response
|
||||
} else {
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
let ResponsesOutput::Complete(response) = responses_route(no_secrets())
|
||||
.execute(call, &host)
|
||||
.await
|
||||
.unwrap()
|
||||
else {
|
||||
panic!()
|
||||
};
|
||||
response
|
||||
};
|
||||
assert_eq!(serde_json::to_value(response).unwrap(), body);
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.url.path(), "/responses");
|
||||
assert_eq!(sent.header("authorization"), Some("Bearer test-key"));
|
||||
assert_eq!(sent.header("x-hook"), Some("called"));
|
||||
assert_eq!(sent.json(), json!({"model":"test-model", "input":"hello"}));
|
||||
assert!(matches!(
|
||||
&host.events.0.lock().unwrap()[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Machine(_),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::direct(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption(
|
||||
call: ResponsesCall,
|
||||
#[case] hosted: bool,
|
||||
) {
|
||||
let body = "event: response.completed\ndata: {\"type\":\"response.completed\"}\n\n";
|
||||
let upstream = upstream([ResponseTemplate::new(200)
|
||||
.insert_header("x-request-id", "response-stream")
|
||||
.set_body_raw(body, "text/event-stream")])
|
||||
.await;
|
||||
let host = RecordingCall::<Responses>::new(ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
optional_params: json!({"stream":true}).as_object().unwrap().clone(),
|
||||
..call
|
||||
});
|
||||
let (headers, bytes) = if hosted {
|
||||
assert_eq!(
|
||||
litellm_host::in_process::run_hosted(
|
||||
responses_route(no_secrets()).machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
HostedCompletion::StreamEnded
|
||||
);
|
||||
(
|
||||
host.head.lock().unwrap().take().unwrap().headers,
|
||||
host.chunks.lock().unwrap().concat(),
|
||||
)
|
||||
} else {
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
let ResponsesOutput::Stream { head, chunks } = responses_route(no_secrets())
|
||||
.execute(call, &host)
|
||||
.await
|
||||
.unwrap()
|
||||
else {
|
||||
panic!()
|
||||
};
|
||||
assert_eq!(host.events.0.lock().unwrap().len(), 1);
|
||||
(
|
||||
head.headers,
|
||||
chunks.try_collect::<Vec<_>>().await.unwrap().concat(),
|
||||
)
|
||||
};
|
||||
assert!(headers.contains(&("x-request-id".into(), "response-stream".into())));
|
||||
assert_eq!(bytes, body.as_bytes());
|
||||
assert!(matches!(
|
||||
&host.events.0.lock().unwrap()[..],
|
||||
[CallEvent::Started { .. }, CallEvent::Succeeded { .. }]
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::http(429, json!({"error":"limited"}))]
|
||||
#[case::invalid_response(200, json!({"unexpected":true}))]
|
||||
#[tokio::test]
|
||||
async fn provider_failures_emit_failure_once(
|
||||
call: ResponsesCall,
|
||||
#[case] status: u16,
|
||||
#[case] body: serde_json::Value,
|
||||
) {
|
||||
let upstream = upstream([ResponseTemplate::new(status).set_body_json(body)]).await;
|
||||
let host = RecordingCall::<Responses>::new(ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
});
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
assert!(
|
||||
responses_route(no_secrets())
|
||||
.execute(call, &host)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
let events = host.events.0.lock().unwrap();
|
||||
assert!(matches!(events.last(), Some(CallEvent::Failed { .. })));
|
||||
assert_eq!(
|
||||
events
|
||||
.iter()
|
||||
.filter(|event| matches!(
|
||||
event,
|
||||
CallEvent::Failed { .. } | CallEvent::Succeeded { .. }
|
||||
))
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::explicit(true)]
|
||||
#[case::from_secrets(false)]
|
||||
#[tokio::test]
|
||||
async fn credentials_and_endpoint_are_resolved_only_when_needed(
|
||||
call: ResponsesCall,
|
||||
#[case] explicit: bool,
|
||||
) {
|
||||
let upstream = upstream([json_response(
|
||||
json!({"id":"response", "model":"test-model", "output":[]}),
|
||||
)])
|
||||
.await;
|
||||
let base = upstream.uri();
|
||||
let key = "resolved-test-key";
|
||||
let secrets = Arc::new(RecordingSecrets::new([
|
||||
("OPENAI_API_KEY", key),
|
||||
("OPENAI_BASE_URL", base.as_str()),
|
||||
]));
|
||||
let call = ResponsesCall {
|
||||
api_key: explicit.then(|| key.into()),
|
||||
api_base: explicit.then(|| base.clone()),
|
||||
..call
|
||||
};
|
||||
responses_route(secrets.clone())
|
||||
.execute(call, &())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("authorization"),
|
||||
Some(format!("Bearer {key}").as_str())
|
||||
);
|
||||
if explicit {
|
||||
assert!(secrets.requested().is_empty());
|
||||
} else {
|
||||
assert!(secrets.requested().contains(&"OPENAI_API_KEY".into()));
|
||||
assert!(secrets.requested().contains(&"OPENAI_BASE_URL".into()));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::provider("other/model", "other")]
|
||||
#[case::conflicting_prefix("other/model", "openai")]
|
||||
#[tokio::test]
|
||||
async fn unsupported_providers_fail_before_secrets_or_transport(
|
||||
call: ResponsesCall,
|
||||
#[case] model: &str,
|
||||
#[case] provider: &str,
|
||||
) {
|
||||
let upstream = upstream([]).await;
|
||||
let secrets = Arc::new(RecordingSecrets::failing());
|
||||
let call = ResponsesCall {
|
||||
model: model.into(),
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
};
|
||||
assert!(
|
||||
responses_route(secrets.clone())
|
||||
.execute(call, &())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
assert!(secrets.requested().is_empty());
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::success(200, json!({"id":"response-1", "model":"test-model", "output":[]}), "success")]
|
||||
#[case::invalid_response(200, json!("private-response-sentinel"), "failure")]
|
||||
#[case::upstream_error(429, json!({"error":"private-response-sentinel"}), "failure")]
|
||||
#[tokio::test]
|
||||
async fn route_tracing_covers_native_and_hosted_outcomes(
|
||||
call: ResponsesCall,
|
||||
traces: TraceCapture,
|
||||
#[case] status: u16,
|
||||
#[case] body: serde_json::Value,
|
||||
#[case] outcome: &str,
|
||||
#[values(false, true)] hosted: bool,
|
||||
) {
|
||||
let upstream = upstream([status_response(status, body)]).await;
|
||||
let host = RecordingCall::<Responses>::new(ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
input: json!("private-prompt-sentinel"),
|
||||
api_key: Some("private-key-sentinel".into()),
|
||||
..call
|
||||
});
|
||||
let route = responses_route(no_secrets());
|
||||
let result = traces
|
||||
.logger()
|
||||
.instrument(async {
|
||||
if hosted {
|
||||
litellm_host::in_process::run_hosted(
|
||||
route.clone().machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.map(|_| ())
|
||||
} else {
|
||||
route
|
||||
.execute(host.request().unwrap(), &())
|
||||
.await
|
||||
.map(|_| ())
|
||||
}
|
||||
})
|
||||
.await;
|
||||
assert_eq!(result.is_ok(), outcome == "success");
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
let [summary] = summaries.as_slice() else {
|
||||
panic!("expected one route summary: {summaries:?}")
|
||||
};
|
||||
assert_eq!(summary["route"], "responses");
|
||||
assert_eq!(summary["model"], "openai/test-model");
|
||||
assert_eq!(summary["resolved_model"], "test-model");
|
||||
assert_eq!(summary["provider"], "openai");
|
||||
assert_eq!(summary["stream"], false);
|
||||
assert_eq!(summary["outcome"], outcome);
|
||||
assert!(summary["duration_ms"].as_f64().unwrap() >= 0.0);
|
||||
let sends = traces.summaries("litellm.provider.send");
|
||||
assert_eq!(sends.len(), 1);
|
||||
assert_eq!(sends[0]["status"], status);
|
||||
assert!(!format!("{:?}", traces.records()).contains("private-"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::exhausted(true, "success")]
|
||||
#[case::dropped(false, "cancelled")]
|
||||
#[tokio::test]
|
||||
async fn stream_trace_survives_handoff_and_closes_before_the_stream_object_is_dropped(
|
||||
call: ResponsesCall,
|
||||
traces: TraceCapture,
|
||||
#[case] exhaust: bool,
|
||||
#[case] outcome: &str,
|
||||
) {
|
||||
let upstream = upstream([ResponseTemplate::new(200).set_body_string("stream-bytes")]).await;
|
||||
let output = traces
|
||||
.logger()
|
||||
.instrument(async {
|
||||
responses_route(no_secrets())
|
||||
.execute(
|
||||
ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
optional_params: json!({"stream":true}).as_object().unwrap().clone(),
|
||||
..call
|
||||
},
|
||||
&(),
|
||||
)
|
||||
.await
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(traces.summaries("litellm.route").is_empty());
|
||||
let ResponsesOutput::Stream { mut chunks, .. } = output else {
|
||||
panic!("expected stream")
|
||||
};
|
||||
let captured = traces.clone();
|
||||
tokio::spawn(async move {
|
||||
if exhaust {
|
||||
while chunks.try_next().await.unwrap().is_some() {}
|
||||
assert_eq!(captured.summaries("litellm.route").len(), 1);
|
||||
assert!(chunks.try_next().await.unwrap().is_none());
|
||||
}
|
||||
drop(chunks);
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["stream"], true);
|
||||
assert_eq!(summaries[0]["outcome"], outcome);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn preparation_failure_is_traced_but_unpolled_builders_are_not(
|
||||
call: ResponsesCall,
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
traces.logger().scope(|| {
|
||||
drop(responses_route(no_secrets()).execute(
|
||||
ResponsesCall {
|
||||
model: "unknown/model".into(),
|
||||
..call
|
||||
},
|
||||
&(),
|
||||
))
|
||||
});
|
||||
assert!(traces.records().is_empty());
|
||||
let result = traces
|
||||
.logger()
|
||||
.instrument(async {
|
||||
responses_route(no_secrets())
|
||||
.execute(
|
||||
ResponsesCall {
|
||||
model: "unknown/model".into(),
|
||||
..self::call()
|
||||
},
|
||||
&(),
|
||||
)
|
||||
.await
|
||||
})
|
||||
.await;
|
||||
assert!(result.is_err());
|
||||
let summaries = traces.summaries("litellm.route");
|
||||
assert_eq!(summaries.len(), 1);
|
||||
assert_eq!(summaries[0]["outcome"], "failure");
|
||||
assert!(traces.summaries("litellm.provider.send").is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn websocket_operations_trace_outcomes_without_capturing_frames_or_credentials(
|
||||
traces: TraceCapture,
|
||||
) {
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use litellm_core::responses::websocket::ResponsesWebSocketConnection;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move {
|
||||
let (socket, _) = listener.accept().await.unwrap();
|
||||
let mut socket = tokio_tungstenite::accept_async(socket).await.unwrap();
|
||||
let message = socket.next().await.unwrap().unwrap();
|
||||
socket.send(message).await.unwrap();
|
||||
let _ = socket.next().await;
|
||||
});
|
||||
traces
|
||||
.logger()
|
||||
.instrument(async {
|
||||
let connection = ResponsesWebSocketConnection::connect_url(
|
||||
&format!("ws://{address}/responses"),
|
||||
&std::collections::HashMap::from([(
|
||||
"authorization".into(),
|
||||
"private-key-sentinel".into(),
|
||||
)]),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
connection
|
||||
.send_text("private-frame-sentinel".into())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
connection.recv_text().await.unwrap().as_deref(),
|
||||
Some("private-frame-sentinel")
|
||||
);
|
||||
connection.close().await.unwrap();
|
||||
assert!(
|
||||
connection
|
||||
.send_text("private-frame-sentinel".into())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
})
|
||||
.await;
|
||||
server.await.unwrap();
|
||||
let sends = traces.summaries("litellm.websocket.send_text");
|
||||
assert_eq!(sends.len(), 2);
|
||||
assert_eq!(sends[0]["outcome"], "success");
|
||||
assert_eq!(sends[1]["outcome"], "failure");
|
||||
assert_eq!(
|
||||
traces.summaries("litellm.websocket.connect_url")[0]["outcome"],
|
||||
"success"
|
||||
);
|
||||
assert_eq!(
|
||||
traces.summaries("litellm.websocket.recv_text")[0]["outcome"],
|
||||
"success"
|
||||
);
|
||||
assert_eq!(
|
||||
traces.summaries("litellm.websocket.close")[0]["outcome"],
|
||||
"success"
|
||||
);
|
||||
assert!(!format!("{:?}", traces.records()).contains("private-"));
|
||||
}
|
||||
|
|
@ -7,7 +7,8 @@ use std::sync::{Arc, Mutex};
|
|||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_http::{
|
||||
HttpClientConfig, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
|
||||
ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution,
|
||||
media::PublicDnsResolver,
|
||||
};
|
||||
use litellm_secrets::{SecretValue, source::SecretSource};
|
||||
use serde_json::Value;
|
||||
|
|
@ -24,6 +25,76 @@ pub fn resources() -> litellm_core::resources::CoreResources {
|
|||
litellm_core::resources::CoreResources::new(Arc::new(http_pool()))
|
||||
}
|
||||
|
||||
pub fn no_secrets() -> Arc<dyn SecretSource> {
|
||||
Arc::new(RecordingSecrets::empty())
|
||||
}
|
||||
|
||||
pub fn provider_http(
|
||||
resources: &litellm_core::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
) -> litellm_http::Client {
|
||||
resources
|
||||
.pool
|
||||
.client(config, ClientVariant::Provider)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
pub fn messages_route(secrets: Arc<dyn SecretSource>) -> litellm_core::messages::MessagesRoute {
|
||||
let resources = resources();
|
||||
litellm_core::messages::MessagesRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute {
|
||||
let resources = resources();
|
||||
litellm_core::chat_completions::ChatCompletionsRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
no_secrets(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn responses_route(secrets: Arc<dyn SecretSource>) -> litellm_core::responses::ResponsesRoute {
|
||||
let resources = resources();
|
||||
litellm_core::responses::ResponsesRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn audio_transcription_route() -> litellm_core::audio_transcription::AudioTranscriptionRoute {
|
||||
let resources = resources();
|
||||
litellm_core::audio_transcription::AudioTranscriptionRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
no_secrets(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_ocr_route(
|
||||
resources: &litellm_core::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
url_policy: litellm_http::media::UrlPolicy,
|
||||
settings: litellm_llms::base_llm::ocr::settings::OcrSettings,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> litellm_core::ocr::OcrRoute {
|
||||
litellm_core::ocr::OcrRoute::new(
|
||||
litellm_llms::base_llm::ocr::handler::OcrClient::new(
|
||||
&resources.pool,
|
||||
config,
|
||||
url_policy,
|
||||
resources.auth.clone(),
|
||||
settings,
|
||||
secrets,
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn http_config() -> HttpClientConfig {
|
||||
Resolution::from(&HttpSettings::default()).config
|
||||
}
|
||||
|
|
@ -168,3 +239,152 @@ impl SecretSource for RecordingSecrets {
|
|||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RecordingCall<P: litellm_host::protocol::Protocol> {
|
||||
pub request: Mutex<Option<P::Request>>,
|
||||
pub events: Arc<CallEvents>,
|
||||
pub chunks: Mutex<Vec<P::Chunk>>,
|
||||
pub head: Mutex<Option<P::StreamHead>>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct CallEvents(pub Mutex<Vec<litellm_host::event::CallEvent>>);
|
||||
|
||||
impl litellm_host::lifecycle::CallObserver for CallEvents {
|
||||
fn observe(&self, event: litellm_host::event::CallEvent) {
|
||||
self.0.lock().unwrap().push(event);
|
||||
}
|
||||
}
|
||||
|
||||
impl<P: litellm_host::protocol::Protocol> RecordingCall<P> {
|
||||
pub fn new(request: P::Request) -> Self {
|
||||
Self {
|
||||
request: Mutex::new(Some(request)),
|
||||
events: Arc::new(CallEvents::default()),
|
||||
chunks: Mutex::new(Vec::new()),
|
||||
head: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<P: litellm_host::protocol::Protocol> litellm_host::hooks::RouteHooks<P::Error>
|
||||
for RecordingCall<P>
|
||||
{
|
||||
fn observer(&self) -> Option<Arc<dyn litellm_host::lifecycle::CallObserver>> {
|
||||
Some(self.events.clone())
|
||||
}
|
||||
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: litellm_host::event::WireRequest,
|
||||
_: litellm_host::event::RequestContext,
|
||||
) -> Result<litellm_host::event::WireRequest, P::Error> {
|
||||
Ok(litellm_host::event::WireRequest {
|
||||
headers: wire
|
||||
.headers
|
||||
.into_iter()
|
||||
.chain([("x-hook".into(), "called".into())])
|
||||
.collect(),
|
||||
..wire
|
||||
})
|
||||
}
|
||||
|
||||
async fn on_event(&self, event: litellm_host::event::MachineEvent) -> Result<(), P::Error> {
|
||||
self.events
|
||||
.0
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(litellm_host::event::CallEvent::Machine(event));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl<P> RecordingCall<P>
|
||||
where
|
||||
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
|
||||
P::Error: From<litellm_host::machine::MachineFault>,
|
||||
{
|
||||
pub fn request(&self) -> Result<P::Request, P::Error> {
|
||||
self.request
|
||||
.lock()
|
||||
.unwrap()
|
||||
.take()
|
||||
.ok_or_else(|| litellm_host::machine::MachineFault::Abandoned.into())
|
||||
}
|
||||
pub fn runtime(&self) -> litellm_host::in_process::Host<'_, (), Self, Self> {
|
||||
litellm_host::in_process::Host {
|
||||
services: &(),
|
||||
hooks: self,
|
||||
stream: self,
|
||||
observer: Some(self),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<P> litellm_host::in_process::StreamConsumer<P> for RecordingCall<P>
|
||||
where
|
||||
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
|
||||
P::Error: From<litellm_host::machine::MachineFault>,
|
||||
{
|
||||
async fn open_stream(
|
||||
&self,
|
||||
head: P::StreamHead,
|
||||
) -> Result<litellm_host::protocol::Demand, P::Error> {
|
||||
*self.head.lock().unwrap() = Some(head);
|
||||
Ok(litellm_host::protocol::Demand::More)
|
||||
}
|
||||
async fn send_chunk(
|
||||
&self,
|
||||
chunk: P::Chunk,
|
||||
) -> Result<litellm_host::protocol::Demand, P::Error> {
|
||||
self.chunks.lock().unwrap().push(chunk);
|
||||
Ok(litellm_host::protocol::Demand::More)
|
||||
}
|
||||
}
|
||||
impl<P> litellm_host::lifecycle::CallObserver for RecordingCall<P>
|
||||
where
|
||||
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
|
||||
P::Error: From<litellm_host::machine::MachineFault>,
|
||||
{
|
||||
fn observe(&self, event: litellm_host::event::CallEvent) {
|
||||
self.events.0.lock().unwrap().push(event.clone());
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct TraceCapture(Arc<Mutex<Vec<Value>>>);
|
||||
|
||||
impl TraceCapture {
|
||||
pub fn logger(&self) -> litellm_tracing::Logger {
|
||||
litellm_tracing::Logger::new(self.clone())
|
||||
}
|
||||
|
||||
pub fn records(&self) -> Vec<Value> {
|
||||
self.0.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
pub fn summaries(&self, name: &str) -> Vec<Value> {
|
||||
self.records()
|
||||
.into_iter()
|
||||
.filter(|record| record["span_name"] == name)
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_tracing::Sink for TraceCapture {
|
||||
fn enabled(&self, metadata: &litellm_tracing::Metadata<'_>) -> bool {
|
||||
metadata.target().starts_with("litellm_core")
|
||||
}
|
||||
|
||||
fn emit(&self, record: &litellm_tracing::Record) {
|
||||
self.0
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(Value::Object(record.fields.clone()));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::fixture]
|
||||
pub fn traces() -> TraceCapture {
|
||||
TraceCapture::default()
|
||||
}
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue