mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
Merge remote-tracking branch 'origin/main' into claude/litellm-pr-43216-gdh4ct
# Conflicts: # tests/unit/litellm_core_utils/test_token_counter.py
This commit is contained in:
commit
2379161eee
189 changed files with 4153 additions and 1247 deletions
|
|
@ -323,6 +323,7 @@ jobs:
|
|||
CHOCOLATEY_CONFIRM_ALL: "true"
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
no_output_timeout: 30m
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
|
|
@ -381,6 +382,7 @@ jobs:
|
|||
uv run --no-sync python -m pytest tests/windows_tests/ -v
|
||||
- run:
|
||||
name: Guard against MAX_PATH-busting packaged wheel paths
|
||||
no_output_timeout: 30m
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
|
|
@ -3327,22 +3329,12 @@ workflows:
|
|||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser]
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- integration_contracts:
|
||||
name: integration-<< matrix.suite >>-replica
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, database]
|
||||
mode: [replica]
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
build_and_test:
|
||||
unless:
|
||||
or:
|
||||
|
|
@ -3350,101 +3342,60 @@ workflows:
|
|||
- not:
|
||||
equal: ["", << pipeline.parameters.routing_parity_base >>]
|
||||
jobs:
|
||||
- using_litellm_on_windows:
|
||||
filters: &main_branches
|
||||
branches:
|
||||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- unit:
|
||||
filters: *main_branches
|
||||
- using_litellm_on_windows
|
||||
- unit
|
||||
- provider_replay_harness
|
||||
- base_sdk_install:
|
||||
filters: *main_branches
|
||||
- local_testing_part1:
|
||||
filters: *main_branches
|
||||
- local_testing_part2:
|
||||
filters: *main_branches
|
||||
- langfuse_logging_unit_tests:
|
||||
filters: *main_branches
|
||||
- litellm_assistants_api_testing:
|
||||
filters: *main_branches
|
||||
- litellm_router_testing:
|
||||
filters: *main_branches
|
||||
- litellm_router_unit_testing:
|
||||
filters: *main_branches
|
||||
- auth_ui_unit_tests:
|
||||
filters: *main_branches
|
||||
- build_docker_database_image:
|
||||
filters: *main_branches
|
||||
- e2e_ui_testing:
|
||||
filters: *main_branches
|
||||
- e2e_ui_testing_server_root_path:
|
||||
filters: *main_branches
|
||||
- base_sdk_install
|
||||
- local_testing_part1
|
||||
- local_testing_part2
|
||||
- langfuse_logging_unit_tests
|
||||
- litellm_assistants_api_testing
|
||||
- litellm_router_testing
|
||||
- litellm_router_unit_testing
|
||||
- auth_ui_unit_tests
|
||||
- build_docker_database_image
|
||||
- e2e_ui_testing
|
||||
- e2e_ui_testing_server_root_path
|
||||
- build_and_test:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- e2e_openai_endpoints:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_logging_guardrails_model_info_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_spend_accuracy_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_multi_instance_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_store_model_in_db_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_build_from_pip_tests:
|
||||
filters: *main_branches
|
||||
- proxy_build_from_pip_tests
|
||||
- proxy_pass_through_endpoint_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- proxy_e2e_anthropic_messages_tests:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- llm_translation_testing:
|
||||
filters: *main_branches
|
||||
- realtime_translation_testing:
|
||||
filters: *main_branches
|
||||
- agent_testing:
|
||||
filters: *main_branches
|
||||
- guardrails_testing:
|
||||
filters: *main_branches
|
||||
- google_generate_content_endpoint_testing:
|
||||
filters: *main_branches
|
||||
- llm_responses_api_testing:
|
||||
filters: *main_branches
|
||||
- ocr_testing:
|
||||
filters: *main_branches
|
||||
- search_testing:
|
||||
filters: *main_branches
|
||||
- batches_testing:
|
||||
filters: *main_branches
|
||||
- litellm_utils_testing:
|
||||
filters: *main_branches
|
||||
- pass_through_unit_testing:
|
||||
filters: *main_branches
|
||||
- image_gen_testing:
|
||||
filters: *main_branches
|
||||
- logging_testing:
|
||||
filters: *main_branches
|
||||
- audio_testing:
|
||||
filters: *main_branches
|
||||
- redis_caching_unit_tests:
|
||||
filters: *main_branches
|
||||
- llm_translation_testing
|
||||
- realtime_translation_testing
|
||||
- agent_testing
|
||||
- guardrails_testing
|
||||
- google_generate_content_endpoint_testing
|
||||
- llm_responses_api_testing
|
||||
- ocr_testing
|
||||
- search_testing
|
||||
- batches_testing
|
||||
- litellm_utils_testing
|
||||
- pass_through_unit_testing
|
||||
- image_gen_testing
|
||||
- logging_testing
|
||||
- audio_testing
|
||||
- redis_caching_unit_tests
|
||||
- upload-coverage:
|
||||
requires:
|
||||
- realtime_translation_testing
|
||||
|
|
@ -3469,18 +3420,12 @@ workflows:
|
|||
- db_migration_disable_update_check:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python_3_13:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python_v2_migration_resolver:
|
||||
filters: *main_branches
|
||||
- installing_litellm_on_python
|
||||
- installing_litellm_on_python_3_13
|
||||
- installing_litellm_on_python_v2_migration_resolver
|
||||
- helm_chart_testing:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
- test_bad_database_url:
|
||||
requires:
|
||||
- build_docker_database_image
|
||||
filters: *main_branches
|
||||
|
|
|
|||
2
.github/workflows/test-unit.yml
vendored
2
.github/workflows/test-unit.yml
vendored
|
|
@ -61,7 +61,7 @@ jobs:
|
|||
|
||||
- shard: core-utils
|
||||
artifact-name: core-utils
|
||||
test-path: "tests/test_litellm/litellm_core_utils"
|
||||
test-path: ""
|
||||
unit-flag: core-utils
|
||||
workers: 2
|
||||
reruns: 1
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from litellm.integrations.email_templates.key_rotated_email import (
|
|||
from litellm.integrations.email_templates.templates import (
|
||||
MAX_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
SOFT_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
TEAM_SOFT_BUDGET_ALERT_EMAIL_TEMPLATE,
|
||||
)
|
||||
from litellm.integrations.email_templates.user_invitation_email import (
|
||||
|
|
@ -48,6 +49,12 @@ from litellm.secret_managers.main import get_secret_bool
|
|||
from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL
|
||||
|
||||
|
||||
def _max_budget_alert_id(user_info: CallInfo) -> str:
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
return f"team_member:{user_info.user_id}:{user_info.team_id}"
|
||||
return user_info.token or user_info.user_id or "default_id"
|
||||
|
||||
|
||||
def _parse_email_list(raw) -> List[str]:
|
||||
"""Parse emails from a list or comma-separated string."""
|
||||
if isinstance(raw, list):
|
||||
|
|
@ -373,17 +380,31 @@ class BaseEmailLogger(CustomLogger):
|
|||
greeting = html.escape(
|
||||
event.user_email or event.key_alias or event.token or ""
|
||||
)
|
||||
email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
recipient_email=greeting,
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
if event.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
email_html_content = TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
member=html.escape(event.user_email or event.user_id or ""),
|
||||
team_alias=html.escape(event.team_alias or event.team_id or ""),
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
else:
|
||||
email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format(
|
||||
email_logo_url=email_params.logo_url,
|
||||
recipient_email=greeting,
|
||||
percentage=percentage,
|
||||
spend=spend_str,
|
||||
max_budget=max_budget_str,
|
||||
alert_threshold=alert_threshold_str,
|
||||
base_url=email_params.base_url,
|
||||
email_support_contact=email_params.support_contact,
|
||||
email_footer=email_params.signature,
|
||||
)
|
||||
await self.send_email(
|
||||
from_email=self.DEFAULT_LITELLM_EMAIL,
|
||||
to_email=recipient_emails,
|
||||
|
|
@ -607,7 +628,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
if user_info.spend < threshold_amount:
|
||||
continue
|
||||
|
||||
_id = user_info.token or user_info.user_id or "default_id"
|
||||
_id = _max_budget_alert_id(user_info)
|
||||
_cache_key = (
|
||||
f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}"
|
||||
)
|
||||
|
|
@ -618,7 +639,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
emails.append(user_info.user_email)
|
||||
if not emails:
|
||||
verbose_proxy_logger.warning(
|
||||
"No recipients for %d%% threshold on key %s, skipping alert",
|
||||
"No recipients for %d%% threshold on %s, skipping alert",
|
||||
threshold_pct,
|
||||
_id,
|
||||
)
|
||||
|
|
@ -633,7 +654,11 @@ class BaseEmailLogger(CustomLogger):
|
|||
if send_count is not None and send_count > 1:
|
||||
continue
|
||||
|
||||
event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
event_message = (
|
||||
f"Team Member Budget Alert - {threshold_pct}% of Team Member Budget Reached"
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER
|
||||
else f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached"
|
||||
)
|
||||
webhook_event = WebhookEvent(
|
||||
event="max_budget_alert",
|
||||
event_message=event_message,
|
||||
|
|
|
|||
11
litellm-rust/Cargo.lock
generated
11
litellm-rust/Cargo.lock
generated
|
|
@ -2911,6 +2911,7 @@ dependencies = [
|
|||
"litellm-cache",
|
||||
"litellm-cache-response",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
|
|
@ -2944,6 +2945,7 @@ dependencies = [
|
|||
"litellm-auth-types",
|
||||
"litellm-cache",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"percent-encoding",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
|
|
@ -2970,6 +2972,7 @@ dependencies = [
|
|||
"futures-util",
|
||||
"litellm-cache",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"qdrant-client",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
|
|
@ -3043,6 +3046,7 @@ dependencies = [
|
|||
"litellm-auth-aws",
|
||||
"litellm-cache",
|
||||
"litellm-cache-testing",
|
||||
"litellm-http",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
|
|
@ -3207,11 +3211,13 @@ dependencies = [
|
|||
"http 1.4.2",
|
||||
"hyper-util",
|
||||
"litellm-core-utils",
|
||||
"rcgen",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"rustls 0.23.42",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"veil",
|
||||
|
|
@ -3311,6 +3317,7 @@ dependencies = [
|
|||
"serde_json",
|
||||
"serde_with",
|
||||
"sha2 0.10.9",
|
||||
"strum",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
|
|
@ -3344,6 +3351,7 @@ dependencies = [
|
|||
"google-cloud-auth",
|
||||
"google-cloud-kms-v1",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-python-compat",
|
||||
"litellm-secrets-aws",
|
||||
"litellm-secrets-azure",
|
||||
|
|
@ -3391,6 +3399,7 @@ dependencies = [
|
|||
"litellm-auth-azure",
|
||||
"litellm-auth-types",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-secrets-types",
|
||||
"percent-encoding",
|
||||
"reqwest 0.12.28",
|
||||
|
|
@ -3410,6 +3419,7 @@ version = "0.1.0"
|
|||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-secrets-types",
|
||||
"litellm-tracing",
|
||||
"moka",
|
||||
|
|
@ -3438,6 +3448,7 @@ dependencies = [
|
|||
"litellm-auth-gcp",
|
||||
"litellm-auth-types",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-secrets-types",
|
||||
"moka",
|
||||
"percent-encoding",
|
||||
|
|
|
|||
|
|
@ -7,4 +7,16 @@ disallowed-methods = [
|
|||
{ path = "pyo3_async_runtimes::tokio::local_future_into_py", reason = "use litellm_host_python::run_async / run_async_value" },
|
||||
{ path = "pyo3_async_runtimes::tokio::run", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
{ path = "pyo3_async_runtimes::tokio::run_until_complete", reason = "use litellm_host_python::run_sync / run_sync_value" },
|
||||
{ path = "reqwest::Client::new", reason = "take litellm_http::Client from HttpClientPool" },
|
||||
{ path = "reqwest::Client::builder", reason = "HttpClientConfig owns client construction" },
|
||||
{ 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" },
|
||||
]
|
||||
|
||||
# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS,
|
||||
# proxy and timeout settings. Only crates/http builds one.
|
||||
disallowed-types = [
|
||||
{ path = "reqwest::Client", reason = "take litellm_http::Client from HttpClientPool; only crates/http builds one" },
|
||||
{ path = "reqwest::ClientBuilder", reason = "HttpClientConfig owns client construction" },
|
||||
]
|
||||
|
|
|
|||
|
|
@ -22,5 +22,6 @@ aws-types = "1.4.0"
|
|||
aws-smithy-runtime-api = "1.13.0"
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
reqwest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -962,7 +962,7 @@ mod tests {
|
|||
&no_env,
|
||||
)
|
||||
.await?;
|
||||
let client = reqwest::Client::new();
|
||||
let client = litellm_http::Client::plain_for_test();
|
||||
let mut failures = Vec::new();
|
||||
|
||||
for region in ["us-west-2", "us-east-1"] {
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
litellm-auth-azure.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
|
|
@ -19,6 +20,7 @@ tokio.workspace = true
|
|||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-response.workspace = true
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ impl<C: CacheCodec> AzureBlobCache<C> {
|
|||
pub async fn connect(
|
||||
account_url: &str,
|
||||
container: &str,
|
||||
http: reqwest::Client,
|
||||
http: litellm_http::Client,
|
||||
codec: C,
|
||||
runtime: Handle,
|
||||
) -> Result<Self, Error> {
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ use azure_core::{
|
|||
use futures_util::TryStreamExt;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct ReqwestTransport(pub reqwest::Client);
|
||||
pub struct ReqwestTransport(pub litellm_http::Client);
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl HttpClient for ReqwestTransport {
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ async fn connect(server: &MockServer) -> AzureBlobCache<JsonCodec<Value>> {
|
|||
None,
|
||||
ClientOptions {
|
||||
transport: Some(Transport::new(Arc::new(ReqwestTransport(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
)))),
|
||||
..ClientOptions::default()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
|
|
@ -15,6 +16,7 @@ reqwest.workspace = true
|
|||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ use litellm_cache::{
|
|||
BaseCache, BatchCache, BatchEntry, CacheCodec, DisconnectCache, Error, ExactCacheContext,
|
||||
FlushCache,
|
||||
};
|
||||
use litellm_http::Client;
|
||||
use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode};
|
||||
use reqwest::Client;
|
||||
|
||||
use crate::{GcpTokenSource, TokenSource};
|
||||
|
||||
|
|
|
|||
|
|
@ -96,7 +96,7 @@ async fn cache_exposes_its_configuration(#[future(awt)] server: MockServer) {
|
|||
path_service_account: Some("/secrets/sa.json".into()),
|
||||
..support::config(&server, Some("folder"))
|
||||
},
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_cache::JsonCodec::<Value>::new(),
|
||||
);
|
||||
assert_eq!(cache.bucket_name(), "bucket");
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ pub fn cache_with_token(
|
|||
) -> JsonGcsCache {
|
||||
GcsCache::with_token_source(
|
||||
config(server, gcs_path),
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
JsonCodec::new(),
|
||||
token,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
qdrant-client = { workspace = true, features = ["serde"] }
|
||||
|
|
@ -17,6 +18,7 @@ tokio.workspace = true
|
|||
uuid.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
futures-executor = "0.3"
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_cache::{Error, semantic::Embedder};
|
||||
use reqwest::Client;
|
||||
use litellm_http::Client;
|
||||
use serde_json::Value;
|
||||
|
||||
pub struct OpenAiEmbedder {
|
||||
|
|
|
|||
|
|
@ -5,6 +5,10 @@ use std::{
|
|||
|
||||
use litellm_cache::{Error, semantic::Embedder};
|
||||
use litellm_cache_qdrant_semantic::{OpenAiEmbedder, OpenAiEmbedderConfig};
|
||||
use litellm_http::{
|
||||
ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution,
|
||||
media::PublicDnsResolver,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
use tokio::{
|
||||
|
|
@ -104,7 +108,7 @@ fn config(base: String, timeout: Option<Duration>) -> OpenAiEmbedderConfig {
|
|||
async fn posts_embeddings_request_and_parses_vector() {
|
||||
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config(
|
||||
format!("{}/", server.base_url()),
|
||||
Some(Duration::from_secs(1)),
|
||||
|
|
@ -156,14 +160,17 @@ async fn status_timeout_and_body_errors_are_unavailable(
|
|||
) {
|
||||
let server =
|
||||
TestHttpServer::response_after(status, body, Duration::from_millis(delay_ms)).await;
|
||||
let embedder = OpenAiEmbedder::new(reqwest::Client::new(), config(server.base_url(), timeout));
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config(server.base_url(), timeout),
|
||||
);
|
||||
assert_eq!(embedder.async_embed("hello", None).await, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn sync_embedding_is_unsupported() {
|
||||
let embedder = OpenAiEmbedder::new(
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
config("http://127.0.0.1:9".to_owned(), None),
|
||||
);
|
||||
assert_eq!(
|
||||
|
|
@ -176,9 +183,12 @@ fn sync_embedding_is_unsupported() {
|
|||
#[tokio::test]
|
||||
async fn uses_the_injected_client() {
|
||||
let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await;
|
||||
let client = reqwest::Client::builder()
|
||||
.user_agent("litellm-embedder-test")
|
||||
.build()
|
||||
let config_with_agent = HttpClientConfig {
|
||||
user_agent: Some("litellm-embedder-test".into()),
|
||||
..Resolution::from(&HttpSettings::default()).config
|
||||
};
|
||||
let client = HttpClientPool::new(Arc::new(PublicDnsResolver))
|
||||
.client(&config_with_agent, ClientVariant::Provider)
|
||||
.unwrap();
|
||||
let embedder = OpenAiEmbedder::new(client, config(server.base_url(), None));
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ license.workspace = true
|
|||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-http.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
litellm-auth-aws.workspace = true
|
||||
aws-sdk-s3 = { version = "1.146.1", default-features = false, features = ["rustls", "rt-tokio"] }
|
||||
|
|
@ -19,6 +20,7 @@ reqwest.workspace = true
|
|||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-cache-testing.workspace = true
|
||||
rstest.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
|
|
|
|||
|
|
@ -42,7 +42,12 @@ pub struct S3Cache<C: CacheCodec> {
|
|||
}
|
||||
|
||||
impl<C: CacheCodec> S3Cache<C> {
|
||||
pub fn new(config: S3CacheConfig, http: reqwest::Client, codec: C, runtime: Handle) -> Self {
|
||||
pub fn new(
|
||||
config: S3CacheConfig,
|
||||
http: litellm_http::Client,
|
||||
codec: C,
|
||||
runtime: Handle,
|
||||
) -> Self {
|
||||
let endpoint_url: Option<String> = config.endpoint.map(|endpoint| endpoint.url);
|
||||
let base = aws_sdk_s3::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ use aws_smithy_runtime_api::client::{
|
|||
use aws_smithy_types::body::SdkBody;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct ReqwestHttpClient(pub(crate) reqwest::Client);
|
||||
pub(crate) struct ReqwestHttpClient(pub(crate) litellm_http::Client);
|
||||
|
||||
impl HttpClient for ReqwestHttpClient {
|
||||
fn http_connector(
|
||||
|
|
|
|||
|
|
@ -32,7 +32,12 @@ pub fn config(endpoint: &str) -> S3CacheConfig {
|
|||
}
|
||||
|
||||
pub fn cache_with(config: S3CacheConfig, runtime: Handle) -> JsonS3Cache {
|
||||
S3Cache::new(config, reqwest::Client::new(), JsonCodec::new(), runtime)
|
||||
S3Cache::new(
|
||||
config,
|
||||
litellm_http::Client::plain_for_test(),
|
||||
JsonCodec::new(),
|
||||
runtime,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn cache(endpoint: &str) -> JsonS3Cache {
|
||||
|
|
|
|||
|
|
@ -1,19 +1,19 @@
|
|||
- Target invariants, not completion claims
|
||||
- Keep this crate 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) plus the kwargs rewrites the wrapper makes on the way in (credential-name inheritance, the budget and retry-count limits)
|
||||
- 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
|
||||
- Rust drives the call; every litellm Python internal it still borrows is a variant of `LegacyPython`, grouped by subsystem (`Wrapper`, `Logging`, `DeploymentHooks`)
|
||||
- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json`
|
||||
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
|
||||
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the legacy path rewrites it (setup, deployment hook, prepare) and the bound request object whose attributes back keywords the caller omitted; routes hand it over through `run_legacy_call` and keep no copy
|
||||
- `setup` reuses a `Logging` the caller passed as `litellm_logging_obj` (the proxy and Router are the live cases) and otherwise builds one through `function_setup`, as `@client` does
|
||||
- Either way every phase calls the same `Logging` method the Python path calls; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; routes hand it over through `run_legacy_call` and keep no copy
|
||||
- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's
|
||||
- Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation
|
||||
- Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view
|
||||
- Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup`
|
||||
- Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only
|
||||
- A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-host`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view
|
||||
- Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts
|
||||
- Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch
|
||||
- Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once
|
||||
- Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view
|
||||
- Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object, resolved through `litellm_host_python::lookup`
|
||||
- Retain body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only
|
||||
- Success and failure handlers receive the exact selected public response or exception
|
||||
- A failure-handler error cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch
|
||||
- Dispatch errors never replay provider work or trigger the opposite outcome; the proxy releases deferred success at most once
|
||||
- Delivery follows the registry, not the callable's type: direct, awaited, executor-submitted, logging-worker and deferred paths stay distinct
|
||||
- Traverse every retained Python edge; `close` is idempotent and restores the correlation context once
|
||||
|
|
|
|||
|
|
@ -6,9 +6,6 @@
|
|||
"start_time",
|
||||
"asynchronous"
|
||||
],
|
||||
"check_limits": [
|
||||
"kwargs"
|
||||
],
|
||||
"finalize": [
|
||||
"response",
|
||||
"logger",
|
||||
|
|
@ -76,11 +73,6 @@
|
|||
],
|
||||
"custom_pricing_fields": [],
|
||||
"is_internal_call": [],
|
||||
"credential_list": [],
|
||||
"warn_unknown_credential": [
|
||||
"name",
|
||||
"loaded"
|
||||
],
|
||||
"before_deployment_call": [
|
||||
"kwargs",
|
||||
"call_type"
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ use serde_json::Value;
|
|||
use crate::{
|
||||
DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger,
|
||||
deferred::{PendingLogging, PendingSuccess},
|
||||
finalize, is_internal_call, prepare,
|
||||
finalize, is_internal_call,
|
||||
python::Streaming,
|
||||
setup,
|
||||
};
|
||||
|
|
@ -117,9 +117,13 @@ 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> {
|
||||
let prepared = prepare(py, self.call.kwargs().bind(py), self.logger()?)?.unbind();
|
||||
self.call.set_kwargs(prepared);
|
||||
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)))
|
||||
}
|
||||
|
||||
|
|
@ -580,8 +584,6 @@ assert prepared['document'] is replacement
|
|||
assert prepared['pages'] is replaced_kwargs['pages']
|
||||
assert prepared['litellm_logging_obj'] is logger
|
||||
assert 'litellm_logging_obj' not in replaced_kwargs
|
||||
[checked] = [value for name, value in logger.calls if name == 'check_limits']
|
||||
assert checked is prepared
|
||||
",
|
||||
);
|
||||
});
|
||||
|
|
@ -616,8 +618,6 @@ kwargs = {'logger': logger, 'vendor_extension': opaque}
|
|||
&locals,
|
||||
c"
|
||||
assert prepared['vendor_extension'] is opaque
|
||||
[checked] = [value for name, value in logger.calls if name == 'check_limits']
|
||||
assert checked['vendor_extension'] is opaque
|
||||
assert hooked == ([opaque] if asynchronous else []), hooked
|
||||
",
|
||||
);
|
||||
|
|
@ -733,45 +733,6 @@ assert all(value is failure for name, value in logger.calls if name.endswith('_h
|
|||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::synchronous(false)]
|
||||
#[case::asynchronous(true)]
|
||||
fn a_limit_rejected_before_the_call_surfaces_as_the_callers_error(#[case] asynchronous: bool) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = namespace(
|
||||
py,
|
||||
c"
|
||||
class BudgetExceeded(Exception):
|
||||
pass
|
||||
|
||||
rejection = BudgetExceeded('over budget')
|
||||
|
||||
class LimitedLogger(StubLogger):
|
||||
def check_limits(self, arguments):
|
||||
raise rejection
|
||||
|
||||
logger = LimitedLogger()
|
||||
logger.hooks = {'pre': lambda kwargs: kwargs}
|
||||
kwargs = {'logger': logger}
|
||||
",
|
||||
);
|
||||
let mut logging = legacy_call(py, &locals, asynchronous);
|
||||
let kwargs = local(&locals, "kwargs")
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap()
|
||||
.unbind();
|
||||
let result = logging.begin(py, kwargs, 0.0).and_then(|step| match step {
|
||||
LifecycleStep::Await(_) => {
|
||||
logging.resume(py, Ok(local(&locals, "kwargs").unbind()))
|
||||
}
|
||||
step => Ok(step),
|
||||
});
|
||||
let error = result.err().unwrap();
|
||||
assert!(error.value(py).is(local(&locals, "rejection")));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@
|
|||
//! this crate holds them.
|
||||
|
||||
use litellm_host::{machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{ProtocolHost, lookup, run_call};
|
||||
use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call};
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
|
|
@ -39,7 +39,8 @@ impl PublicCall {
|
|||
}
|
||||
|
||||
/// The keyword view the legacy path currently reads: the caller's copy until
|
||||
/// `function_setup`, then each rewrite (setup, deployment hook, prepare) in turn.
|
||||
/// `function_setup`, then each rewrite (setup, deployment hook, the driver's preflight)
|
||||
/// in turn.
|
||||
pub(crate) fn kwargs(&self) -> &Py<PyDict> {
|
||||
&self.kwargs
|
||||
}
|
||||
|
|
@ -64,13 +65,15 @@ impl PublicCall {
|
|||
}
|
||||
|
||||
/// Runs one native call under the legacy `Logging` contract: the protocol host projects from
|
||||
/// the keyword view the contract prepares, and the contract observes the call.
|
||||
/// the keyword view the contract prepares and `preflight` rewrites, and the contract
|
||||
/// observes the call.
|
||||
pub fn run_legacy_call<H, M>(
|
||||
py: Python<'_>,
|
||||
surface: LegacySurface,
|
||||
call: PublicCall,
|
||||
machine: M,
|
||||
host: H,
|
||||
preflight: Preflight,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
|
|
@ -83,6 +86,7 @@ where
|
|||
machine,
|
||||
host,
|
||||
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
|
||||
preflight,
|
||||
arguments,
|
||||
asynchronous,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
//! 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, the deferred
|
||||
//! proxy release, and the kwargs rewrites the wrapper makes on the way in (credential-name
|
||||
//! inheritance, budget and retry-count limits). All of it sits behind one
|
||||
//! sync and async callback registries it fans out to, the deployment hooks and the deferred
|
||||
//! proxy release. All of it sits behind one
|
||||
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
|
||||
//! core never learn which Python object is on the other end.
|
||||
//! 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.
|
||||
//!
|
||||
//! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`]
|
||||
//! is where those objects live, and [`run_legacy_call`] is how a route hands them over
|
||||
|
|
@ -14,14 +15,12 @@ mod call;
|
|||
mod callbacks;
|
||||
mod deferred;
|
||||
mod logger;
|
||||
mod preparation;
|
||||
mod python;
|
||||
pub(crate) use adapter::LegacyLogging;
|
||||
pub use adapter::{LegacySurface, PassThroughStream};
|
||||
pub use call::{PublicCall, run_legacy_call};
|
||||
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
|
||||
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
|
||||
pub(crate) use preparation::prepare;
|
||||
|
||||
#[cfg(test)]
|
||||
mod test_support {
|
||||
|
|
@ -77,7 +76,6 @@ FAKES = {
|
|||
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
|
||||
kwargs=kwargs,
|
||||
),
|
||||
'check_limits': lambda arguments: arguments['logger'].check_limits(arguments),
|
||||
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
|
||||
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -104,8 +102,6 @@ FAKES = {
|
|||
'restore_context': lambda logger: logger.record('restore', None),
|
||||
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
|
||||
'is_internal_call': lambda: legacy.is_internal.get(),
|
||||
'credential_list': lambda: [],
|
||||
'warn_unknown_credential': lambda name, loaded: None,
|
||||
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
|
||||
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
|
||||
'success', response, call_type
|
||||
|
|
@ -162,9 +158,6 @@ class StubLogger:
|
|||
self.record(phase + '_hook', call_type)
|
||||
return self.hooks.get(phase, lambda value: 'awaitable')(value)
|
||||
|
||||
def check_limits(self, arguments):
|
||||
self.record('check_limits', arguments)
|
||||
|
||||
def failure_handler(self, error, trace, start, end):
|
||||
self.record('failure_handler', error)
|
||||
|
||||
|
|
|
|||
|
|
@ -19,18 +19,12 @@ pub(crate) enum LegacyPython {
|
|||
Streaming(Streaming),
|
||||
}
|
||||
|
||||
/// The `@client` wrapper around the call: `function_setup`, limits, credentials,
|
||||
/// response metadata and the correlation context.
|
||||
/// The `@client` wrapper around the call: `function_setup`, response metadata and the
|
||||
/// correlation context.
|
||||
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
|
||||
pub(crate) enum Wrapper {
|
||||
#[strum(serialize = "setup")]
|
||||
Setup,
|
||||
#[strum(serialize = "check_limits")]
|
||||
CheckLimits,
|
||||
#[strum(serialize = "credential_list")]
|
||||
CredentialList,
|
||||
#[strum(serialize = "warn_unknown_credential")]
|
||||
WarnUnknownCredential,
|
||||
#[strum(serialize = "is_internal_call")]
|
||||
IsInternalCall,
|
||||
#[strum(serialize = "finalize")]
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ url.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-llms = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,13 +0,0 @@
|
|||
use std::{sync::OnceLock, time::Duration};
|
||||
|
||||
use crate::constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS;
|
||||
|
||||
pub(super) fn http_client() -> &'static reqwest::Client {
|
||||
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
|
||||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new())
|
||||
})
|
||||
}
|
||||
|
|
@ -1,10 +1,16 @@
|
|||
use litellm_http::request::truncate_error_body;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_http::{Client, request::truncate_error_body};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, client::http_client};
|
||||
use crate::audio_transcription::types::ProviderAudioTranscriptionRequest;
|
||||
use super::Error;
|
||||
use crate::{
|
||||
audio_transcription::types::ProviderAudioTranscriptionRequest,
|
||||
constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS,
|
||||
};
|
||||
|
||||
pub async fn execute_audio_transcription_provider_call(
|
||||
http: &Client,
|
||||
request: ProviderAudioTranscriptionRequest,
|
||||
) -> Result<Value, Error> {
|
||||
let response = crate::outbound::outbound_request::<Error>(
|
||||
|
|
@ -12,11 +18,15 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
request.url.clone(),
|
||||
request.upstream_headers.clone(),
|
||||
&request.body,
|
||||
request.timeout,
|
||||
Some(
|
||||
request
|
||||
.timeout
|
||||
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
|
||||
),
|
||||
&request.optional_params,
|
||||
)
|
||||
.await?
|
||||
.send(http_client())
|
||||
.send(http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
|
|
|
|||
|
|
@ -1,16 +1,21 @@
|
|||
mod error;
|
||||
pub mod types;
|
||||
pub use error::Error;
|
||||
mod client;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub use handler::execute_audio_transcription_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool};
|
||||
pub use prepare::prepare_audio_transcription_provider_call;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::audio_transcription::types::AudioTranscriptionRequest;
|
||||
|
||||
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
|
||||
.await
|
||||
pub async fn audio_transcription(
|
||||
pool: &HttpClientPool,
|
||||
config: &HttpClientConfig,
|
||||
request: AudioTranscriptionRequest<'_>,
|
||||
) -> Result<Value, Error> {
|
||||
let request = prepare_audio_transcription_provider_call(request)?;
|
||||
let http = pool.client(config, ClientVariant::Provider)?;
|
||||
execute_audio_transcription_provider_call(&http, request).await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,14 +0,0 @@
|
|||
use std::{sync::OnceLock, time::Duration};
|
||||
|
||||
use crate::constants::{CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS, CHAT_COMPLETIONS_TIMEOUT_SECS};
|
||||
|
||||
pub(super) fn http_client() -> &'static reqwest::Client {
|
||||
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
|
||||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS))
|
||||
.connect_timeout(Duration::from_secs(CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new())
|
||||
})
|
||||
}
|
||||
|
|
@ -1,20 +1,24 @@
|
|||
use litellm_http::{outbound::OutboundRequest, request::truncate_error_body};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body};
|
||||
use litellm_llms::base_llm::chat::transformation::ProviderChatResponseData;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, client::http_client, prepare::prepare_provider_request};
|
||||
use crate::chat_completions::types::{
|
||||
ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
|
||||
use super::{Error, prepare::prepare_provider_request};
|
||||
use crate::{
|
||||
chat_completions::types::{ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest},
|
||||
constants::CHAT_COMPLETIONS_TIMEOUT_SECS,
|
||||
};
|
||||
|
||||
pub(super) async fn execute_chat_completions_provider_call(
|
||||
http: &Client,
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let request = prepare_provider_request(request)?;
|
||||
let outbound = outbound_request(&request).await?;
|
||||
|
||||
let response = outbound.send(http_client()).await.map_err(|err| {
|
||||
let response = outbound.send(http).await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
|
|
@ -72,7 +76,11 @@ pub(super) async fn outbound_request(
|
|||
request.url.clone(),
|
||||
request.upstream_headers.clone(),
|
||||
&request.body,
|
||||
request.timeout,
|
||||
Some(
|
||||
request
|
||||
.timeout
|
||||
.unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)),
|
||||
),
|
||||
&request.optional_params,
|
||||
)
|
||||
.await
|
||||
|
|
|
|||
|
|
@ -9,11 +9,11 @@
|
|||
mod error;
|
||||
pub mod types;
|
||||
pub use error::Error;
|
||||
mod client;
|
||||
mod common_utils;
|
||||
pub(crate) mod handler;
|
||||
mod prepare;
|
||||
use handler::execute_chat_completions_provider_call;
|
||||
use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool};
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use prepare::{parse_messages, resolve_provider_config, resolve_request};
|
||||
use serde_json::{Map, Value};
|
||||
|
|
@ -21,9 +21,13 @@ use serde_json::{Map, Value};
|
|||
use crate::chat_completions::types::ChatCompletionsRequest;
|
||||
|
||||
pub async fn chat_completions(
|
||||
pool: &HttpClientPool,
|
||||
config: &HttpClientConfig,
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
execute_chat_completions_provider_call(resolve_request(request)?).await
|
||||
let request = resolve_request(request)?;
|
||||
let http = pool.client(config, ClientVariant::Provider)?;
|
||||
execute_chat_completions_provider_call(&http, request).await
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
|
|
|
|||
|
|
@ -5,9 +5,6 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
|
|||
/// timeout from the caller still overrides this on the request builder.
|
||||
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// Connect timeout for Anthropic Messages provider calls, in seconds.
|
||||
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
/// Provider name used for Anthropic Messages when a deployment's provider model
|
||||
/// does not carry an explicit provider prefix.
|
||||
pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
|
||||
|
|
@ -16,9 +13,6 @@ pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
|
|||
/// seconds. Mirrors the Python chat completions default.
|
||||
pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// Connect timeout for chat completions provider calls, in seconds.
|
||||
pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
|
||||
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// `object` field every non-streaming chat completion response carries.
|
||||
|
|
|
|||
|
|
@ -1,14 +0,0 @@
|
|||
use std::{sync::OnceLock, time::Duration};
|
||||
|
||||
use crate::constants::{MESSAGES_CONNECT_TIMEOUT_SECS, MESSAGES_TIMEOUT_SECS};
|
||||
|
||||
pub(super) fn http_client() -> &'static reqwest::Client {
|
||||
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
|
||||
CLIENT.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(MESSAGES_TIMEOUT_SECS))
|
||||
.connect_timeout(Duration::from_secs(MESSAGES_CONNECT_TIMEOUT_SECS))
|
||||
.build()
|
||||
.unwrap_or_else(|_| reqwest::Client::new())
|
||||
})
|
||||
}
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
use litellm_http::request::string_headers as shared_string_headers;
|
||||
pub(super) use litellm_http::request::truncate_error_body;
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -17,6 +17,8 @@ pub enum Error {
|
|||
#[error(transparent)]
|
||||
Auth(#[from] litellm_auth::Error),
|
||||
#[error(transparent)]
|
||||
Client(#[from] litellm_http::Error),
|
||||
#[error(transparent)]
|
||||
Transport(#[from] litellm_http::transport::Error),
|
||||
#[error(transparent)]
|
||||
Headers(#[from] litellm_http::request::HeaderError),
|
||||
|
|
|
|||
|
|
@ -5,13 +5,15 @@ use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMes
|
|||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{Error, client::http_client, common_utils::truncate_error_body};
|
||||
use super::{Error, common_utils::truncate_error_body};
|
||||
use crate::constants::MESSAGES_TIMEOUT_SECS;
|
||||
|
||||
pub(super) fn network(error: reqwest::Error) -> Error {
|
||||
Error::Transport(TransportError::Network(error.to_string()))
|
||||
}
|
||||
|
||||
pub(super) async fn send(
|
||||
http: &litellm_http::Client,
|
||||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
body: &Value,
|
||||
|
|
@ -20,13 +22,11 @@ pub(super) async fn send(
|
|||
let encoded = serde_json::to_vec(body)
|
||||
.map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?;
|
||||
let builder = headers.iter().fold(
|
||||
http_client().post(url).body(encoded),
|
||||
http.post(url)
|
||||
.body(encoded)
|
||||
.timeout(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))),
|
||||
|builder, (key, value)| builder.header(key, value),
|
||||
);
|
||||
let builder = match timeout {
|
||||
Some(duration) => builder.timeout(duration),
|
||||
None => builder,
|
||||
};
|
||||
http_request(builder).await.map_err(network)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -7,13 +7,13 @@
|
|||
mod error;
|
||||
pub mod types;
|
||||
pub use error::Error;
|
||||
mod client;
|
||||
mod common_utils;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool};
|
||||
use litellm_secrets::source::EnvironmentSecrets;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
|
||||
|
|
@ -21,7 +21,11 @@ use serde_json::Value;
|
|||
|
||||
use crate::messages::types::MessagesRequest;
|
||||
|
||||
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
|
||||
pub async fn messages(
|
||||
pool: &HttpClientPool,
|
||||
config: &HttpClientConfig,
|
||||
request: MessagesRequest<'_>,
|
||||
) -> Result<AnthropicMessagesResponse, Error> {
|
||||
let Value::Object(body) = request.body else {
|
||||
return Err(Error::InvalidRequest(
|
||||
"messages body must be an object".into(),
|
||||
|
|
@ -38,8 +42,15 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
|
|||
timeout: request.timeout,
|
||||
shaping: request.shaping,
|
||||
};
|
||||
let secrets = Arc::new(EnvironmentSecrets::python_compatible());
|
||||
match litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await? {
|
||||
let secrets = Arc::new(EnvironmentSecrets::python_compatible(
|
||||
pool.client(config, ClientVariant::Provider)?,
|
||||
));
|
||||
match litellm_host::run::run(
|
||||
messages_machine(pool, config, secrets)?,
|
||||
&LocalMessagesHost::new(call),
|
||||
)
|
||||
.await?
|
||||
{
|
||||
MessagesOutput::Message(message) => Ok(*message),
|
||||
MessagesOutput::Streamed => Err(Error::Unsupported(
|
||||
"streamed responses need a streaming host",
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ use litellm_core_utils::{
|
|||
settings::Lookup,
|
||||
};
|
||||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request,
|
||||
anthropic::messages::handler::shape_anthropic_messages_request,
|
||||
base_llm::anthropic_messages::transformation::{
|
||||
BaseAnthropicMessagesConfig, MessagesTransformContext,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ use litellm_host::{
|
|||
machine::{CallMachine, HostChannel, MachineFault},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_http::{Client, ClientVariant, HttpClientConfig, HttpClientPool};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
|
||||
|
|
@ -108,12 +109,20 @@ impl Host<Messages> for LocalMessagesHost {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
CallMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
|
||||
pub fn messages_machine(
|
||||
pool: &HttpClientPool,
|
||||
config: &HttpClientConfig,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesMachine, litellm_http::Error> {
|
||||
let http = pool.client(config, ClientVariant::Provider)?;
|
||||
Ok(CallMachine::new(move |host| {
|
||||
Box::pin(execute(host, http.clone(), secrets.clone()))
|
||||
}))
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
host: MessagesHost,
|
||||
http: Client,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
let call = host.project().await?;
|
||||
|
|
@ -164,7 +173,7 @@ async fn execute(
|
|||
context,
|
||||
)
|
||||
.await?;
|
||||
let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?;
|
||||
let response = send(&http, &wire.url, &wire.headers, &wire.body, request.timeout).await?;
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -110,7 +110,10 @@ mod tests {
|
|||
}
|
||||
|
||||
fn client() -> OcrClient {
|
||||
OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new())
|
||||
OcrClient::for_test(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_http::Client::no_redirect_for_test(),
|
||||
)
|
||||
}
|
||||
|
||||
fn request(model: &str, base: &str, document: Value, options: Value) -> LiteLLMOcrRequest {
|
||||
|
|
|
|||
|
|
@ -10,6 +10,10 @@ use support::*;
|
|||
|
||||
const MODEL: &str = "mistral.voxtral-mini-3b-2507";
|
||||
|
||||
async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
audio_transcription(&http_pool(), &http_config(), request).await
|
||||
}
|
||||
|
||||
fn transcript_response(text: &str) -> ResponseTemplate {
|
||||
json_response(json!({"output": {"message": {"content": [{"text": text}]}}}))
|
||||
}
|
||||
|
|
@ -47,7 +51,7 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region(
|
|||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let response = audio_transcription(AudioTranscriptionRequest {
|
||||
let response = transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
optional_params: aws_params(region),
|
||||
..request
|
||||
|
|
@ -79,7 +83,7 @@ async fn the_provider_can_come_from_the_model_prefix(request: AudioTranscription
|
|||
let base = upstream.uri();
|
||||
let model = format!("bedrock/{MODEL}");
|
||||
|
||||
audio_transcription(AudioTranscriptionRequest {
|
||||
transcribe(AudioTranscriptionRequest {
|
||||
model: &model,
|
||||
custom_llm_provider: None,
|
||||
api_base: Some(&base),
|
||||
|
|
@ -110,7 +114,7 @@ async fn audio_and_transcription_params_reach_the_converse_body(
|
|||
])
|
||||
.collect();
|
||||
|
||||
audio_transcription(AudioTranscriptionRequest {
|
||||
transcribe(AudioTranscriptionRequest {
|
||||
audio: json!({"data": "AQI=", "format": format}),
|
||||
api_base: Some(&base),
|
||||
optional_params,
|
||||
|
|
@ -142,7 +146,7 @@ async fn invalid_audio_is_rejected_before_sending(
|
|||
let upstream = upstream([transcript_response("hello")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = audio_transcription(AudioTranscriptionRequest {
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
audio,
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
|
|
@ -174,7 +178,7 @@ async fn unsupported_providers_are_rejected_before_sending(
|
|||
#[case] provider: Option<&'static str>,
|
||||
#[case] reported: &str,
|
||||
) {
|
||||
let error = audio_transcription(AudioTranscriptionRequest {
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
model,
|
||||
custom_llm_provider: provider,
|
||||
api_base: Some(UNREACHABLE_BASE),
|
||||
|
|
@ -189,7 +193,7 @@ async fn unsupported_providers_are_rejected_before_sending(
|
|||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_non_string_extra_header_is_rejected(request: AudioTranscriptionRequest<'static>) {
|
||||
let error = audio_transcription(AudioTranscriptionRequest {
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
extra_headers: Some(Map::from_iter([("x-count".to_string(), json!(3))])),
|
||||
api_base: Some(UNREACHABLE_BASE),
|
||||
..request
|
||||
|
|
@ -212,7 +216,7 @@ async fn an_upstream_error_keeps_its_status_and_body(
|
|||
upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = audio_transcription(AudioTranscriptionRequest {
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
|
|
@ -239,7 +243,7 @@ async fn an_unreadable_success_body_is_an_invalid_response(
|
|||
let upstream = upstream([response]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = audio_transcription(AudioTranscriptionRequest {
|
||||
let error = transcribe(AudioTranscriptionRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ use litellm_core::chat_completions::{
|
|||
Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest,
|
||||
};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
|
|
@ -13,6 +14,10 @@ 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(&http_pool(), &http_config(), request).await
|
||||
}
|
||||
|
||||
fn object(value: Value) -> Map<String, Value> {
|
||||
let Value::Object(map) = value else {
|
||||
panic!("expected a json object, got {value}");
|
||||
|
|
@ -50,7 +55,7 @@ async fn anthropic_round_trip_translates_the_conversation_and_normalizes_the_res
|
|||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let response = chat_completions(ChatCompletionsRequest {
|
||||
let response = complete(ChatCompletionsRequest {
|
||||
messages: json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
|
|
@ -90,7 +95,7 @@ async fn the_deployment_key_replaces_a_caller_supplied_x_api_key(
|
|||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
chat_completions(ChatCompletionsRequest {
|
||||
complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
extra_headers: Some(object(
|
||||
json!({"x-api-key": "caller-key", "x-trace": "kept"}),
|
||||
|
|
@ -116,7 +121,7 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq
|
|||
.await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let response = chat_completions(ChatCompletionsRequest {
|
||||
let response = complete(ChatCompletionsRequest {
|
||||
model: "bedrock/anthropic.claude-sonnet-4-5",
|
||||
optional_params: object(json!({
|
||||
"aws_access_key_id": "access-key",
|
||||
|
|
@ -167,7 +172,7 @@ async fn a_response_it_cannot_normalize_is_reported_as_already_sent(
|
|||
let upstream = upstream([anthropic_response(body)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = chat_completions(ChatCompletionsRequest {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
|
|
@ -188,7 +193,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
|
|||
let upstream = upstream([ResponseTemplate::new(status).set_body_string("slow down")]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = chat_completions(ChatCompletionsRequest {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
})
|
||||
|
|
@ -210,7 +215,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body(
|
|||
async fn a_connection_that_is_never_established_declines_instead_of_failing(
|
||||
request: ChatCompletionsRequest<'static>,
|
||||
) {
|
||||
let error = chat_completions(ChatCompletionsRequest {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(UNREACHABLE_BASE),
|
||||
..request
|
||||
})
|
||||
|
|
@ -232,7 +237,7 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline(
|
|||
upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = chat_completions(ChatCompletionsRequest {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
api_base: Some(&base),
|
||||
timeout: Some(Duration::from_millis(100)),
|
||||
..request
|
||||
|
|
@ -307,7 +312,7 @@ async fn a_declined_request_fails_the_call_before_sending(
|
|||
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
|
||||
let base = upstream.uri();
|
||||
|
||||
let error = chat_completions(ChatCompletionsRequest {
|
||||
let error = complete(ChatCompletionsRequest {
|
||||
optional_params: object(json!({"stream": true})),
|
||||
api_base: Some(&base),
|
||||
..request
|
||||
|
|
|
|||
|
|
@ -78,7 +78,7 @@ impl Host<Messages> for RecordingHost {
|
|||
}
|
||||
|
||||
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
}
|
||||
|
||||
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
|
|
|
|||
|
|
@ -2,9 +2,11 @@ use std::{sync::Arc, time::Duration};
|
|||
|
||||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
|
||||
route::{LocalMessagesHost, MessagesCall, MessagesMachine, MessagesOutput, messages_machine},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
|
||||
use rstest::fixture;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
|
@ -75,11 +77,16 @@ fn headers<'a>(pairs: impl IntoIterator<Item = (&'a str, &'a str)>) -> Option<Ma
|
|||
)
|
||||
}
|
||||
|
||||
fn machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
|
||||
messages_machine(&http_pool(), &http_config(), secrets)
|
||||
.expect("default HTTP settings build a client")
|
||||
}
|
||||
|
||||
async fn run_with(
|
||||
secrets: Arc<RecordingSecrets>,
|
||||
call: MessagesCall,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await
|
||||
litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await
|
||||
}
|
||||
|
||||
/// Runs the route with a secret source that knows nothing, so no environment leaks in.
|
||||
|
|
|
|||
|
|
@ -190,29 +190,40 @@ fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> {
|
|||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn the_facade_runs_the_route_in_process() {
|
||||
async fn the_facade_sends_through_the_injected_http_pool_configuration() {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let base = upstream.uri();
|
||||
let settings = HttpSettings {
|
||||
user_agent: Some("host-owned/1".into()),
|
||||
..HttpSettings::default()
|
||||
};
|
||||
|
||||
let message = messages(facade_request(
|
||||
json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}),
|
||||
&base,
|
||||
))
|
||||
let message = messages(
|
||||
&http_pool(),
|
||||
&Resolution::from(&settings).config,
|
||||
facade_request(
|
||||
json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}),
|
||||
&base,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("messages request succeeds");
|
||||
|
||||
assert_eq!(message.id, "msg_1");
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("x-api-key"),
|
||||
Some("sk-ant")
|
||||
);
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.header("x-api-key"), Some("sk-ant"));
|
||||
assert_eq!(sent.header("user-agent"), Some("host-owned/1"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn the_facade_rejects_a_body_that_is_not_an_object() {
|
||||
let error = messages(facade_request(json!([]), UNREACHABLE_BASE))
|
||||
.await
|
||||
.expect_err("a non-object body is rejected");
|
||||
let error = messages(
|
||||
&http_pool(),
|
||||
&http_config(),
|
||||
facade_request(json!([]), UNREACHABLE_BASE),
|
||||
)
|
||||
.await
|
||||
.expect_err("a non-object body is rejected");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ fn sse_response() -> ResponseTemplate {
|
|||
}
|
||||
|
||||
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ use litellm_core::ocr::{
|
|||
types::LiteLLMOcrRequest,
|
||||
wire::{OcrWireRequest, decode_request},
|
||||
};
|
||||
use litellm_http::Client;
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::OcrClient,
|
||||
|
|
@ -37,11 +38,7 @@ fn object(value: Value) -> Map<String, Value> {
|
|||
}
|
||||
|
||||
fn ocr_client() -> OcrClient {
|
||||
let document_http = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("test document client builds");
|
||||
OcrClient::for_test(reqwest::Client::new(), document_http)
|
||||
OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test())
|
||||
}
|
||||
|
||||
async fn perform(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
|
||||
|
|
|
|||
|
|
@ -1,10 +1,7 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth_gcp::VertexAuth;
|
||||
use litellm_http::{
|
||||
HttpClientPool, HttpSettings, Resolution,
|
||||
media::{PublicDnsResolver, UrlPolicy},
|
||||
};
|
||||
use litellm_http::{HttpSettings, Resolution, media::UrlPolicy};
|
||||
use litellm_llms::{
|
||||
base_llm::ocr::{
|
||||
settings::OcrSettings,
|
||||
|
|
@ -192,12 +189,16 @@ async fn the_client_uses_the_injected_http_pool_configuration() {
|
|||
..HttpSettings::default()
|
||||
};
|
||||
let client = OcrClient::new(
|
||||
&HttpClientPool::new(Arc::new(PublicDnsResolver)),
|
||||
&http_pool(),
|
||||
&Resolution::from(&settings).config,
|
||||
UrlPolicy::default(),
|
||||
VertexAuth::default(),
|
||||
OcrSettings::default(),
|
||||
Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
|
||||
Arc::new(
|
||||
litellm_secrets::source::EnvironmentSecrets::python_compatible(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
),
|
||||
),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
|
|
|
|||
|
|
@ -3,9 +3,12 @@
|
|||
|
||||
#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset
|
||||
|
||||
use std::sync::Mutex;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_http::{
|
||||
HttpClientConfig, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
|
||||
};
|
||||
use litellm_secrets::{SecretValue, source::SecretSource};
|
||||
use serde_json::Value;
|
||||
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
|
||||
|
|
@ -13,6 +16,14 @@ use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
|
|||
/// A port nothing listens on, for calls that must fail before any request is sent.
|
||||
pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1";
|
||||
|
||||
pub fn http_pool() -> HttpClientPool {
|
||||
HttpClientPool::new(Arc::new(PublicDnsResolver))
|
||||
}
|
||||
|
||||
pub fn http_config() -> HttpClientConfig {
|
||||
Resolution::from(&HttpSettings::default()).config
|
||||
}
|
||||
|
||||
/// Starts an upstream that answers its n-th request with the n-th response and 404s after.
|
||||
pub async fn upstream(responses: impl IntoIterator<Item = ResponseTemplate>) -> MockServer {
|
||||
let server = MockServer::start().await;
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits
|
||||
- No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features
|
||||
- The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business
|
||||
- `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance)
|
||||
- `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, rewritten in place by the route's `Preflight`, not the caller's dict; a protocol host that projects from it inherits the adapter's rewrites (for the legacy adapter: setup, deployment hooks) and the preflight's (credential inheritance)
|
||||
- A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is
|
||||
- A failing `classify` is raised with the native error's text as its `__context__`, never swallowed
|
||||
- Use standard PyO3 ownership and conversion APIs
|
||||
|
|
|
|||
|
|
@ -9,6 +9,12 @@ pub fn missing_state() -> PyErr {
|
|||
PyRuntimeError::new_err("missing native call state")
|
||||
}
|
||||
|
||||
/// The SDK's request policy, run by the driver on the keyword view `begin` returned and
|
||||
/// before the protocol host projects from it. It rewrites that view in place, so the
|
||||
/// lifecycle that returned it sees the rewrite too; a rejection fails the call as a host
|
||||
/// failure, so the lifecycle still observes it.
|
||||
pub type Preflight = fn(Python<'_>, &Bound<'_, PyDict>) -> PyResult<()>;
|
||||
|
||||
/// What an adapter step produced: either the value the driver asked for, or a Python
|
||||
/// awaitable the driver hands back to the caller's task before asking again.
|
||||
pub enum LifecycleStep {
|
||||
|
|
|
|||
|
|
@ -14,7 +14,8 @@ use pyo3::types::PyDict;
|
|||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::adapter::{
|
||||
InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state,
|
||||
InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle,
|
||||
missing_state,
|
||||
};
|
||||
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
|
||||
use crate::handle::{Execution, ExecutionBody, ExecutionStep};
|
||||
|
|
@ -83,6 +84,7 @@ where
|
|||
{
|
||||
host: H,
|
||||
adapter: Box<dyn PythonLifecycle>,
|
||||
preflight: Preflight,
|
||||
machine: Option<Arc<Mutex<MachineState<M>>>>,
|
||||
arguments: Option<Py<PyDict>>,
|
||||
started_at: f64,
|
||||
|
|
@ -95,12 +97,14 @@ where
|
|||
}
|
||||
|
||||
/// Runs one native call for Python: synchronously, or as a coroutine that awaits every
|
||||
/// host suspension inline in the caller's task.
|
||||
/// host suspension inline in the caller's task. `preflight` runs once, on the keyword view
|
||||
/// the adapter's `begin` returned, before the host projects from it.
|
||||
pub fn run_call<H, M>(
|
||||
py: Python<'_>,
|
||||
machine: M,
|
||||
host: H,
|
||||
adapter: Box<dyn PythonLifecycle>,
|
||||
preflight: Preflight,
|
||||
arguments: Py<PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
|
|
@ -111,6 +115,7 @@ where
|
|||
let mut driver = PythonDriver {
|
||||
host,
|
||||
adapter,
|
||||
preflight,
|
||||
machine: Some(Arc::new(Mutex::new(MachineState {
|
||||
machine,
|
||||
result: None,
|
||||
|
|
@ -213,6 +218,9 @@ where
|
|||
match (expect, step) {
|
||||
(Expect::Started, LifecycleStep::Done) => self.begin(py),
|
||||
(Expect::Arguments, LifecycleStep::Arguments(arguments)) => {
|
||||
if let Err(error) = (self.preflight)(py, arguments.bind(py)) {
|
||||
return self.adapter_failed(py, error);
|
||||
}
|
||||
self.arguments = Some(arguments);
|
||||
self.stage = Stage::Call;
|
||||
self.resume_machine(py, None)
|
||||
|
|
@ -869,6 +877,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
host: SyntheticHost,
|
||||
script: AdapterScript,
|
||||
asynchronous: bool,
|
||||
) -> (PyResult<Py<PyAny>>, Vec<String>) {
|
||||
run_preflighted(py, machine, host, script, no_preflight, asynchronous)
|
||||
}
|
||||
|
||||
fn no_preflight(_: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn run_preflighted(
|
||||
py: Python<'_>,
|
||||
machine: CallMachine<Synthetic>,
|
||||
host: SyntheticHost,
|
||||
script: AdapterScript,
|
||||
preflight: Preflight,
|
||||
asynchronous: bool,
|
||||
) -> (PyResult<Py<PyAny>>, Vec<String>) {
|
||||
let log = Log(host.log.0.clone());
|
||||
let adapter = SyntheticAdapter {
|
||||
|
|
@ -882,6 +905,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
machine,
|
||||
host,
|
||||
Box::new(adapter),
|
||||
preflight,
|
||||
arguments.unbind(),
|
||||
asynchronous,
|
||||
);
|
||||
|
|
@ -1088,6 +1112,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
streaming_machine(),
|
||||
StreamingHost,
|
||||
Box::new(adapter),
|
||||
no_preflight,
|
||||
PyDict::new(py).unbind(),
|
||||
asynchronous,
|
||||
)
|
||||
|
|
@ -1291,6 +1316,87 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
});
|
||||
}
|
||||
|
||||
/// The rejection a preflight raised, kept so a test can check the caller receives that
|
||||
/// exact object. A `Preflight` is a plain `fn`, so it cannot capture one itself.
|
||||
static REJECTION: Mutex<Option<Py<PyBaseException>>> = Mutex::new(None);
|
||||
|
||||
fn rejecting_preflight(py: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> {
|
||||
let error = PyValueError::new_err("over budget");
|
||||
*REJECTION.lock().unwrap() = Some(error.value(py).clone().unbind());
|
||||
Err(error)
|
||||
}
|
||||
|
||||
fn inheriting_preflight(_: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> {
|
||||
arguments.set_item("api_key", "inherited")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_preflight_rejection_is_the_callers_error_and_the_machine_never_starts() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
crate::initialize_python();
|
||||
Python::attach(|py| {
|
||||
install_lifecycle_module(py);
|
||||
for asynchronous in [false, true] {
|
||||
let (result, log) = run_preflighted(
|
||||
py,
|
||||
success_machine(),
|
||||
SyntheticHost {
|
||||
log: Log::default(),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
},
|
||||
AdapterScript::Plain,
|
||||
rejecting_preflight,
|
||||
asynchronous,
|
||||
);
|
||||
let error = result.unwrap_err();
|
||||
let raised = REJECTION.lock().unwrap().take().unwrap();
|
||||
assert!(error.value(py).is(&raised));
|
||||
assert_eq!(
|
||||
log,
|
||||
[
|
||||
"started",
|
||||
"begin",
|
||||
"failed:Host:over budget",
|
||||
"adapter.close",
|
||||
"host.close"
|
||||
]
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_host_projects_from_the_keyword_view_the_preflight_rewrote() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
crate::initialize_python();
|
||||
Python::attach(|py| {
|
||||
install_lifecycle_module(py);
|
||||
for asynchronous in [false, true] {
|
||||
let (result, _) = run_preflighted(
|
||||
py,
|
||||
success_machine(),
|
||||
SyntheticHost {
|
||||
log: Log::default(),
|
||||
op: OpScript::Answer,
|
||||
classifier_fails: false,
|
||||
},
|
||||
AdapterScript::Plain,
|
||||
inheriting_preflight,
|
||||
asynchronous,
|
||||
);
|
||||
assert_eq!(
|
||||
result.unwrap().extract::<String>(py).unwrap(),
|
||||
"project:2|sign|rewritten"
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_adapters_finalized_response_is_what_the_call_returns_and_reports() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
|
|
@ -1412,6 +1518,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
success_machine(),
|
||||
host,
|
||||
Box::new(adapter),
|
||||
no_preflight,
|
||||
PyDict::new(py).unbind(),
|
||||
false,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -15,7 +15,8 @@ mod handle;
|
|||
mod marshal;
|
||||
|
||||
pub use adapter::{
|
||||
InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state,
|
||||
InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle,
|
||||
missing_state,
|
||||
};
|
||||
pub use argument::lookup;
|
||||
pub use callable::wrap_failure;
|
||||
|
|
|
|||
|
|
@ -22,5 +22,7 @@ veil.workspace = true
|
|||
webpki-roots.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rcgen = "0.14.10"
|
||||
tempfile.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
38
litellm-rust/crates/http/src/client.rs
Normal file
38
litellm-rust/crates/http/src/client.rs
Normal file
|
|
@ -0,0 +1,38 @@
|
|||
use std::ops::Deref;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct Client(reqwest::Client);
|
||||
|
||||
impl Client {
|
||||
pub(crate) fn new(client: reqwest::Client) -> Self {
|
||||
Self(client)
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub fn for_test(client: reqwest::Client) -> Self {
|
||||
Self(client)
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub fn plain_for_test() -> Self {
|
||||
Self(reqwest::Client::new())
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub fn no_redirect_for_test() -> Self {
|
||||
Self(
|
||||
reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("a client without TLS or proxy settings builds"),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for Client {
|
||||
type Target = reqwest::Client;
|
||||
|
||||
fn deref(&self) -> &reqwest::Client {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
|
@ -18,10 +18,16 @@ pub enum Verify {
|
|||
BuiltInRoots,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
|
||||
pub enum ClientIdentity {
|
||||
Pem(PathBuf),
|
||||
Split { certificate: PathBuf, key: PathBuf },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
|
||||
pub struct HttpClientConfig {
|
||||
pub verify: Verify,
|
||||
pub client_certificate: Option<PathBuf>,
|
||||
pub client_certificate: Option<ClientIdentity>,
|
||||
pub key_exchange_group: Option<KeyExchangeGroup>,
|
||||
pub tls12_cipher_suites: Option<Vec<Tls12CipherSuite>>,
|
||||
pub force_ipv4: bool,
|
||||
|
|
@ -67,7 +73,7 @@ impl From<&HttpSettings> for Resolution {
|
|||
Self {
|
||||
config: HttpClientConfig {
|
||||
verify: Verify::from(settings),
|
||||
client_certificate: settings.ssl_certificate.clone(),
|
||||
client_certificate: settings.ssl_certificate.clone().map(ClientIdentity::Pem),
|
||||
key_exchange_group: curve.clone().ok().flatten(),
|
||||
tls12_cipher_suites: ciphers.tls12_cipher_suites,
|
||||
force_ipv4: settings.force_ipv4,
|
||||
|
|
@ -276,7 +282,7 @@ mod tests {
|
|||
config,
|
||||
HttpClientConfig {
|
||||
verify: Verify::BuiltInRoots,
|
||||
client_certificate: Some("/client.pem".into()),
|
||||
client_certificate: Some(ClientIdentity::Pem("/client.pem".into())),
|
||||
key_exchange_group: None,
|
||||
tls12_cipher_suites: None,
|
||||
force_ipv4: true,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,10 @@
|
|||
#![allow(
|
||||
clippy::disallowed_types,
|
||||
clippy::disallowed_methods,
|
||||
reason = "this crate is the one place reqwest clients are built"
|
||||
)]
|
||||
|
||||
mod client;
|
||||
mod config;
|
||||
mod error;
|
||||
pub mod media;
|
||||
|
|
@ -9,7 +16,8 @@ mod settings;
|
|||
mod tls;
|
||||
pub mod transport;
|
||||
|
||||
pub use config::{HttpClientConfig, Resolution, Verify};
|
||||
pub use client::Client;
|
||||
pub use config::{ClientIdentity, HttpClientConfig, Resolution, Verify};
|
||||
pub use error::{Error, TlsSource};
|
||||
pub use pool::{ClientVariant, HttpClientPool};
|
||||
pub use proxy::EnvironmentProxies;
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ use reqwest::{
|
|||
dns::{Addrs, Name, Resolve, Resolving},
|
||||
};
|
||||
|
||||
use crate::{ClientVariant, HttpClientConfig, HttpClientPool};
|
||||
use crate::{Client, ClientVariant, HttpClientConfig, HttpClientPool};
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
|
|
@ -93,8 +93,8 @@ type ProxyMatch = Arc<dyn Fn(&Url) -> bool + Send + Sync>;
|
|||
|
||||
#[derive(Clone)]
|
||||
pub struct MediaFetcher {
|
||||
pinned: reqwest::Client,
|
||||
unpinned: reqwest::Client,
|
||||
pinned: Client,
|
||||
unpinned: Client,
|
||||
uses_proxy: ProxyMatch,
|
||||
address_resolver: Arc<dyn AddressResolver>,
|
||||
url_policy: UrlPolicy,
|
||||
|
|
@ -154,7 +154,7 @@ impl MediaFetcher {
|
|||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub fn for_test(client: reqwest::Client) -> Self {
|
||||
pub fn for_test(client: Client) -> Self {
|
||||
Self {
|
||||
pinned: client.clone(),
|
||||
unpinned: client,
|
||||
|
|
@ -230,7 +230,7 @@ impl MediaFetcher {
|
|||
}
|
||||
}
|
||||
|
||||
async fn client_for(&self, url: &Url) -> Result<&reqwest::Client, Error> {
|
||||
async fn client_for(&self, url: &Url) -> Result<&Client, Error> {
|
||||
if !self.url_policy.validate {
|
||||
return Ok(&self.unpinned);
|
||||
}
|
||||
|
|
@ -520,10 +520,7 @@ mod tests {
|
|||
b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf; charset=binary\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc",
|
||||
)
|
||||
.await;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("test client builds");
|
||||
let client = Client::no_redirect_for_test();
|
||||
let media = MediaFetcher::for_test(client)
|
||||
.fetch(url, policy(3, 0))
|
||||
.await
|
||||
|
|
@ -539,10 +536,7 @@ mod tests {
|
|||
b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc",
|
||||
)
|
||||
.await;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("test client builds");
|
||||
let client = Client::no_redirect_for_test();
|
||||
let error = MediaFetcher::for_test(client)
|
||||
.fetch(url, policy(2, 0))
|
||||
.await
|
||||
|
|
@ -557,10 +551,7 @@ mod tests {
|
|||
b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n2\r\nab\r\n2\r\ncd\r\n0\r\n\r\n",
|
||||
)
|
||||
.await;
|
||||
let client = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("test client builds");
|
||||
let client = Client::no_redirect_for_test();
|
||||
let error = MediaFetcher::for_test(client)
|
||||
.fetch(url, policy(3, 0))
|
||||
.await
|
||||
|
|
|
|||
|
|
@ -107,7 +107,7 @@ impl OutboundRequest {
|
|||
self.timeout
|
||||
}
|
||||
|
||||
pub async fn send(self, client: &reqwest::Client) -> Result<reqwest::Response, reqwest::Error> {
|
||||
pub async fn send(self, client: &crate::Client) -> Result<reqwest::Response, reqwest::Error> {
|
||||
let builder = with_headers(
|
||||
client.post(&self.url).body(self.body),
|
||||
&self.headers,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ use std::{
|
|||
|
||||
use reqwest::dns::Resolve;
|
||||
|
||||
use crate::{config::HttpClientConfig, error::Error, proxy::EnvironmentProxies};
|
||||
use crate::{client::Client, config::HttpClientConfig, error::Error, proxy::EnvironmentProxies};
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
|
||||
pub enum ClientVariant {
|
||||
|
|
@ -48,7 +48,7 @@ impl HttpClientPool {
|
|||
&self,
|
||||
config: &HttpClientConfig,
|
||||
variant: ClientVariant,
|
||||
) -> Result<reqwest::Client, Error> {
|
||||
) -> Result<Client, Error> {
|
||||
let effective = match variant {
|
||||
ClientVariant::Media => HttpClientConfig {
|
||||
client_certificate: None,
|
||||
|
|
@ -65,7 +65,7 @@ impl HttpClientPool {
|
|||
if let Some(pooled) = self.lock().get(&key)
|
||||
&& pooled.built_at.elapsed() < self.ttl
|
||||
{
|
||||
return Ok(pooled.client.clone());
|
||||
return Ok(Client::new(pooled.client.clone()));
|
||||
}
|
||||
let client = self
|
||||
.apply(variant, reqwest::ClientBuilder::try_from(&key.0)?)
|
||||
|
|
@ -77,7 +77,7 @@ impl HttpClientPool {
|
|||
built_at: Instant::now(),
|
||||
},
|
||||
);
|
||||
Ok(client)
|
||||
Ok(Client::new(client))
|
||||
}
|
||||
|
||||
fn lock(&self) -> MutexGuard<'_, Clients> {
|
||||
|
|
@ -116,7 +116,7 @@ mod tests {
|
|||
};
|
||||
|
||||
use super::*;
|
||||
use crate::{HttpSettings, Resolution, Verify};
|
||||
use crate::{ClientIdentity, HttpSettings, Resolution, Verify};
|
||||
|
||||
struct FixedResolver(SocketAddr);
|
||||
|
||||
|
|
@ -288,7 +288,9 @@ mod tests {
|
|||
fn media_variant_never_loads_the_client_certificate() {
|
||||
let pool = pool();
|
||||
let with_identity = HttpClientConfig {
|
||||
client_certificate: Some(std::env::temp_dir().join("litellm-http-absent-client.pem")),
|
||||
client_certificate: Some(ClientIdentity::Pem(
|
||||
std::env::temp_dir().join("litellm-http-absent-client.pem"),
|
||||
)),
|
||||
..config("a")
|
||||
};
|
||||
assert!(
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ use rustls::{
|
|||
};
|
||||
|
||||
use crate::{
|
||||
config::{HttpClientConfig, Verify},
|
||||
config::{ClientIdentity, HttpClientConfig, Verify},
|
||||
error::{Error, TlsSource},
|
||||
};
|
||||
|
||||
|
|
@ -203,11 +203,15 @@ impl TryFrom<&HttpClientConfig> for ClientConfig {
|
|||
};
|
||||
let mut tls = match &config.client_certificate {
|
||||
None => verified.with_no_client_auth(),
|
||||
Some(path) => {
|
||||
let (chain, key) = identity(path, TlsSource::ClientIdentity)?;
|
||||
Some(identity) => {
|
||||
let (certificate, key) = match identity {
|
||||
ClientIdentity::Pem(path) => (path, path),
|
||||
ClientIdentity::Split { certificate, key } => (certificate, key),
|
||||
};
|
||||
let (chain, private_key) = client_identity(certificate, key)?;
|
||||
verified
|
||||
.with_client_auth_cert(chain, key)
|
||||
.map_err(|error| invalid_pem(path, TlsSource::ClientIdentity, error))?
|
||||
.with_client_auth_cert(chain, private_key)
|
||||
.map_err(|error| invalid_pem(key, TlsSource::ClientIdentity, error))?
|
||||
}
|
||||
};
|
||||
tls.alpn_protocols = if config.http2 {
|
||||
|
|
@ -233,17 +237,18 @@ fn bundle_roots(path: &Path, source: TlsSource) -> Result<RootCertStore, Error>
|
|||
Ok(store)
|
||||
}
|
||||
|
||||
fn identity(
|
||||
path: &Path,
|
||||
source: TlsSource,
|
||||
fn client_identity(
|
||||
certificate: &Path,
|
||||
key: &Path,
|
||||
) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>), Error> {
|
||||
let chain = certificates(path, source)?;
|
||||
let source = TlsSource::ClientIdentity;
|
||||
let chain = certificates(certificate, source)?;
|
||||
if chain.is_empty() {
|
||||
return Err(invalid_pem(path, source, "no certificates found"));
|
||||
return Err(invalid_pem(certificate, source, "no certificates found"));
|
||||
}
|
||||
let key = PrivateKeyDer::from_pem_slice(&read(path, source)?)
|
||||
.map_err(|error| invalid_pem(path, source, error))?;
|
||||
Ok((chain, key))
|
||||
let private_key = PrivateKeyDer::from_pem_slice(&read(key, source)?)
|
||||
.map_err(|error| invalid_pem(key, source, error))?;
|
||||
Ok((chain, private_key))
|
||||
}
|
||||
|
||||
fn certificates(path: &Path, source: TlsSource) -> Result<Vec<CertificateDer<'static>>, Error> {
|
||||
|
|
@ -405,7 +410,7 @@ mod tests {
|
|||
)
|
||||
.unwrap();
|
||||
let result = ClientConfig::try_from(&HttpClientConfig {
|
||||
client_certificate: Some(path.clone()),
|
||||
client_certificate: Some(ClientIdentity::Pem(path.clone())),
|
||||
..config(HttpSettings::default())
|
||||
})
|
||||
.map(drop);
|
||||
|
|
@ -419,4 +424,29 @@ mod tests {
|
|||
}) if reported == path
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_client_identity_reads_the_key_from_its_own_file() {
|
||||
let identity = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let certificate = directory.path().join("client.crt");
|
||||
let key = directory.path().join("client.key");
|
||||
std::fs::write(&certificate, identity.cert.pem()).unwrap();
|
||||
std::fs::write(&key, identity.signing_key.serialize_pem()).unwrap();
|
||||
|
||||
let split = ClientConfig::try_from(&HttpClientConfig {
|
||||
client_certificate: Some(ClientIdentity::Split {
|
||||
certificate: certificate.clone(),
|
||||
key,
|
||||
}),
|
||||
..config(HttpSettings::default())
|
||||
});
|
||||
let combined = ClientConfig::try_from(&HttpClientConfig {
|
||||
client_certificate: Some(ClientIdentity::Pem(certificate)),
|
||||
..config(HttpSettings::default())
|
||||
});
|
||||
|
||||
assert!(split.unwrap().client_auth_cert_resolver.has_certs());
|
||||
assert!(combined.is_err());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ tokio = { workspace = true, features = ["sync"] }
|
|||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
aws-smithy-eventstream = "=0.61.4"
|
||||
aws-smithy-types = "1.6.1"
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ use time::OffsetDateTime;
|
|||
use url::Url;
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::transformation::resolve_anthropic_api_base,
|
||||
anthropic::messages::transformation::resolve_anthropic_api_base,
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ use litellm_types::{
|
|||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::streaming_iterator::{
|
||||
anthropic::messages::streaming_iterator::{
|
||||
AnthropicContentBlock, AnthropicContentBlockDelta, AnthropicMessagesStreamEvent,
|
||||
AnthropicStreamUsage,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -11,9 +11,7 @@ use serde_json::{Map, Value, json};
|
|||
use crate::{
|
||||
anthropic::{
|
||||
ANTHROPIC_OAUTH_TOKEN_PREFIX,
|
||||
experimental_pass_through::messages::transformation::{
|
||||
complete_anthropic_url, resolve_anthropic_api_key,
|
||||
},
|
||||
messages::transformation::{complete_anthropic_url, resolve_anthropic_api_key},
|
||||
},
|
||||
base_llm::chat::transformation::{
|
||||
BaseConfig, Error, ProviderChatRequestData, ProviderChatResponseData, RequestAuth,
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
pub mod messages;
|
||||
|
|
@ -1,7 +1,8 @@
|
|||
pub mod common_utils;
|
||||
|
||||
pub mod batches;
|
||||
pub mod chat;
|
||||
pub mod common_utils;
|
||||
pub mod count_tokens;
|
||||
pub mod experimental_pass_through;
|
||||
pub mod messages;
|
||||
|
||||
pub const ANTHROPIC_OAUTH_TOKEN_PREFIX: &str = "sk-ant-oat";
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ use litellm_types::llms::anthropic_messages::{
|
|||
};
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::transformation::{
|
||||
anthropic::messages::transformation::{
|
||||
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
|
||||
},
|
||||
base_llm::{
|
||||
|
|
|
|||
|
|
@ -450,7 +450,7 @@ fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result<i64, E
|
|||
}
|
||||
|
||||
async fn read_operation_response(
|
||||
http_client: &reqwest::Client,
|
||||
http_client: &litellm_http::Client,
|
||||
response: reqwest::Response,
|
||||
original_url: &str,
|
||||
headers: &[(String, String)],
|
||||
|
|
@ -489,7 +489,7 @@ async fn read_operation_response(
|
|||
}
|
||||
|
||||
async fn poll_operation(
|
||||
http_client: &reqwest::Client,
|
||||
http_client: &litellm_http::Client,
|
||||
url: Url,
|
||||
headers: &[(String, String)],
|
||||
connection: &OcrConnection,
|
||||
|
|
|
|||
|
|
@ -4,8 +4,7 @@ use litellm_types::llms::anthropic_messages::{
|
|||
};
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::thinking::ThinkingContext,
|
||||
base_llm::chat::transformation::Error,
|
||||
anthropic::messages::thinking::ThinkingContext, base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
pub type Headers = Vec<(String, String)>;
|
||||
|
|
|
|||
|
|
@ -184,16 +184,21 @@ mod tests {
|
|||
reqwest::header::AUTHORIZATION,
|
||||
reqwest::header::HeaderValue::from_static("Bearer provider-secret"),
|
||||
);
|
||||
let provider_http = reqwest::Client::builder()
|
||||
.default_headers(provider_headers)
|
||||
.build()
|
||||
.unwrap();
|
||||
let document_http = reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.unwrap();
|
||||
let client =
|
||||
crate::base_llm::ocr::handler::OcrClient::for_test(provider_http, document_http);
|
||||
#[expect(
|
||||
clippy::disallowed_methods,
|
||||
clippy::disallowed_types,
|
||||
reason = "the pool has no default-header setting to stand in for provider credentials"
|
||||
)]
|
||||
let provider_http = litellm_http::Client::for_test(
|
||||
reqwest::Client::builder()
|
||||
.default_headers(provider_headers)
|
||||
.build()
|
||||
.unwrap(),
|
||||
);
|
||||
let client = crate::base_llm::ocr::handler::OcrClient::for_test(
|
||||
provider_http,
|
||||
litellm_http::Client::no_redirect_for_test(),
|
||||
);
|
||||
let converted = inline_remote_document(
|
||||
client.document_fetcher(),
|
||||
OcrDocument::ImageUrl {
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ use futures_util::future::BoxFuture;
|
|||
use litellm_auth_gcp::VertexAuth;
|
||||
use litellm_host::event::WireRequest;
|
||||
use litellm_http::{
|
||||
ClientVariant, HttpClientConfig, HttpClientPool,
|
||||
Client, ClientVariant, HttpClientConfig, HttpClientPool,
|
||||
media::{MediaFetcher, UrlPolicy},
|
||||
outbound::{OutboundRequest, RequestSigner},
|
||||
transport,
|
||||
|
|
@ -33,8 +33,8 @@ pub trait CallHooks<E>: Send + Sync {
|
|||
|
||||
#[derive(Clone)]
|
||||
pub struct OcrClient {
|
||||
provider_http: reqwest::Client,
|
||||
polling_http: reqwest::Client,
|
||||
provider_http: Client,
|
||||
polling_http: Client,
|
||||
document_fetcher: MediaFetcher,
|
||||
vertex_auth: VertexAuth,
|
||||
settings: OcrSettings,
|
||||
|
|
@ -60,11 +60,11 @@ impl OcrClient {
|
|||
})
|
||||
}
|
||||
|
||||
pub fn provider_http(&self) -> &reqwest::Client {
|
||||
pub fn provider_http(&self) -> &Client {
|
||||
&self.provider_http
|
||||
}
|
||||
|
||||
pub fn polling_http(&self) -> &reqwest::Client {
|
||||
pub fn polling_http(&self) -> &Client {
|
||||
&self.polling_http
|
||||
}
|
||||
|
||||
|
|
@ -85,17 +85,18 @@ impl OcrClient {
|
|||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-support"))]
|
||||
pub fn for_test(provider_http: reqwest::Client, document_http: reqwest::Client) -> Self {
|
||||
pub fn for_test(provider_http: Client, no_redirect_http: Client) -> Self {
|
||||
Self {
|
||||
secrets: Arc::new(
|
||||
litellm_secrets::source::EnvironmentSecrets::python_compatible(
|
||||
provider_http.clone(),
|
||||
),
|
||||
),
|
||||
provider_http,
|
||||
polling_http: reqwest::Client::builder()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.expect("test polling client builds"),
|
||||
document_fetcher: MediaFetcher::for_test(document_http),
|
||||
polling_http: no_redirect_http.clone(),
|
||||
document_fetcher: MediaFetcher::for_test(no_redirect_http),
|
||||
vertex_auth: VertexAuth::default(),
|
||||
settings: OcrSettings::default(),
|
||||
secrets: Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -311,7 +312,7 @@ mod tests {
|
|||
let _connection = listener.accept().await.unwrap();
|
||||
tokio::time::sleep(Duration::from_secs(1)).await;
|
||||
});
|
||||
let error = reqwest::Client::new()
|
||||
let error = litellm_http::Client::plain_for_test()
|
||||
.get(format!("http://{address}"))
|
||||
.timeout(Duration::from_millis(10))
|
||||
.send()
|
||||
|
|
|
|||
|
|
@ -564,7 +564,10 @@ mod tests {
|
|||
let params = ReductoParseV3Config
|
||||
.map_ocr_params(&overrides, "parse-v3")
|
||||
.unwrap();
|
||||
let client = OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new());
|
||||
let client = OcrClient::for_test(
|
||||
litellm_http::Client::plain_for_test(),
|
||||
litellm_http::Client::no_redirect_for_test(),
|
||||
);
|
||||
let connection = OcrConnection::default();
|
||||
let document = serde_json::from_value(
|
||||
json!({"type":"document_url","document_url":"reducto://ready.pdf"}),
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ async fn read_bounded(response: String, limit: usize) -> Result<bytes::Bytes, Er
|
|||
socket.write_all(response.as_bytes()).await.unwrap();
|
||||
std::future::pending::<()>().await;
|
||||
});
|
||||
let response = reqwest::Client::new()
|
||||
let response = litellm_http::Client::plain_for_test()
|
||||
.get(format!("http://{address}"))
|
||||
.send()
|
||||
.await
|
||||
|
|
|
|||
|
|
@ -42,6 +42,9 @@ pub struct ModelInfo {
|
|||
pub cache_creation_input_token_cost_above_200k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_200k_tokens_batches: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_256k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -78,6 +81,9 @@ pub struct ModelInfo {
|
|||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_200k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_200k_tokens_batches: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_200k_tokens_priority: Option<f64>,
|
||||
|
|
@ -113,6 +119,10 @@ pub struct ModelInfo {
|
|||
pub code_interpreter_cost_per_session: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub comment: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub computer_use_input_cost_per_1k_tokens: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub computer_use_output_cost_per_1k_tokens: Option<f64>,
|
||||
/// Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub default_reasoning_effort: Option<ReasoningEffort>,
|
||||
|
|
@ -120,6 +130,10 @@ pub struct ModelInfo {
|
|||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub deprecation_date: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub file_search_cost_per_1k_calls: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub file_search_cost_per_gb_per_day: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub gemini_audio_only_live: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub gemini_native_audio: Option<bool>,
|
||||
|
|
@ -174,6 +188,9 @@ pub struct ModelInfo {
|
|||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_200k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_200k_tokens_batches: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_200k_tokens_priority: Option<f64>,
|
||||
|
|
@ -265,6 +282,26 @@ pub struct ModelInfo {
|
|||
pub output_cost_per_image_1536: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_image_512: Option<f64>,
|
||||
#[serde(
|
||||
rename = "output_cost_per_image_0.5K",
|
||||
skip_serializing_if = "Option::is_none"
|
||||
)]
|
||||
pub output_cost_per_image_0_5k: Option<f64>,
|
||||
#[serde(
|
||||
rename = "output_cost_per_image_1K",
|
||||
skip_serializing_if = "Option::is_none"
|
||||
)]
|
||||
pub output_cost_per_image_1k: Option<f64>,
|
||||
#[serde(
|
||||
rename = "output_cost_per_image_2K",
|
||||
skip_serializing_if = "Option::is_none"
|
||||
)]
|
||||
pub output_cost_per_image_2k: Option<f64>,
|
||||
#[serde(
|
||||
rename = "output_cost_per_image_4K",
|
||||
skip_serializing_if = "Option::is_none"
|
||||
)]
|
||||
pub output_cost_per_image_4k: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_image_token: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
@ -297,6 +334,9 @@ pub struct ModelInfo {
|
|||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_200k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_200k_tokens_batches: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_200k_tokens_priority: Option<f64>,
|
||||
|
|
@ -357,6 +397,8 @@ pub struct ModelInfo {
|
|||
/// Provider default requests-per-minute limit.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub rpm: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub rules: Option<Vec<Value>>,
|
||||
/// USD cost per web search query, keyed by search context size.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub search_context_cost_per_query: Option<SearchContextCostPerQuery>,
|
||||
|
|
@ -475,6 +517,8 @@ pub struct ModelInfo {
|
|||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub uses_embed_content: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub vector_store_cost_per_gb_per_day: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub vertex_ai_audio_api: Option<VertexAiAudioApi>,
|
||||
/// Whether web search is billed per query or per prompt.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
|
|
|||
|
|
@ -55,12 +55,14 @@ pyo3-async-runtimes.workspace = true
|
|||
reqwest.workspace = true
|
||||
redis = { version = "1.7.0", features = ["tls-rustls"] }
|
||||
serde_json.workspace = true
|
||||
strum.workspace = true
|
||||
veil.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["rt", "sync"] }
|
||||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
litellm-secrets-aws.workspace = true
|
||||
serde.workspace = true
|
||||
serde_with.workspace = true
|
||||
|
|
|
|||
10
litellm-rust/crates/python-bridge/preflight_contract.json
Normal file
10
litellm-rust/crates/python-bridge/preflight_contract.json
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
{
|
||||
"credential_list": [],
|
||||
"warn_unknown_credential": [
|
||||
"name",
|
||||
"loaded"
|
||||
],
|
||||
"check_limits": [
|
||||
"kwargs"
|
||||
]
|
||||
}
|
||||
|
|
@ -9,10 +9,10 @@ use super::{
|
|||
cache_error,
|
||||
config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig},
|
||||
embedder::PythonEmbedder,
|
||||
host_client,
|
||||
native::NativeResponseCache,
|
||||
};
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use crate::http::host_client;
|
||||
|
||||
fn declined(reason: UnsupportedCacheConfig) -> PyErr {
|
||||
RustBridgeDeclined::new_err(reason.message())
|
||||
|
|
|
|||
|
|
@ -1511,7 +1511,7 @@ mod tests {
|
|||
path_service_account: Some("credentials.json".into()),
|
||||
endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(),
|
||||
},
|
||||
reqwest::Client::new(),
|
||||
litellm_http::Client::plain_for_test(),
|
||||
Some("token".into()),
|
||||
);
|
||||
let matching_config = NativeCacheConfig {
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::http::host_client;
|
||||
use crate::logger::run_sync_value;
|
||||
use litellm_auth_aws::AwsAuthConfig;
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
|
|
@ -19,7 +20,6 @@ use super::{
|
|||
config::{QdrantSemanticCacheConfig, project_redis_semantic},
|
||||
embedder::PythonEmbedder,
|
||||
facade::FacadeGuard,
|
||||
host_client,
|
||||
native::NativeResponseCache,
|
||||
request::duration,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -13,11 +13,9 @@ mod resolver;
|
|||
mod semantic;
|
||||
|
||||
use litellm_cache::Error;
|
||||
use litellm_http::ClientVariant;
|
||||
use pyo3::{
|
||||
exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
};
|
||||
|
||||
pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver};
|
||||
|
|
@ -29,11 +27,3 @@ fn cache_error(error: Error) -> PyErr {
|
|||
_ => PyRuntimeError::new_err(error.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// The host's pooled HTTP client, configured from the proxy's HTTP settings.
|
||||
fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult<reqwest::Client> {
|
||||
let http_config = crate::http::call_config(py, &PyDict::new(py), true)?;
|
||||
crate::http::pool()
|
||||
.client(&http_config, variant)
|
||||
.map_err(crate::http::client_error)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -92,7 +92,7 @@ impl NativeResponseCache {
|
|||
))
|
||||
}
|
||||
|
||||
pub async fn s3(config: S3CacheConfig, http: reqwest::Client) -> Self {
|
||||
pub async fn s3(config: S3CacheConfig, http: litellm_http::Client) -> Self {
|
||||
let runtime = tokio::runtime::Handle::current();
|
||||
let backend = S3Cache::new(config, http, ResponseCacheCodec, runtime);
|
||||
let identity = BackendIdentity::S3 {
|
||||
|
|
@ -112,7 +112,7 @@ impl NativeResponseCache {
|
|||
Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity))
|
||||
}
|
||||
|
||||
pub fn gcs(config: GcsConfig, client: reqwest::Client, token: Option<String>) -> Self {
|
||||
pub fn gcs(config: GcsConfig, client: litellm_http::Client, token: Option<String>) -> Self {
|
||||
let backend = match token {
|
||||
Some(token) => GcsCache::with_token_source(
|
||||
config,
|
||||
|
|
@ -133,7 +133,7 @@ impl NativeResponseCache {
|
|||
pub async fn azure_blob(
|
||||
account_url: &str,
|
||||
container: &str,
|
||||
http: reqwest::Client,
|
||||
http: litellm_http::Client,
|
||||
) -> Result<Self, Error> {
|
||||
let backend = AzureBlobCache::connect(
|
||||
account_url,
|
||||
|
|
@ -242,7 +242,7 @@ impl NativeResponseCache {
|
|||
|
||||
pub async fn qdrant_semantic(
|
||||
config: QdrantSemanticCacheConfig,
|
||||
client: reqwest::Client,
|
||||
client: litellm_http::Client,
|
||||
runtime: tokio::runtime::Handle,
|
||||
) -> Result<Self, Error> {
|
||||
let qdrant = qdrant_client::Qdrant::from_url(&config.grpc_url)
|
||||
|
|
|
|||
|
|
@ -6,8 +6,8 @@ use std::{
|
|||
|
||||
use litellm_core_utils::settings::ProcessEnvironment;
|
||||
use litellm_http::{
|
||||
HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, Resolution, SslVerify,
|
||||
TlsSource, Unsupported,
|
||||
Client, ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer,
|
||||
Resolution, SslVerify, TlsSource, Unsupported,
|
||||
media::{PublicDnsResolver, UrlPolicy},
|
||||
};
|
||||
use pyo3::{
|
||||
|
|
@ -97,7 +97,10 @@ pub(crate) fn call_config(
|
|||
let settings = HttpSettings::from_layers([
|
||||
for_call(call_ssl_verify(kwargs)?, asynchronous),
|
||||
HttpSettingsLayer::from_environment(&ProcessEnvironment),
|
||||
configured(&PythonSettings::Http.read(py)?)?,
|
||||
match PythonSettings::Http.read_or_unset(py)? {
|
||||
Some(snapshot) => configured(&snapshot)?,
|
||||
None => HttpSettingsLayer::default(),
|
||||
},
|
||||
])
|
||||
.without_missing_files(&|path: &Path| path.exists());
|
||||
let resolution = Resolution::from(&settings);
|
||||
|
|
@ -107,6 +110,11 @@ pub(crate) fn call_config(
|
|||
Ok(resolution.config)
|
||||
}
|
||||
|
||||
pub(crate) fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult<Client> {
|
||||
let config = call_config(py, &PyDict::new(py), true)?;
|
||||
pool().client(&config, variant).map_err(client_error)
|
||||
}
|
||||
|
||||
pub(crate) fn client_error(error: litellm_http::Error) -> PyErr {
|
||||
match error {
|
||||
litellm_http::Error::Read {
|
||||
|
|
@ -143,7 +151,10 @@ fn unreported(
|
|||
}
|
||||
|
||||
pub(crate) fn url_policy(py: Python<'_>) -> PyResult<UrlPolicy> {
|
||||
project_url_policy(&PythonSettings::UrlPolicy.read(py)?)
|
||||
match PythonSettings::UrlPolicy.read_or_unset(py)? {
|
||||
Some(snapshot) => project_url_policy(&snapshot),
|
||||
None => Ok(UrlPolicy::default()),
|
||||
}
|
||||
}
|
||||
|
||||
fn project_url_policy(snapshot: &Snapshot<'_>) -> PyResult<UrlPolicy> {
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ mod errors;
|
|||
mod http;
|
||||
mod logger;
|
||||
mod marshal;
|
||||
mod preflight;
|
||||
mod python_settings;
|
||||
mod routes;
|
||||
mod secrets;
|
||||
|
|
|
|||
|
|
@ -1,9 +1,51 @@
|
|||
//! The SDK's request policy the driver runs on every route's keyword view before the host
|
||||
//! projects from it: credential-name inheritance from `litellm.credential_list`, then the
|
||||
//! budget and retry-count limits. It is the `@client` prologue after `function_setup` and the
|
||||
//! deployment hook, and belongs to no callback contract.
|
||||
|
||||
use pyo3::{
|
||||
prelude::*,
|
||||
types::{PyDict, PyList},
|
||||
};
|
||||
use strum::{IntoStaticStr, VariantArray};
|
||||
|
||||
use crate::python::Wrapper;
|
||||
const MODULE: &str = "litellm.rust_bridge.preflight";
|
||||
|
||||
/// The litellm globals the preflight still reads through Python. `preflight_contract.json`
|
||||
/// pins each function's parameters on both sides.
|
||||
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
|
||||
pub(crate) enum PythonPreflight {
|
||||
#[strum(serialize = "credential_list")]
|
||||
CredentialList,
|
||||
#[strum(serialize = "warn_unknown_credential")]
|
||||
WarnUnknownCredential,
|
||||
#[strum(serialize = "check_limits")]
|
||||
CheckLimits,
|
||||
}
|
||||
|
||||
impl PythonPreflight {
|
||||
fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult<Bound<'py, PyAny>>
|
||||
where
|
||||
A: pyo3::call::PyCallArgs<'py>,
|
||||
{
|
||||
py.import(MODULE)?.getattr(<&str>::from(self))?.call1(args)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../preflight_contract.json");
|
||||
|
||||
/// Rewrites `arguments` in place, in the order the Python wrapper runs: credentials first,
|
||||
/// so the limits see the same view the provider request is built from.
|
||||
pub(crate) fn sdk_preflight(py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> {
|
||||
inherit_credentials(py, arguments, || {
|
||||
Ok(PythonPreflight::CredentialList
|
||||
.call(py, ())?
|
||||
.cast_into::<PyList>()?)
|
||||
})?;
|
||||
PythonPreflight::CheckLimits.call(py, (arguments,))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
struct CredentialEntry<'py>(Bound<'py, PyAny>);
|
||||
|
||||
|
|
@ -17,22 +59,6 @@ impl<'py> CredentialEntry<'py> {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn prepare<'py>(
|
||||
py: Python<'py>,
|
||||
kwargs: &Bound<'py, PyDict>,
|
||||
logger: &crate::PythonLogger,
|
||||
) -> PyResult<Bound<'py, PyDict>> {
|
||||
let arguments = kwargs.copy()?;
|
||||
arguments.set_item("litellm_logging_obj", logger.object(py))?;
|
||||
inherit_credentials(py, &arguments, || {
|
||||
Ok(Wrapper::CredentialList
|
||||
.call(py, ())?
|
||||
.cast_into::<PyList>()?)
|
||||
})?;
|
||||
Wrapper::CheckLimits.call(py, (&arguments,))?;
|
||||
Ok(arguments)
|
||||
}
|
||||
|
||||
fn inherit_credentials<'py>(
|
||||
py: Python<'py>,
|
||||
arguments: &Bound<'py, PyDict>,
|
||||
|
|
@ -54,7 +80,7 @@ fn inherit_credentials<'py>(
|
|||
.map(|credential| CredentialEntry(credential).name())
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
let Some(index) = names.iter().position(|name| *name == requested) else {
|
||||
Wrapper::WarnUnknownCredential.call(py, (requested, names.len()))?;
|
||||
PythonPreflight::WarnUnknownCredential.call(py, (requested, names.len()))?;
|
||||
return Ok(());
|
||||
};
|
||||
let selected = CredentialEntry(credentials.get_item(index)?);
|
||||
|
|
@ -71,7 +97,42 @@ fn inherit_credentials<'py>(
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeSet;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use super::*;
|
||||
use strum::VariantArray;
|
||||
|
||||
/// Tests share one interpreter, and the stub module below is global state, so the
|
||||
/// tests that install it run one at a time.
|
||||
static PREFLIGHT_MODULE: Mutex<()> = Mutex::new(());
|
||||
|
||||
/// A fresh stand-in for `litellm.rust_bridge.preflight` that records every call, then
|
||||
/// `script` run against it with the module bound as `preflight`.
|
||||
fn preflight_stubs<'py>(py: Python<'py>, script: &std::ffi::CStr) -> Bound<'py, PyDict> {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"
|
||||
import sys
|
||||
import types
|
||||
|
||||
for name in ('litellm', 'litellm.rust_bridge'):
|
||||
sys.modules.setdefault(name, types.ModuleType(name))
|
||||
preflight = types.ModuleType('litellm.rust_bridge.preflight')
|
||||
preflight.warnings = []
|
||||
preflight.checked = []
|
||||
preflight.credential_list = lambda: []
|
||||
preflight.warn_unknown_credential = lambda name, loaded: preflight.warnings.append((name, loaded))
|
||||
preflight.check_limits = lambda kwargs: preflight.checked.append(kwargs)
|
||||
sys.modules['litellm.rust_bridge.preflight'] = preflight
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
py.run(script, Some(&locals), Some(&locals)).unwrap();
|
||||
locals
|
||||
}
|
||||
|
||||
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
|
||||
let locals = PyDict::new(py);
|
||||
|
|
@ -312,4 +373,110 @@ arguments = {'litellm_credential_name': 'ocr-test'}
|
|||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_borrowed_function_is_in_the_python_contract() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let contract = litellm_host_python::json_loads(py, PYTHON_CONTRACT.as_bytes()).unwrap();
|
||||
let declared: BTreeSet<String> = contract
|
||||
.bind(py)
|
||||
.cast::<PyDict>()
|
||||
.unwrap()
|
||||
.keys()
|
||||
.extract()
|
||||
.map(|names: Vec<String>| names.into_iter().collect())
|
||||
.unwrap();
|
||||
let called: BTreeSet<String> = PythonPreflight::VARIANTS
|
||||
.iter()
|
||||
.map(|&function| <&str>::from(function).to_owned())
|
||||
.collect();
|
||||
assert_eq!(
|
||||
called.len(),
|
||||
PythonPreflight::VARIANTS.len(),
|
||||
"a function is borrowed twice"
|
||||
);
|
||||
assert_eq!(called, declared);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn an_unknown_name_is_reported_with_the_loaded_count_and_leaves_the_arguments_alone() {
|
||||
let _guard = PREFLIGHT_MODULE
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = preflight_stubs(
|
||||
py,
|
||||
c"
|
||||
class Credential:
|
||||
credential_name = 'listed'
|
||||
credential_values = {'api_key': 'listed-key'}
|
||||
preflight.credential_list = lambda: [Credential(), Credential()]
|
||||
arguments = {'litellm_credential_name': 'missing'}
|
||||
",
|
||||
);
|
||||
sdk_preflight(py, &argument_dict(&locals)).unwrap();
|
||||
py.run(
|
||||
c"
|
||||
assert arguments == {'litellm_credential_name': 'missing'}, arguments
|
||||
assert preflight.warnings == [('missing', 2)], preflight.warnings
|
||||
assert preflight.checked == [arguments]
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn limits_are_checked_on_the_arguments_after_credentials_are_inherited() {
|
||||
let _guard = PREFLIGHT_MODULE
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = preflight_stubs(
|
||||
py,
|
||||
c"
|
||||
class Credential:
|
||||
credential_name = 'ocr-test'
|
||||
credential_values = {'api_key': 'inherited'}
|
||||
preflight.credential_list = lambda: [Credential()]
|
||||
rejection = RuntimeError('Max retries per request hit!')
|
||||
def check_limits(arguments):
|
||||
preflight.checked.append(dict(arguments))
|
||||
raise rejection
|
||||
preflight.check_limits = check_limits
|
||||
arguments = {'litellm_credential_name': 'ocr-test'}
|
||||
",
|
||||
);
|
||||
let error = sdk_preflight(py, &argument_dict(&locals)).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
.is(locals.get_item("rejection").unwrap().unwrap())
|
||||
);
|
||||
py.run(
|
||||
c"
|
||||
assert preflight.checked == [{'litellm_credential_name': 'ocr-test', 'api_key': 'inherited'}], preflight.checked
|
||||
assert arguments['api_key'] == 'inherited'
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
fn argument_dict<'py>(locals: &Bound<'py, PyDict>) -> Bound<'py, PyDict> {
|
||||
locals
|
||||
.get_item("arguments")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap()
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
use pyo3::prelude::*;
|
||||
use pyo3::{exceptions::PyModuleNotFoundError, prelude::*};
|
||||
|
||||
use crate::coercion::{FieldSpec, ProjectionError};
|
||||
|
||||
|
|
@ -40,15 +40,45 @@ impl PythonSettings {
|
|||
Ok(Snapshot { group: self, value })
|
||||
}
|
||||
|
||||
/// Reads the accessor, or `None` when the litellm package is not installed
|
||||
/// (a bare extension module), meaning there are no configured values.
|
||||
pub(crate) fn read_or_unset(self, py: Python<'_>) -> PyResult<Option<Snapshot<'_>>> {
|
||||
match self.read(py) {
|
||||
Ok(snapshot) => Ok(Some(snapshot)),
|
||||
Err(error) => {
|
||||
if missing_module(py, &error, "litellm")? {
|
||||
Ok(None)
|
||||
} else {
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn snapshot(self, value: Bound<'_, PyAny>) -> Snapshot<'_> {
|
||||
Snapshot { group: self, value }
|
||||
}
|
||||
}
|
||||
|
||||
fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult<bool> {
|
||||
if !error.is_instance_of::<PyModuleNotFoundError>(py) {
|
||||
return Ok(false);
|
||||
}
|
||||
Ok(error
|
||||
.value(py)
|
||||
.getattr("name")?
|
||||
.extract::<Option<String>>()?
|
||||
.is_some_and(|name| name == expected))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use pyo3::{exceptions::PyRuntimeError, prelude::*, types::PyDict};
|
||||
use pyo3::{
|
||||
exceptions::{PyImportError, PyModuleNotFoundError, PyRuntimeError},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
};
|
||||
|
||||
use super::PythonSettings;
|
||||
use crate::coercion::FieldSpec;
|
||||
|
|
@ -140,4 +170,153 @@ values = (Descriptor(), SimpleNamespace(flag=Truth()))
|
|||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn read_or_unset_returns_none_when_litellm_is_missing() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"
|
||||
import sys
|
||||
class MissingLitellm:
|
||||
def find_spec(self, fullname, path=None, target=None):
|
||||
if fullname == 'litellm':
|
||||
raise ModuleNotFoundError('No module named litellm', name='litellm')
|
||||
finder = MissingLitellm()
|
||||
previous_litellm = sys.modules.get('litellm')
|
||||
had_litellm = 'litellm' in sys.modules
|
||||
sys.meta_path.insert(0, finder)
|
||||
sys.modules.pop('litellm', None)
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let result = PythonSettings::Http.read_or_unset(py);
|
||||
assert!(result.unwrap().is_none());
|
||||
py.run(
|
||||
c"
|
||||
sys.meta_path.remove(finder)
|
||||
if had_litellm:
|
||||
sys.modules['litellm'] = previous_litellm
|
||||
else:
|
||||
sys.modules.pop('litellm', None)
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn read_or_unset_propagates_nested_module_not_found_errors() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"
|
||||
import sys
|
||||
import types
|
||||
previous_modules = {
|
||||
name: sys.modules[name]
|
||||
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.settings')
|
||||
if name in sys.modules
|
||||
}
|
||||
litellm = types.ModuleType('litellm')
|
||||
litellm.__path__ = []
|
||||
rust_bridge = types.ModuleType('litellm.rust_bridge')
|
||||
rust_bridge.__path__ = []
|
||||
settings = types.ModuleType('litellm.rust_bridge.settings')
|
||||
def http_settings():
|
||||
raise ModuleNotFoundError('No module named certifi', name='certifi')
|
||||
settings.http_settings = http_settings
|
||||
litellm.rust_bridge = rust_bridge
|
||||
rust_bridge.settings = settings
|
||||
sys.modules['litellm'] = litellm
|
||||
sys.modules['litellm.rust_bridge'] = rust_bridge
|
||||
sys.modules['litellm.rust_bridge.settings'] = settings
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let error = match PythonSettings::Http.read_or_unset(py) {
|
||||
Ok(_) => panic!("nested module errors must propagate"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(error.is_instance_of::<PyModuleNotFoundError>(py));
|
||||
assert_eq!(
|
||||
error
|
||||
.value(py)
|
||||
.getattr("name")
|
||||
.unwrap()
|
||||
.extract::<String>()
|
||||
.unwrap(),
|
||||
"certifi"
|
||||
);
|
||||
py.run(
|
||||
c"
|
||||
for name in ('litellm.rust_bridge.settings', 'litellm.rust_bridge', 'litellm'):
|
||||
sys.modules.pop(name, None)
|
||||
sys.modules.update(previous_modules)
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn read_or_unset_propagates_import_errors() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"
|
||||
import sys
|
||||
import types
|
||||
previous_modules = {
|
||||
name: sys.modules[name]
|
||||
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.settings')
|
||||
if name in sys.modules
|
||||
}
|
||||
litellm = types.ModuleType('litellm')
|
||||
litellm.__path__ = []
|
||||
rust_bridge = types.ModuleType('litellm.rust_bridge')
|
||||
rust_bridge.__path__ = []
|
||||
settings = types.ModuleType('litellm.rust_bridge.settings')
|
||||
def http_settings():
|
||||
raise ImportError('cannot import name setting')
|
||||
settings.http_settings = http_settings
|
||||
litellm.rust_bridge = rust_bridge
|
||||
rust_bridge.settings = settings
|
||||
sys.modules['litellm'] = litellm
|
||||
sys.modules['litellm.rust_bridge'] = rust_bridge
|
||||
sys.modules['litellm.rust_bridge.settings'] = settings
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let error = match PythonSettings::Http.read_or_unset(py) {
|
||||
Ok(_) => panic!("import errors must propagate"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(error.is_instance_of::<PyImportError>(py));
|
||||
assert_eq!(error.to_string(), "ImportError: cannot import name setting");
|
||||
py.run(
|
||||
c"
|
||||
for name in ('litellm.rust_bridge.settings', 'litellm.rust_bridge', 'litellm'):
|
||||
sys.modules.pop(name, None)
|
||||
sys.modules.update(previous_modules)
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,8 @@ use litellm_core::audio_transcription::{
|
|||
Error, audio_transcription as run_audio_transcription, types::AudioTranscriptionRequest,
|
||||
};
|
||||
use litellm_host_python::from_py_argument;
|
||||
use pyo3::prelude::*;
|
||||
use litellm_http::HttpClientConfig;
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
|
|
@ -12,6 +13,7 @@ use crate::{
|
|||
};
|
||||
|
||||
async fn execute(
|
||||
config: HttpClientConfig,
|
||||
audio: Value,
|
||||
optional_params: Map<String, Value>,
|
||||
options: RouteOptions,
|
||||
|
|
@ -24,16 +26,20 @@ async fn execute(
|
|||
extra_headers,
|
||||
timeout,
|
||||
} = options;
|
||||
run_audio_transcription(AudioTranscriptionRequest {
|
||||
model: &model,
|
||||
audio,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers,
|
||||
optional_params,
|
||||
timeout,
|
||||
})
|
||||
run_audio_transcription(
|
||||
crate::http::pool(),
|
||||
&config,
|
||||
AudioTranscriptionRequest {
|
||||
model: &model,
|
||||
audio,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers,
|
||||
optional_params,
|
||||
timeout,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
|
|
@ -62,9 +68,10 @@ pub(crate) fn transcription(
|
|||
extra_headers,
|
||||
timeout: optional_timeout(timeout_seconds),
|
||||
};
|
||||
let config = crate::http::call_config(py, &PyDict::new(py), false)?;
|
||||
run_sync(
|
||||
py,
|
||||
execute(audio, optional_params.unwrap_or_default(), options),
|
||||
execute(config, audio, optional_params.unwrap_or_default(), options),
|
||||
audio_transcription_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
|
@ -94,9 +101,10 @@ pub(crate) fn atranscription<'py>(
|
|||
extra_headers,
|
||||
timeout: optional_timeout(timeout_seconds),
|
||||
};
|
||||
let config = crate::http::call_config(py, &PyDict::new(py), true)?;
|
||||
run_async(
|
||||
py,
|
||||
execute(audio, optional_params.unwrap_or_default(), options),
|
||||
execute(config, audio, optional_params.unwrap_or_default(), options),
|
||||
audio_transcription_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ use litellm_core::chat_completions::{
|
|||
types::ChatCompletionsRequest,
|
||||
};
|
||||
use litellm_host_python::from_py_argument;
|
||||
use litellm_http::HttpClientConfig;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
|
|
@ -20,6 +21,7 @@ use crate::{
|
|||
};
|
||||
|
||||
async fn execute(
|
||||
config: HttpClientConfig,
|
||||
messages: Vec<Value>,
|
||||
optional_params: Map<String, Value>,
|
||||
options: RouteOptions,
|
||||
|
|
@ -32,16 +34,20 @@ async fn execute(
|
|||
extra_headers,
|
||||
timeout,
|
||||
} = options;
|
||||
run_chat_completions(ChatCompletionsRequest {
|
||||
model: &model,
|
||||
messages: Value::Array(messages),
|
||||
optional_params,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers,
|
||||
timeout,
|
||||
})
|
||||
run_chat_completions(
|
||||
crate::http::pool(),
|
||||
&config,
|
||||
ChatCompletionsRequest {
|
||||
model: &model,
|
||||
messages: Value::Array(messages),
|
||||
optional_params,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers,
|
||||
timeout,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
|
|
@ -87,9 +93,15 @@ pub(crate) fn chat_completions(
|
|||
extra_headers,
|
||||
timeout: optional_timeout(timeout_seconds),
|
||||
};
|
||||
let config = crate::http::call_config(py, &PyDict::new(py), false)?;
|
||||
run_sync(
|
||||
py,
|
||||
execute(messages, optional_params.unwrap_or_default(), options),
|
||||
execute(
|
||||
config,
|
||||
messages,
|
||||
optional_params.unwrap_or_default(),
|
||||
options,
|
||||
),
|
||||
chat_completions_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
|
@ -119,9 +131,15 @@ pub(crate) fn achat_completions<'py>(
|
|||
extra_headers,
|
||||
timeout: optional_timeout(timeout_seconds),
|
||||
};
|
||||
let config = crate::http::call_config(py, &PyDict::new(py), true)?;
|
||||
run_async(
|
||||
py,
|
||||
execute(messages, optional_params.unwrap_or_default(), options),
|
||||
execute(
|
||||
config,
|
||||
messages,
|
||||
optional_params.unwrap_or_default(),
|
||||
options,
|
||||
),
|
||||
chat_completions_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -27,12 +27,16 @@ fn run_messages(
|
|||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let secrets = crate::secrets::source(py)?;
|
||||
let config = crate::http::call_config(py, &kwargs, asynchronous)?;
|
||||
let machine = messages_machine(crate::http::pool(), &config, secrets)
|
||||
.map_err(crate::http::client_error)?;
|
||||
run_legacy_call(
|
||||
py,
|
||||
SURFACE,
|
||||
PublicCall::capture(&request, &args, &kwargs)?,
|
||||
crate::logger::LoggedMachine::new(messages_machine(secrets)),
|
||||
crate::logger::LoggedMachine::new(machine),
|
||||
MessagesPythonHost::new(request.unbind()),
|
||||
crate::preflight::sdk_preflight,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -70,6 +70,7 @@ fn run_ocr(
|
|||
PublicCall::capture(&request, &args, &kwargs)?,
|
||||
crate::logger::LoggedMachine::new(ocr_machine(client)),
|
||||
OcrPythonHost::new(request.unbind()),
|
||||
crate::preflight::sdk_preflight,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue