mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_responses_queued_id_encryption
# Conflicts: # type-discipline-budget.json
This commit is contained in:
commit
2f5d9ae194
482 changed files with 75631 additions and 3618 deletions
5
.github/ci-coverage-allowlist.yml
vendored
5
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -4,6 +4,11 @@ description: >-
|
|||
by a job nor listed here, so every entry below is a decision on the record.
|
||||
|
||||
test_paths:
|
||||
- reason: >-
|
||||
The Rust/Python parity harness is run manually through its local CLI. Recorded replay,
|
||||
fixture generation, and harness checks are intentionally outside pull request CI
|
||||
paths:
|
||||
- tests/rust-python-harness
|
||||
- reason: >-
|
||||
What is left of the caching suite in tests/local_testing that runs nowhere. Every job that
|
||||
globs that directory either deselects it (local_testing_part1 and part2 carry `-k "... and
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 5601
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15288
|
||||
"limit": 15287
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -105,10 +105,10 @@
|
|||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38324
|
||||
"limit": 38323
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19625
|
||||
"limit": 19624
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29861
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.63"
|
||||
version = "0.1.64"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.63"
|
||||
version = "0.1.64"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import shutil
|
|||
import subprocess
|
||||
import tempfile
|
||||
import time
|
||||
from dataclasses import dataclass, replace
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
|
|
@ -45,6 +46,38 @@ _MIGRATION_TS_RE = re.compile(r"^(\d{14})_")
|
|||
|
||||
_MIGRATION_DEADLOCK_MARKER = "deadlock detected"
|
||||
|
||||
MAX_MIGRATE_DEPLOY_ATTEMPTS = 4
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _MigrateAttemptBudget:
|
||||
"""Retries left, and the recoveries already run.
|
||||
|
||||
A recovery that lands something new costs nothing, so a database full of
|
||||
objects `prisma db push` created works through them one per pass. Anything
|
||||
that made no progress spends an attempt, so a stuck run still gives up.
|
||||
"""
|
||||
|
||||
attempts_left: int
|
||||
recoveries: frozenset[str] = frozenset()
|
||||
|
||||
@property
|
||||
def exhausted(self) -> bool:
|
||||
return self.attempts_left <= 0
|
||||
|
||||
@property
|
||||
def attempt_number(self) -> int:
|
||||
return MAX_MIGRATE_DEPLOY_ATTEMPTS - self.attempts_left + 1
|
||||
|
||||
def spend(self) -> "_MigrateAttemptBudget":
|
||||
return replace(self, attempts_left=self.attempts_left - 1)
|
||||
|
||||
def after_recovery(self, recovery: str) -> "_MigrateAttemptBudget":
|
||||
if recovery in self.recoveries:
|
||||
return self.spend()
|
||||
return replace(self, recoveries=self.recoveries | {recovery})
|
||||
|
||||
|
||||
_SPEND_LOGS_ALTER_RE = re.compile(r'^ALTER\s+TABLE\s+"LiteLLM_SpendLogs"\s', re.IGNORECASE)
|
||||
_SPEND_LOGS_ARTIFACT_DROP_RE = re.compile(
|
||||
r'^DROP\s+TABLE\s+"LiteLLM_SpendLogs_[^"]*"', re.IGNORECASE
|
||||
|
|
@ -716,6 +749,9 @@ class ProxyExtrasDBManager:
|
|||
Ahead-of-HEAD state (DB has migrations newer than this build ships)
|
||||
is logged as a warning, not a fatal error — users whose DBs got into
|
||||
weird shapes from the old thrashing should still be able to start.
|
||||
|
||||
The retry budget only counts attempts that made no progress: see
|
||||
_MigrateAttemptBudget.
|
||||
"""
|
||||
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
|
||||
migrations_dir = ProxyExtrasDBManager._get_prisma_dir()
|
||||
|
|
@ -749,8 +785,9 @@ class ProxyExtrasDBManager:
|
|||
original_dir = os.getcwd()
|
||||
os.chdir(migrations_dir)
|
||||
deploy_timeout = prisma_migrate_deploy_timeout()
|
||||
budget = _MigrateAttemptBudget(attempts_left=MAX_MIGRATE_DEPLOY_ATTEMPTS)
|
||||
try:
|
||||
for attempt in range(4):
|
||||
while not budget.exhausted:
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[_get_prisma_command(), "migrate", "deploy"],
|
||||
|
|
@ -767,168 +804,155 @@ class ProxyExtrasDBManager:
|
|||
logger.warning(
|
||||
"prisma migrate deploy attempt %s timed out after %ss, retrying. "
|
||||
"Raise %s if this database needs longer to apply its pending migrations.",
|
||||
attempt + 1,
|
||||
budget.attempt_number,
|
||||
deploy_timeout,
|
||||
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR,
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
continue
|
||||
next_budget = budget.spend()
|
||||
|
||||
except subprocess.CalledProcessError as e:
|
||||
stderr = e.stderr or ""
|
||||
next_budget = ProxyExtrasDBManager._budget_after_deploy_failure(
|
||||
e, budget, schema_path
|
||||
)
|
||||
|
||||
if "P3005" in stderr and "database schema is not empty" in stderr:
|
||||
logger.info(
|
||||
"Schema exists but no migrations ledger — creating baseline"
|
||||
)
|
||||
ProxyExtrasDBManager._create_baseline_migration(schema_path)
|
||||
continue
|
||||
|
||||
if "P3009" in stderr:
|
||||
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
|
||||
if (
|
||||
migration_match
|
||||
and ProxyExtrasDBManager._is_idempotent_error(stderr)
|
||||
):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} failed idempotently — marking applied and retrying"
|
||||
)
|
||||
try:
|
||||
ProxyExtrasDBManager._roll_back_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
):
|
||||
pass # may already be rolled-back
|
||||
try:
|
||||
ProxyExtrasDBManager._resolve_specific_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
) as resolve_err:
|
||||
# We're already inside the outer
|
||||
# `except CalledProcessError` handler —
|
||||
# re-raising CalledProcessError from here
|
||||
# would escape as itself, bypassing
|
||||
# proxy_cli.py's `except RuntimeError`.
|
||||
raise RuntimeError(
|
||||
f"Failed to mark migration {name} as applied "
|
||||
f"after idempotent recovery. Manual "
|
||||
f"intervention may be required.\n\n"
|
||||
f"Detail: {resolve_err}"
|
||||
) from resolve_err
|
||||
continue
|
||||
if migration_match:
|
||||
migration_name = migration_match.group(1)
|
||||
ledger_logs = ProxyExtrasDBManager._failed_migration_logs(migration_name)
|
||||
if ledger_logs is not None and (
|
||||
ledger_logs == "" or _MIGRATION_DEADLOCK_MARKER in ledger_logs
|
||||
):
|
||||
logger.info(
|
||||
"Migration %s failed in a concurrent migrate deploy "
|
||||
"deadlock race, rolling its ledger row back and retrying",
|
||||
migration_name,
|
||||
)
|
||||
ProxyExtrasDBManager._roll_back_migration_best_effort(migration_name)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
continue
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
if "P3018" in stderr:
|
||||
if ProxyExtrasDBManager._is_permission_error(stderr):
|
||||
raise RuntimeError(
|
||||
"Database migration failed due to insufficient "
|
||||
"permissions. Please grant the required privileges "
|
||||
f"and retry.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
migration_match = re.search(
|
||||
r"Migration name: (\d+_\S+)", stderr
|
||||
)
|
||||
if (
|
||||
migration_match
|
||||
and ProxyExtrasDBManager._is_idempotent_error(stderr)
|
||||
):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
|
||||
)
|
||||
try:
|
||||
ProxyExtrasDBManager._roll_back_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
):
|
||||
pass # may already be rolled-back
|
||||
try:
|
||||
ProxyExtrasDBManager._resolve_specific_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
) as resolve_err:
|
||||
raise RuntimeError(
|
||||
f"Failed to mark migration {name} as applied "
|
||||
f"after idempotent recovery. Manual "
|
||||
f"intervention may be required.\n\n"
|
||||
f"Detail: {resolve_err}"
|
||||
) from resolve_err
|
||||
continue
|
||||
|
||||
if migration_match and _MIGRATION_DEADLOCK_MARKER in stderr:
|
||||
logger.info(
|
||||
"Migration %s deadlocked against a concurrent "
|
||||
"migrate deploy, rolling its ledger row back "
|
||||
"and retrying",
|
||||
migration_match.group(1),
|
||||
)
|
||||
ProxyExtrasDBManager._roll_back_migration_best_effort(
|
||||
migration_match.group(1)
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
continue
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
|
||||
if _MIGRATION_DEADLOCK_MARKER in stderr:
|
||||
logger.info(
|
||||
"prisma migrate deploy attempt %s deadlocked against "
|
||||
"a concurrent migrate deploy, retrying",
|
||||
attempt + 1,
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
continue
|
||||
|
||||
if "P1002" in stderr and "advisory lock" in stderr:
|
||||
logger.info(
|
||||
"prisma migrate deploy attempt %s timed out waiting for "
|
||||
"the advisory lock a concurrent migrate deploy holds, retrying",
|
||||
attempt + 1,
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
continue
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from e
|
||||
if next_budget.attempts_left < budget.attempts_left:
|
||||
time.sleep(random.randrange(5, 15))
|
||||
budget = next_budget # rebind-ok: the loop carries the budget from one migrate deploy pass to the next
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed after 4 attempts (retry loop "
|
||||
"exhausted by timeouts, deadlock retries, or repeated "
|
||||
"idempotent-recovery continues). Check database connectivity, "
|
||||
f"Database migration failed after {MAX_MIGRATE_DEPLOY_ATTEMPTS} "
|
||||
"attempts that made no progress (timeouts, deadlock retries, or a "
|
||||
"recovery that had already run once). Check database connectivity, "
|
||||
"load, and _prisma_migrations ledger state, and raise "
|
||||
f"{PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR} if the attempts timed out."
|
||||
)
|
||||
finally:
|
||||
os.chdir(original_dir)
|
||||
|
||||
@staticmethod
|
||||
def _budget_after_deploy_failure(
|
||||
error: subprocess.CalledProcessError,
|
||||
budget: "_MigrateAttemptBudget",
|
||||
schema_path: str,
|
||||
) -> "_MigrateAttemptBudget":
|
||||
"""Recover from one failed `prisma migrate deploy`, and price the pass.
|
||||
|
||||
Returns the budget the next pass runs under, or raises when the failure
|
||||
is not one this resolver knows how to recover from.
|
||||
"""
|
||||
stderr = error.stderr or ""
|
||||
|
||||
if "P3005" in stderr and "database schema is not empty" in stderr:
|
||||
logger.info("Schema exists but no migrations ledger — creating baseline")
|
||||
if ProxyExtrasDBManager._create_baseline_migration(schema_path):
|
||||
return budget.after_recovery("baseline")
|
||||
return budget.spend()
|
||||
|
||||
if "P3009" in stderr:
|
||||
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
|
||||
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} failed idempotently — marking applied and retrying"
|
||||
)
|
||||
ProxyExtrasDBManager._mark_migration_applied(name)
|
||||
return budget.after_recovery(f"resolved:{name}")
|
||||
if migration_match:
|
||||
migration_name = migration_match.group(1)
|
||||
ledger_logs = ProxyExtrasDBManager._failed_migration_logs(migration_name)
|
||||
if ledger_logs is not None and (
|
||||
ledger_logs == "" or _MIGRATION_DEADLOCK_MARKER in ledger_logs
|
||||
):
|
||||
logger.info(
|
||||
"Migration %s failed in a concurrent migrate deploy "
|
||||
"deadlock race, rolling its ledger row back and retrying",
|
||||
migration_name,
|
||||
)
|
||||
ProxyExtrasDBManager._roll_back_migration_best_effort(migration_name)
|
||||
return budget.spend()
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
if "P3018" in stderr:
|
||||
if ProxyExtrasDBManager._is_permission_error(stderr):
|
||||
raise RuntimeError(
|
||||
"Database migration failed due to insufficient "
|
||||
"permissions. Please grant the required privileges "
|
||||
f"and retry.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
migration_match = re.search(r"Migration name: (\d+_\S+)", stderr)
|
||||
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
|
||||
name = migration_match.group(1)
|
||||
logger.info(
|
||||
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
|
||||
)
|
||||
ProxyExtrasDBManager._mark_migration_applied(name)
|
||||
return budget.after_recovery(f"resolved:{name}")
|
||||
|
||||
if migration_match and _MIGRATION_DEADLOCK_MARKER in stderr:
|
||||
logger.info(
|
||||
"Migration %s deadlocked against a concurrent "
|
||||
"migrate deploy, rolling its ledger row back "
|
||||
"and retrying",
|
||||
migration_match.group(1),
|
||||
)
|
||||
ProxyExtrasDBManager._roll_back_migration_best_effort(
|
||||
migration_match.group(1)
|
||||
)
|
||||
return budget.spend()
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
if _MIGRATION_DEADLOCK_MARKER in stderr:
|
||||
logger.info(
|
||||
"prisma migrate deploy attempt %s deadlocked against "
|
||||
"a concurrent migrate deploy, retrying",
|
||||
budget.attempt_number,
|
||||
)
|
||||
return budget.spend()
|
||||
|
||||
if "P1002" in stderr and "advisory lock" in stderr:
|
||||
logger.info(
|
||||
"prisma migrate deploy attempt %s timed out waiting for "
|
||||
"the advisory lock a concurrent migrate deploy holds, retrying",
|
||||
budget.attempt_number,
|
||||
)
|
||||
return budget.spend()
|
||||
|
||||
raise RuntimeError(
|
||||
"Database migration failed and cannot be auto-recovered. "
|
||||
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
|
||||
) from error
|
||||
|
||||
@staticmethod
|
||||
def _mark_migration_applied(name: str) -> None:
|
||||
"""Roll a failed ledger row back if it is still there, then mark it applied."""
|
||||
try:
|
||||
ProxyExtrasDBManager._roll_back_migration(name)
|
||||
except (subprocess.CalledProcessError, subprocess.TimeoutExpired):
|
||||
pass # may already be rolled-back
|
||||
try:
|
||||
ProxyExtrasDBManager._resolve_specific_migration(name)
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
subprocess.TimeoutExpired,
|
||||
) as resolve_err:
|
||||
# We're called from inside an `except CalledProcessError` handler —
|
||||
# re-raising CalledProcessError from here would escape as itself,
|
||||
# bypassing proxy_cli.py's `except RuntimeError`.
|
||||
raise RuntimeError(
|
||||
f"Failed to mark migration {name} as applied "
|
||||
f"after idempotent recovery. Manual "
|
||||
f"intervention may be required.\n\n"
|
||||
f"Detail: {resolve_err}"
|
||||
) from resolve_err
|
||||
|
||||
@staticmethod
|
||||
def apply_replica_identity_full_if_requested() -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.92"
|
||||
version = "0.4.93"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.92"
|
||||
version = "0.4.93"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ rand = "0.8"
|
|||
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls", "http2", "stream"] }
|
||||
rstest = "0.26.1"
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
serde_json = "1.0"
|
||||
serde_json = { version = "1.0", features = ["float_roundtrip"] }
|
||||
sha2 = "0.10"
|
||||
subtle = "2"
|
||||
thiserror = "2.0"
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ const AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGE
|
|||
const AZURE_DOCUMENT_INTELLIGENCE_API_VERSION: &str = "2024-11-30";
|
||||
const AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI: i64 = 96;
|
||||
|
||||
const AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS: &[&str] = &["pages"];
|
||||
const AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS: &[&str] = &["pages", "features"];
|
||||
|
||||
pub struct AzureAiOcrConfig;
|
||||
pub struct AzureDocumentIntelligenceOcrConfig;
|
||||
|
|
@ -192,6 +192,46 @@ fn normalize_pages_param(pages: &Value) -> Result<Option<String>, Error> {
|
|||
}
|
||||
}
|
||||
|
||||
fn feature_token_is_valid(token: &str) -> bool {
|
||||
let Some((first, rest)) = token.as_bytes().split_first() else {
|
||||
return false;
|
||||
};
|
||||
first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric)
|
||||
}
|
||||
|
||||
fn invalid_features_error(features: &Value) -> Error {
|
||||
Error::InvalidRequest(format!(
|
||||
"Invalid `features` for Azure Document Intelligence: {features:?}. Expected a list of feature names or a comma-separated string like 'keyValuePairs' or 'keyValuePairs,languages'."
|
||||
))
|
||||
}
|
||||
|
||||
fn normalize_features_param(features: &Value) -> Result<Option<String>, Error> {
|
||||
let normalized = match features {
|
||||
Value::String(value) => value
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.collect::<Vec<_>>()
|
||||
.join(","),
|
||||
Value::Array(values) if values.is_empty() => return Ok(None),
|
||||
Value::Array(values) => values
|
||||
.iter()
|
||||
.map(Value::as_str)
|
||||
.collect::<Option<Vec<_>>>()
|
||||
.ok_or_else(|| invalid_features_error(features))?
|
||||
.into_iter()
|
||||
.map(str::trim)
|
||||
.collect::<Vec<_>>()
|
||||
.join(","),
|
||||
_ => return Err(invalid_features_error(features)),
|
||||
};
|
||||
|
||||
if normalized.split(',').all(feature_token_is_valid) {
|
||||
Ok(Some(normalized))
|
||||
} else {
|
||||
Err(invalid_features_error(features))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn complete_document_intelligence_url(
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
|
|
@ -213,6 +253,13 @@ pub fn complete_document_intelligence_url(
|
|||
url.push_str(&normalized);
|
||||
}
|
||||
|
||||
if let Some(features) = optional_params.get("features")
|
||||
&& let Some(normalized) = normalize_features_param(features)?
|
||||
{
|
||||
url.push_str("&features=");
|
||||
url.push_str(&normalized);
|
||||
}
|
||||
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
|
|
@ -475,6 +522,103 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_normalizes_features() {
|
||||
let params = serde_json::Map::from_iter([(
|
||||
"features".to_string(),
|
||||
json!("keyValuePairs, languages"),
|
||||
)]);
|
||||
let url = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect("url builds");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&features=keyValuePairs,languages"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_combines_pages_and_feature_list() {
|
||||
let params = serde_json::Map::from_iter([
|
||||
("pages".to_string(), json!([0, 1, 2])),
|
||||
(
|
||||
"features".to_string(),
|
||||
json!([" keyValuePairs ", "languages"]),
|
||||
),
|
||||
]);
|
||||
let url = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect("url builds");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&pages=1,2,3&features=keyValuePairs,languages"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_omits_empty_feature_list() {
|
||||
let params = serde_json::Map::from_iter([("features".to_string(), json!([]))]);
|
||||
let url = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect("url builds");
|
||||
|
||||
assert_eq!(
|
||||
url,
|
||||
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_url_rejects_invalid_features() {
|
||||
for features in [
|
||||
json!("keyValuePairs&pages=9"),
|
||||
json!(""),
|
||||
json!(["keyValuePairs", 1]),
|
||||
json!({"feature": "keyValuePairs"}),
|
||||
] {
|
||||
let params = serde_json::Map::from_iter([("features".to_string(), features.clone())]);
|
||||
let error = complete_document_intelligence_url(
|
||||
Some("https://example.cognitiveservices.azure.com"),
|
||||
"prebuilt-layout",
|
||||
¶ms,
|
||||
&|_| None,
|
||||
)
|
||||
.expect_err("invalid features must fail");
|
||||
|
||||
assert!(
|
||||
matches!(error, Error::InvalidRequest(message) if message.contains("Invalid `features`")),
|
||||
"features={features:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_maps_features() {
|
||||
let params = Map::from_iter([
|
||||
("features".to_string(), json!(["keyValuePairs"])),
|
||||
("unsupported".to_string(), json!(true)),
|
||||
]);
|
||||
|
||||
assert_eq!(
|
||||
AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG.map_ocr_params(¶ms),
|
||||
Map::from_iter([("features".to_string(), json!(["keyValuePairs"]))])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_intelligence_request_uses_base64_source_for_data_uri() {
|
||||
let body = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG
|
||||
|
|
|
|||
|
|
@ -59,3 +59,41 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
|||
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
|
||||
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
|
||||
}
|
||||
|
||||
pub(crate) fn ocr_error_to_pyerr(err: Error) -> PyErr {
|
||||
match err {
|
||||
Error::MissingField("document_url" | "image_url") => {
|
||||
PyValueError::new_err("Document URL is required")
|
||||
}
|
||||
Error::Http { status, body } => RustUpstreamError::new_err((status, body)),
|
||||
other => core_error_to_pyerr(other),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod ocr_error_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn ocr_errors_preserve_python_validation_and_provider_details() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
for field in ["document_url", "image_url"] {
|
||||
let mapped = ocr_error_to_pyerr(Error::MissingField(field));
|
||||
assert!(mapped.is_instance_of::<PyValueError>(py));
|
||||
assert_eq!(mapped.value(py).to_string(), "Document URL is required");
|
||||
}
|
||||
let mapped = ocr_error_to_pyerr(Error::Http {
|
||||
status: 429,
|
||||
body: r#"{"message":"rate limited"}"#.to_string(),
|
||||
});
|
||||
assert!(mapped.is_instance_of::<RustUpstreamError>(py));
|
||||
let args: (u16, String) = mapped
|
||||
.value(py)
|
||||
.getattr("args")
|
||||
.and_then(|args| args.extract())
|
||||
.expect("OCR failures retain status and unprefixed provider message");
|
||||
assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string()));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
|
|||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::errors::ocr_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
|
||||
|
||||
fn prepare_ocr(
|
||||
|
|
@ -69,5 +69,5 @@ bridge_route! {
|
|||
timeout_seconds: Option<f64>,
|
||||
},
|
||||
prepare = prepare_ocr,
|
||||
errors = core_error_to_pyerr,
|
||||
errors = ocr_error_to_pyerr,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -932,7 +932,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode)
|
||||
|
||||
def _convert_content_str_to_input_text(self, content: str, role: str) -> dict[str, object]:
|
||||
if role == "user" or role == "system" or role == "tool":
|
||||
if role in ("user", "system", "developer", "tool"):
|
||||
return {"type": "input_text", "text": content}
|
||||
else:
|
||||
return {"type": "output_text", "text": content}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_ran
|
|||
|
||||
DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
|
||||
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
|
||||
AZURE_OPENAI_AUDIO_PROVIDERS: Final = frozenset({"azure", "azure_ai"})
|
||||
ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5))
|
||||
ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS: Final = 2000
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS: Final[frozenset[str]] = frozenset(
|
||||
|
|
@ -39,6 +40,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
|
|||
"router_general_settings",
|
||||
"ignore_invalid_deployments",
|
||||
"fallback_access_check",
|
||||
"heuristic_v2_router_limit",
|
||||
}
|
||||
)
|
||||
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
||||
|
|
@ -1450,6 +1452,7 @@ SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY: Final = "_session_deployment_affin
|
|||
CONSUMED_REQUEST_TAGS_METADATA_KEY: Final = "_consumed_request_tags"
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY: Final = "internal_call_origin"
|
||||
SESSION_ID_GENERATED_METADATA_KEY: Final = "litellm_session_id_generated"
|
||||
SESSION_ID_OMITTED_METADATA_KEY: Final = "litellm_session_id_omitted"
|
||||
LITELLM_TRUNCATED_PAYLOAD_FIELD: Final = "litellm_truncated"
|
||||
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE: Final = (
|
||||
"Truncation is a DB storage safeguard. "
|
||||
|
|
|
|||
|
|
@ -18,9 +18,11 @@ caller's identity metadata, minus two things that must never be forwarded as-is:
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY, NON_INFERENCE_CALL_TYPES
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import initialize_standard_callback_dynamic_params
|
||||
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, InternalCallOrigin
|
||||
|
||||
BUDGET_RESERVATION_METADATA_KEYS: Final = frozenset({"user_api_key_budget_reservation"})
|
||||
|
|
@ -142,6 +144,19 @@ def forwarded_internal_call_metadata(
|
|||
}
|
||||
|
||||
|
||||
def parent_session_kwargs(request_kwargs: Mapping[str, object] | None) -> Mapping[str, str]:
|
||||
kwargs: Final = request_kwargs or MappingProxyType({})
|
||||
return MappingProxyType(
|
||||
{k: v for k in ("litellm_session_id", "litellm_trace_id") if isinstance(v := kwargs.get(k), str)}
|
||||
)
|
||||
|
||||
|
||||
def effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | None) -> bool | None:
|
||||
return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else None).get(
|
||||
"turn_off_message_logging"
|
||||
)
|
||||
|
||||
|
||||
def sanitized_forwardable_call_metadata(
|
||||
parent_metadata: Mapping[str, object],
|
||||
call_origin: InternalCallOrigin,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ Helper utilities for tracking the cost of built-in tools.
|
|||
from collections.abc import Mapping
|
||||
from typing import Final, Literal
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.constants import OPENAI_FILE_SEARCH_COST_PER_1K_CALLS
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
|
|
@ -13,6 +15,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
|||
from litellm.types.llms.openai import (
|
||||
FileSearchTool,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesToolUsage,
|
||||
WebSearchOptions,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -32,6 +35,17 @@ def _output_item_type(output_item: object) -> str | None:
|
|||
return item_type if isinstance(item_type, str) else None
|
||||
|
||||
|
||||
def _reported_web_search_requests(response_object: ResponsesAPIResponse) -> int | None:
|
||||
tool_usage: Final = getattr(response_object, "tool_usage", None)
|
||||
if tool_usage is None:
|
||||
return None
|
||||
try:
|
||||
web_search: Final = ResponsesToolUsage.model_validate(tool_usage).web_search
|
||||
except ValidationError:
|
||||
return None
|
||||
return None if web_search is None else web_search.num_requests
|
||||
|
||||
|
||||
def _usage_reports_server_side_web_search_calls(usage: Usage) -> bool:
|
||||
details: Final = getattr(usage, "server_side_tool_usage_details", None)
|
||||
if not isinstance(details, Mapping):
|
||||
|
|
@ -182,15 +196,19 @@ class StandardBuiltInToolCostTracking:
|
|||
|
||||
Providers that report a request count in usage (gemini, anthropic, xai, vertex) are handled by
|
||||
get_cost_for_web_search_request and never reach here. This path prices per call, so it must count
|
||||
the web_search_call items. Chat-completions responses only expose url_citation annotations with no
|
||||
count, so they floor to a single billable search.
|
||||
the web_search_call items, unless the response reports the billable count itself
|
||||
(Bedrock's tool_usage.web_search.num_requests, which excludes open_page fetches). Chat-completions
|
||||
responses only expose url_citation annotations with no count, so they floor to a single billable search.
|
||||
"""
|
||||
if isinstance(response_object, ResponsesAPIResponse):
|
||||
count = sum(
|
||||
1 for output_item in response_object.output if _output_item_type(output_item) == "web_search_call"
|
||||
)
|
||||
return max(count, 1)
|
||||
return 1
|
||||
if not isinstance(response_object, ResponsesAPIResponse):
|
||||
return 1
|
||||
reported: Final = _reported_web_search_requests(response_object)
|
||||
if reported is not None:
|
||||
return reported
|
||||
count: Final = sum(
|
||||
1 for output_item in response_object.output if _output_item_type(output_item) == "web_search_call"
|
||||
)
|
||||
return max(count, 1)
|
||||
|
||||
@staticmethod
|
||||
def _handle_file_search_cost(
|
||||
|
|
|
|||
|
|
@ -428,7 +428,7 @@ def _coerce_off_peak_rate(value: object, default: float) -> float:
|
|||
return default
|
||||
|
||||
|
||||
def _apply_off_peak_pricing(
|
||||
def apply_off_peak_pricing(
|
||||
model_info: ModelInfo,
|
||||
current_time: datetime | None,
|
||||
prompt_base_cost: float,
|
||||
|
|
@ -462,7 +462,7 @@ def _apply_off_peak_to_base_costs(
|
|||
has no field for them.
|
||||
"""
|
||||
prompt, completion, cache_creation, cache_creation_above_1hr, cache_read = base_costs
|
||||
off_peak_prompt, off_peak_completion, off_peak_cache_read = _apply_off_peak_pricing(
|
||||
off_peak_prompt, off_peak_completion, off_peak_cache_read = apply_off_peak_pricing(
|
||||
model_info, current_time, prompt, completion, cache_read
|
||||
)
|
||||
return (off_peak_prompt, off_peak_completion, cache_creation, cache_creation_above_1hr, off_peak_cache_read)
|
||||
|
|
|
|||
|
|
@ -1554,6 +1554,22 @@ def with_prompt_cache_breakpoint(target: _MarkedT, marker: object) -> _MarkedT:
|
|||
return cast(_MarkedT, marked) # cast-ok: same block shape as the input plus the marker key
|
||||
|
||||
|
||||
LITELLM_INTERNAL_MESSAGE_FIELDS: Final = frozenset({"thinking_blocks", "reasoning_content", "provider_specific_fields"})
|
||||
|
||||
|
||||
def strip_litellm_internal_message_fields(message: AllMessageValues) -> AllMessageValues:
|
||||
"""Drop the fields litellm attaches to assistant messages (e.g. when translating Anthropic thinking
|
||||
blocks) that OpenAI-compatible endpoints with strict schemas reject as extra inputs."""
|
||||
if LITELLM_INTERNAL_MESSAGE_FIELDS.isdisjoint(message):
|
||||
return message
|
||||
return cast( # cast-ok: same TypedDict minus internal keys
|
||||
AllMessageValues,
|
||||
{ # mutable-ok: provider transforms mutate message dicts in place downstream
|
||||
key: value for key, value in message.items() if key not in LITELLM_INTERNAL_MESSAGE_FIELDS
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def filter_value_from_dict(dictionary: dict, key: str, depth: int = 0) -> Any:
|
||||
"""
|
||||
Filters a value from a dictionary
|
||||
|
|
|
|||
|
|
@ -2337,6 +2337,9 @@ class CustomStreamWrapper:
|
|||
else:
|
||||
self.sent_last_chunk = True
|
||||
processed_chunk: Final = self.finish_reason_handler()
|
||||
if self.stream_options is None:
|
||||
usage: Final = calculate_total_usage(chunks=self.chunks)
|
||||
processed_chunk._hidden_params["usage"] = usage # pyright: ignore[reportPrivateUsage] # sync parity
|
||||
# see sync __next__'s sibling branch: deliberately do NOT restore
|
||||
# here - this chunk is still this call's own data, and restoring
|
||||
# before returning it would corrupt the caller's own log
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from litellm.exceptions import UnsupportedParamsError
|
|||
from litellm.llms.openai.chat.gpt_5_transformation import (
|
||||
OpenAIGPT5Config,
|
||||
_get_effort_level,
|
||||
is_gpt_reasoning_series_name,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
|
@ -35,26 +36,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
|
|||
|
||||
@classmethod
|
||||
def is_model_gpt_5_model(cls, model: str) -> bool:
|
||||
"""Check if the Azure model string refers to a gpt-5 variant.
|
||||
|
||||
Accepts both explicit gpt-5 model names and the ``gpt5_series/`` prefix
|
||||
used for manual routing.
|
||||
"""
|
||||
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
|
||||
# …) are regular chat models: they support temperature and tool_choice but NOT
|
||||
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
|
||||
#
|
||||
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
|
||||
# models and must stay on the GPT-5 path. The distinguishing feature is that
|
||||
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
|
||||
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
|
||||
# number (i.e. "gpt-5.<digit>-chat").
|
||||
#
|
||||
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
|
||||
# than a substring check) makes this boundary explicit and avoids any ambiguity
|
||||
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
|
||||
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "azure/"
|
||||
return ("gpt-5" in model and not _normalized.startswith("gpt-5-chat")) or "gpt5_series" in model
|
||||
return is_gpt_reasoning_series_name(model) or "gpt5_series" in model
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[str]:
|
||||
"""Get supported parameters for Azure OpenAI GPT-5 models.
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
convert_to_azure_openai_messages,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import GPT_REASONING_SERIES_MARKERS
|
||||
from litellm.types.llms.azure import (
|
||||
API_VERSION_MONTH_SUPPORTED_RESPONSE_FORMAT,
|
||||
API_VERSION_YEAR_SUPPORTED_RESPONSE_FORMAT,
|
||||
|
|
@ -139,7 +140,7 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
name family needs the rename, including the ``gpt-5-chat*`` models that are excluded from
|
||||
the reasoning path by https://github.com/BerriAI/litellm/issues/13781.
|
||||
"""
|
||||
return "gpt-5" in model or "gpt5_series" in model
|
||||
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) or "gpt5_series" in model
|
||||
|
||||
def _is_response_format_supported_model(self, model: str) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
filter_value_from_dict,
|
||||
)
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
|
||||
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
|
||||
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
|
||||
from litellm.llms.openai.openai import OpenAIConfig
|
||||
|
|
@ -207,20 +208,18 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
message["content"] = texts
|
||||
return stripped_messages
|
||||
|
||||
def _is_azure_openai_model(self, model: str, api_base: str | None) -> bool:
|
||||
try:
|
||||
if "/" in model:
|
||||
model = model.split("/", 1)[1]
|
||||
if (
|
||||
model in litellm.open_ai_chat_completion_models
|
||||
or model in litellm.open_ai_text_completion_models
|
||||
or model in litellm.open_ai_embedding_models
|
||||
):
|
||||
return True
|
||||
def _is_foundry_model_inference_base(self, api_base: str) -> bool:
|
||||
return is_foundry_model_inference_base(api_base)
|
||||
|
||||
except Exception:
|
||||
def _is_azure_openai_model(self, model: str, api_base: str | None) -> bool:
|
||||
if api_base is None or self._is_foundry_model_inference_base(api_base):
|
||||
return False
|
||||
return False
|
||||
stripped_model: Final = model.split("/", 1)[1] if "/" in model else model
|
||||
return (
|
||||
stripped_model in litellm.open_ai_chat_completion_models
|
||||
or stripped_model in litellm.open_ai_text_completion_models
|
||||
or stripped_model in litellm.open_ai_embedding_models
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
|
||||
|
|
@ -10,6 +11,14 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
AzureAIApiKeyHeader = Literal["Authorization", "api-key", "Api-Key", "Ocp-Apim-Subscription-Key"]
|
||||
|
||||
|
||||
def is_foundry_model_inference_base(api_base: str) -> bool:
|
||||
parsed: Final = urlparse(api_base)
|
||||
host: Final = parsed.hostname
|
||||
if host is None or not host.endswith(".services.ai.azure.com"):
|
||||
return False
|
||||
return "/openai/deployments" not in parsed.path
|
||||
|
||||
|
||||
def get_azure_ai_entra_token(litellm_params: Mapping[str, object] | None = None) -> str | None:
|
||||
"""
|
||||
Resolve an Entra ID / OAuth access token for an Azure AI Foundry deployment.
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
from typing import Final
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
from openai import OpenAI
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
|
|
@ -16,6 +18,16 @@ from litellm.utils import convert_to_model_response_object
|
|||
from .cohere_transformation import AzureAICohereConfig
|
||||
|
||||
|
||||
def _foundry_models_route_base(api_base: str | None) -> str | None:
|
||||
if api_base is None or not is_foundry_model_inference_base(api_base):
|
||||
return api_base
|
||||
parts: Final = urlsplit(api_base)
|
||||
path: Final = parts.path.rstrip("/")
|
||||
if path.endswith("/models"):
|
||||
return api_base
|
||||
return urlunsplit((parts.scheme, parts.netloc, f"{path}/models", parts.query, parts.fragment))
|
||||
|
||||
|
||||
class AzureAIEmbedding(OpenAIChatCompletion):
|
||||
def _process_response(
|
||||
self,
|
||||
|
|
@ -214,6 +226,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
|
||||
assemble result in-order, and return
|
||||
"""
|
||||
resolved_api_base: Final = _foundry_models_route_base(api_base)
|
||||
if aembedding is True:
|
||||
return self.async_embedding(
|
||||
model,
|
||||
|
|
@ -223,7 +236,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
model_response,
|
||||
optional_params,
|
||||
api_key,
|
||||
api_base,
|
||||
resolved_api_base,
|
||||
client,
|
||||
)
|
||||
|
||||
|
|
@ -245,7 +258,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
model_response=model_response,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_base=resolved_api_base,
|
||||
client=client,
|
||||
)
|
||||
|
||||
|
|
@ -262,7 +275,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
model_response,
|
||||
optional_params,
|
||||
api_key,
|
||||
api_base,
|
||||
resolved_api_base,
|
||||
client=(client if client is not None and isinstance(client, OpenAI) else None),
|
||||
aembedding=aembedding,
|
||||
shared_session=shared_session,
|
||||
|
|
|
|||
|
|
@ -48,7 +48,9 @@ _BASE_SUFFIXES_TO_STRIP: Final = (
|
|||
)
|
||||
|
||||
# Per Bedrock Mantle Responses API validation errors.
|
||||
_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES = frozenset({"function", "mcp", "custom", "namespace", "tool_search"})
|
||||
_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES: Final = frozenset(
|
||||
{"function", "mcp", "custom", "namespace", "tool_search", "web_search"}
|
||||
)
|
||||
|
||||
_BEDROCK_MANTLE_SUPPORTED_SERVICE_TIERS: Final = frozenset({"auto", "default"})
|
||||
|
||||
|
|
|
|||
|
|
@ -99,9 +99,12 @@ from litellm.types.containers.main import (
|
|||
)
|
||||
from litellm.types.files import StreamingMediaUploadConfig, TwoStepFileUploadConfig
|
||||
from litellm.types.integrations.custom_logger import (
|
||||
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
|
||||
AgenticLoopPlan,
|
||||
AgenticLoopRequestPatch,
|
||||
AgenticLoopSafetyError,
|
||||
converted_stream_requested,
|
||||
is_interception_internal_key,
|
||||
)
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
|
|
@ -2760,6 +2763,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
if self._has_agentic_completion_hook(logging_obj):
|
||||
agentic_kwargs: Final = dict(litellm_params) # mutable-ok: agentic hooks mutate kwargs in place
|
||||
final_response: Final = run_async_function(
|
||||
self._call_agentic_completion_hooks,
|
||||
response=initial_response,
|
||||
|
|
@ -2770,10 +2774,19 @@ class BaseLLMHTTPHandler:
|
|||
logging_obj=logging_obj,
|
||||
stream=False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=dict(litellm_params),
|
||||
kwargs=agentic_kwargs,
|
||||
api_surface="responses",
|
||||
)
|
||||
return final_response if final_response is not None else initial_response
|
||||
result: Final = final_response if final_response is not None else initial_response
|
||||
if converted_stream_requested(agentic_kwargs) and not agentic_kwargs.get("_agentic_loop_depth"):
|
||||
return self._wrap_responses_response_as_fake_stream(
|
||||
result=result,
|
||||
model=model,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
return result
|
||||
|
||||
return initial_response
|
||||
|
||||
|
|
@ -2939,6 +2952,7 @@ class BaseLLMHTTPHandler:
|
|||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
agentic_kwargs: Final = dict(litellm_params) # mutable-ok: agentic hooks mutate kwargs in place
|
||||
final_response: Final = await self._call_agentic_completion_hooks(
|
||||
response=initial_response,
|
||||
model=model,
|
||||
|
|
@ -2948,15 +2962,12 @@ class BaseLLMHTTPHandler:
|
|||
logging_obj=logging_obj,
|
||||
stream=False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=dict(litellm_params),
|
||||
kwargs=agentic_kwargs,
|
||||
api_surface="responses",
|
||||
)
|
||||
|
||||
result: Final = final_response if final_response is not None else initial_response
|
||||
interception_converted_stream: Final = litellm_params.get(
|
||||
"_code_interpreter_interception_converted_stream"
|
||||
) or litellm_params.get("_websearch_interception_converted_stream")
|
||||
if interception_converted_stream and not litellm_params.get("_agentic_loop_depth"):
|
||||
if converted_stream_requested(agentic_kwargs) and not agentic_kwargs.get("_agentic_loop_depth"):
|
||||
return self._wrap_responses_response_as_fake_stream(
|
||||
result=result,
|
||||
model=model,
|
||||
|
|
@ -5420,8 +5431,7 @@ class BaseLLMHTTPHandler:
|
|||
kwargs_for_followup: Final = {
|
||||
k: v
|
||||
for k, v in kwargs.items()
|
||||
if not k.startswith("_websearch_interception")
|
||||
and not k.startswith("_compression_interception")
|
||||
if not is_interception_internal_key(k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES)
|
||||
and k != "_code_interpreter_interception_converted_stream"
|
||||
and k not in internal_keys
|
||||
and k not in optional_params
|
||||
|
|
|
|||
|
|
@ -7,11 +7,13 @@ cached, cache-creation, output, reasoning) is billed at that one tier's rate.
|
|||
See https://help.aliyun.com/zh/model-studio/billing-for-model-studio
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import select_tier_for_input, tier_rate
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
apply_off_peak_pricing,
|
||||
parse_completion_tokens_details,
|
||||
parse_prompt_tokens_details,
|
||||
)
|
||||
|
|
@ -32,6 +34,19 @@ class TokenBreakdown:
|
|||
return self.text_tokens + self.cached_tokens + self.cache_creation_tokens
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TokenRates:
|
||||
input_rate: float
|
||||
cache_read_rate: float
|
||||
cache_creation_rate: float
|
||||
output_rate: float
|
||||
reasoning_rate: float | None
|
||||
|
||||
@property
|
||||
def billed_reasoning_rate(self) -> float:
|
||||
return self.output_rate if self.reasoning_rate is None else self.reasoning_rate
|
||||
|
||||
|
||||
def _extract_token_breakdown(usage: Usage) -> TokenBreakdown:
|
||||
prompt_details: Final = parse_prompt_tokens_details(usage)
|
||||
cached_tokens: Final = prompt_details["cache_hit_tokens"]
|
||||
|
|
@ -57,69 +72,75 @@ def _flat_rate(model_info: ModelInfo, cost_key: str, fallback_cost_key: str) ->
|
|||
return float(value)
|
||||
|
||||
|
||||
def _calculate_prompt_cost(
|
||||
breakdown: TokenBreakdown,
|
||||
model_info: ModelInfo,
|
||||
tier: dict | None,
|
||||
) -> float:
|
||||
if tier is not None:
|
||||
return (
|
||||
(breakdown.text_tokens * tier_rate(tier, "input_cost_per_token"))
|
||||
+ (breakdown.cached_tokens * tier_rate(tier, "cache_read_input_token_cost", "input_cost_per_token"))
|
||||
+ (
|
||||
breakdown.cache_creation_tokens
|
||||
* tier_rate(tier, "cache_creation_input_token_cost", "input_cost_per_token")
|
||||
)
|
||||
)
|
||||
|
||||
input_cost: Final = float(model_info.get("input_cost_per_token") or 0.0)
|
||||
cache_read_cost: Final = _flat_rate(model_info, "cache_read_input_token_cost", "input_cost_per_token")
|
||||
cache_creation_cost: Final = _flat_rate(model_info, "cache_creation_input_token_cost", "input_cost_per_token")
|
||||
|
||||
return (
|
||||
(breakdown.text_tokens * input_cost)
|
||||
+ (breakdown.cached_tokens * cache_read_cost)
|
||||
+ (breakdown.cache_creation_tokens * cache_creation_cost)
|
||||
def _flat_rates(model_info: ModelInfo) -> TokenRates:
|
||||
reasoning_rate: Final = model_info.get("output_cost_per_reasoning_token")
|
||||
return TokenRates(
|
||||
input_rate=float(model_info.get("input_cost_per_token") or 0.0),
|
||||
cache_read_rate=_flat_rate(model_info, "cache_read_input_token_cost", "input_cost_per_token"),
|
||||
cache_creation_rate=_flat_rate(model_info, "cache_creation_input_token_cost", "input_cost_per_token"),
|
||||
output_rate=float(model_info.get("output_cost_per_token") or 0.0),
|
||||
reasoning_rate=None if reasoning_rate is None else float(reasoning_rate),
|
||||
)
|
||||
|
||||
|
||||
def _calculate_completion_cost(
|
||||
breakdown: TokenBreakdown,
|
||||
model_info: ModelInfo,
|
||||
tier: dict | None,
|
||||
) -> float:
|
||||
def _tier_rates(model_info: ModelInfo, tier: dict) -> TokenRates:
|
||||
# A tier that declares output rates keeps the request on them, all-or-nothing. A tier table
|
||||
# spelling out only input rates would serve every completion for free, so there the model's
|
||||
# own output rates stand in
|
||||
tier_declares_output: Final = tier is not None and "output_cost_per_token" in tier
|
||||
output_cost: Final = (
|
||||
tier_rate(tier, "output_cost_per_token")
|
||||
if tier_declares_output
|
||||
else float(model_info.get("output_cost_per_token") or 0.0)
|
||||
)
|
||||
tier_declares_reasoning: Final = tier is not None and "output_cost_per_reasoning_token" in tier
|
||||
model_reasoning_rate: Final = None if tier_declares_output else model_info.get("output_cost_per_reasoning_token")
|
||||
reasoning_cost: Final = (
|
||||
tier_rate(tier, "output_cost_per_reasoning_token", "output_cost_per_token")
|
||||
if tier_declares_reasoning
|
||||
else float(model_reasoning_rate)
|
||||
if model_reasoning_rate is not None
|
||||
else output_cost
|
||||
flat_rates: Final = _flat_rates(model_info)
|
||||
tier_declares_output: Final = "output_cost_per_token" in tier
|
||||
tier_declares_reasoning: Final = "output_cost_per_reasoning_token" in tier
|
||||
return TokenRates(
|
||||
input_rate=tier_rate(tier, "input_cost_per_token"),
|
||||
cache_read_rate=tier_rate(tier, "cache_read_input_token_cost", "input_cost_per_token"),
|
||||
cache_creation_rate=tier_rate(tier, "cache_creation_input_token_cost", "input_cost_per_token"),
|
||||
output_rate=tier_rate(tier, "output_cost_per_token") if tier_declares_output else flat_rates.output_rate,
|
||||
reasoning_rate=(
|
||||
tier_rate(tier, "output_cost_per_reasoning_token")
|
||||
if tier_declares_reasoning
|
||||
else None
|
||||
if tier_declares_output
|
||||
else flat_rates.reasoning_rate
|
||||
),
|
||||
)
|
||||
|
||||
return (breakdown.completion_tokens * output_cost) + (breakdown.reasoning_tokens * reasoning_cost)
|
||||
|
||||
def _off_peak_rates(model_info: ModelInfo, current_time: datetime | None, rates: TokenRates) -> TokenRates:
|
||||
input_rate, output_rate, cache_read_rate = apply_off_peak_pricing(
|
||||
model_info, current_time, rates.input_rate, rates.output_rate, rates.cache_read_rate
|
||||
)
|
||||
return replace(rates, input_rate=input_rate, output_rate=output_rate, cache_read_rate=cache_read_rate)
|
||||
|
||||
|
||||
def cost_per_token(model: str, usage: Usage, custom_llm_provider: str = "dashscope") -> tuple[float, float]:
|
||||
def _bill(breakdown: TokenBreakdown, rates: TokenRates) -> tuple[float, float]:
|
||||
prompt_cost: Final = (
|
||||
(breakdown.text_tokens * rates.input_rate)
|
||||
+ (breakdown.cached_tokens * rates.cache_read_rate)
|
||||
+ (breakdown.cache_creation_tokens * rates.cache_creation_rate)
|
||||
)
|
||||
completion_cost: Final = (breakdown.completion_tokens * rates.output_rate) + (
|
||||
breakdown.reasoning_tokens * rates.billed_reasoning_rate
|
||||
)
|
||||
return prompt_cost, completion_cost
|
||||
|
||||
|
||||
def cost_per_token(
|
||||
model: str,
|
||||
usage: Usage,
|
||||
custom_llm_provider: str = "dashscope",
|
||||
current_time: datetime | None = None,
|
||||
) -> tuple[float, float]:
|
||||
"""
|
||||
Calculate cost per token for Dashscope models.
|
||||
|
||||
Supports both tiered and flat pricing with cached and reasoning tokens.
|
||||
Supports both tiered and flat pricing with cached and reasoning tokens, and swaps in the
|
||||
model's off_peak_pricing rates while one of its windows is open.
|
||||
|
||||
Args:
|
||||
model: Model name without provider prefix
|
||||
usage: LiteLLM Usage block
|
||||
custom_llm_provider: The provider id the request resolved to; dashscope or one of its brand aliases
|
||||
current_time: The moment the request is billed at; defaults to now, UTC
|
||||
|
||||
Returns:
|
||||
Tuple[float, float] - (prompt_cost_in_usd, completion_cost_in_usd)
|
||||
|
|
@ -133,8 +154,7 @@ def cost_per_token(model: str, usage: Usage, custom_llm_provider: str = "dashsco
|
|||
if tiered_pricing
|
||||
else None
|
||||
)
|
||||
standard_rates: Final = _flat_rates(model_info) if tier is None else _tier_rates(model_info, tier)
|
||||
rates: Final = _off_peak_rates(model_info, current_time, standard_rates)
|
||||
|
||||
prompt_cost: Final = _calculate_prompt_cost(breakdown=breakdown, model_info=model_info, tier=tier)
|
||||
completion_cost: Final = _calculate_completion_cost(breakdown=breakdown, model_info=model_info, tier=tier)
|
||||
|
||||
return prompt_cost, completion_cost
|
||||
return _bill(breakdown, rates)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to Databricks' `/chat/completion
|
|||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import AsyncIterator, Coroutine, Iterator
|
||||
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload
|
||||
|
||||
import httpx
|
||||
|
|
@ -15,6 +15,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
_should_convert_tool_call_to_json_mode,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
strip_litellm_internal_message_fields,
|
||||
strip_name_from_message,
|
||||
)
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
|
|
@ -55,6 +56,14 @@ from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
|||
from ..common_utils import DatabricksBase, DatabricksException
|
||||
|
||||
|
||||
def _is_bare_assistant_message(message_dict: Mapping[str, object]) -> bool:
|
||||
"""Databricks rejects assistant messages with neither content nor tool calls, e.g. a replayed
|
||||
thinking-only turn once its `thinking_blocks` are stripped."""
|
||||
return message_dict.get("role") == "assistant" and not any(
|
||||
message_dict.get(key) for key in ("content", "tool_calls", "function_call")
|
||||
)
|
||||
|
||||
|
||||
def _sanitize_empty_content(message_dict: dict[str, Any]) -> None:
|
||||
"""
|
||||
Remove or filter content so empty text blocks are not sent.
|
||||
|
|
@ -423,6 +432,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
|
|||
"""
|
||||
Databricks does not support:
|
||||
- 'name' in user message.
|
||||
- litellm's internal `thinking_blocks` / `reasoning_content` on assistant messages.
|
||||
"""
|
||||
new_messages = []
|
||||
for idx, message in enumerate(messages):
|
||||
|
|
@ -431,10 +441,13 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
|
|||
else:
|
||||
_message = message
|
||||
_message = strip_name_from_message(_message, allowed_name_roles=["user"])
|
||||
_message = strip_litellm_internal_message_fields(_message)
|
||||
# Move message-level cache_control into a content block when content is a string.
|
||||
if "cache_control" in _message and isinstance(_message.get("content"), str):
|
||||
_message = self._move_cache_control_into_string_content_block(_message)
|
||||
_sanitize_empty_content(cast(dict[str, Any], _message))
|
||||
if _is_bare_assistant_message(_message):
|
||||
continue
|
||||
new_messages.append(_message)
|
||||
|
||||
if "claude" not in model:
|
||||
|
|
|
|||
|
|
@ -61,6 +61,14 @@ def _get_effort_level(value: str | dict | None) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
GPT_REASONING_SERIES_MARKERS: Final = ("gpt-5", "gpt-6")
|
||||
|
||||
|
||||
def is_gpt_reasoning_series_name(model: str) -> bool:
|
||||
normalized: Final = model.split("/")[-1]
|
||||
return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and not normalized.startswith("gpt-5-chat")
|
||||
|
||||
|
||||
class OpenAIGPT5Config(OpenAIGPTConfig):
|
||||
"""Configuration for gpt-5 models including GPT-5-Codex variants.
|
||||
|
||||
|
|
@ -73,21 +81,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
|
|||
|
||||
@classmethod
|
||||
def is_model_gpt_5_model(cls, model: str) -> bool:
|
||||
# The gpt-5-chat* family (gpt-5-chat, gpt-5-chat-latest, gpt-5-chat-2025-08-07,
|
||||
# …) are regular chat models: they support temperature and tool_choice but NOT
|
||||
# reasoning_effort. They must NOT be routed through the GPT-5 reasoning path.
|
||||
#
|
||||
# Versioned chat models such as gpt-5.3-chat and gpt-5.1-chat ARE reasoning
|
||||
# models and must stay on the GPT-5 path. The distinguishing feature is that
|
||||
# the gpt-5-chat family has a literal "-chat" immediately after "gpt-5"
|
||||
# (i.e. "gpt-5-chat…"), while versioned chat models interpose a minor version
|
||||
# number (i.e. "gpt-5.<digit>-chat").
|
||||
#
|
||||
# Using a startswith("gpt-5-chat") prefix check on the normalized name (rather
|
||||
# than a substring check) makes this boundary explicit and avoids any ambiguity
|
||||
# if future model names coincidentally contain "gpt-5-chat" as an interior run.
|
||||
_normalized: Final = model.split("/")[-1] # strip provider prefix, e.g. "openai/"
|
||||
return "gpt-5" in model and not _normalized.startswith("gpt-5-chat")
|
||||
return is_gpt_reasoning_series_name(model)
|
||||
|
||||
@classmethod
|
||||
def is_model_gpt_5_search_model(cls, model: str) -> bool:
|
||||
|
|
@ -122,6 +116,8 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
|
|||
def is_model_gpt_5_4_plus_model(cls, model: str) -> bool:
|
||||
"""Check if the model is gpt-5.4 or newer (5.4, 5.5, 5.6, etc., including pro)."""
|
||||
model_name: Final = model.split("/")[-1]
|
||||
if model_name.startswith("gpt-6"):
|
||||
return True
|
||||
if not model_name.startswith("gpt-5."):
|
||||
return False
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import time
|
|||
import uuid
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Optional
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
|
|
@ -43,6 +44,14 @@ _OPENAI_INIT_PARAMS: Final[tuple[str, ...]] = _get_client_init_params(OpenAI)
|
|||
_AZURE_OPENAI_INIT_PARAMS: Final[tuple[str, ...]] = _get_client_init_params(AzureOpenAI)
|
||||
|
||||
|
||||
_OPENAI_API_HOST: Final[str] = "api.openai.com"
|
||||
|
||||
|
||||
def is_openai_backed_api_base(api_base: str) -> bool:
|
||||
hostname: Final = urlsplit(api_base).hostname
|
||||
return hostname is not None and (hostname == _OPENAI_API_HOST or hostname.endswith(f".{_OPENAI_API_HOST}"))
|
||||
|
||||
|
||||
class OpenAIError(BaseLLMException):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -82,8 +82,8 @@ class GPTImageGenerationConfig(BaseImageGenerationConfig):
|
|||
)
|
||||
|
||||
# set optional params
|
||||
image_response.size = optional_params.get("size", "1024x1024") # default is always 1024x1024
|
||||
image_response.quality = optional_params.get("quality", "high") # always hd for dall-e-3
|
||||
image_response.output_format = optional_params.get("response_format", "png") # always png for dall-e-3
|
||||
image_response.size = image_response.size or optional_params.get("size", "1024x1024")
|
||||
image_response.quality = image_response.quality or optional_params.get("quality", "high")
|
||||
image_response.output_format = image_response.output_format or optional_params.get("output_format", "png")
|
||||
|
||||
return image_response
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import time
|
|||
import types
|
||||
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -55,6 +54,7 @@ from .common_utils import (
|
|||
OpenAIError,
|
||||
build_output_token_limit_response,
|
||||
drop_params_from_unprocessable_entity_error,
|
||||
is_openai_backed_api_base,
|
||||
is_output_token_limit_error,
|
||||
)
|
||||
from .workload_identity import resolve_openai_workload_identity_config
|
||||
|
|
@ -1190,10 +1190,8 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
"""
|
||||
if stream_options is not None:
|
||||
return {"stream_options": stream_options}
|
||||
else:
|
||||
# by default litellm will include usage for openai endpoints
|
||||
if api_base is None or urlparse(api_base).hostname == "api.openai.com":
|
||||
return {"stream_options": {"include_usage": True}}
|
||||
if api_base is None or is_openai_backed_api_base(api_base):
|
||||
return {"stream_options": {"include_usage": True}}
|
||||
return {}
|
||||
|
||||
# Embedding
|
||||
|
|
|
|||
|
|
@ -33,8 +33,9 @@ import time
|
|||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import accumulate
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Union, cast
|
||||
|
||||
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
|
@ -42,6 +43,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
OpenAiResponsesToChatCompletionStreamIterator,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
||||
|
|
@ -74,6 +76,7 @@ from litellm.types.llms.openai import (
|
|||
OutputTextDoneEvent,
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
|
|
@ -115,6 +118,199 @@ class ResponsesStreamChunk(TypedDict, total=False):
|
|||
content_index: ReadOnly[int]
|
||||
|
||||
|
||||
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{"function_call_output": "output", "message": "content"}
|
||||
)
|
||||
|
||||
_EMPTY_RESPONSES_REQUEST: Final[ResponsesAPIOptionalRequestParams] = {}
|
||||
|
||||
|
||||
def _item_rewrite_field(item: Mapping[str, object]) -> str | None:
|
||||
item_type: Final = item.get("type")
|
||||
if item_type is None:
|
||||
return "content" if "content" in item else None
|
||||
if not isinstance(item_type, str):
|
||||
return None
|
||||
return _PATCHABLE_ITEM_FIELDS.get(item_type)
|
||||
|
||||
|
||||
def _rewritten_input_item(item: Mapping[str, object], rewritten: object) -> Mapping[str, object] | None:
|
||||
field: Final = _item_rewrite_field(item)
|
||||
if field is None or not isinstance(rewritten, Mapping):
|
||||
return None
|
||||
rewritten_content: Final = rewritten.get("content")
|
||||
if isinstance(item.get(field), str) and isinstance(rewritten_content, str):
|
||||
return {**item, field: rewritten_content} # mutable-ok: request input items must stay JSON-plain dicts
|
||||
rewritten_row: Final = cast("AllMessageValues", rewritten) # cast-ok: guardrails hand back chat-shaped rows
|
||||
converted_items, _ = LiteLLMResponsesTransformationHandler().convert_chat_completion_messages_to_responses_api(
|
||||
[rewritten_row] # mutable-ok: converter signature takes a list
|
||||
)
|
||||
if len(converted_items) != 1 or not isinstance(converted_items[0], Mapping):
|
||||
return None
|
||||
first_converted: Final = cast("Mapping[str, object]", converted_items[0]) # cast-ok: isinstance-checked above
|
||||
converted_value: Final = first_converted.get(field)
|
||||
if converted_value is None:
|
||||
return None
|
||||
return {**item, field: converted_value} # mutable-ok: request input items must stay JSON-plain dicts
|
||||
|
||||
|
||||
def _is_function_call_item(item: object) -> bool:
|
||||
return isinstance(item, Mapping) and item.get("type") in ("function_call", "custom_tool_call")
|
||||
|
||||
|
||||
def _last_message_role(messages: Sequence[object]) -> str | None:
|
||||
if not messages:
|
||||
return None
|
||||
last: Final = messages[-1]
|
||||
role: Final = last.get("role") if isinstance(last, Mapping) else getattr(last, "role", None)
|
||||
return role if isinstance(role, str) else None
|
||||
|
||||
|
||||
def _provenance_unit_bounds(
|
||||
raw_input: Sequence[object],
|
||||
solo_conversions: Sequence[Sequence[object]],
|
||||
) -> tuple[tuple[int, int], ...]:
|
||||
trailing_roles: Final = tuple(
|
||||
accumulate(
|
||||
(_last_message_role(messages) for messages in solo_conversions),
|
||||
lambda previous, current: current if current is not None else previous,
|
||||
)
|
||||
)
|
||||
start_indexes: Final = tuple(
|
||||
index
|
||||
for index in range(len(raw_input))
|
||||
if index == 0 or not (_is_function_call_item(raw_input[index]) and trailing_roles[index - 1] == "assistant")
|
||||
)
|
||||
return tuple(zip(start_indexes, (*start_indexes[1:], len(raw_input))))
|
||||
|
||||
|
||||
def _input_item_provenance(
|
||||
raw_input: Sequence[object],
|
||||
expected_messages: Sequence[object],
|
||||
) -> tuple[Mapping[int, int], frozenset[int]] | None:
|
||||
if not all(isinstance(item, Mapping) for item in raw_input):
|
||||
return None
|
||||
solo_conversions: Final = tuple(
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=cast("ResponseInputParam", [item]), # cast-ok: items checked as Mappings above
|
||||
responses_api_request=_EMPTY_RESPONSES_REQUEST,
|
||||
)
|
||||
for item in raw_input
|
||||
)
|
||||
full_conversion: Final = tuple(
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=cast("ResponseInputParam", list(raw_input)), # cast-ok: items checked as Mappings above
|
||||
responses_api_request=_EMPTY_RESPONSES_REQUEST,
|
||||
)
|
||||
)
|
||||
if full_conversion != tuple(expected_messages):
|
||||
return None
|
||||
units: Final = _provenance_unit_bounds(raw_input, solo_conversions)
|
||||
unit_messages: Final = tuple(
|
||||
tuple(solo_conversions[start])
|
||||
if end - start == 1
|
||||
else tuple(
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=cast("ResponseInputParam", list(raw_input[start:end])), # cast-ok: checked as Mappings above
|
||||
responses_api_request=_EMPTY_RESPONSES_REQUEST,
|
||||
)
|
||||
)
|
||||
for start, end in units
|
||||
)
|
||||
if tuple(message for messages in unit_messages for message in messages) != full_conversion:
|
||||
return None
|
||||
boundaries: Final = tuple(accumulate((len(messages) for messages in unit_messages), initial=0))
|
||||
item_for_message: Final = MappingProxyType(
|
||||
{
|
||||
message_index: start
|
||||
for unit_index, (start, end) in enumerate(units)
|
||||
if end - start == 1
|
||||
for message_index in range(boundaries[unit_index], boundaries[unit_index + 1])
|
||||
}
|
||||
)
|
||||
tainted: Final = frozenset(
|
||||
message_index
|
||||
for unit_index, (start, end) in enumerate(units)
|
||||
if end - start > 1
|
||||
for message_index in range(boundaries[unit_index], boundaries[unit_index + 1])
|
||||
)
|
||||
return item_for_message, tainted
|
||||
|
||||
|
||||
class _RequestFields(NamedTuple):
|
||||
input: tuple[object, ...]
|
||||
instructions: str | None
|
||||
|
||||
|
||||
class _ExtractedInputs(NamedTuple):
|
||||
inputs: GenericGuardrailAPIInputs
|
||||
task_mappings: tuple[tuple[int, int | None], ...]
|
||||
|
||||
|
||||
def _patched_request_fields(
|
||||
raw_input: object,
|
||||
instructions: object,
|
||||
original_messages: Sequence[object],
|
||||
structured_messages: Sequence[object],
|
||||
) -> _RequestFields | None:
|
||||
if not isinstance(raw_input, list) or len(original_messages) != len(structured_messages):
|
||||
return None
|
||||
offset: Final = 1 if instructions else 0
|
||||
provenance: Final = _input_item_provenance(raw_input, tuple(original_messages)[offset:])
|
||||
if provenance is None:
|
||||
return None
|
||||
item_for_message, tainted = provenance
|
||||
changed: Final = tuple(
|
||||
(index, rewritten)
|
||||
for index, (original, rewritten) in enumerate(zip(original_messages, structured_messages))
|
||||
if original != rewritten
|
||||
)
|
||||
instruction_rewrites: Final = tuple(rewritten for index, rewritten in changed if index < offset)
|
||||
rewritten_instructions: Final = (
|
||||
instruction_rewrites[0].get("content")
|
||||
if instruction_rewrites and isinstance(instruction_rewrites[0], Mapping)
|
||||
else instructions
|
||||
)
|
||||
instructions_value: Final = rewritten_instructions if isinstance(rewritten_instructions, str) else None
|
||||
if rewritten_instructions is not None and instructions_value is None:
|
||||
return None
|
||||
body_changes: Final = tuple((index - offset, rewritten) for index, rewritten in changed if index >= offset)
|
||||
if any(message_index in tainted or message_index not in item_for_message for message_index, _ in body_changes):
|
||||
return None
|
||||
replacements: Final = MappingProxyType(
|
||||
{
|
||||
item_for_message[message_index]: _rewritten_input_item(
|
||||
cast("Mapping[str, object]", raw_input[item_for_message[message_index]]), # cast-ok: checked Mappings
|
||||
rewritten,
|
||||
)
|
||||
for message_index, rewritten in body_changes
|
||||
}
|
||||
)
|
||||
if len(replacements) != len(body_changes) or any(item is None for item in replacements.values()):
|
||||
return None
|
||||
return _RequestFields(
|
||||
input=tuple(replacements.get(index, item) for index, item in enumerate(raw_input)),
|
||||
instructions=instructions_value,
|
||||
)
|
||||
|
||||
|
||||
def _patch_or_convert_request_fields(
|
||||
raw_input: object,
|
||||
instructions: object,
|
||||
original_messages: Sequence[object],
|
||||
structured_messages: Sequence[AllMessageValues],
|
||||
) -> _RequestFields | None:
|
||||
if not isinstance(structured_messages, list):
|
||||
return None
|
||||
patched: Final = _patched_request_fields(raw_input, instructions, original_messages, structured_messages)
|
||||
if patched is not None:
|
||||
return patched
|
||||
input_items, converted_instructions = (
|
||||
LiteLLMResponsesTransformationHandler().convert_chat_completion_messages_to_responses_api(structured_messages)
|
||||
)
|
||||
return _RequestFields(input=tuple(input_items), instructions=converted_instructions)
|
||||
|
||||
|
||||
def _next_stream_sequence_number(responses_so_far: Sequence[Any] | None) -> int:
|
||||
sequence_numbers: Final = (
|
||||
item.get("sequence_number") if isinstance(item, dict) else getattr(item, "sequence_number", None)
|
||||
|
|
@ -162,9 +358,8 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
Handles both string input and list of message objects.
|
||||
"""
|
||||
input_data: Final[str | ResponseInputParam | None] = data.get("input")
|
||||
if input_data is None:
|
||||
if not isinstance(input_data, (str, list)):
|
||||
return data
|
||||
|
||||
structured_messages: Final = self.get_structured_messages(data)
|
||||
raw_tools: Final = data.get("tools")
|
||||
original_tools: Final[tuple[Mapping[str, object], ...]] = (
|
||||
|
|
@ -173,94 +368,93 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
flattened_tool_groups: Final = tuple(
|
||||
form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools)
|
||||
)
|
||||
flattened_tools: Final = tuple(
|
||||
cast(ChatCompletionToolParam, tool) # cast-ok: mcp tools ride along in the guardrail's tool list
|
||||
for group in flattened_tool_groups
|
||||
for tool in group
|
||||
)
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
|
||||
copy.deepcopy(flattened_tools)
|
||||
)
|
||||
|
||||
# Handle simple string input
|
||||
if isinstance(input_data, str):
|
||||
inputs = GenericGuardrailAPIInputs(texts=[input_data])
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = structured_messages
|
||||
# Include model information if available
|
||||
model = data.get("model")
|
||||
if model:
|
||||
inputs["model"] = model
|
||||
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data
|
||||
self._apply_guardrailed_tools_to_data(
|
||||
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
|
||||
)
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed string input")
|
||||
return data
|
||||
|
||||
# Handle list input (ResponseInputParam)
|
||||
if not isinstance(input_data, list):
|
||||
extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups)
|
||||
if not extracted.inputs.get("texts"):
|
||||
return data
|
||||
if structured_messages:
|
||||
extracted.inputs["structured_messages"] = structured_messages
|
||||
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=extracted.inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
self._apply_guardrailed_tools_to_data(
|
||||
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
|
||||
)
|
||||
written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs)
|
||||
if written_back is not None:
|
||||
data["input"] = list(written_back.input) # mutable-ok: JSON body
|
||||
if written_back.instructions is None:
|
||||
data.pop("instructions", None)
|
||||
else:
|
||||
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
|
||||
elif isinstance(input_data, str):
|
||||
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
|
||||
else:
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=guardrailed_inputs.get("texts") or (),
|
||||
task_mappings=extracted.task_mappings,
|
||||
)
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
|
||||
return data
|
||||
|
||||
def _extract_guardrail_inputs(
|
||||
self,
|
||||
data: Mapping[str, object],
|
||||
input_data: "str | ResponseInputParam",
|
||||
flattened_tool_groups: Sequence[Sequence[Mapping[str, object]]],
|
||||
) -> _ExtractedInputs:
|
||||
texts_to_check: Final[list[str]] = []
|
||||
images_to_check: Final[list[str]] = []
|
||||
task_mappings: Final[list[tuple[int, int | None]]] = []
|
||||
|
||||
# Step 1: Extract all text content, images, and tools
|
||||
for msg_idx, message in enumerate(input_data):
|
||||
self._extract_input_text_and_images(
|
||||
message=message,
|
||||
msg_idx=msg_idx,
|
||||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
task_mappings=task_mappings,
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
|
||||
copy.deepcopy(
|
||||
tuple(
|
||||
cast(ChatCompletionToolParam, tool) # cast-ok: mcp tools ride along in the guardrail's tool list
|
||||
for group in flattened_tool_groups
|
||||
for tool in group
|
||||
)
|
||||
)
|
||||
)
|
||||
if isinstance(input_data, str):
|
||||
texts_to_check.append(input_data)
|
||||
else:
|
||||
for msg_idx, message in enumerate(input_data):
|
||||
self._extract_input_text_and_images(
|
||||
message=message,
|
||||
msg_idx=msg_idx,
|
||||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
task_mappings=task_mappings,
|
||||
)
|
||||
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
model: Final = data.get("model")
|
||||
if isinstance(model, str):
|
||||
inputs["model"] = model
|
||||
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings))
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = structured_messages
|
||||
# Include model information if available
|
||||
model = data.get("model")
|
||||
if model:
|
||||
inputs["model"] = model
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
self._apply_guardrailed_tools_to_data(
|
||||
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
|
||||
)
|
||||
|
||||
# Step 3: Map guardrail responses back to original input structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=guardrailed_texts,
|
||||
task_mappings=task_mappings,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", input_data)
|
||||
|
||||
return data
|
||||
@staticmethod
|
||||
def _written_back_request_fields(
|
||||
data: Mapping[str, object],
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> _RequestFields | None:
|
||||
guardrailed: Final = guardrailed_inputs.get("structured_messages")
|
||||
if guardrailed is None or guardrailed is structured_messages:
|
||||
return None
|
||||
return _patch_or_convert_request_fields(
|
||||
data.get("input"),
|
||||
data.get("instructions"),
|
||||
structured_messages or (),
|
||||
guardrailed,
|
||||
)
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> list[str]:
|
||||
"""Extract tool names from Responses API request (tools[].name for function
|
||||
|
|
@ -331,8 +525,8 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
async def _apply_guardrail_responses_to_input(
|
||||
self,
|
||||
messages: Any, # Can be List[Dict[str, Any]] or ResponseInputParam
|
||||
responses: list[str],
|
||||
task_mappings: list[tuple[int, int | None]],
|
||||
responses: Sequence[str],
|
||||
task_mappings: Sequence[tuple[int, int | None]],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrail responses back to input messages.
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
)
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.llms.openai.chat.gpt_5_transformation import is_gpt_reasoning_series_name
|
||||
from litellm.responses.litellm_completion_transformation.custom_tools import TOOL_CALL_ITEM_ID_PREFIX_BY_TYPE
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import *
|
||||
|
|
@ -88,7 +89,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
parts: Final = model.split("/")
|
||||
if len(parts) > 1 and parts[0] not in ("openai",):
|
||||
return False
|
||||
return "gpt-5" in model and "gpt-5-chat" not in model
|
||||
return is_gpt_reasoning_series_name(model)
|
||||
|
||||
@staticmethod
|
||||
def _supports_reasoning_effort_none(model: str) -> bool:
|
||||
|
|
|
|||
|
|
@ -998,6 +998,16 @@ def replace_project_and_location_in_route(requested_route: str, vertex_project:
|
|||
return modified_route
|
||||
|
||||
|
||||
def _api_version_for_route(requested_route: str) -> Literal["v1", "v1beta1"]:
|
||||
return "v1beta1" if "cachedContent" in requested_route else "v1"
|
||||
|
||||
|
||||
def _with_api_version(requested_route: str) -> str:
|
||||
if not requested_route.startswith("/projects/"):
|
||||
return requested_route
|
||||
return f"/{_api_version_for_route(requested_route)}{requested_route}"
|
||||
|
||||
|
||||
def construct_target_url(
|
||||
base_url: str,
|
||||
requested_route: str,
|
||||
|
|
@ -1017,18 +1027,19 @@ def construct_target_url(
|
|||
|
||||
new_base_url: Final = httpx.URL(base_url)
|
||||
if "locations" in requested_route: # contains the target project id + location
|
||||
if vertex_project and vertex_location:
|
||||
requested_route = replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
|
||||
return new_base_url.copy_with(path=requested_route)
|
||||
targeted_route: Final = (
|
||||
replace_project_and_location_in_route(requested_route, vertex_project, vertex_location)
|
||||
if vertex_project and vertex_location
|
||||
else requested_route
|
||||
)
|
||||
return new_base_url.copy_with(path=_with_api_version(targeted_route))
|
||||
|
||||
"""
|
||||
- Add endpoint version (e.g. v1beta for cachedContent, v1 for rest)
|
||||
- Add default project id
|
||||
- Add default location
|
||||
"""
|
||||
vertex_version: Literal["v1", "v1beta1"] = "v1"
|
||||
if "cachedContent" in requested_route:
|
||||
vertex_version = "v1beta1"
|
||||
vertex_version: Literal["v1", "v1beta1"] = _api_version_for_route(requested_route)
|
||||
|
||||
# Check if the requested route starts with a version
|
||||
# e.g. /v1beta1/publishers/google/models/gemini-3-pro-preview:streamGenerateContent
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from copy import deepcopy
|
|||
from functools import partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, Union, cast, get_args
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from litellm._logging import _redact_string
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -60,6 +61,7 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
from litellm.constants import (
|
||||
AZURE_OPENAI_AUDIO_PROVIDERS,
|
||||
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
|
||||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
|
||||
)
|
||||
|
|
@ -984,6 +986,12 @@ def mock_completion(
|
|||
|
||||
|
||||
_OPENAI_DEFAULT_API_BASE: Final = "https://api.openai.com/v1"
|
||||
_OPENAI_API_HOST: Final = "api.openai.com"
|
||||
|
||||
|
||||
def _is_openai_backed_api_base(api_base: str) -> bool:
|
||||
hostname: Final = urlsplit(api_base).hostname
|
||||
return hostname is not None and (hostname == _OPENAI_API_HOST or hostname.endswith(f".{_OPENAI_API_HOST}"))
|
||||
|
||||
|
||||
def _resolve_openai_api_base(api_base: str | None) -> str:
|
||||
|
|
@ -1053,7 +1061,7 @@ def responses_api_bridge_check(
|
|||
# natively by Chat Completions with reasoning on, so custom-only requests stay on
|
||||
# chat and keep their native custom tool_call response shape.
|
||||
# - The UNSET-effort arm only fires against endpoints known to enforce that
|
||||
# constraint (the default OpenAI endpoint, or Azure OpenAI where api_base is
|
||||
# constraint (any api.openai.com host, or Azure OpenAI where api_base is
|
||||
# always set): chat-only OpenAI-compatible backends registered under the openai
|
||||
# provider with a custom api_base and gpt-5.4+ model names serve tools without
|
||||
# reasoning fine and have no /responses route, so they keep pre-existing
|
||||
|
|
@ -1068,14 +1076,15 @@ def responses_api_bridge_check(
|
|||
reasoning_active = reasoning_effort.get("effort") != "none" or reasoning_effort.get("summary") is not None
|
||||
else:
|
||||
reasoning_active = reasoning_effort != "none"
|
||||
# The reasoning+tools constraint is enforced only by the real OpenAI endpoint (and Azure OpenAI).
|
||||
# Resolve the effective base arg>global>env>default exactly as the chat handler does, so a custom
|
||||
# base set via litellm.api_base or OPENAI_BASE_URL/OPENAI_API_BASE isn't misread as the default and
|
||||
# bridged to a /responses route it lacks. A whitespace-only base collapses to the default too.
|
||||
resolved_api_base: Final = _resolve_openai_api_base(api_base)
|
||||
on_constraint_enforcing_endpoint: Final = custom_llm_provider == "azure" or resolved_api_base.strip() in (
|
||||
"",
|
||||
_OPENAI_DEFAULT_API_BASE,
|
||||
# The reasoning+tools constraint is enforced by the real OpenAI backend behind any api.openai.com
|
||||
# host (the default URL or a PrivateLink hostname such as <region>.privatelink.api.openai.com) and
|
||||
# by Azure OpenAI. Resolve the effective base arg>global>env>default exactly as the chat handler
|
||||
# does, so a custom base set via litellm.api_base or OPENAI_BASE_URL/OPENAI_API_BASE isn't misread
|
||||
# as the default and bridged to a /responses route it lacks. A whitespace-only base collapses to
|
||||
# the default too.
|
||||
resolved_api_base: Final = _resolve_openai_api_base(api_base).strip()
|
||||
on_constraint_enforcing_endpoint: Final = (
|
||||
custom_llm_provider == "azure" or resolved_api_base == "" or _is_openai_backed_api_base(resolved_api_base)
|
||||
)
|
||||
if (
|
||||
custom_llm_provider in ("openai", "azure")
|
||||
|
|
@ -7769,7 +7778,7 @@ def transcription(
|
|||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if custom_llm_provider == "azure" and provider_config is None:
|
||||
if custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS and provider_config is None:
|
||||
# azure configs
|
||||
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
|
|
@ -8056,7 +8065,10 @@ def speech(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
response: HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent] | None = None
|
||||
if custom_llm_provider == "openai" or custom_llm_provider in litellm.openai_compatible_providers:
|
||||
if custom_llm_provider == "openai" or (
|
||||
custom_llm_provider in litellm.openai_compatible_providers
|
||||
and custom_llm_provider not in AZURE_OPENAI_AUDIO_PROVIDERS
|
||||
):
|
||||
if voice is None or not (isinstance(voice, str)):
|
||||
raise litellm.BadRequestError(
|
||||
message="'voice' is required to be passed as a string for OpenAI TTS",
|
||||
|
|
@ -8110,7 +8122,7 @@ def speech(
|
|||
aspeech=aspeech,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
elif custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS:
|
||||
# Check if this is Azure Speech Service (Cognitive Services TTS)
|
||||
if model.startswith("speech/"):
|
||||
from litellm.llms.azure.text_to_speech.transformation import (
|
||||
|
|
|
|||
|
|
@ -29277,6 +29277,75 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": true
|
||||
},
|
||||
"gpt-6-astra": {
|
||||
"cache_creation_input_token_cost": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_priority": 5e-05,
|
||||
"cache_creation_input_token_cost_flex": 6.25e-06,
|
||||
"cache_creation_input_token_cost_priority": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": 4e-06,
|
||||
"cache_read_input_token_cost_flex": 5e-07,
|
||||
"cache_read_input_token_cost_priority": 2e-06,
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens_priority": 4e-05,
|
||||
"input_cost_per_token_batches": 5e-06,
|
||||
"input_cost_per_token_flex": 5e-06,
|
||||
"input_cost_per_token_priority": 2e-05,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 7.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 3.75e-05,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 0.00015,
|
||||
"output_cost_per_token_batches": 2.5e-05,
|
||||
"output_cost_per_token_flex": 2.5e-05,
|
||||
"output_cost_per_token_priority": 0.0001,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
"regional_processing_uplift_multiplier_us": 1.1,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_native_streaming": true,
|
||||
"supports_none_reasoning_effort": false,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_cache_breakpoint": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
|
|
@ -52911,6 +52980,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -52933,7 +53007,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-terra": {
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
|
|
@ -52944,6 +53019,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 4.4e-07,
|
||||
"output_cost_per_token": 1.32e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 1.98e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -52966,7 +53046,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-cyber": {
|
||||
"input_cost_per_token": 1.375e-05,
|
||||
|
|
@ -53005,6 +53086,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 4.4e-08,
|
||||
"output_cost_per_token": 1.32e-06,
|
||||
"output_cost_per_token_above_272k_tokens": 1.98e-06,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -53027,7 +53113,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"us.openai.gpt-5.6-sol": {
|
||||
"input_cost_per_token": 4.4e-06,
|
||||
|
|
@ -53192,6 +53279,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -53213,7 +53305,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.4": {
|
||||
"input_cost_per_token": 2.75e-06,
|
||||
|
|
@ -53222,6 +53315,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 5.5e-07,
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 2.475e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -53243,7 +53341,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/google.gemma-4-31b": {
|
||||
"input_cost_per_token": 1.4e-07,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from collections.abc import Mapping, Sequence
|
|||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
from typing import TYPE_CHECKING, Final, Literal, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
from starlette.datastructures import Headers
|
||||
|
|
@ -305,6 +305,12 @@ def _admission_failure_fallback(
|
|||
raise exc
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MCPServerAccess:
|
||||
server_ids: tuple[str, ...]
|
||||
scope: Literal["unscoped", "scoped", "unresolved"] = "unscoped"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DcrBridgeTarget:
|
||||
"""The single DCR-bridge server a request targets, paired with the exact name the caller
|
||||
|
|
@ -1456,6 +1462,18 @@ class MCPRequestHandler:
|
|||
*,
|
||||
keyless_source: bool = False,
|
||||
) -> list[str]:
|
||||
access: Final = await MCPRequestHandler.get_mcp_server_access(
|
||||
user_api_key_auth,
|
||||
keyless_source=keyless_source,
|
||||
)
|
||||
return list(access.server_ids)
|
||||
|
||||
@staticmethod
|
||||
async def get_mcp_server_access(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
*,
|
||||
keyless_source: bool = False,
|
||||
) -> MCPServerAccess:
|
||||
"""
|
||||
Get list of allowed MCP servers for the given user/key based on permissions.
|
||||
|
||||
|
|
@ -1478,13 +1496,17 @@ class MCPRequestHandler:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
|
||||
try:
|
||||
# A keyless admitted subject resolves per source BEFORE any single-source rule here. Ordering
|
||||
# matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above
|
||||
# this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is
|
||||
# independent; an opt-out silences only its own source, inside the recursive call).
|
||||
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
|
||||
return await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)
|
||||
return MCPServerAccess(
|
||||
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
|
||||
)
|
||||
|
||||
# Get allowed servers from key and team
|
||||
allowed_mcp_servers_for_key = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
|
||||
|
|
@ -1492,7 +1514,7 @@ class MCPRequestHandler:
|
|||
# The key explicitly opted out of every MCP server. This overrides
|
||||
# team inheritance and additive grants (mirrors no-default-models).
|
||||
if SpecialMCPServerNames.no_mcp_servers.value in allowed_mcp_servers_for_key:
|
||||
return []
|
||||
return MCPServerAccess(server_ids=(), scope="scoped")
|
||||
|
||||
allowed_mcp_servers_for_team = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
|
||||
|
||||
|
|
@ -1572,7 +1594,7 @@ class MCPRequestHandler:
|
|||
"require_end_user_mcp_access_defined=True and end_user %s has no MCP permissions - blocking MCP access",
|
||||
user_api_key_auth.end_user_id,
|
||||
)
|
||||
return []
|
||||
return MCPServerAccess(server_ids=(), scope="scoped")
|
||||
|
||||
#########################################################
|
||||
# Check agent permissions if agent_id is set on the key
|
||||
|
|
@ -1601,14 +1623,22 @@ class MCPRequestHandler:
|
|||
#########################################################
|
||||
# Apply org-level ceiling if org_id is set
|
||||
#########################################################
|
||||
allowed_mcp_servers = await MCPRequestHandler._apply_primary_org_ceiling(
|
||||
allowed_mcp_servers, org_restricts = await MCPRequestHandler._apply_primary_org_ceiling(
|
||||
allowed_mcp_servers,
|
||||
user_api_key_auth,
|
||||
has_lower_level_mcp_restrictions,
|
||||
keyless_source=keyless_source,
|
||||
)
|
||||
|
||||
return list(set(allowed_mcp_servers))
|
||||
declares_key_mcp_scope: Final = getattr(key_object_permission, "mcp_servers", None) is not None
|
||||
return MCPServerAccess(
|
||||
server_ids=tuple(set(allowed_mcp_servers)),
|
||||
scope=(
|
||||
"scoped"
|
||||
if has_lower_level_mcp_restrictions or org_restricts or declares_key_mcp_scope
|
||||
else "unscoped"
|
||||
),
|
||||
)
|
||||
except Exception as e:
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
# A ceiling we KNOW exists and cannot read. Denying is the only answer that does not
|
||||
|
|
@ -1616,7 +1646,10 @@ class MCPRequestHandler:
|
|||
verbose_logger.warning("Denying MCP access, entitlement unreadable: %s", e)
|
||||
else:
|
||||
verbose_logger.warning("Failed to get allowed MCP servers: %s", e)
|
||||
return []
|
||||
return MCPServerAccess(
|
||||
server_ids=(),
|
||||
scope="scoped" if getattr(key_object_permission, "mcp_servers", None) is not None else "unresolved",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _apply_primary_org_ceiling(
|
||||
|
|
@ -1624,7 +1657,7 @@ class MCPRequestHandler:
|
|||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
has_lower_level_mcp_restrictions: bool,
|
||||
keyless_source: bool = False,
|
||||
) -> list[str]:
|
||||
) -> tuple[list[str], bool]:
|
||||
"""Cap the resolved server list by this caller's org ceiling: an explicit org list intersects
|
||||
lower-level restrictions (else becomes the ceiling); no org or an empty list leaves it unchanged.
|
||||
|
||||
|
|
@ -1638,7 +1671,7 @@ class MCPRequestHandler:
|
|||
cannot be read raises out of ``_get_allowed_mcp_servers_for_org`` and never arrives here as
|
||||
``None``, so key auth cannot silently shed a ceiling an operator did configure."""
|
||||
if not (user_api_key_auth and user_api_key_auth.org_id):
|
||||
return allowed_mcp_servers
|
||||
return allowed_mcp_servers, False
|
||||
allowed_mcp_servers_for_org: Final = await MCPRequestHandler._get_allowed_mcp_servers_for_org(user_api_key_auth)
|
||||
if allowed_mcp_servers_for_org is None:
|
||||
verbose_logger.warning(
|
||||
|
|
@ -1646,9 +1679,9 @@ class MCPRequestHandler:
|
|||
user_api_key_auth.org_id,
|
||||
"denying (keyless admitted subject)" if keyless_source else "leaving uncapped (key auth)",
|
||||
)
|
||||
return [] if keyless_source else allowed_mcp_servers
|
||||
return ([] if keyless_source else allowed_mcp_servers), False
|
||||
if len(allowed_mcp_servers_for_org) == 0:
|
||||
return allowed_mcp_servers
|
||||
return allowed_mcp_servers, False
|
||||
if has_lower_level_mcp_restrictions or keyless_source:
|
||||
# Org can only cap lower-level restrictions. A keyless admitted source ALWAYS takes this
|
||||
# arm: its model unions GRANTS, so an org list may only narrow a source, never become one.
|
||||
|
|
@ -1657,7 +1690,7 @@ class MCPRequestHandler:
|
|||
# No lower-level restrictions → org list becomes the ceiling.
|
||||
capped = allowed_mcp_servers_for_org
|
||||
verbose_logger.debug("Applied org ceiling filter. Final allowed servers: %s", capped)
|
||||
return capped
|
||||
return capped, True
|
||||
|
||||
@staticmethod
|
||||
def _scoped_source_auth(
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
|
|||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
MCPServerAccess,
|
||||
_is_mcp_admitted_user_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
|
|
@ -2958,7 +2959,13 @@ class MCPServerManager:
|
|||
return None
|
||||
return user_api_key_auth.mcp_session_resource_server_id
|
||||
|
||||
async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth | None = None) -> list[str]:
|
||||
async def get_allowed_mcp_servers(
|
||||
self,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
*,
|
||||
access: MCPServerAccess | None = None,
|
||||
general_settings: Mapping[str, object] | None = None,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Get the allowed MCP Servers for the user.
|
||||
|
||||
|
|
@ -2967,6 +2974,9 @@ class MCPServerManager:
|
|||
2. If admin and no object_permission, return all servers
|
||||
3. Otherwise, use standard permission checks
|
||||
"""
|
||||
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
|
||||
|
||||
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
|
||||
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
|
||||
|
||||
# A keyless admitted subject is resolved per grant source, and channel decisions that are
|
||||
|
|
@ -3007,11 +3017,16 @@ class MCPServerManager:
|
|||
# whole registry, for keys AND admitted session subjects alike (one predicate owns the
|
||||
# question). Seeded into the union rather than returned early so the session resource
|
||||
# scope below still bounds a per-server envelope held by an admin.
|
||||
combined_servers: Final = (
|
||||
set(self.get_registry().keys())
|
||||
if await MCPRequestHandler.admin_view_unscoped(user_api_key_auth)
|
||||
else set(await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth))
|
||||
admin_unscoped: Final = await MCPRequestHandler.admin_view_unscoped(user_api_key_auth)
|
||||
resolved_access: Final = (
|
||||
MCPServerAccess(server_ids=())
|
||||
if admin_unscoped
|
||||
else access or await MCPRequestHandler.get_mcp_server_access(user_api_key_auth)
|
||||
)
|
||||
resolved_server_ids: Final = (
|
||||
set(self.get_registry().keys()) if admin_unscoped else set(resolved_access.server_ids)
|
||||
)
|
||||
combined_servers: Final = set(resolved_server_ids)
|
||||
verbose_logger.debug("Allowed MCP Servers for user api key auth: %s", combined_servers)
|
||||
combined_servers.update(
|
||||
await self.operator_open_server_ids(
|
||||
|
|
@ -3052,6 +3067,18 @@ class MCPServerManager:
|
|||
]
|
||||
combined_servers.update(delegate_server_ids)
|
||||
|
||||
restrict_allow_all: Final = (
|
||||
resolved_general_settings.get("mcp_allow_all_keys_respects_mcp_scope", False)
|
||||
and user_api_key_auth is not None
|
||||
and user_api_key_auth.via_virtual_key
|
||||
and resolved_access.scope != "unscoped"
|
||||
)
|
||||
if restrict_allow_all:
|
||||
combined_servers.difference_update(
|
||||
set(allow_all_server_ids)
|
||||
- resolved_server_ids
|
||||
- (set(submitted_server_ids) if resolved_access.scope != "unresolved" else set())
|
||||
)
|
||||
if len(combined_servers) == 0:
|
||||
verbose_logger.debug("No allowed MCP Servers found for user api key auth.")
|
||||
scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth)
|
||||
|
|
|
|||
|
|
@ -1216,9 +1216,9 @@ class GenerateKeyRequest(KeyRequestBase):
|
|||
organization_id: str | None = None
|
||||
project_id: str | None = None
|
||||
|
||||
@field_validator("team_id", mode="before")
|
||||
@field_validator("team_id", "organization_id", "project_id", mode="before")
|
||||
@classmethod
|
||||
def treat_cleared_team_id_as_unset(cls, v: object) -> object:
|
||||
def treat_cleared_id_as_unset(cls, v: object) -> object:
|
||||
if v == "":
|
||||
return None
|
||||
return v
|
||||
|
|
@ -1278,6 +1278,13 @@ class UpdateKeyRequest(KeyRequestBase):
|
|||
rotation_interval: str | None = None
|
||||
organization_id: str | None = None
|
||||
|
||||
@field_validator("organization_id", mode="before")
|
||||
@classmethod
|
||||
def treat_cleared_organization_id_as_unset(cls, v: object) -> object:
|
||||
if v == "":
|
||||
return None
|
||||
return v
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_temp_budget(self) -> "UpdateKeyRequest":
|
||||
if self.temp_budget_increase is not None or self.temp_budget_expiry is not None:
|
||||
|
|
@ -1923,6 +1930,13 @@ class NewTeamRequest(TeamBase):
|
|||
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
@field_validator("team_id", mode="before")
|
||||
@classmethod
|
||||
def treat_blank_team_id_as_unset(cls, v: object) -> object:
|
||||
if isinstance(v, str) and not v.strip():
|
||||
return None
|
||||
return v
|
||||
|
||||
|
||||
class GlobalEndUsersSpend(LiteLLMPydanticObjectBase):
|
||||
api_key: str | None = None
|
||||
|
|
@ -2594,9 +2608,9 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.",
|
||||
)
|
||||
missing_session_id: Literal["generate", "reject"] | None = Field(
|
||||
missing_session_id: Literal["generate", "reject", "omit"] | None = Field(
|
||||
None,
|
||||
description="What to do with LLM API requests that carry no session id (x-litellm-session-id header, metadata.session_id, etc.). 'generate' stamps one id into litellm_session_id, litellm_trace_id and metadata.session_id so SpendLogs and logging callbacks agree; 'reject' returns 400. Unset keeps the legacy behavior where SpendLogs falls back to the trace id while callbacks get no session id.",
|
||||
description="What to do with LLM API requests that carry no session id (x-litellm-session-id header, metadata.session_id, etc.). 'generate' stamps one id into litellm_session_id, litellm_trace_id and metadata.session_id so SpendLogs and logging callbacks agree; 'reject' returns 400; 'omit' leaves SpendLogs.session_id null, matching callbacks such as Langfuse that only record a client-established metadata.session_id. Unset keeps the legacy behavior where SpendLogs falls back to the trace id while callbacks get no session id.",
|
||||
)
|
||||
enable_public_model_hub: bool = Field(
|
||||
default=False,
|
||||
|
|
@ -4225,6 +4239,8 @@ class TeamAccessGroupModelGrant(LiteLLMPydanticObjectBase):
|
|||
access_group_id: str
|
||||
access_group_name: str
|
||||
models: tuple[str, ...]
|
||||
mcp_server_ids: tuple[str, ...] = ()
|
||||
agent_ids: tuple[str, ...] = ()
|
||||
|
||||
|
||||
class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable):
|
||||
|
|
|
|||
|
|
@ -939,35 +939,32 @@ async def make_agent_public(
|
|||
if agent is None:
|
||||
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found")
|
||||
|
||||
if litellm.public_agent_groups is None:
|
||||
litellm.public_agent_groups = []
|
||||
# handle duplicates
|
||||
if not AGENT_REGISTRY.ids_for_agent(agent.agent_id).isdisjoint(litellm.public_agent_groups):
|
||||
config: Final = await proxy_config.get_config()
|
||||
|
||||
current_public_agent_groups: Final = list(litellm.public_agent_groups or [])
|
||||
if not AGENT_REGISTRY.ids_for_agent(agent.agent_id).isdisjoint(current_public_agent_groups):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Agent with name {agent.agent_name} already in public agent groups",
|
||||
)
|
||||
litellm.public_agent_groups.append(agent.agent_id)
|
||||
updated_public_agent_groups: Final = [*current_public_agent_groups, agent.agent_id]
|
||||
|
||||
# Load existing config
|
||||
config: Final = await proxy_config.get_config()
|
||||
|
||||
# Update config with new settings
|
||||
if "litellm_settings" not in config or config["litellm_settings"] is None:
|
||||
config["litellm_settings"] = {}
|
||||
|
||||
config["litellm_settings"]["public_agent_groups"] = litellm.public_agent_groups
|
||||
config["litellm_settings"]["public_agent_groups"] = updated_public_agent_groups
|
||||
|
||||
# Save the updated config
|
||||
await proxy_config.save_config(new_config=config)
|
||||
|
||||
litellm.public_agent_groups = updated_public_agent_groups
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Updated public agent groups to: %s by user: %s", litellm.public_agent_groups, user_api_key_dict.user_id
|
||||
"Updated public agent groups to: %s by user: %s", updated_public_agent_groups, user_api_key_dict.user_id
|
||||
)
|
||||
|
||||
return {
|
||||
"message": "Successfully updated public agent groups",
|
||||
"public_agent_groups": litellm.public_agent_groups,
|
||||
"public_agent_groups": updated_public_agent_groups,
|
||||
"updated_by": user_api_key_dict.user_id,
|
||||
}
|
||||
except HTTPException:
|
||||
|
|
|
|||
|
|
@ -16,6 +16,10 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy._types import EnterpriseLicenseData
|
||||
|
||||
|
||||
AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router"
|
||||
HEURISTIC_V2_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit."
|
||||
|
||||
|
||||
class LicenseCheck:
|
||||
"""
|
||||
- Check if license in env
|
||||
|
|
@ -149,6 +153,19 @@ class LicenseCheck:
|
|||
return False
|
||||
return team_count > _max_teams_in_license
|
||||
|
||||
def heuristic_v2_router_limit(self) -> int | None:
|
||||
"""
|
||||
How many heuristic_v2 auto-routers this proxy may hold: unlimited (None) only when the
|
||||
signed license lists the auto_router feature, otherwise one. A license verified through
|
||||
the API carries no feature list, so it does not lift the limit either.
|
||||
"""
|
||||
if self.airgapped_license_data is None:
|
||||
return 1
|
||||
allowed_features: Final = self.airgapped_license_data.get("allowed_features")
|
||||
if isinstance(allowed_features, list) and AUTO_ROUTER_LICENSE_FEATURE in allowed_features:
|
||||
return None
|
||||
return 1
|
||||
|
||||
def verify_license_without_api_request(self, public_key, license_key):
|
||||
try:
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
|
|
@ -179,19 +196,21 @@ class LicenseCheck:
|
|||
# Decode and parse the data
|
||||
license_data: Final = json.loads(message.decode())
|
||||
|
||||
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
|
||||
|
||||
# debug information provided in license data
|
||||
verbose_proxy_logger.debug("License data: %s", license_data)
|
||||
|
||||
# Check expiration date
|
||||
expiration_date: Final = datetime.strptime(license_data["expiration_date"], "%Y-%m-%d")
|
||||
if expiration_date < datetime.now():
|
||||
self.airgapped_license_data = None
|
||||
return False, "License has expired"
|
||||
|
||||
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
|
||||
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
self.airgapped_license_data = None
|
||||
verbose_proxy_logger.debug(
|
||||
"litellm.proxy.auth.litellm_license.py::verify_license_without_api_request - Unable to verify License locally. - %s",
|
||||
e,
|
||||
|
|
|
|||
|
|
@ -489,7 +489,7 @@ lite codex exec "summarize the repo"
|
|||
|
||||
Each command resolves your LiteLLM key (logging in via SSO when none is stored and you are at a terminal; otherwise it expects `LITELLM_PROXY_API_KEY` or `--api-key`), checks the key against the proxy so bad credentials fail immediately instead of deep inside the agent, exports the environment variables the agent reads, then replaces itself with the agent process.
|
||||
|
||||
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol).
|
||||
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. It also gets `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` (again unless you already set it) so Claude Code v2.1.129+ fills its `/model` picker from the proxy's `/v1/models`; Claude Code only lists entries whose id contains `claude` or `anthropic`, and older versions ignore the variable. Export it as `0` to turn discovery off. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol).
|
||||
|
||||
Options (these belong to the wrapper, so put them before the agent's own flags):
|
||||
|
||||
|
|
@ -505,7 +505,7 @@ The credential is short-lived by design (default 24h, configurable via `LITELLM_
|
|||
|
||||
### Route Every Claude Code Session Through the Proxy
|
||||
|
||||
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` when that key is missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
|
||||
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` to `1` when those keys are missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
|
||||
|
||||
Two things need to already be true: you've run `lite login` (or `lite login --pkce`, whose key the helper renews on its own), since the apiKeyHelper depends on that stored token, and the proxy is already reachable, since `lite up` does not start one for you.
|
||||
|
||||
|
|
@ -529,7 +529,7 @@ Cursor is not supported: it has no equivalent file-based config to hot-patch thi
|
|||
lite --base-url https://your-proxy.example.com login --config-claude
|
||||
```
|
||||
|
||||
It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
|
||||
It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY`, and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
|
||||
|
||||
Because the credential is reached through `apiKeyHelper` rather than copied into the file, a later `lite login` refreshes it with no further action: Claude Code re-runs the helper on every request and picks up whatever token the most recent login stored. Nothing secret is written to `settings.json`.
|
||||
|
||||
|
|
|
|||
|
|
@ -16,6 +16,8 @@ ANTHROPIC_AUTH_TOKEN_ENV: Final = "ANTHROPIC_AUTH_TOKEN"
|
|||
ANTHROPIC_API_KEY_ENV: Final = "ANTHROPIC_API_KEY"
|
||||
ENABLE_TOOL_SEARCH_ENV: Final = "ENABLE_TOOL_SEARCH"
|
||||
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
|
||||
ENABLE_GATEWAY_MODEL_DISCOVERY_ENV: Final = "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"
|
||||
ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE: Final = "1"
|
||||
OPENAI_BASE_URL_ENV: Final = "OPENAI_BASE_URL"
|
||||
OPENAI_API_KEY_ENV: Final = "OPENAI_API_KEY"
|
||||
|
||||
|
|
@ -67,7 +69,9 @@ def build_agent_env(
|
|||
Anthropic key cannot win over the bearer token we set. ENABLE_TOOL_SEARCH
|
||||
defaults to true because Claude Code turns tool search off when
|
||||
ANTHROPIC_BASE_URL is not a first-party Anthropic host; a value already in
|
||||
the environment is left alone.
|
||||
the environment is left alone. CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY
|
||||
defaults to 1 so Claude Code (v2.1.129+) fills its /model picker from the
|
||||
proxy's /v1/models; likewise left alone when already set.
|
||||
"""
|
||||
env: Final = dict(base_env)
|
||||
root: Final = base_url.rstrip("/")
|
||||
|
|
@ -77,6 +81,8 @@ def build_agent_env(
|
|||
env.pop(ANTHROPIC_API_KEY_ENV, None)
|
||||
if ENABLE_TOOL_SEARCH_ENV not in env:
|
||||
env[ENABLE_TOOL_SEARCH_ENV] = ENABLE_TOOL_SEARCH_VALUE
|
||||
if ENABLE_GATEWAY_MODEL_DISCOVERY_ENV not in env:
|
||||
env[ENABLE_GATEWAY_MODEL_DISCOVERY_ENV] = ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE
|
||||
if PROFILE_OPENAI in profiles:
|
||||
env[OPENAI_BASE_URL_ENV] = root + "/v1"
|
||||
env[OPENAI_API_KEY_ENV] = api_key
|
||||
|
|
|
|||
|
|
@ -26,6 +26,8 @@ ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
|
|||
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
|
||||
ENABLE_TOOL_SEARCH_KEY: Final = "ENABLE_TOOL_SEARCH"
|
||||
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
|
||||
ENABLE_GATEWAY_MODEL_DISCOVERY_KEY: Final = "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"
|
||||
ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE: Final = "1"
|
||||
|
||||
CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json"
|
||||
BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json"
|
||||
|
|
@ -77,13 +79,16 @@ def merge_claude_settings(
|
|||
stray env.ANTHROPIC_API_KEY is dropped so it cannot outrank the helper-issued
|
||||
token (same reasoning as build_agent_env in agents.py). ENABLE_TOOL_SEARCH
|
||||
defaults to true because Claude Code turns tool search off when
|
||||
ANTHROPIC_BASE_URL is not a first-party Anthropic host; an existing value is
|
||||
left alone. Every other key is preserved untouched.
|
||||
ANTHROPIC_BASE_URL is not a first-party Anthropic host, and
|
||||
CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY defaults to 1 so the /model picker
|
||||
is filled from the proxy's /v1/models; existing values of both are left
|
||||
alone. Every other key is preserved untouched.
|
||||
"""
|
||||
raw_env: Final = settings.get(ENV_KEY, {})
|
||||
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
|
||||
env: Final = {
|
||||
ENABLE_TOOL_SEARCH_KEY: ENABLE_TOOL_SEARCH_VALUE,
|
||||
ENABLE_GATEWAY_MODEL_DISCOVERY_KEY: ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE,
|
||||
**{key: value for key, value in base_env.items() if key != ANTHROPIC_API_KEY_KEY},
|
||||
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
|
||||
}
|
||||
|
|
@ -156,6 +161,8 @@ __all__ = (
|
|||
"AUTOROUTE_BACKUP_PATH",
|
||||
"BACKUP_PATH",
|
||||
"CLAUDE_SETTINGS_PATH",
|
||||
"ENABLE_GATEWAY_MODEL_DISCOVERY_KEY",
|
||||
"ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE",
|
||||
"ENABLE_TOOL_SEARCH_KEY",
|
||||
"ENABLE_TOOL_SEARCH_VALUE",
|
||||
"ENV_KEY",
|
||||
|
|
|
|||
|
|
@ -135,7 +135,7 @@ async def get_credentials(
|
|||
]
|
||||
return {"success": True, "credentials": masked_credentials}
|
||||
except Exception as e:
|
||||
return handle_exception_on_proxy(e)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -239,13 +239,18 @@ async def delete_credential(
|
|||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
await CredentialsRepository(prisma_client).delete_by_name(credential_name)
|
||||
deleted: Final = await CredentialsRepository(prisma_client).delete_by_name(credential_name)
|
||||
if deleted is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail="Credential not found. Got credential name: " + credential_name,
|
||||
)
|
||||
|
||||
## DELETE FROM LITELLM ##
|
||||
litellm.credential_list = [cred for cred in litellm.credential_list if cred.credential_name != credential_name]
|
||||
return {"success": True, "message": "Credential deleted successfully"}
|
||||
except Exception as e:
|
||||
return handle_exception_on_proxy(e)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
def update_db_credential(
|
||||
|
|
|
|||
|
|
@ -124,6 +124,42 @@ def add_missing_query_params(url: str, params: Mapping[str, str | int | float])
|
|||
return urllib.parse.urlunsplit(parsed._replace(query=query))
|
||||
|
||||
|
||||
LIBPQ_VERIFY_SSLMODES: Final[frozenset[str]] = frozenset({"verify-ca", "verify-full"})
|
||||
|
||||
|
||||
def translate_libpq_ssl_params(url: str) -> str:
|
||||
"""Rewrite libpq's certificate-verification params into Prisma's dialect.
|
||||
|
||||
Prisma's engine only knows ``sslmode=disable|prefer|require``, ``sslcert``
|
||||
(the CA bundle) and ``sslaccept=strict``. It silently discards
|
||||
``sslrootcert`` and downgrades ``sslmode=verify-ca`` / ``verify-full`` to
|
||||
``prefer``, so a URL copied from libpq / RDS docs connects over TLS with no
|
||||
certificate check at all. ``verify-ca`` and ``verify-full`` both become
|
||||
``require`` (Prisma has no CA-only mode), ``sslrootcert`` becomes
|
||||
``sslcert``, and either one turns on ``sslaccept=strict`` (chain and
|
||||
hostname), matching libpq where a root cert makes ``require`` verify.
|
||||
Prisma params the operator pinned themselves win; anything else is left
|
||||
untouched.
|
||||
"""
|
||||
parsed: Final = urllib.parse.urlsplit(url)
|
||||
pairs: Final = tuple(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True))
|
||||
keys: Final = frozenset(key for key, _ in pairs)
|
||||
wants_verify: Final = any(key == "sslmode" and value in LIBPQ_VERIFY_SSLMODES for key, value in pairs)
|
||||
if not wants_verify and "sslrootcert" not in keys:
|
||||
return url
|
||||
translated: Final = tuple(
|
||||
("sslmode", "require") if key == "sslmode" and value in LIBPQ_VERIFY_SSLMODES else (key, value)
|
||||
for key, value in pairs
|
||||
if key != "sslrootcert"
|
||||
)
|
||||
root_cert: Final = tuple(
|
||||
("sslcert", value) for key, value in pairs if key == "sslrootcert" and "sslcert" not in keys
|
||||
)
|
||||
strict: Final = () if "sslaccept" in keys else (("sslaccept", "strict"),)
|
||||
query: Final = urllib.parse.urlencode(translated + root_cert + strict)
|
||||
return urllib.parse.urlunsplit(parsed._replace(query=query))
|
||||
|
||||
|
||||
def reader_shareable_params(params: Mapping[str, str | int | float]) -> Mapping[str, str | int | float]:
|
||||
"""Return the subset of ``params`` the read replica is allowed to inherit."""
|
||||
return MappingProxyType({key: value for key, value in params.items() if key in CONNECTION_PARAM_KEYS})
|
||||
|
|
@ -403,6 +439,11 @@ class DatabaseURLSettings(BaseSettings):
|
|||
self._raise_for_unsupported_scheme()
|
||||
wrote_writer: Final = self.apply_writer_url_to_env()
|
||||
|
||||
for env_var in ("DATABASE_URL", "DIRECT_URL"):
|
||||
url = os.environ.get(env_var)
|
||||
if url:
|
||||
os.environ[env_var] = translate_libpq_ssl_params(url)
|
||||
|
||||
# DATABASE_DISABLE_PREPARED_STATEMENTS maps to Prisma's `pgbouncer=true`
|
||||
# URL param, same as the CLI's `database_disable_prepared_statements`
|
||||
# config key. An explicit `pgbouncer` value already on the URL wins.
|
||||
|
|
@ -418,7 +459,7 @@ class DatabaseURLSettings(BaseSettings):
|
|||
reader_url: Final = self.build_reader_url() or self.database_url_read_replica
|
||||
if reader_url is not None:
|
||||
os.environ["DATABASE_URL_READ_REPLICA"] = add_missing_query_params(
|
||||
reader_url,
|
||||
translate_libpq_ssl_params(reader_url),
|
||||
connection_params_from_url(os.environ.get("DATABASE_URL", "")),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -916,9 +916,10 @@ class CompresrGuardrail(CustomGuardrail):
|
|||
def _mirror_texts_channel(input_texts: object, applied: _CompressionResult) -> list[object] | None:
|
||||
"""Compressed content mirrored into the Responses `texts` channel.
|
||||
|
||||
The chat/Anthropic handlers round-trip ``structured_messages``; the
|
||||
Responses translation cannot rebuild its input from chat messages and
|
||||
instead writes back through ``texts``. This matches by value, so a
|
||||
The chat/Anthropic/Responses handlers round-trip
|
||||
``structured_messages``; translations without that round-trip write
|
||||
back through ``texts``, so the compressed content is mirrored there
|
||||
too. This matches by value, so a
|
||||
replacement is applied only when it is unambiguous: one compression per
|
||||
text, and every occurrence in ``texts`` accounted for by a compressed
|
||||
target. Anything else is left uncompressed rather than risk a wrong or
|
||||
|
|
|
|||
|
|
@ -50,6 +50,9 @@ if TYPE_CHECKING:
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
BYPASS_HEADER: Final = "x-headroom-bypass"
|
||||
_STREAM_CONVERTIBLE_CALL_TYPES: Final = frozenset(
|
||||
(CallTypes.completion, CallTypes.acompletion, CallTypes.responses, CallTypes.aresponses)
|
||||
)
|
||||
HEADROOM_RETRIEVE_TOOL_NAME: Final = "headroom_retrieve"
|
||||
_HASH_PATTERN: Final = re.compile(r"hash=([a-f0-9]{24})")
|
||||
_HASH_CACHE_TTL_SECONDS: Final = 15 * 60
|
||||
|
|
@ -725,6 +728,10 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.debug("Headroom: %s header set; skipping compression", BYPASS_HEADER)
|
||||
return inputs
|
||||
|
||||
if request_data.get("background"):
|
||||
verbose_proxy_logger.debug("Headroom: background request; skipping compression")
|
||||
return inputs
|
||||
|
||||
structured_messages: Final = inputs.get("structured_messages")
|
||||
if not _is_object_list(structured_messages) or not structured_messages:
|
||||
return inputs
|
||||
|
|
@ -826,9 +833,9 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
) -> dict[str, Any] | None: # mutable-ok: overrides CustomLogger hook whose contract is a plain dict
|
||||
base_result: Final = await super().async_pre_call_deployment_hook(kwargs, call_type)
|
||||
effective: Final = base_result if base_result is not None else kwargs
|
||||
if call_type not in (CallTypes.completion, CallTypes.acompletion):
|
||||
if call_type not in _STREAM_CONVERTIBLE_CALL_TYPES:
|
||||
return base_result
|
||||
if not effective.get("stream"):
|
||||
if not effective.get("stream") or effective.get("background"):
|
||||
return base_result
|
||||
if not has_headroom_retrieve_tool(effective.get("tools")):
|
||||
return base_result
|
||||
|
|
|
|||
|
|
@ -168,11 +168,8 @@ class _ProxyDBLogger(CustomLogger):
|
|||
"custom_llm_provider"
|
||||
) or request_data.get("custom_llm_provider", "")
|
||||
|
||||
# Propagate standard_logging_object and litellm_trace_id from the
|
||||
# Logging instance so that _get_session_id_for_spend_log uses the same
|
||||
# trace_id that Langfuse received (via async_failure_handler).
|
||||
# Without this, the DB session_id would be a random UUID that doesn't
|
||||
# match the Langfuse trace_id, making failed requests unsearchable.
|
||||
# Propagate standard_logging_object and litellm_trace_id from the Logging
|
||||
# instance so the failure row carries the same trace_id Langfuse received.
|
||||
_litellm_logging_obj: Final = request_data.get("litellm_logging_obj")
|
||||
if _litellm_logging_obj is not None:
|
||||
if not request_data.get("standard_logging_object"):
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from litellm.constants import (
|
|||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
SESSION_ID_GENERATED_METADATA_KEY,
|
||||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
)
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
||||
|
|
@ -733,12 +734,18 @@ def apply_missing_session_id_policy(
|
|||
general_settings: Mapping[str, object] | None,
|
||||
request: Request,
|
||||
) -> None:
|
||||
for metadata_key in ("metadata", "litellm_metadata"):
|
||||
if isinstance(client_metadata := data.get(metadata_key), dict):
|
||||
client_metadata.pop(SESSION_ID_OMITTED_METADATA_KEY, None)
|
||||
metadata: Final = data.get(_metadata_variable_name)
|
||||
policy: Final = general_settings.get("missing_session_id") if general_settings else None
|
||||
if policy is None or not _is_llm_inference_route(request):
|
||||
return
|
||||
metadata: Final = data.get(_metadata_variable_name)
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
if policy == "omit":
|
||||
metadata[SESSION_ID_OMITTED_METADATA_KEY] = True
|
||||
return
|
||||
if data.get("litellm_session_id") or metadata.get("session_id"):
|
||||
return
|
||||
match policy:
|
||||
|
|
@ -760,7 +767,8 @@ def apply_missing_session_id_policy(
|
|||
)
|
||||
case _:
|
||||
verbose_proxy_logger.warning(
|
||||
"Ignoring unknown general_settings.missing_session_id=%r; expected 'generate' or 'reject'", policy
|
||||
"Ignoring unknown general_settings.missing_session_id=%r; expected 'generate', 'reject' or 'omit'",
|
||||
policy,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import json
|
|||
import traceback
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Literal, Protocol, cast, overload
|
||||
|
||||
import fastapi
|
||||
|
|
@ -77,6 +78,7 @@ from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
|
|||
BulkUpdateUserRequest,
|
||||
BulkUpdateUserResponse,
|
||||
UserListResponse,
|
||||
UserSearchWhere,
|
||||
UserUpdateResult,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.scim_v2 import (
|
||||
|
|
@ -426,6 +428,11 @@ async def add_new_user_to_default_team(
|
|||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
|
||||
async def _fetch_user_team_ids(user_id: str, prisma_client: "PrismaClient") -> tuple[str, ...]:
|
||||
user_row: Final = await _user_table(prisma_client).find_unique(where={"user_id": user_id})
|
||||
return tuple(user_row.teams) if user_row is not None else ()
|
||||
|
||||
|
||||
@router.post(
|
||||
"/user/new",
|
||||
tags=["Internal User management"],
|
||||
|
|
@ -580,6 +587,11 @@ async def new_user(
|
|||
)
|
||||
|
||||
user_id: Final = cast(str | None, response.get("user_id", None))
|
||||
attached_team_ids: Final = (
|
||||
await _fetch_user_team_ids(user_id=user_id, prisma_client=prisma_client)
|
||||
if user_id is not None and (_team_id is not None or teams is not None)
|
||||
else None
|
||||
)
|
||||
|
||||
if organization_ids is not None and user_id is not None:
|
||||
await _add_user_to_organizations(
|
||||
|
|
@ -596,6 +608,8 @@ async def new_user(
|
|||
response_dict[key] = value
|
||||
|
||||
response_dict["key"] = response.get("token", "")
|
||||
if attached_team_ids is not None:
|
||||
response_dict["teams"] = list(attached_team_ids)
|
||||
|
||||
new_user_response: Final = NewUserResponse.model_validate(response_dict)
|
||||
|
||||
|
|
@ -2068,6 +2082,22 @@ async def _authorize_user_list_request(
|
|||
return ",".join(allowed_org_ids)
|
||||
|
||||
|
||||
_NO_SEARCH_WHERE: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _user_search_where(search: str | None) -> Mapping[str, object]:
|
||||
"""Prisma predicate for `/user/list?search=`: user_id or user_email contains it, case-insensitive."""
|
||||
if not search:
|
||||
return _NO_SEARCH_WHERE
|
||||
search_where: Final[UserSearchWhere] = {
|
||||
"OR": (
|
||||
{"user_id": {"contains": search, "mode": "insensitive"}},
|
||||
{"user_email": {"contains": search, "mode": "insensitive"}},
|
||||
)
|
||||
}
|
||||
return search_where
|
||||
|
||||
|
||||
@router.get(
|
||||
"/user/list",
|
||||
tags=["Internal User management"],
|
||||
|
|
@ -2079,6 +2109,10 @@ async def get_users(
|
|||
user_ids: str | None = fastapi.Query(default=None, description="Get list of users by user_ids"),
|
||||
sso_user_ids: str | None = fastapi.Query(default=None, description="Get list of users by sso_user_id"),
|
||||
user_email: str | None = fastapi.Query(default=None, description="Filter users by partial email match"),
|
||||
search: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description="Combined search: matches users whose 'user_id' or 'user_email' contains the value (case-insensitive).",
|
||||
),
|
||||
team: str | None = fastapi.Query(default=None, description="Filter users by team id"),
|
||||
page: int = fastapi.Query(default=1, ge=1, description="Page number"),
|
||||
page_size: int = fastapi.Query(default=25, ge=1, le=100, description="Number of items per page"),
|
||||
|
|
@ -2109,6 +2143,8 @@ async def get_users(
|
|||
Get list of users by sso_ids. Comma separated list of sso_ids.
|
||||
user_email: Optional[str]
|
||||
Filter users by partial email match
|
||||
search: Optional[str]
|
||||
Combined search: matches users whose user_id or user_email contains the value (case-insensitive)
|
||||
team: Optional[str]
|
||||
Filter users by team id. Will match if user has this team in their teams array.
|
||||
page: int
|
||||
|
|
@ -2185,7 +2221,11 @@ async def get_users(
|
|||
where_conditions["organization_memberships"] = {"some": {"organization_id": {"in": org_id_list}}}
|
||||
|
||||
## Filter any none fastapi.Query params - e.g. where_conditions: {'user_email': {'contains': Query(None), 'mode': 'insensitive'}, 'teams': {'has': Query(None)}}
|
||||
where_conditions = {k: v for k, v in where_conditions.items() if v is not None}
|
||||
where: Final[Mapping[str, object]] = {
|
||||
key: value
|
||||
for key, value in (*where_conditions.items(), *_user_search_where(search).items())
|
||||
if value is not None
|
||||
}
|
||||
|
||||
# Build order_by conditions
|
||||
|
||||
|
|
@ -2194,14 +2234,14 @@ async def get_users(
|
|||
)
|
||||
|
||||
users: Final[Sequence[prisma_models.LiteLLM_UserTable]] = await UserRepository(prisma_client).table.find_many(
|
||||
where=where_conditions,
|
||||
where=where,
|
||||
skip=skip,
|
||||
take=page_size,
|
||||
order=(order_by if order_by else {"created_at": "desc"}), # Default to created_at desc if no sort specified
|
||||
)
|
||||
|
||||
# Get total count of user rows
|
||||
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where_conditions)
|
||||
total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where)
|
||||
|
||||
# Get key count for each user
|
||||
user_key_counts: Final = await get_user_key_counts(prisma_client, [user.user_id for user in users])
|
||||
|
|
|
|||
|
|
@ -13,10 +13,11 @@ model/{model_id}/update - PATCH endpoint for model update.
|
|||
import asyncio
|
||||
import datetime
|
||||
import json
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from json import JSONDecodeError
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
|
||||
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeVar, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator
|
||||
|
|
@ -49,6 +50,7 @@ from litellm.proxy._types import (
|
|||
TeamModelDeleteRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import (
|
||||
coordination_redis_cache,
|
||||
|
|
@ -68,6 +70,10 @@ from litellm.proxy.management_endpoints.team_endpoints import (
|
|||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
update_team as _legacy_update_team,
|
||||
)
|
||||
from litellm.proxy.management_helpers.access_group_model_sync import (
|
||||
sync_access_groups_for_deleted_model,
|
||||
sync_access_groups_for_renamed_model,
|
||||
)
|
||||
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import (
|
||||
PTU_COST_ATTRIBUTION_ENV_VAR,
|
||||
|
|
@ -92,6 +98,9 @@ from litellm.router_strategy.complexity_router import (
|
|||
from litellm.router_utils.auto_router_model_naming import (
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
carries_complexity_router_settings,
|
||||
count_heuristic_v2_routers,
|
||||
heuristic_v2_limit_violation,
|
||||
uses_heuristic_v2_classifier,
|
||||
validate_complexity_router_config_placement,
|
||||
validate_complexity_router_config_write,
|
||||
validate_strategy_router_model_write,
|
||||
|
|
@ -149,6 +158,8 @@ class _ProxyModelTable(Protocol):
|
|||
|
||||
def find_many(self, *, where: Mapping[str, object]) -> Awaitable[Sequence[_ProxyModelRow]]: ...
|
||||
|
||||
def create(self, *, data: Mapping[str, object]) -> Awaitable[_ProxyModelRow]: ...
|
||||
|
||||
def update(
|
||||
self, *, where: Mapping[str, object], data: Mapping[str, object]
|
||||
) -> Awaitable[_ProxyModelRow | None]: ...
|
||||
|
|
@ -162,6 +173,9 @@ class _TxModelTables(Protocol):
|
|||
litellm_proxymodeltable: _ProxyModelTable
|
||||
|
||||
|
||||
_RowT = TypeVar("_RowT")
|
||||
|
||||
|
||||
class _ExistingModelRow(Protocol):
|
||||
@property
|
||||
def litellm_params(self) -> Mapping[str, object]: ...
|
||||
|
|
@ -265,6 +279,66 @@ def _raise_on_strategy_router_write_violation(
|
|||
)
|
||||
|
||||
|
||||
HEURISTIC_V2_SLOT_LOCK_KEY: Final = 5_872_301
|
||||
_HEURISTIC_V2_LOCK_SQL: Final = "SELECT 1 AS locked FROM pg_advisory_xact_lock($1)"
|
||||
_HEURISTIC_V2_DB_ROWS_SQL: Final = """
|
||||
SELECT count(*)::int AS held FROM "LiteLLM_ProxyModelTable"
|
||||
WHERE model_id <> $1
|
||||
AND (CASE jsonb_typeof(litellm_params) WHEN 'string' THEN (litellm_params #>> '{}')::jsonb ELSE litellm_params END)
|
||||
-> 'complexity_router_config' ->> 'classifier_type' = 'heuristic_v2'
|
||||
"""
|
||||
|
||||
|
||||
def _effective_complexity_router_config(
|
||||
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
|
||||
) -> object:
|
||||
"""The complexity config a write leaves on the row: the incoming one when the write carries it, else the stored one."""
|
||||
incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config
|
||||
if incoming is not None or existing_params is None:
|
||||
return incoming
|
||||
return existing_params.complexity_router_config
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _heuristic_v2_slot(
|
||||
prisma_client: PrismaClient, *, effective_config: object, model_id: str | None
|
||||
) -> AsyncGenerator[_ProxyModelTable, None]:
|
||||
"""Hand out the model table to write through while the row's claim on a heuristic_v2 slot is settled.
|
||||
|
||||
A write that leaves the row on classifier_type heuristic_v2 under a limited license runs
|
||||
inside one transaction that takes an advisory lock in its own statement before counting
|
||||
(a statement's snapshot predates anything it locks), so pods cannot both pass the count:
|
||||
the DB rows (any pod, either JSON shape) plus this proxy's config.yaml routers are judged
|
||||
against the license limit and the write is refused with a 403 before it happens. The row
|
||||
being edited keeps its own slot through ``model_id``. Every other write, and every write on
|
||||
an unlimited license, goes through the repository table with no lock. Only the row write
|
||||
itself may run inside: anything that needs a second connection (the team model bookkeeping)
|
||||
must wait until the transaction has committed and the lock is released. The transaction
|
||||
writes bypass the repository's publish-on-write, so the config change is published once
|
||||
after commit, the way delete_team_models does.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import _license_check, llm_router
|
||||
|
||||
limit: Final = _license_check.heuristic_v2_router_limit()
|
||||
if limit is None or not uses_heuristic_v2_classifier(effective_config):
|
||||
yield _proxy_model_table(prisma_client)
|
||||
return
|
||||
async with prisma_client.db.tx() as tx_ctx:
|
||||
tables: Final[_TxModelTables] = tx_ctx
|
||||
await tx_ctx.query_raw(_HEURISTIC_V2_LOCK_SQL, HEURISTIC_V2_SLOT_LOCK_KEY)
|
||||
rows: Sequence[Mapping[str, object]] = await tx_ctx.query_raw(_HEURISTIC_V2_DB_ROWS_SQL, model_id or "")
|
||||
db_held: Final = rows[0].get("held") if rows else 0
|
||||
config_rows: Final = () if llm_router is None else tuple(llm_router.config_deployments())
|
||||
held: Final = (db_held if isinstance(db_held, int) else 0) + count_heuristic_v2_routers(config_rows)
|
||||
violation: Final = heuristic_v2_limit_violation(held=held + 1, limit=limit)
|
||||
if violation is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail=f"{violation} {HEURISTIC_V2_LICENSE_REMEDY}"
|
||||
)
|
||||
yield tables.litellm_proxymodeltable
|
||||
await publish_config_change(redis_cache=coordination_redis_cache(), object_type="litellm_proxymodeltable")
|
||||
|
||||
|
||||
ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add"
|
||||
_REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm")
|
||||
|
||||
|
|
@ -715,22 +789,30 @@ async def patch_model(
|
|||
existing_params=db_model.litellm_params,
|
||||
)
|
||||
|
||||
requested_model_name: Final = patch_data.model_name
|
||||
stored_model_name: str | None = None
|
||||
|
||||
async def write_row(update_data: PrismaCompatibleUpdateDBModel) -> _ProxyModelRow | None:
|
||||
nonlocal stored_model_name
|
||||
stored_model_name = update_data.get("model_name")
|
||||
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
update_data["updated_at"] = cast(str, get_utc_datetime())
|
||||
async with _heuristic_v2_slot(
|
||||
prisma_client,
|
||||
effective_config=_effective_complexity_router_config(
|
||||
patch_data.litellm_params, db_model.litellm_params
|
||||
),
|
||||
model_id=model_id,
|
||||
) as table:
|
||||
return await table.update(where={"model_id": model_id}, data=update_data)
|
||||
|
||||
# Handle team model updates with proper alias management
|
||||
update_data: Final = await _update_team_model_in_db(
|
||||
updated_model: Final = await _update_team_model_in_db(
|
||||
db_model=db_model,
|
||||
patch_data=patch_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
# Add metadata about update
|
||||
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
update_data["updated_at"] = cast(str, get_utc_datetime())
|
||||
|
||||
# Perform partial update
|
||||
updated_model: Final = await _proxy_model_table(prisma_client).update(
|
||||
where={"model_id": model_id},
|
||||
data=update_data,
|
||||
write_row=write_row,
|
||||
)
|
||||
|
||||
if updated_model is None:
|
||||
|
|
@ -741,6 +823,19 @@ async def patch_model(
|
|||
param=None,
|
||||
)
|
||||
|
||||
if (
|
||||
stored_model_name is not None
|
||||
and stored_model_name == requested_model_name
|
||||
and stored_model_name != db_model.model_name
|
||||
):
|
||||
await sync_access_groups_for_renamed_model(
|
||||
prisma_client=prisma_client,
|
||||
model_id=model_id,
|
||||
old_name=db_model.model_name,
|
||||
new_name=stored_model_name,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
|
||||
live_before_reload: Final = live_model_ids_snapshot()
|
||||
reload_outcome: Final = await clear_cache()
|
||||
|
|
@ -961,7 +1056,8 @@ async def _add_model_to_db(
|
|||
prisma_client: PrismaClient,
|
||||
new_encryption_key: str | None = None,
|
||||
should_create_model_in_db: bool = True,
|
||||
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
|
||||
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
|
||||
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
|
||||
# encrypt litellm params #
|
||||
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
|
||||
_original_litellm_model_name: Final = model_params.litellm_params.model
|
||||
|
|
@ -979,18 +1075,20 @@ async def _add_model_to_db(
|
|||
if model_params.model_info.id is not None:
|
||||
_data["model_id"] = model_params.model_info.id
|
||||
_create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above
|
||||
if should_create_model_in_db:
|
||||
model_response = await ModelRepository(prisma_client).table.create(data=_create_data)
|
||||
else:
|
||||
model_response = LiteLLM_ProxyModelTable(**_data)
|
||||
return model_response
|
||||
if not should_create_model_in_db:
|
||||
return LiteLLM_ProxyModelTable(**_data)
|
||||
if slot is None:
|
||||
return await _proxy_model_table(prisma_client).create(data=_create_data)
|
||||
async with slot as table:
|
||||
return await table.create(data=_create_data)
|
||||
|
||||
|
||||
async def _add_team_model_to_db(
|
||||
model_params: Deployment,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
|
||||
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
|
||||
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
|
||||
"""
|
||||
If 'team_id' is provided,
|
||||
|
||||
|
|
@ -1021,6 +1119,7 @@ async def _add_team_model_to_db(
|
|||
model_params=model_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
slot=slot,
|
||||
)
|
||||
|
||||
if original_model_name:
|
||||
|
|
@ -1041,7 +1140,8 @@ async def _update_team_model_in_db(
|
|||
patch_data: updateDeployment,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
) -> PrismaCompatibleUpdateDBModel:
|
||||
write_row: Callable[[PrismaCompatibleUpdateDBModel], Awaitable[_RowT]],
|
||||
) -> _RowT:
|
||||
"""
|
||||
Handle team model updates with proper alias management.
|
||||
|
||||
|
|
@ -1049,6 +1149,9 @@ async def _update_team_model_in_db(
|
|||
- Creates unique internal model_name and team alias
|
||||
- Adds model to team object
|
||||
- Preserves team_public_model_name for external reference
|
||||
|
||||
The row is written through ``write_row`` before the team's model list is touched, so a
|
||||
refused or failed write leaves the team as it was (the create path orders itself the same way).
|
||||
"""
|
||||
# Validate team_id if present in patch_data
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
|
@ -1060,9 +1163,7 @@ async def _update_team_model_in_db(
|
|||
premium_user=premium_user,
|
||||
)
|
||||
|
||||
# Validated before any write, beside the premium check the create path already runs
|
||||
# here. The team ACL is updated below and autocommits, so a validator that raises
|
||||
# further down would leave the team mutated and the deployment row never written.
|
||||
# Validated before the row write, beside the premium check the create path already runs here.
|
||||
#
|
||||
# The merged view is what gets stored, so that is what has to satisfy the invariants.
|
||||
# Validating the patch alone rejected a partial edit of an already valid deployment:
|
||||
|
|
@ -1082,7 +1183,7 @@ async def _update_team_model_in_db(
|
|||
|
||||
# No team_id in patch, proceed with standard update
|
||||
if patch_team_id is None:
|
||||
return update_db_model(db_model=db_model, updated_patch=patch_data)
|
||||
return await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
|
||||
|
||||
# Determine public model name
|
||||
public_model_name: Final = _get_public_model_name(
|
||||
|
|
@ -1101,11 +1202,14 @@ async def _update_team_model_in_db(
|
|||
db_team_id: Final = db_model.model_info.team_id if db_model.model_info else None
|
||||
is_new_team_assignment: Final = db_team_id != patch_team_id
|
||||
|
||||
# Team rows keep their internal UUID-based model_name; the public name lives in model_info
|
||||
patch_data.model_name = f"model_name_{patch_team_id}_{uuid.uuid4()}" if is_new_team_assignment else None
|
||||
row: Final = await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
|
||||
|
||||
if is_new_team_assignment:
|
||||
await _setup_new_team_model_assignment(
|
||||
team_id=patch_team_id,
|
||||
public_model_name=public_model_name,
|
||||
patch_data=patch_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
else:
|
||||
|
|
@ -1113,12 +1217,11 @@ async def _update_team_model_in_db(
|
|||
team_id=patch_team_id,
|
||||
public_model_name=public_model_name,
|
||||
db_model=db_model,
|
||||
patch_data=patch_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
return update_db_model(db_model=db_model, updated_patch=patch_data)
|
||||
return row
|
||||
|
||||
|
||||
def _get_public_model_name(
|
||||
|
|
@ -1170,13 +1273,9 @@ def _get_public_model_name(
|
|||
async def _setup_new_team_model_assignment(
|
||||
team_id: str,
|
||||
public_model_name: str,
|
||||
patch_data: updateDeployment,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
"""Set up a new team model with unique name and team membership."""
|
||||
unique_model_name: Final = f"model_name_{team_id}_{uuid.uuid4()}"
|
||||
patch_data.model_name = unique_model_name
|
||||
|
||||
"""Register a newly team-assigned model's public name on the team."""
|
||||
await team_model_add(
|
||||
data=TeamModelAddRequest(
|
||||
team_id=team_id,
|
||||
|
|
@ -1366,7 +1465,6 @@ async def _update_existing_team_model_assignment(
|
|||
team_id: str,
|
||||
public_model_name: str,
|
||||
db_model: Deployment,
|
||||
patch_data: updateDeployment,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient | None,
|
||||
) -> None:
|
||||
|
|
@ -1390,9 +1488,6 @@ async def _update_existing_team_model_assignment(
|
|||
old_public_name: Final = db_model.model_info.team_public_model_name if db_model.model_info else None
|
||||
|
||||
if old_public_name and public_model_name != old_public_name:
|
||||
# Clear user-supplied public name from patch before any early return so the
|
||||
# caller does not overwrite the internal UUID-based model_name in the DB.
|
||||
patch_data.model_name = None
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"prisma_client not initialized; skipping public name update entirely to avoid orphaned entries"
|
||||
|
|
@ -1440,10 +1535,6 @@ async def _update_existing_team_model_assignment(
|
|||
# else: old_public_name == public_model_name (no rename needed)
|
||||
# No team_model_add/delete calls required; public name is already registered
|
||||
|
||||
# Always clear patch_data.model_name to prevent caller from overwriting
|
||||
# the internal UUID-based model_name in the DB with the user-supplied public name
|
||||
patch_data.model_name = None
|
||||
|
||||
|
||||
class ModelManagementAuthChecks:
|
||||
"""
|
||||
|
|
@ -1673,6 +1764,12 @@ async def delete_model(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
await sync_access_groups_for_deleted_model(
|
||||
prisma_client=prisma_client,
|
||||
model_id=model_info.id,
|
||||
model_name=model_params.model_name,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
## CREATE AUDIT LOG ##
|
||||
asyncio.create_task(
|
||||
|
|
@ -1853,18 +1950,19 @@ async def add_new_model(
|
|||
reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None)
|
||||
try:
|
||||
_original_litellm_model_name: Final = model_params.model_name
|
||||
if model_params.model_info.team_id is None:
|
||||
model_response = await _add_model_to_db(
|
||||
model_params=priced_model_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
else:
|
||||
model_response = await _add_team_model_to_db(
|
||||
model_params=priced_model_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
add_model: Final = (
|
||||
_add_model_to_db if model_params.model_info.team_id is None else _add_team_model_to_db
|
||||
)
|
||||
model_response = await add_model(
|
||||
model_params=priced_model_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
slot=_heuristic_v2_slot(
|
||||
prisma_client,
|
||||
effective_config=priced_model_params.litellm_params.complexity_router_config,
|
||||
model_id=priced_model_params.model_info.id,
|
||||
),
|
||||
)
|
||||
reload_outcome = await proxy_config.add_deployment(
|
||||
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
|
|
@ -1878,6 +1976,8 @@ async def add_new_model(
|
|||
passed_model_info=priced_model_params.model_info,
|
||||
)
|
||||
except Exception as e:
|
||||
if isinstance(e, HTTPException):
|
||||
raise
|
||||
verbose_proxy_logger.exception("Exception in add_new_model: %s", e)
|
||||
|
||||
else:
|
||||
|
|
@ -2027,25 +2127,43 @@ async def update_model(
|
|||
model_params.litellm_params[k] = encrypted_value
|
||||
|
||||
### MERGE WITH EXISTING DATA ###
|
||||
merged_dictionary: Final = {}
|
||||
_mp: Final[dict[str, object]] = model_params.litellm_params.dict()
|
||||
merged_dictionary: Final = {
|
||||
key: _existing_litellm_params_dict[key] if value is None else value
|
||||
for key, value in _mp.items()
|
||||
if value is not None or _existing_litellm_params_dict.get(key) is not None
|
||||
}
|
||||
|
||||
for key, value in _mp.items():
|
||||
if value is not None:
|
||||
merged_dictionary[key] = value
|
||||
elif key in _existing_litellm_params_dict and _existing_litellm_params_dict[key] is not None:
|
||||
merged_dictionary[key] = _existing_litellm_params_dict[key]
|
||||
else:
|
||||
pass
|
||||
|
||||
renamed_to: Final = (
|
||||
model_params.model_name
|
||||
if model_params.model_name not in (None, deployment.model_name)
|
||||
and deployment.model_info.team_id is None
|
||||
else None
|
||||
)
|
||||
_data: Final[dict[str, str]] = {
|
||||
"litellm_params": json.dumps(merged_dictionary),
|
||||
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
||||
**({} if renamed_to is None else {"model_name": renamed_to}),
|
||||
}
|
||||
model_response: Final = await _proxy_model_table(prisma_client).update(
|
||||
where={"model_id": _model_id},
|
||||
data=_data,
|
||||
)
|
||||
async with _heuristic_v2_slot(
|
||||
prisma_client,
|
||||
effective_config=_effective_complexity_router_config(
|
||||
model_params.litellm_params, deployment.litellm_params
|
||||
),
|
||||
model_id=_model_id,
|
||||
) as table:
|
||||
model_response: Final = await table.update(
|
||||
where={"model_id": _model_id},
|
||||
data=_data,
|
||||
)
|
||||
if renamed_to is not None:
|
||||
await sync_access_groups_for_renamed_model(
|
||||
prisma_client=prisma_client,
|
||||
model_id=_model_id,
|
||||
old_name=deployment.model_name,
|
||||
new_name=renamed_to,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
|
||||
live_before_reload: Final = live_model_ids_snapshot()
|
||||
|
|
|
|||
|
|
@ -4318,6 +4318,8 @@ async def _resolve_team_access_group_resources(
|
|||
access_group_id=group.access_group_id,
|
||||
access_group_name=group.access_group_name,
|
||||
models=tuple(group.access_model_names or ()),
|
||||
mcp_server_ids=tuple(group.access_mcp_server_ids or ()),
|
||||
agent_ids=tuple(group.access_agent_ids or ()),
|
||||
)
|
||||
for group in resolved_groups
|
||||
),
|
||||
|
|
|
|||
119
litellm/proxy/management_helpers/access_group_model_sync.py
Normal file
119
litellm/proxy/management_helpers/access_group_model_sync.py
Normal file
|
|
@ -0,0 +1,119 @@
|
|||
"""
|
||||
Keep `litellm_accessgrouptable.access_model_names` pointing at deployment names that still exist.
|
||||
|
||||
Unified access groups store model names, not ids, so a deployment rename or delete that leaves
|
||||
the arrays alone strands every group on a name nothing serves any more.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Final, Protocol
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
|
||||
from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_caches
|
||||
from litellm.repositories.table_repositories import AccessGroupRepository
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
class _TouchedGroupRow(BaseModel):
|
||||
access_group_id: str
|
||||
|
||||
|
||||
class _DeploymentCountRow(BaseModel):
|
||||
deployment_count: int
|
||||
|
||||
|
||||
class _RawExecutor(Protocol):
|
||||
async def query_raw(self, query: str, *args: str) -> Sequence[object]: ...
|
||||
|
||||
|
||||
_BACKING_DEPLOYMENTS_SQL: Final = (
|
||||
'SELECT COUNT(*)::int AS deployment_count FROM "LiteLLM_ProxyModelTable" WHERE "model_name" = $1'
|
||||
)
|
||||
|
||||
_REPLACE_MODEL_NAME_SQL: Final = (
|
||||
'UPDATE "LiteLLM_AccessGroupTable" '
|
||||
'SET "access_model_names" = array_replace(array_remove("access_model_names", $2), $1, $2) '
|
||||
'WHERE $1 = ANY("access_model_names") '
|
||||
'RETURNING "access_group_id"'
|
||||
)
|
||||
|
||||
_APPEND_MODEL_NAME_SQL: Final = (
|
||||
'UPDATE "LiteLLM_AccessGroupTable" '
|
||||
'SET "access_model_names" = array_append("access_model_names", $2) '
|
||||
'WHERE $1 = ANY("access_model_names") AND NOT ($2 = ANY("access_model_names")) '
|
||||
'RETURNING "access_group_id"'
|
||||
)
|
||||
|
||||
_REMOVE_MODEL_NAME_SQL: Final = (
|
||||
'UPDATE "LiteLLM_AccessGroupTable" '
|
||||
'SET "access_model_names" = array_remove("access_model_names", $1) '
|
||||
'WHERE $1 = ANY("access_model_names") '
|
||||
'RETURNING "access_group_id"'
|
||||
)
|
||||
|
||||
|
||||
def _raw_executor(prisma_client: object) -> _RawExecutor:
|
||||
db: Final = AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
|
||||
return WriterPinnedClient(db).db # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
|
||||
|
||||
|
||||
def _config_sourced_sibling(llm_router: Router, deployment_id: str, model_id: str) -> bool:
|
||||
if deployment_id == model_id:
|
||||
return False
|
||||
deployment: Final = llm_router.get_deployment(model_id=deployment_id)
|
||||
return deployment is not None and not deployment.model_info.db_model
|
||||
|
||||
|
||||
def _served_by_a_config_deployment(llm_router: Router | None, model_name: str, model_id: str) -> bool:
|
||||
if llm_router is None:
|
||||
return False
|
||||
return any(
|
||||
_config_sourced_sibling(llm_router, deployment_id, model_id)
|
||||
for deployment_id in llm_router.get_model_ids(model_name=model_name)
|
||||
)
|
||||
|
||||
|
||||
async def _still_backed(executor: _RawExecutor, llm_router: Router | None, model_name: str, model_id: str) -> bool:
|
||||
if _served_by_a_config_deployment(llm_router, model_name, model_id):
|
||||
return True
|
||||
count_rows: Final = await executor.query_raw(_BACKING_DEPLOYMENTS_SQL, model_name)
|
||||
return any(_DeploymentCountRow.model_validate(row).deployment_count > 0 for row in count_rows)
|
||||
|
||||
|
||||
async def _rewrite_groups(executor: _RawExecutor, sql: str, *names: str) -> None:
|
||||
touched_rows: Final = await executor.query_raw(sql, *names)
|
||||
await invalidate_access_group_caches(
|
||||
tuple(_TouchedGroupRow.model_validate(row).access_group_id for row in touched_rows)
|
||||
)
|
||||
|
||||
|
||||
async def sync_access_groups_for_renamed_model(
|
||||
prisma_client: object,
|
||||
*,
|
||||
model_id: str,
|
||||
old_name: str,
|
||||
new_name: str,
|
||||
llm_router: Router | None,
|
||||
) -> None:
|
||||
if old_name == new_name:
|
||||
return
|
||||
executor: Final = _raw_executor(prisma_client)
|
||||
old_name_still_backed: Final = await _still_backed(executor, llm_router, old_name, model_id)
|
||||
await _rewrite_groups(
|
||||
executor, _APPEND_MODEL_NAME_SQL if old_name_still_backed else _REPLACE_MODEL_NAME_SQL, old_name, new_name
|
||||
)
|
||||
|
||||
|
||||
async def sync_access_groups_for_deleted_model(
|
||||
prisma_client: object,
|
||||
*,
|
||||
model_id: str,
|
||||
model_name: str,
|
||||
llm_router: Router | None,
|
||||
) -> None:
|
||||
executor: Final = _raw_executor(prisma_client)
|
||||
if await _still_backed(executor, llm_router, model_name, model_id):
|
||||
return
|
||||
await _rewrite_groups(executor, _REMOVE_MODEL_NAME_SQL, model_name)
|
||||
|
|
@ -1730,6 +1730,16 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict:
|
|||
return headers
|
||||
|
||||
|
||||
def _is_vertex_anthropic_count_tokens_route(endpoint: str) -> bool:
|
||||
return endpoint.rsplit("/", 1)[-1].split(":", 1)[0] == "count-tokens"
|
||||
|
||||
|
||||
def _upstream_headers_for_vertex_route(endpoint: str, headers: Mapping[str, str]) -> Mapping[str, str]:
|
||||
if not _is_vertex_anthropic_count_tokens_route(endpoint):
|
||||
return headers
|
||||
return MappingProxyType({name: value for name, value in headers.items() if name.lower() != "anthropic-beta"})
|
||||
|
||||
|
||||
def get_vertex_pass_through_handler(
|
||||
call_type: Literal["discovery", "aiplatform"], # noqa: UP037 # ruff reports quoted Literal values here
|
||||
) -> BaseVertexAIPassThroughHandler:
|
||||
|
|
@ -2128,7 +2138,7 @@ async def _base_vertex_proxy_route(
|
|||
endpoint_func: Final = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
target=target,
|
||||
custom_headers=headers,
|
||||
custom_headers=_upstream_headers_for_vertex_route(endpoint, headers),
|
||||
is_streaming_request=is_streaming_request,
|
||||
) # dynamically construct pass-through endpoint based on incoming path
|
||||
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm._uuid import uuid
|
||||
from litellm.constants import (
|
||||
MAXIMUM_TRACEBACK_LINES_TO_LOG,
|
||||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -581,8 +582,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
)
|
||||
|
||||
# Set internal keys after merging client-supplied metadata so a request
|
||||
# body that mirrors them cannot clobber the authenticated key or the
|
||||
# real parent span.
|
||||
# body that mirrors them cannot clobber the authenticated key, the real
|
||||
# parent span, or the proxy's own session-id decision.
|
||||
_metadata.pop(SESSION_ID_OMITTED_METADATA_KEY, None)
|
||||
_metadata["user_api_key"] = user_api_key_dict.api_key
|
||||
_metadata["litellm_parent_otel_span"] = user_api_key_dict.parent_otel_span
|
||||
_metadata["user_api_key_budget_reservation"] = user_api_key_dict.budget_reservation
|
||||
|
|
|
|||
|
|
@ -29,6 +29,12 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.auth_utils import is_request_body_safe
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.path_utils import safe_filename
|
||||
from litellm.proxy.prompts.prompt_registry import (
|
||||
DEFAULT_PROMPT_ENVIRONMENT,
|
||||
get_base_prompt_id,
|
||||
get_version_number,
|
||||
prompt_environment_or_default,
|
||||
)
|
||||
from litellm.repositories.table_repositories import PromptRepository
|
||||
from litellm.types.prompts.init_prompts import (
|
||||
ListPromptsResponse,
|
||||
|
|
@ -102,165 +108,20 @@ def _prompt_table(prisma_client: "PrismaClient") -> _PromptTableActions:
|
|||
return PromptRepository(prisma_client).table
|
||||
|
||||
|
||||
def get_base_prompt_id(prompt_id: str) -> str:
|
||||
"""
|
||||
Extract the base prompt ID by stripping the version suffix if present.
|
||||
|
||||
Args:
|
||||
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v1" or "jack_success_v1")
|
||||
|
||||
Returns:
|
||||
Base prompt ID without version suffix (e.g., "jack_success")
|
||||
|
||||
Examples:
|
||||
>>> get_base_prompt_id("jack_success.v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success_v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success")
|
||||
"jack_success"
|
||||
"""
|
||||
# Try dot separator first (.v)
|
||||
if ".v" in prompt_id:
|
||||
return prompt_id.split(".v")[0]
|
||||
# Try underscore separator (_v)
|
||||
if "_v" in prompt_id:
|
||||
return prompt_id.split("_v")[0]
|
||||
return prompt_id
|
||||
|
||||
|
||||
def get_version_number(prompt_id: str) -> int:
|
||||
"""
|
||||
Extract the version number from a versioned prompt ID.
|
||||
|
||||
Args:
|
||||
prompt_id: Prompt ID that may include version suffix (e.g., "jack_success.v2" or "jack_success_v2")
|
||||
|
||||
Returns:
|
||||
Version number (defaults to 1 if no version suffix or invalid format)
|
||||
|
||||
Examples:
|
||||
>>> get_version_number("jack_success.v2")
|
||||
2
|
||||
>>> get_version_number("jack_success_v2")
|
||||
2
|
||||
>>> get_version_number("jack_success")
|
||||
1
|
||||
"""
|
||||
# Try dot separator first (.v)
|
||||
if ".v" in prompt_id:
|
||||
version_str = prompt_id.split(".v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Try underscore separator (_v)
|
||||
if "_v" in prompt_id:
|
||||
version_str = prompt_id.split("_v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return 1
|
||||
|
||||
|
||||
def construct_versioned_prompt_id(prompt_id: str, version: int | None = None) -> str:
|
||||
"""
|
||||
Construct a versioned prompt ID from a base prompt_id and version number.
|
||||
|
||||
Args:
|
||||
prompt_id: Base prompt ID (e.g., "jack_success")
|
||||
version: Version number (if None, returns the base prompt_id unchanged)
|
||||
|
||||
Returns:
|
||||
Versioned prompt ID (e.g., "jack_success.v4")
|
||||
|
||||
Examples:
|
||||
>>> construct_versioned_prompt_id("jack_success", 4)
|
||||
"jack_success.v4"
|
||||
>>> construct_versioned_prompt_id("jack_success", None)
|
||||
"jack_success"
|
||||
>>> construct_versioned_prompt_id("jack_success.v2", 4)
|
||||
"jack_success.v4"
|
||||
"""
|
||||
if version is None:
|
||||
return prompt_id
|
||||
|
||||
# Strip any existing version suffix first
|
||||
base_id: Final = get_base_prompt_id(prompt_id)
|
||||
return f"{base_id}.v{version}"
|
||||
|
||||
|
||||
def get_latest_version_prompt_id(prompt_id: str, all_prompt_ids: Mapping[str, object]) -> str:
|
||||
"""
|
||||
Find the latest version of a prompt from available prompt IDs.
|
||||
|
||||
Args:
|
||||
prompt_id: Base prompt ID or versioned prompt ID (e.g., "jack_success" or "jack_success.v2")
|
||||
all_prompt_ids: Dictionary of all available prompt IDs (keys are prompt IDs)
|
||||
|
||||
Returns:
|
||||
The prompt ID with the highest version number, or the original prompt_id if no versions exist
|
||||
|
||||
Examples:
|
||||
>>> all_ids = {"jack.v1": {}, "jack.v2": {}, "jack.v3": {}}
|
||||
>>> get_latest_version_prompt_id("jack", all_ids)
|
||||
"jack.v3"
|
||||
>>> get_latest_version_prompt_id("jack.v1", all_ids)
|
||||
"jack.v3"
|
||||
>>> all_ids = {"simple": {}}
|
||||
>>> get_latest_version_prompt_id("simple", all_ids)
|
||||
"simple"
|
||||
"""
|
||||
base_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
|
||||
# Find all versions of this prompt
|
||||
matching_versions: Final = []
|
||||
for stored_prompt_id in all_prompt_ids:
|
||||
if get_base_prompt_id(prompt_id=stored_prompt_id) == base_id:
|
||||
version_num = get_version_number(prompt_id=stored_prompt_id)
|
||||
matching_versions.append((version_num, stored_prompt_id))
|
||||
|
||||
# Use the highest version number
|
||||
if matching_versions:
|
||||
matching_versions.sort(reverse=True)
|
||||
return matching_versions[0][1]
|
||||
else:
|
||||
# No versioned prompts found, use the base ID as-is
|
||||
return prompt_id
|
||||
|
||||
|
||||
def get_latest_prompt_versions(prompts: list[PromptSpec]) -> list[PromptSpec]:
|
||||
"""
|
||||
Filter a list of prompts to return only the latest version of each unique prompt.
|
||||
|
||||
Args:
|
||||
prompts: List of PromptSpec objects
|
||||
|
||||
Returns:
|
||||
List of PromptSpec objects with only the latest version of each prompt
|
||||
Filter prompts down to the latest version per (base prompt id, environment).
|
||||
"""
|
||||
latest_prompts: Final[dict[str, PromptSpec]] = {}
|
||||
|
||||
for prompt in prompts:
|
||||
base_id = get_base_prompt_id(prompt_id=prompt.prompt_id)
|
||||
version = get_version_number(prompt_id=prompt.prompt_id)
|
||||
|
||||
# Keep the prompt with the highest version number
|
||||
if base_id not in latest_prompts:
|
||||
latest_prompts[base_id] = prompt
|
||||
else:
|
||||
existing_version = get_version_number(prompt_id=latest_prompts[base_id].prompt_id)
|
||||
if version > existing_version:
|
||||
latest_prompts[base_id] = prompt
|
||||
|
||||
sorted_prompts: Final = sorted(prompts, key=lambda prompt: get_version_number(prompt_id=prompt.prompt_id))
|
||||
latest_prompts: Final = {
|
||||
(get_base_prompt_id(prompt_id=prompt.prompt_id), prompt_environment_or_default(prompt.environment)): prompt
|
||||
for prompt in sorted_prompts
|
||||
}
|
||||
return list(latest_prompts.values())
|
||||
|
||||
|
||||
async def get_next_version_for_prompt(
|
||||
prisma_client: "PrismaClient", prompt_id: str, environment: str = "development"
|
||||
prisma_client: "PrismaClient", prompt_id: str, environment: str = DEFAULT_PROMPT_ENVIRONMENT
|
||||
) -> int:
|
||||
"""
|
||||
Get the next version number for a prompt in a specific environment.
|
||||
|
|
@ -403,11 +264,14 @@ async def list_prompts(
|
|||
if key_metadata is not None:
|
||||
prompts: Final = cast(list[str] | None, key_metadata.get("prompts", None))
|
||||
if prompts is not None:
|
||||
all_prompts = [
|
||||
IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS[prompt_id]
|
||||
for prompt_id in prompts
|
||||
if prompt_id in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS
|
||||
allowed_prompt_ids: Final = frozenset(prompts)
|
||||
allowed_prompts: Final = [
|
||||
spec
|
||||
for spec in IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.values()
|
||||
if spec.prompt_id in allowed_prompt_ids
|
||||
or get_base_prompt_id(prompt_id=spec.prompt_id) in allowed_prompt_ids
|
||||
]
|
||||
all_prompts = get_latest_prompt_versions(prompts=allowed_prompts)
|
||||
if environment:
|
||||
all_prompts = [p for p in all_prompts if p.environment == environment]
|
||||
prompt_list: Final = []
|
||||
|
|
@ -576,7 +440,7 @@ def _get_prompt_template(prompt_spec: PromptSpec, base_prompt_id: str) -> Prompt
|
|||
metadata=parsed.get("metadata"),
|
||||
)
|
||||
else:
|
||||
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(prompt_spec.prompt_id)
|
||||
prompt_callback: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
|
||||
if prompt_callback is not None:
|
||||
integration_name: Final = prompt_callback.integration_name
|
||||
if integration_name == "dotprompt":
|
||||
|
|
@ -690,15 +554,10 @@ async def get_prompt_info(
|
|||
if env_prompts:
|
||||
prompt_spec = create_versioned_prompt_spec(db_prompt=env_prompts[0])
|
||||
|
||||
# Fallback: use in-memory registry (no environment filter)
|
||||
if prompt_spec is None and environment is None:
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if prompt_spec is None:
|
||||
latest_prompt_id: Final = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
|
||||
if prompt_spec is None:
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
|
||||
prompt_id, version=requested_version, environment=environment
|
||||
)
|
||||
|
||||
if prompt_spec is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -785,7 +644,7 @@ async def create_prompt(
|
|||
environment: Final = (
|
||||
request.prompt_info.environment
|
||||
if request.prompt_info and request.prompt_info.environment
|
||||
else "development"
|
||||
else DEFAULT_PROMPT_ENVIRONMENT
|
||||
)
|
||||
|
||||
# Get next version number
|
||||
|
|
@ -885,7 +744,7 @@ async def update_prompt(
|
|||
environment: Final = (
|
||||
request.prompt_info.environment
|
||||
if request.prompt_info and request.prompt_info.environment
|
||||
else "development"
|
||||
else DEFAULT_PROMPT_ENVIRONMENT
|
||||
)
|
||||
|
||||
# Check if any version of this prompt exists (in any environment)
|
||||
|
|
@ -897,9 +756,7 @@ async def update_prompt(
|
|||
detail=f"Prompt with ID {base_prompt_id} not found",
|
||||
)
|
||||
|
||||
# Check if it's a config prompt
|
||||
existing_in_memory: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
if existing_in_memory and existing_in_memory.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
|
|
@ -988,40 +845,26 @@ async def delete_prompt(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
try:
|
||||
# Try to get prompt directly first
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(prompt_id)
|
||||
|
||||
# If not found, try to find the latest version
|
||||
if existing_prompt is None:
|
||||
latest_prompt_id: Final = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
existing_prompt = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(latest_prompt_id)
|
||||
# Use the resolved prompt_id for deletion
|
||||
prompt_id = latest_prompt_id
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(prompt_id, environment=environment)
|
||||
|
||||
if existing_prompt is None:
|
||||
raise HTTPException(status_code=404, detail=f"Prompt with ID {prompt_id} not found")
|
||||
|
||||
if existing_prompt.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot delete config prompts.",
|
||||
)
|
||||
|
||||
# Get the base prompt ID (without version suffix) for database deletion
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
|
||||
# Build delete filter; scope to environment if provided
|
||||
delete_where: Final[dict[str, str]] = {"prompt_id": base_prompt_id}
|
||||
if environment:
|
||||
delete_where["environment"] = environment
|
||||
|
||||
# Delete versions from the database (scoped to environment if provided)
|
||||
delete_where: Final[dict[str, str]] = {
|
||||
"prompt_id": base_prompt_id,
|
||||
**({"environment": environment} if environment else {}),
|
||||
}
|
||||
await _prompt_table(prisma_client).delete_many(where=delete_where)
|
||||
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(base_prompt_id, environment=environment or None)
|
||||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id(
|
||||
base_prompt_id=base_prompt_id, environment=environment or None
|
||||
)
|
||||
|
||||
env_msg: Final = f" from {environment}" if environment else ""
|
||||
return {"message": f"Prompt {base_prompt_id} deleted successfully{env_msg}"}
|
||||
|
|
@ -1093,7 +936,7 @@ async def patch_prompt(
|
|||
try:
|
||||
# Resolve the target row: find the latest version in the given environment
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
env: Final = environment or "development"
|
||||
env: Final = prompt_environment_or_default(environment)
|
||||
requested_version: Final = get_version_number(prompt_id=prompt_id) if prompt_id != base_prompt_id else None
|
||||
|
||||
# Build query to find the exact row by composite unique key
|
||||
|
|
@ -1117,11 +960,7 @@ async def patch_prompt(
|
|||
|
||||
target_row: Final = db_rows[0]
|
||||
|
||||
# Check if prompt exists in memory
|
||||
versioned_id: Final = f"{base_prompt_id}.v{target_row.version}"
|
||||
existing_prompt: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(versioned_id)
|
||||
|
||||
if existing_prompt and existing_prompt.prompt_info.prompt_type == "config":
|
||||
if IN_MEMORY_PROMPT_REGISTRY.has_config_prompt(base_prompt_id=base_prompt_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot update config prompts.",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import importlib
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -14,6 +14,87 @@ from litellm.types.prompts.init_prompts import (
|
|||
|
||||
prompt_initializer_registry = {}
|
||||
|
||||
DEFAULT_PROMPT_ENVIRONMENT: Final = "development"
|
||||
PROMPT_ENVIRONMENT_SERVE_PRECEDENCE: Final = ("production", "staging", "development")
|
||||
|
||||
|
||||
def get_base_prompt_id(prompt_id: str) -> str:
|
||||
"""
|
||||
Extract the base prompt ID by stripping the version suffix if present.
|
||||
|
||||
Examples:
|
||||
>>> get_base_prompt_id("jack_success.v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success_v1")
|
||||
"jack_success"
|
||||
>>> get_base_prompt_id("jack_success")
|
||||
"jack_success"
|
||||
"""
|
||||
if ".v" in prompt_id:
|
||||
return prompt_id.split(".v")[0]
|
||||
if "_v" in prompt_id:
|
||||
return prompt_id.split("_v")[0]
|
||||
return prompt_id
|
||||
|
||||
|
||||
def get_version_number(prompt_id: str) -> int:
|
||||
"""
|
||||
Extract the version number from a versioned prompt ID (defaults to 1).
|
||||
|
||||
Examples:
|
||||
>>> get_version_number("jack_success.v2")
|
||||
2
|
||||
>>> get_version_number("jack_success_v2")
|
||||
2
|
||||
>>> get_version_number("jack_success")
|
||||
1
|
||||
"""
|
||||
if ".v" in prompt_id:
|
||||
version_str = prompt_id.split(".v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if "_v" in prompt_id:
|
||||
version_str = prompt_id.split("_v")[1]
|
||||
try:
|
||||
return int(version_str)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return 1
|
||||
|
||||
|
||||
def prompt_environment_or_default(environment: str | None) -> str:
|
||||
return environment or DEFAULT_PROMPT_ENVIRONMENT
|
||||
|
||||
|
||||
def registry_key_for_prompt(prompt: PromptSpec) -> str:
|
||||
return f"{prompt.prompt_id}::{prompt_environment_or_default(prompt.environment)}"
|
||||
|
||||
|
||||
def parse_prompt_version(raw_version: object) -> int | None:
|
||||
if isinstance(raw_version, bool):
|
||||
return None
|
||||
if isinstance(raw_version, int):
|
||||
return raw_version
|
||||
if isinstance(raw_version, str) and raw_version.isdigit():
|
||||
return int(raw_version)
|
||||
return None
|
||||
|
||||
|
||||
def _spec_version(prompt: PromptSpec) -> int:
|
||||
return prompt.version if prompt.version is not None else get_version_number(prompt_id=prompt.prompt_id)
|
||||
|
||||
|
||||
def _default_serve_environment(prompts: Sequence[PromptSpec]) -> str:
|
||||
present: Final = frozenset(prompt_environment_or_default(prompt.environment) for prompt in prompts)
|
||||
ladder_pick: Final = next((env for env in PROMPT_ENVIRONMENT_SERVE_PRECEDENCE if env in present), None)
|
||||
if ladder_pick is not None:
|
||||
return ladder_pick
|
||||
return min(present) if present else DEFAULT_PROMPT_ENVIRONMENT
|
||||
|
||||
|
||||
def get_prompt_initializer_from_integrations():
|
||||
"""
|
||||
|
|
@ -113,17 +194,16 @@ class InMemoryPromptRegistry:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
prompt_id: Final = prompt.prompt_id
|
||||
if prompt_id in self.IN_MEMORY_PROMPTS:
|
||||
verbose_proxy_logger.debug("prompt_id already exists in IN_MEMORY_PROMPTS")
|
||||
return self.IN_MEMORY_PROMPTS[prompt_id]
|
||||
registry_key: Final = registry_key_for_prompt(prompt)
|
||||
if registry_key in self.IN_MEMORY_PROMPTS:
|
||||
verbose_proxy_logger.debug("prompt already exists in IN_MEMORY_PROMPTS")
|
||||
return self.IN_MEMORY_PROMPTS[registry_key]
|
||||
|
||||
parsed_prompt, custom_prompt_callback = self._build_prompt_callback(prompt=prompt)
|
||||
litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback)
|
||||
|
||||
# store references to the prompt in memory
|
||||
self.IN_MEMORY_PROMPTS[prompt_id] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[prompt_id] = custom_prompt_callback
|
||||
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[registry_key] = custom_prompt_callback
|
||||
|
||||
return parsed_prompt
|
||||
|
||||
|
|
@ -166,68 +246,93 @@ class InMemoryPromptRegistry:
|
|||
import litellm
|
||||
|
||||
parsed_prompt, new_callback = self._build_prompt_callback(prompt=prompt)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt.prompt_id, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(prompt.prompt_id, None)
|
||||
registry_key: Final = registry_key_for_prompt(parsed_prompt)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
|
||||
if stale_callback is not None:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
|
||||
litellm.logging_callback_manager.add_litellm_callback(new_callback)
|
||||
self.IN_MEMORY_PROMPTS[prompt.prompt_id] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[prompt.prompt_id] = new_callback
|
||||
self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
|
||||
self.prompt_id_to_custom_prompt[registry_key] = new_callback
|
||||
return parsed_prompt
|
||||
|
||||
def sync_prompt_from_db(self, prompt: PromptSpec) -> PromptSpec | None:
|
||||
existing: Final = self.IN_MEMORY_PROMPTS.get(prompt.prompt_id)
|
||||
existing: Final = self.IN_MEMORY_PROMPTS.get(registry_key_for_prompt(prompt))
|
||||
if existing is None:
|
||||
return self.initialize_prompt(prompt=prompt)
|
||||
if existing.litellm_params == prompt.litellm_params and existing.prompt_info == prompt.prompt_info:
|
||||
return existing
|
||||
return self.reload_prompt(prompt=prompt)
|
||||
|
||||
def get_prompt_by_id(self, prompt_id: str) -> PromptSpec | None:
|
||||
def resolve_prompt_spec(
|
||||
self,
|
||||
prompt_id: str,
|
||||
version: int | None = None,
|
||||
environment: str | None = None,
|
||||
) -> PromptSpec | None:
|
||||
"""
|
||||
Get a prompt by its ID from memory
|
||||
"""
|
||||
return self.IN_MEMORY_PROMPTS.get(prompt_id)
|
||||
Resolve a prompt spec by base prompt id, optional version, and optional environment.
|
||||
|
||||
def get_prompt_callback_by_id(self, prompt_id: str) -> CustomPromptManagement | None:
|
||||
With no environment, resolves within the default serve environment
|
||||
(production > staging > development > alphabetical first present).
|
||||
With no version, resolves to the highest version in the chosen environment.
|
||||
"""
|
||||
Get a prompt callback by its ID from memory
|
||||
"""
|
||||
return self.prompt_id_to_custom_prompt.get(prompt_id)
|
||||
base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id)
|
||||
base_matches: Final = tuple(
|
||||
spec
|
||||
for spec in self.IN_MEMORY_PROMPTS.values()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
)
|
||||
if not base_matches:
|
||||
return None
|
||||
resolved_environment: Final = (
|
||||
environment if environment is not None else _default_serve_environment(base_matches)
|
||||
)
|
||||
env_matches: Final = tuple(
|
||||
spec for spec in base_matches if prompt_environment_or_default(spec.environment) == resolved_environment
|
||||
)
|
||||
if not env_matches:
|
||||
return None
|
||||
if version is not None:
|
||||
return next((spec for spec in env_matches if _spec_version(spec) == version), None)
|
||||
return max(env_matches, key=_spec_version)
|
||||
|
||||
def remove_prompt(self, prompt_id: str) -> None:
|
||||
def get_prompt_callback_for_prompt(self, prompt: PromptSpec) -> CustomPromptManagement | None:
|
||||
return self.prompt_id_to_custom_prompt.get(registry_key_for_prompt(prompt))
|
||||
|
||||
def has_config_prompt(self, base_prompt_id: str) -> bool:
|
||||
return any(
|
||||
spec.prompt_info.prompt_type == "config"
|
||||
for spec in self.IN_MEMORY_PROMPTS.values()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
)
|
||||
|
||||
def remove_prompt(self, registry_key: str) -> None:
|
||||
import litellm
|
||||
|
||||
self.IN_MEMORY_PROMPTS.pop(prompt_id, None)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(prompt_id, None)
|
||||
self.IN_MEMORY_PROMPTS.pop(registry_key, None)
|
||||
stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None)
|
||||
if stale_callback is not None:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback)
|
||||
|
||||
def delete_prompts_by_base_id(self, base_prompt_id: str, environment: str | None = None) -> list[str]:
|
||||
"""
|
||||
Delete all prompts matching the given base prompt ID from memory, along with their
|
||||
registered callbacks; scoped to one environment when given.
|
||||
Delete matching prompts from memory, along with their registered callbacks,
|
||||
scoped to one environment when given.
|
||||
|
||||
Args:
|
||||
base_prompt_id: The base prompt ID (without version suffix)
|
||||
environment: When set, only delete prompts deployed to this environment
|
||||
|
||||
Returns:
|
||||
List of prompt IDs that were deleted
|
||||
Returns the registry keys that were deleted.
|
||||
"""
|
||||
from litellm.proxy.prompts.prompt_endpoints import get_base_prompt_id
|
||||
|
||||
prompts_to_delete: Final = [
|
||||
pid
|
||||
for pid, prompt in self.IN_MEMORY_PROMPTS.items()
|
||||
if get_base_prompt_id(prompt_id=pid) == base_prompt_id
|
||||
and (environment is None or prompt.environment == environment)
|
||||
keys_to_delete: Final = [
|
||||
key
|
||||
for key, spec in self.IN_MEMORY_PROMPTS.items()
|
||||
if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id
|
||||
and (environment is None or prompt_environment_or_default(spec.environment) == environment)
|
||||
]
|
||||
|
||||
for pid in prompts_to_delete:
|
||||
self.remove_prompt(prompt_id=pid)
|
||||
for key in keys_to_delete:
|
||||
self.remove_prompt(registry_key=key)
|
||||
|
||||
return prompts_to_delete
|
||||
return keys_to_delete
|
||||
|
||||
|
||||
IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()
|
||||
|
|
|
|||
|
|
@ -1228,6 +1228,7 @@ def run_server(
|
|||
add_missing_query_params,
|
||||
idle_lifetime_params,
|
||||
reader_shareable_params,
|
||||
translate_libpq_ssl_params,
|
||||
unsupported_db_scheme,
|
||||
unsupported_db_scheme_message,
|
||||
)
|
||||
|
|
@ -1275,11 +1276,15 @@ def run_server(
|
|||
writer_url,
|
||||
connection_url_params,
|
||||
)
|
||||
os.environ["DATABASE_URL"] = add_missing_query_params(modified_url, lifetime_params)
|
||||
os.environ["DATABASE_URL"] = translate_libpq_ssl_params(
|
||||
add_missing_query_params(modified_url, lifetime_params)
|
||||
)
|
||||
if os.getenv("DIRECT_URL", None) is not None:
|
||||
database_url = os.getenv("DIRECT_URL")
|
||||
modified_url = append_query_params(database_url, connection_url_params)
|
||||
os.environ["DIRECT_URL"] = add_missing_query_params(modified_url, lifetime_params)
|
||||
os.environ["DIRECT_URL"] = translate_libpq_ssl_params(
|
||||
add_missing_query_params(modified_url, lifetime_params)
|
||||
)
|
||||
# The reader pool is a real pool against the same configured cap, so it
|
||||
# gets the allowlisted pool params. Schema-affecting ones, including any
|
||||
# the operator smuggled in through database_extra_connection_params, stay
|
||||
|
|
@ -1292,14 +1297,16 @@ def run_server(
|
|||
db_statement_timeout,
|
||||
db_lock_timeout,
|
||||
)
|
||||
os.environ["DATABASE_URL_READ_REPLICA"] = add_missing_query_params(
|
||||
os.environ["DATABASE_URL_READ_REPLICA"] = translate_libpq_ssl_params(
|
||||
add_missing_query_params(
|
||||
_with_query_value(read_replica_url, "options", reader_options)
|
||||
if reader_options
|
||||
else read_replica_url,
|
||||
reader_shareable_params(connection_url_params),
|
||||
),
|
||||
lifetime_params,
|
||||
add_missing_query_params(
|
||||
_with_query_value(read_replica_url, "options", reader_options)
|
||||
if reader_options
|
||||
else read_replica_url,
|
||||
reader_shareable_params(connection_url_params),
|
||||
),
|
||||
lifetime_params,
|
||||
)
|
||||
)
|
||||
subprocess.run(["prisma"], capture_output=True)
|
||||
is_prisma_runnable = True
|
||||
|
|
|
|||
|
|
@ -120,6 +120,8 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
from litellm.router_utils.auto_router_model_naming import (
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
carries_complexity_router_settings,
|
||||
count_heuristic_v2_routers,
|
||||
heuristic_v2_limit_violation,
|
||||
validate_complexity_router_config_placement,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -301,7 +303,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
)
|
||||
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.litellm_license import LicenseCheck
|
||||
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY, LicenseCheck
|
||||
from litellm.proxy.auth.model_checks import (
|
||||
expand_wildcard_deployments_for_model_info,
|
||||
get_all_fallbacks,
|
||||
|
|
@ -4316,6 +4318,19 @@ def validate_deployment_complexity_router_placement(model: Mapping[str, object])
|
|||
raise ValueError(f"model {model.get('model_name', '')!r}: {violation}")
|
||||
|
||||
|
||||
def validate_heuristic_v2_router_limit(model_list: Sequence[Mapping[str, object]], *, limit: int | None) -> None:
|
||||
"""
|
||||
Refuse to start when config.yaml defines more heuristic_v2 auto-routers than the license allows.
|
||||
|
||||
Checked here rather than left to router registration for the same reason as the two
|
||||
validators above: the proxy builds its router with `ignore_invalid_deployments=True`, so
|
||||
the router's own refusal would turn the extra router into a silently missing model.
|
||||
"""
|
||||
violation: Final = heuristic_v2_limit_violation(held=count_heuristic_v2_routers(model_list), limit=limit)
|
||||
if violation is not None:
|
||||
raise ValueError(f"config.yaml model_list: {violation} {HEURISTIC_V2_LICENSE_REMEDY}")
|
||||
|
||||
|
||||
def pin_complexity_router_model_id(model: dict) -> None: # mutable-ok: out-param, model_info is stamped in place
|
||||
"""
|
||||
Stamps `model_info.id` from the raw litellm_params before plugin resolution swaps
|
||||
|
|
@ -5721,6 +5736,7 @@ class ProxyConfig:
|
|||
model_list: Final = config.get("model_list", None)
|
||||
if model_list:
|
||||
router_params["model_list"] = model_list
|
||||
validate_heuristic_v2_router_limit(model_list, limit=_license_check.heuristic_v2_router_limit())
|
||||
print( # noqa: T201
|
||||
"\033[32mLiteLLM: Proxy initialized with Config, Set models:\033[0m"
|
||||
)
|
||||
|
|
@ -5810,6 +5826,7 @@ class ProxyConfig:
|
|||
),
|
||||
ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
|
||||
)
|
||||
|
||||
if redis_usage_cache is not None and router.cache.redis_cache is None:
|
||||
|
|
@ -6270,6 +6287,7 @@ class ProxyConfig:
|
|||
search_tools=search_tools,
|
||||
ignore_invalid_deployments=True,
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
|
||||
)
|
||||
verbose_proxy_logger.debug("updated llm_router: %s", llm_router)
|
||||
else:
|
||||
|
|
@ -7571,7 +7589,7 @@ class ProxyConfig:
|
|||
return create_versioned_prompt_spec(db_prompt=db_prompt)
|
||||
|
||||
async def _init_prompts_in_db(self, prisma_client: PrismaClient):
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY, registry_key_for_prompt
|
||||
from litellm.types.prompts.init_prompts import PromptSpec
|
||||
|
||||
def parse_row(db_prompt: object) -> PromptSpec | None:
|
||||
|
|
@ -7586,21 +7604,12 @@ class ProxyConfig:
|
|||
return None
|
||||
|
||||
try:
|
||||
prompt_ids_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
|
||||
registry_keys_loaded_before_db_read: Final = frozenset(IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS)
|
||||
prompts_in_db: Final[Sequence[object]] = await PromptRepository(prisma_client).table.find_many()
|
||||
parsed_specs: Final[tuple[PromptSpec, ...]] = tuple(
|
||||
spec for row in prompts_in_db if (spec := parse_row(row)) is not None
|
||||
)
|
||||
newest_spec_per_id: Final[Mapping[str, PromptSpec]] = MappingProxyType(
|
||||
{
|
||||
spec.prompt_id: spec
|
||||
for spec in sorted(
|
||||
parsed_specs,
|
||||
key=lambda s: s.updated_at.timestamp() if s.updated_at else float("-inf"),
|
||||
)
|
||||
}
|
||||
)
|
||||
for prompt_spec in newest_spec_per_id.values():
|
||||
for prompt_spec in parsed_specs:
|
||||
try:
|
||||
IN_MEMORY_PROMPT_REGISTRY.sync_prompt_from_db(prompt=prompt_spec)
|
||||
except Exception as prompt_sync_error: # noqa: BLE001 # one poisoned row must not block syncing the remaining prompts
|
||||
|
|
@ -7612,15 +7621,16 @@ class ProxyConfig:
|
|||
# An unparsable row still exists in the DB, so skip the sweep rather than unload its in-memory copy
|
||||
every_row_parsed: Final = len(parsed_specs) == len(prompts_in_db)
|
||||
if every_row_parsed:
|
||||
deleted_db_prompt_ids: Final = tuple(
|
||||
prompt_id
|
||||
for prompt_id in prompt_ids_loaded_before_db_read
|
||||
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(prompt_id)) is not None
|
||||
db_registry_keys: Final = frozenset(registry_key_for_prompt(spec) for spec in parsed_specs)
|
||||
deleted_db_registry_keys: Final = tuple(
|
||||
registry_key
|
||||
for registry_key in registry_keys_loaded_before_db_read
|
||||
if (loaded_spec := IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS.get(registry_key)) is not None
|
||||
and loaded_spec.prompt_info.prompt_type == "db"
|
||||
and prompt_id not in newest_spec_per_id
|
||||
and registry_key not in db_registry_keys
|
||||
)
|
||||
for deleted_prompt_id in deleted_db_prompt_ids:
|
||||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id=deleted_prompt_id)
|
||||
for deleted_registry_key in deleted_db_registry_keys:
|
||||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(registry_key=deleted_registry_key)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("litellm.proxy.proxy_server.py::ProxyConfig:_init_prompts_in_db - %s", e)
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import json
|
|||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Annotated,
|
||||
|
|
@ -12,6 +13,7 @@ from typing import (
|
|||
Literal,
|
||||
NamedTuple,
|
||||
Protocol,
|
||||
TypeAlias,
|
||||
TypedDict,
|
||||
TypeVar,
|
||||
cast, # noqa: TID251 # prisma group_by returns untyped aggregate mappings
|
||||
|
|
@ -19,6 +21,7 @@ from typing import (
|
|||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
|
|
@ -55,9 +58,21 @@ router: Final = APIRouter()
|
|||
|
||||
SPEND_LOGS_PAGINATION_COUNT_CAP: Final = 10000
|
||||
|
||||
_SESSION_GROUP_KEY_SQL: Final = "COALESCE(NULLIF(session_id, ''), request_id), api_key"
|
||||
_SESSION_KEY_EXPR: Final = "COALESCE(NULLIF(session_id, ''), request_id)"
|
||||
_SESSION_GROUP_KEY_SQL: Final = f"{_SESSION_KEY_EXPR}, api_key"
|
||||
_MCP_CALL_TYPES_SQL: Final = "('call_mcp_tool', 'list_mcp_tools')"
|
||||
_AGENT_CALL_TYPE_SQL: Final = "'asend_message'"
|
||||
_SPEND_LOG_LIST_COLUMNS: Final = """
|
||||
request_id, call_type, api_key, spend, total_tokens,
|
||||
prompt_tokens, completion_tokens, "startTime", "endTime",
|
||||
"completionStartTime", model, model_id, model_group,
|
||||
custom_llm_provider, api_base, "user", metadata,
|
||||
cache_hit, cache_key, request_tags, team_id,
|
||||
organization_id, end_user, requester_ip_address,
|
||||
session_id, status, mcp_namespaced_tool_name, agent_id,
|
||||
COALESCE(request_duration_ms,
|
||||
(EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms
|
||||
"""
|
||||
|
||||
_INTERNAL_HEALTH_CHECK_API_KEYS: Final = (
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
|
|
@ -158,6 +173,32 @@ class _SessionSpendRow(TypedDict):
|
|||
session_cache_hit_count: ReadOnly[int]
|
||||
session_llm_count: ReadOnly[int]
|
||||
session_agent_count: ReadOnly[int]
|
||||
session_total_prompt_tokens: ReadOnly[int]
|
||||
session_total_completion_tokens: ReadOnly[int]
|
||||
session_total_tokens: ReadOnly[int]
|
||||
session_models: ReadOnly[Sequence[str]]
|
||||
|
||||
|
||||
_SESSION_MODELS_LIMIT: Final = 10
|
||||
_SESSION_MODEL_NAME_MAX_LEN: Final = 256
|
||||
|
||||
|
||||
class _SessionSpendStats(NamedTuple):
|
||||
session_total_count: int
|
||||
session_total_spend: float
|
||||
mcp_tool_call_count: int
|
||||
mcp_tool_call_spend: float
|
||||
session_cache_hit_count: int
|
||||
session_llm_count: int
|
||||
session_agent_count: int
|
||||
session_total_prompt_tokens: int
|
||||
session_total_completion_tokens: int
|
||||
session_total_tokens: int
|
||||
session_models: Sequence[str]
|
||||
session_models_truncated: bool
|
||||
|
||||
|
||||
_SessionSpendMap: TypeAlias = Mapping[tuple[str, str], _SessionSpendStats]
|
||||
|
||||
|
||||
class _SpendSumAggregate(TypedDict, total=False):
|
||||
|
|
@ -2281,6 +2322,13 @@ async def ui_view_spend_logs(
|
|||
default=False,
|
||||
description="Paginate over sessions instead of raw logs: one representative row per session, total counts sessions",
|
||||
),
|
||||
session_cursor: str | None = fastapi.Query(
|
||||
default=None,
|
||||
description=(
|
||||
"Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. "
|
||||
"UI route only, honored when sorting by startTime"
|
||||
),
|
||||
),
|
||||
):
|
||||
"""
|
||||
View spend logs with pagination support.
|
||||
|
|
@ -2614,6 +2662,18 @@ async def ui_view_spend_logs(
|
|||
sql_params.append(f"%{error_message}%")
|
||||
p += 1
|
||||
|
||||
if group_by_session is True and not is_v2 and not is_request_id_lookup and sort_by == "startTime":
|
||||
return await _ui_session_grouped_spend_logs(
|
||||
prisma_client=prisma_client,
|
||||
sql_conditions=sql_conditions,
|
||||
sql_params=sql_params,
|
||||
next_param_index=p,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
sort_desc=order_direction != "asc",
|
||||
session_cursor=session_cursor,
|
||||
)
|
||||
|
||||
# Build the ORDER BY expression. ttft_ms is computed from
|
||||
# completionStartTime - startTime; non-streaming rows (where
|
||||
# completionStartTime is null or equals endTime) yield NULL, so we
|
||||
|
|
@ -2655,19 +2715,11 @@ async def ui_view_spend_logs(
|
|||
total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP
|
||||
total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total
|
||||
|
||||
select_columns: Final = """request_id, call_type, api_key, spend, total_tokens,
|
||||
prompt_tokens, completion_tokens, "startTime", "endTime",
|
||||
"completionStartTime", model, model_id, model_group,
|
||||
custom_llm_provider, api_base, "user", metadata,
|
||||
cache_hit, cache_key, request_tags, team_id,
|
||||
organization_id, end_user, requester_ip_address,
|
||||
session_id, status, mcp_namespaced_tool_name, agent_id,
|
||||
COALESCE(request_duration_ms, (EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms"""
|
||||
sql_query: Final = (
|
||||
f"""
|
||||
SELECT * FROM (
|
||||
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
|
||||
{select_columns}
|
||||
{_SPEND_LOG_LIST_COLUMNS}
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {joined_conditions}
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
|
||||
|
|
@ -2678,7 +2730,7 @@ async def ui_view_spend_logs(
|
|||
if session_grouping
|
||||
else f"""
|
||||
SELECT
|
||||
{select_columns}
|
||||
{_SPEND_LOG_LIST_COLUMNS}
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {joined_conditions}
|
||||
ORDER BY {_order_expr} {_sql_dir}{_nulls_clause}
|
||||
|
|
@ -2711,6 +2763,162 @@ async def ui_view_spend_logs(
|
|||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
class _SessionPageRow(TypedDict):
|
||||
session_key: ReadOnly[str]
|
||||
api_key: ReadOnly[str]
|
||||
last_activity: ReadOnly[str]
|
||||
|
||||
|
||||
def _parse_session_cursor(session_cursor: str | None) -> tuple[str, str, str] | None:
|
||||
if session_cursor is None or session_cursor.count("|") < 2:
|
||||
return None
|
||||
last_activity, _, rest = session_cursor.partition("|")
|
||||
api_key, _, session_key = rest.partition("|")
|
||||
if not last_activity or not session_key:
|
||||
return None
|
||||
return (last_activity, session_key, api_key)
|
||||
|
||||
|
||||
async def _fetch_session_representatives(
|
||||
prisma_client: "PrismaClient",
|
||||
where_clause: str,
|
||||
sql_params: Sequence[object],
|
||||
next_param_index: int,
|
||||
session_keys: Sequence[tuple[str, str]],
|
||||
) -> list[dict[str, object]]: # mutable-ok: _build_ui_spend_logs_response writes session counts onto each row
|
||||
"""Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
|
||||
rep_query: Final = f"""
|
||||
SELECT * FROM (
|
||||
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
|
||||
{_SPEND_LOG_LIST_COLUMNS}
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {where_clause}
|
||||
AND ({_SESSION_GROUP_KEY_SQL}) IN (
|
||||
SELECT * FROM unnest(${next_param_index}::text[], ${next_param_index + 1}::text[])
|
||||
)
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
|
||||
) AS session_representatives
|
||||
"""
|
||||
rep_rows: Final[Sequence[dict[str, object]]] = await _query_raw( # mutable-ok: rows are enriched in place
|
||||
prisma_client,
|
||||
rep_query,
|
||||
*sql_params,
|
||||
[session_key for session_key, _ in session_keys], # mutable-ok: prisma serializes array params from a list
|
||||
[api_key for _, api_key in session_keys], # mutable-ok: prisma serializes array params from a list
|
||||
)
|
||||
rep_by_key: Final[Mapping[tuple[str, str], dict[str, object]]] = MappingProxyType( # mutable-ok: same rows
|
||||
{(str(row["session_id"] or row["request_id"]), str(row["api_key"])): row for row in rep_rows}
|
||||
)
|
||||
return [rep_by_key[key] for key in session_keys if key in rep_by_key] # mutable-ok: rows are enriched in place
|
||||
|
||||
|
||||
async def _ui_session_grouped_spend_logs(
|
||||
prisma_client: "PrismaClient",
|
||||
sql_conditions: Sequence[str],
|
||||
sql_params: Sequence[object],
|
||||
next_param_index: int,
|
||||
page: int,
|
||||
page_size: int,
|
||||
sort_desc: bool,
|
||||
session_cursor: str | None,
|
||||
) -> Mapping[str, object]:
|
||||
"""
|
||||
One row per session, keyset-paginated by session last activity.
|
||||
|
||||
Sessions are derived on the fly from ``LiteLLM_SpendLogs`` (no extra
|
||||
table): rows sharing a ``session_id`` and ``api_key`` form a session, rows
|
||||
without a session id are singletons keyed by ``request_id``. A page is the
|
||||
next ``page_size`` sessions ordered by ``(MAX(startTime), session_key,
|
||||
api_key)``, resumed from the ``session_cursor`` keyset
|
||||
``'<last_activity>|<api_key>|<session_key>'`` instead of an OFFSET, so
|
||||
page depth does not degrade the query plan. Each session is represented
|
||||
by its newest non-MCP row, enriched by ``_build_ui_spend_logs_response``
|
||||
exactly like the flat listing, and the response carries
|
||||
``next_session_cursor`` / ``has_more`` while ``total`` counts sessions
|
||||
(capped like the flat total).
|
||||
"""
|
||||
where_clause: Final = " AND ".join(sql_conditions) if sql_conditions else "TRUE"
|
||||
cmp_op: Final = "<" if sort_desc else ">"
|
||||
direction: Final = "DESC" if sort_desc else "ASC"
|
||||
|
||||
cursor: Final = _parse_session_cursor(session_cursor)
|
||||
having_clause: Final = (
|
||||
f'HAVING (MAX("startTime"), {_SESSION_GROUP_KEY_SQL}) {cmp_op} '
|
||||
f"(${next_param_index}::timestamp, ${next_param_index + 1}, ${next_param_index + 2})"
|
||||
if cursor
|
||||
else ""
|
||||
)
|
||||
cursor_params: Final[tuple[object, ...]] = cursor if cursor else ()
|
||||
limit_index: Final = next_param_index + len(cursor_params)
|
||||
|
||||
page_query: Final = f"""
|
||||
SELECT {_SESSION_KEY_EXPR} AS session_key,
|
||||
api_key,
|
||||
MAX("startTime")::text AS last_activity
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {where_clause}
|
||||
GROUP BY {_SESSION_GROUP_KEY_SQL}
|
||||
{having_clause}
|
||||
ORDER BY MAX("startTime") {direction}, {_SESSION_KEY_EXPR} {direction}, api_key {direction}
|
||||
LIMIT ${limit_index}
|
||||
"""
|
||||
page_rows: Final[Sequence[_SessionPageRow]] = await _query_raw(
|
||||
prisma_client, page_query, *sql_params, *cursor_params, page_size + 1
|
||||
)
|
||||
|
||||
has_more: Final = len(page_rows) > page_size
|
||||
visible_rows: Final = page_rows[:page_size]
|
||||
next_cursor: Final = (
|
||||
f"{visible_rows[-1]['last_activity']}|{visible_rows[-1]['api_key']}|{visible_rows[-1]['session_key']}"
|
||||
if has_more and visible_rows
|
||||
else None
|
||||
)
|
||||
|
||||
count_query: Final = f"""
|
||||
SELECT COUNT(*) AS total_count
|
||||
FROM (
|
||||
SELECT 1
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {where_clause}
|
||||
GROUP BY {_SESSION_GROUP_KEY_SQL}
|
||||
LIMIT ${next_param_index}
|
||||
) AS bounded_sessions
|
||||
"""
|
||||
count_rows: Final[Sequence[_SpendLogsCountRow]] = await _query_raw(
|
||||
prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1
|
||||
)
|
||||
raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0
|
||||
total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP
|
||||
total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total
|
||||
|
||||
session_keys: Final = tuple((row["session_key"], row["api_key"]) for row in visible_rows)
|
||||
data: Final[list[dict[str, object]]] = ( # mutable-ok: _build_ui_spend_logs_response writes onto each row
|
||||
await _fetch_session_representatives(
|
||||
prisma_client=prisma_client,
|
||||
where_clause=where_clause,
|
||||
sql_params=sql_params,
|
||||
next_param_index=next_param_index,
|
||||
session_keys=session_keys,
|
||||
)
|
||||
if session_keys
|
||||
else [] # mutable-ok: downstream enrichment mutates rows in place
|
||||
)
|
||||
_hydrate_spend_log_metadata(data)
|
||||
|
||||
total_pages: Final = (total_records + page_size - 1) // page_size
|
||||
response: Final[Mapping[str, object]] = await _build_ui_spend_logs_response(
|
||||
prisma_client,
|
||||
data,
|
||||
total_records,
|
||||
page,
|
||||
page_size,
|
||||
total_pages,
|
||||
enrich_session_counts=True,
|
||||
total_is_capped=total_is_capped,
|
||||
)
|
||||
return {**response, "next_session_cursor": next_cursor, "has_more": has_more} # mutable-ok: FastAPI response body
|
||||
|
||||
|
||||
class RequestResponsePayload(NamedTuple):
|
||||
messages: str | list | dict | None
|
||||
response: str | list | dict | None
|
||||
|
|
@ -4080,13 +4288,13 @@ async def _build_ui_spend_logs_response(
|
|||
total_pages: int,
|
||||
enrich_session_counts: bool = True,
|
||||
total_is_capped: bool = False,
|
||||
) -> dict:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Build the paginated response for the UI spend-logs endpoint.
|
||||
|
||||
When ``enrich_session_counts`` is ``True`` (the default for the v1/UI
|
||||
endpoint), each row is enriched with ``session_total_count`` plus spend
|
||||
and call-type aggregates so the frontend knows which sessions are
|
||||
endpoint), each row is enriched with ``session_total_count`` plus spend,
|
||||
token and call-type aggregates so the frontend knows which sessions are
|
||||
expandable (multi-call sessions). One ``GROUP BY (session_id, api_key)``
|
||||
query serves every referenced session, keyed per api key so two callers
|
||||
reusing a session id never see each other's totals. Rows without a
|
||||
|
|
@ -4121,7 +4329,7 @@ async def _build_ui_spend_logs_response(
|
|||
}
|
||||
)
|
||||
|
||||
session_spend_map: dict[tuple[str, str], dict[str, int | float]] = {}
|
||||
session_spend_map: _SessionSpendMap = {}
|
||||
if enrich_session_counts and session_ids:
|
||||
from prisma.errors import PrismaError
|
||||
|
||||
|
|
@ -4139,40 +4347,66 @@ async def _build_ui_spend_logs_response(
|
|||
rows: Final[Sequence[_SessionSpendRow]] = await _query_raw(
|
||||
prisma_client,
|
||||
f"""
|
||||
SELECT session_id, api_key,
|
||||
COUNT(*)::int AS session_total_count,
|
||||
COALESCE(SUM(spend), 0)::double precision AS session_total_spend,
|
||||
COUNT(*) FILTER (
|
||||
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
|
||||
)::int AS mcp_tool_call_count,
|
||||
COALESCE(SUM(spend) FILTER (
|
||||
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
|
||||
), 0)::double precision AS mcp_tool_call_spend,
|
||||
COUNT(*) FILTER (WHERE LOWER(cache_hit) = 'true')::int AS session_cache_hit_count,
|
||||
COUNT(*) FILTER (
|
||||
WHERE call_type NOT IN {_MCP_CALL_TYPES_SQL} AND call_type != {_AGENT_CALL_TYPE_SQL}
|
||||
)::int AS session_llm_count,
|
||||
COUNT(*) FILTER (WHERE call_type = {_AGENT_CALL_TYPE_SQL})::int AS session_agent_count
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE session_id = ANY($1::text[])
|
||||
AND api_key = ANY($2::text[])
|
||||
GROUP BY session_id, api_key
|
||||
SELECT s.*, COALESCE(m.session_models, ARRAY[]::text[]) AS session_models
|
||||
FROM (
|
||||
SELECT session_id, api_key,
|
||||
COUNT(*)::int AS session_total_count,
|
||||
COALESCE(SUM(spend), 0)::double precision AS session_total_spend,
|
||||
COUNT(*) FILTER (
|
||||
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
|
||||
)::int AS mcp_tool_call_count,
|
||||
COALESCE(SUM(spend) FILTER (
|
||||
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
|
||||
), 0)::double precision AS mcp_tool_call_spend,
|
||||
COUNT(*) FILTER (WHERE LOWER(cache_hit) = 'true')::int AS session_cache_hit_count,
|
||||
COUNT(*) FILTER (
|
||||
WHERE call_type NOT IN {_MCP_CALL_TYPES_SQL} AND call_type != {_AGENT_CALL_TYPE_SQL}
|
||||
)::int AS session_llm_count,
|
||||
COUNT(*) FILTER (WHERE call_type = {_AGENT_CALL_TYPE_SQL})::int AS session_agent_count,
|
||||
COALESCE(SUM(prompt_tokens), 0)::bigint AS session_total_prompt_tokens,
|
||||
COALESCE(SUM(completion_tokens), 0)::bigint AS session_total_completion_tokens,
|
||||
COALESCE(SUM(total_tokens), 0)::bigint AS session_total_tokens
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE session_id = ANY($1::text[])
|
||||
AND api_key = ANY($2::text[])
|
||||
GROUP BY session_id, api_key
|
||||
) s
|
||||
LEFT JOIN LATERAL (
|
||||
SELECT ARRAY_AGG(d.model ORDER BY d.model) AS session_models
|
||||
FROM (
|
||||
SELECT DISTINCT LEFT(model, $3::int) AS model
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE session_id = s.session_id
|
||||
AND api_key = s.api_key
|
||||
AND model IS NOT NULL AND model <> ''
|
||||
ORDER BY 1
|
||||
LIMIT $4::int
|
||||
) d
|
||||
) m ON TRUE
|
||||
""",
|
||||
session_ids,
|
||||
authorized_api_keys,
|
||||
_SESSION_MODEL_NAME_MAX_LEN,
|
||||
_SESSION_MODELS_LIMIT + 1,
|
||||
)
|
||||
session_spend_map = {
|
||||
(row["session_id"], row["api_key"]): {
|
||||
"session_total_count": int(row.get("session_total_count") or 0),
|
||||
"session_total_spend": float(row.get("session_total_spend") or 0.0),
|
||||
"mcp_tool_call_count": int(row.get("mcp_tool_call_count") or 0),
|
||||
"mcp_tool_call_spend": float(row.get("mcp_tool_call_spend") or 0.0),
|
||||
"session_cache_hit_count": int(row.get("session_cache_hit_count") or 0),
|
||||
"session_llm_count": int(row.get("session_llm_count") or 0),
|
||||
"session_agent_count": int(row.get("session_agent_count") or 0),
|
||||
}
|
||||
(row["session_id"], row["api_key"]): _SessionSpendStats(
|
||||
session_total_count=int(row.get("session_total_count") or 0),
|
||||
session_total_spend=float(row.get("session_total_spend") or 0.0),
|
||||
mcp_tool_call_count=int(row.get("mcp_tool_call_count") or 0),
|
||||
mcp_tool_call_spend=float(row.get("mcp_tool_call_spend") or 0.0),
|
||||
session_cache_hit_count=int(row.get("session_cache_hit_count") or 0),
|
||||
session_llm_count=int(row.get("session_llm_count") or 0),
|
||||
session_agent_count=int(row.get("session_agent_count") or 0),
|
||||
session_total_prompt_tokens=int(row.get("session_total_prompt_tokens") or 0),
|
||||
session_total_completion_tokens=int(row.get("session_total_completion_tokens") or 0),
|
||||
session_total_tokens=int(row.get("session_total_tokens") or 0),
|
||||
session_models=models[:_SESSION_MODELS_LIMIT],
|
||||
session_models_truncated=len(models) > _SESSION_MODELS_LIMIT,
|
||||
)
|
||||
for row in rows
|
||||
if row.get("session_id") and row.get("api_key") is not None
|
||||
for models in (TypeAdapter(list[str]).validate_python(row.get("session_models") or ()),)
|
||||
}
|
||||
except PrismaError:
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -4187,15 +4421,20 @@ async def _build_ui_spend_logs_response(
|
|||
sid = row_dict.get("session_id")
|
||||
row_api_key = row_dict.get("api_key")
|
||||
session_stats = session_spend_map.get((sid, row_api_key)) if sid and row_api_key is not None else None
|
||||
row_dict["session_total_count"] = int(session_stats["session_total_count"]) if session_stats else 1
|
||||
row_dict["session_total_count"] = session_stats.session_total_count if session_stats else 1
|
||||
if session_stats:
|
||||
row_dict["session_total_spend"] = session_stats["session_total_spend"]
|
||||
if session_stats["mcp_tool_call_count"]:
|
||||
row_dict["mcp_tool_call_count"] = session_stats["mcp_tool_call_count"]
|
||||
row_dict["mcp_tool_call_spend"] = session_stats["mcp_tool_call_spend"]
|
||||
row_dict["session_cache_hit_count"] = session_stats["session_cache_hit_count"]
|
||||
row_dict["session_llm_count"] = session_stats["session_llm_count"]
|
||||
row_dict["session_agent_count"] = session_stats["session_agent_count"]
|
||||
row_dict["session_total_spend"] = session_stats.session_total_spend
|
||||
if session_stats.mcp_tool_call_count:
|
||||
row_dict["mcp_tool_call_count"] = session_stats.mcp_tool_call_count
|
||||
row_dict["mcp_tool_call_spend"] = session_stats.mcp_tool_call_spend
|
||||
row_dict["session_cache_hit_count"] = session_stats.session_cache_hit_count
|
||||
row_dict["session_llm_count"] = session_stats.session_llm_count
|
||||
row_dict["session_agent_count"] = session_stats.session_agent_count
|
||||
row_dict["session_total_prompt_tokens"] = session_stats.session_total_prompt_tokens
|
||||
row_dict["session_total_completion_tokens"] = session_stats.session_total_completion_tokens
|
||||
row_dict["session_total_tokens"] = session_stats.session_total_tokens
|
||||
row_dict["session_models"] = session_stats.session_models
|
||||
row_dict["session_models_truncated"] = session_stats.session_models_truncated
|
||||
enriched.append(row_dict)
|
||||
response_data: list = enriched
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.constants import (
|
|||
LITELLM_TRUNCATED_PAYLOAD_FIELD,
|
||||
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE,
|
||||
REDACTED_BY_LITELM_STRING,
|
||||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
)
|
||||
from litellm.constants import (
|
||||
MAX_STRING_LENGTH_PROMPT_IN_DB as DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB,
|
||||
|
|
@ -578,7 +579,9 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
|
|||
),
|
||||
session_id=_get_session_id_for_spend_log(
|
||||
kwargs=kwargs,
|
||||
metadata=metadata,
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
omit_when_missing=_omits_session_id_when_missing(metadata),
|
||||
),
|
||||
request_duration_ms=_get_request_duration_ms(start_time, end_time),
|
||||
status=_get_status_for_spend_log(
|
||||
|
|
@ -602,26 +605,39 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
|
|||
raise e
|
||||
|
||||
|
||||
def _omits_session_id_when_missing(metadata: Mapping[str, object] | None) -> bool:
|
||||
"""The pre-call stamp pins `omit` on for the requests that carry it, so a config reload between pre-call and spend
|
||||
logging cannot fabricate a session. `apply_missing_session_id_policy` drops any client-supplied copy of the key
|
||||
from both metadata buckets before stamping, which the merge of `litellm_metadata` into `metadata` makes
|
||||
necessary, so a caller cannot forge it. Requests that never reach the pre-call helper, router-model
|
||||
passthrough among them, carry no stamp, so they fall back to the configured policy and `omit` still covers their
|
||||
spend logs."""
|
||||
if metadata is not None and metadata.get(SESSION_ID_OMITTED_METADATA_KEY):
|
||||
return True
|
||||
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
return general_settings.get("missing_session_id") == "omit"
|
||||
|
||||
|
||||
def _get_session_id_for_spend_log(
|
||||
kwargs: dict,
|
||||
kwargs: Mapping[str, object],
|
||||
metadata: Mapping[str, object] | None,
|
||||
standard_logging_payload: StandardLoggingPayload | None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the session id for the spend log.
|
||||
omit_when_missing: bool,
|
||||
) -> str | None:
|
||||
"""Under `omit` only `metadata.session_id`, the key Langfuse reads, counts as a session; `litellm_session_id` may
|
||||
be a copied trace id."""
|
||||
if omit_when_missing:
|
||||
session_id: Final = metadata.get("session_id") if metadata else None
|
||||
return str(session_id) if session_id else None
|
||||
|
||||
This ensures each spend log is associated with a unique session id.
|
||||
|
||||
"""
|
||||
from litellm._uuid import uuid
|
||||
|
||||
if standard_logging_payload is not None and standard_logging_payload.get("trace_id") is not None:
|
||||
return str(standard_logging_payload.get("trace_id"))
|
||||
|
||||
# Users can dynamically set the trace_id for each request by passing `litellm_trace_id` in kwargs
|
||||
if kwargs.get("litellm_trace_id") is not None:
|
||||
return str(kwargs.get("litellm_trace_id"))
|
||||
|
||||
# Ensure we always have a session id, if none is provided
|
||||
return str(uuid.uuid4())
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1478,28 +1478,27 @@ class ProxyLogging:
|
|||
) -> None:
|
||||
"""Process prompt template if applicable."""
|
||||
|
||||
from litellm.proxy.prompts.prompt_endpoints import (
|
||||
construct_versioned_prompt_id,
|
||||
get_latest_version_prompt_id,
|
||||
)
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.utils import get_non_default_completion_params
|
||||
|
||||
if prompt_version is None:
|
||||
lookup_prompt_id = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
else:
|
||||
lookup_prompt_id = construct_versioned_prompt_id(prompt_id=prompt_id, version=prompt_version)
|
||||
|
||||
custom_logger: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(lookup_prompt_id)
|
||||
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id)
|
||||
raw_prompt_environment: Final = data.get("prompt_environment", None)
|
||||
prompt_environment: Final = raw_prompt_environment if isinstance(raw_prompt_environment, str) else None
|
||||
prompt_spec: Final = IN_MEMORY_PROMPT_REGISTRY.resolve_prompt_spec(
|
||||
prompt_id,
|
||||
version=prompt_version,
|
||||
environment=prompt_environment,
|
||||
)
|
||||
custom_logger: Final = (
|
||||
IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_for_prompt(prompt=prompt_spec)
|
||||
if prompt_spec is not None
|
||||
else None
|
||||
)
|
||||
litellm_prompt_id: str | None = None
|
||||
if prompt_spec is not None:
|
||||
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
|
||||
data.pop("prompt_id", None)
|
||||
data.pop("prompt_environment", None)
|
||||
|
||||
if custom_logger and prompt_spec is not None:
|
||||
is_responses_call: Final = call_type == "aresponses"
|
||||
|
|
@ -1542,6 +1541,7 @@ class ProxyLogging:
|
|||
data.pop("prompt_variables", None)
|
||||
data.pop("prompt_label", None)
|
||||
data.pop("prompt_version", None)
|
||||
data.pop("prompt_environment", None)
|
||||
|
||||
def _process_guardrail_metadata(self, data: dict) -> None:
|
||||
"""Process guardrails from metadata and add to applied_guardrails."""
|
||||
|
|
@ -1750,7 +1750,6 @@ class ProxyLogging:
|
|||
|
||||
litellm_logging_obj: Final = cast(Optional["LiteLLMLoggingObj"], data.get("litellm_logging_obj", None))
|
||||
prompt_id: Final[str | None] = data.get("prompt_id", None)
|
||||
prompt_version: Final[int | None] = data.get("prompt_version", None)
|
||||
|
||||
## PROMPT TEMPLATE CHECK ##
|
||||
|
||||
|
|
@ -1760,11 +1759,13 @@ class ProxyLogging:
|
|||
and prompt_id is not None
|
||||
and (call_type == "completion" or call_type == "acompletion" or call_type == "aresponses")
|
||||
):
|
||||
from litellm.proxy.prompts.prompt_registry import parse_prompt_version
|
||||
|
||||
await self._process_prompt_template(
|
||||
data=data,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
prompt_id=prompt_id,
|
||||
prompt_version=prompt_version,
|
||||
prompt_version=parse_prompt_version(data.get("prompt_version", None)),
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
|
|
@ -7531,6 +7532,9 @@ def create_model_info_response(
|
|||
max_input_tokens = configured_input
|
||||
if configured_output is not None:
|
||||
max_output_tokens = configured_output
|
||||
configured_mode: Final = llm_router.get_configured_mode(model_id)
|
||||
if isinstance(configured_mode, str):
|
||||
base["mode"] = configured_mode
|
||||
|
||||
if max_input_tokens is not None:
|
||||
base["max_input_tokens"] = max_input_tokens
|
||||
|
|
|
|||
|
|
@ -30,7 +30,10 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
|
||||
from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
can_user_access_vector_store,
|
||||
filter_listable_vector_stores,
|
||||
)
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -390,11 +393,10 @@ async def list_vector_stores(
|
|||
|
||||
# Filter vector stores based on access control
|
||||
accessible_vector_stores: Final = []
|
||||
for vs in vector_store_map.values():
|
||||
if await _check_vector_store_access(vs, user_api_key_dict):
|
||||
redacted = LiteLLM_ManagedVectorStore(**vs)
|
||||
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
|
||||
accessible_vector_stores.append(redacted)
|
||||
for vs in await filter_listable_vector_stores(vector_store_map.values(), user_api_key_dict):
|
||||
redacted = LiteLLM_ManagedVectorStore(**vs)
|
||||
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
|
||||
accessible_vector_stores.append(redacted)
|
||||
|
||||
total_count: Final = len(accessible_vector_stores)
|
||||
total_pages: Final = (total_count + page_size - 1) // page_size
|
||||
|
|
|
|||
|
|
@ -1,11 +1,17 @@
|
|||
import json
|
||||
import re
|
||||
from collections.abc import Iterable
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
is_ui_session_credential,
|
||||
resolve_ui_session_team_ids,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LitellmUserRoles,
|
||||
|
|
@ -160,10 +166,16 @@ async def can_user_access_vector_store(
|
|||
if _is_proxy_admin(user_api_key_dict):
|
||||
return True
|
||||
|
||||
vector_store_team_id: Final = vector_store.get("team_id")
|
||||
if vector_store_team_id is None:
|
||||
if vector_store.get("team_id") is None:
|
||||
return True
|
||||
|
||||
return await _is_vector_store_granted(vector_store, user_api_key_dict)
|
||||
|
||||
|
||||
async def _is_vector_store_granted(
|
||||
vector_store: LiteLLM_ManagedVectorStore,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> bool:
|
||||
vector_store_id: Final = vector_store.get("vector_store_id") or ""
|
||||
|
||||
key_object_permission = user_api_key_dict.object_permission
|
||||
|
|
@ -178,12 +190,70 @@ async def can_user_access_vector_store(
|
|||
if _object_permission_allows_vector_store(team_object_permission, vector_store_id):
|
||||
return True
|
||||
|
||||
if user_api_key_dict.team_id is not None and user_api_key_dict.team_id == vector_store_team_id:
|
||||
return True
|
||||
return user_api_key_dict.team_id is not None and user_api_key_dict.team_id == vector_store.get("team_id")
|
||||
|
||||
|
||||
async def _team_auth_context(team_id: str, user_api_key_dict: UserAPIKeyAuth) -> UserAPIKeyAuth:
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
team: Final = await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
return user_api_key_dict.model_copy(
|
||||
update=MappingProxyType(
|
||||
{
|
||||
"team_id": team_id,
|
||||
"team_object_permission": team.object_permission,
|
||||
"team_object_permission_id": team.object_permission_id,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _vector_store_listing_auth_contexts(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[UserAPIKeyAuth, ...]:
|
||||
if not is_ui_session_credential(user_api_key_dict):
|
||||
return (user_api_key_dict,)
|
||||
session_key_context: Final = user_api_key_dict.model_copy(
|
||||
update=MappingProxyType({"team_id": None, "team_object_permission": None, "team_object_permission_id": None})
|
||||
)
|
||||
team_ids: Final = await resolve_ui_session_team_ids(user_api_key_dict)
|
||||
team_contexts: Final = tuple([await _team_auth_context(team_id, user_api_key_dict) for team_id in team_ids])
|
||||
return (session_key_context, *team_contexts)
|
||||
|
||||
|
||||
async def _is_vector_store_granted_to_any(
|
||||
vector_store: LiteLLM_ManagedVectorStore,
|
||||
auth_contexts: tuple[UserAPIKeyAuth, ...],
|
||||
) -> bool:
|
||||
for auth_context in auth_contexts:
|
||||
if await _is_vector_store_granted(vector_store, auth_context):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
async def filter_listable_vector_stores(
|
||||
vector_stores: Iterable[LiteLLM_ManagedVectorStore],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[LiteLLM_ManagedVectorStore, ...]:
|
||||
"""Non-admins only see stores their key, one of their teams' object_permission, or team ownership grants."""
|
||||
if _is_proxy_admin(user_api_key_dict):
|
||||
return tuple(vector_stores)
|
||||
|
||||
auth_contexts: Final = await _vector_store_listing_auth_contexts(user_api_key_dict)
|
||||
return tuple([vs for vs in vector_stores if await _is_vector_store_granted_to_any(vs, auth_contexts)])
|
||||
|
||||
|
||||
async def get_litellm_managed_vector_store(
|
||||
vector_store_id: str,
|
||||
) -> LiteLLM_ManagedVectorStore | None:
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from typing import Any, Final, Literal, cast
|
|||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
AZURE_OPENAI_AUDIO_PROVIDERS,
|
||||
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
||||
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
request_timeout,
|
||||
|
|
@ -400,7 +401,7 @@ async def _arealtime(
|
|||
litellm_metadata=_build_litellm_metadata(kwargs),
|
||||
query_params=query_params,
|
||||
)
|
||||
elif _custom_llm_provider == "azure":
|
||||
elif _custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS:
|
||||
api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
# set API KEY
|
||||
api_key = dynamic_api_key or litellm.api_key or litellm.openai_key or get_secret_str("AZURE_API_KEY")
|
||||
|
|
|
|||
|
|
@ -562,6 +562,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
hidden_params: Final = getattr(chunk, "_hidden_params", None)
|
||||
if hidden_params is not None:
|
||||
chunk_dict["_hidden_params"] = dict(hidden_params) if isinstance(hidden_params, dict) else hidden_params
|
||||
if (
|
||||
chunk_dict.get("usage") is None
|
||||
and isinstance(hidden_params, dict)
|
||||
and hidden_params.get("usage") is not None
|
||||
):
|
||||
chunk_dict["usage"] = hidden_params["usage"]
|
||||
return chunk_dict
|
||||
|
||||
def create_reasoning_summary_text_done_event(
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ import time
|
|||
import traceback
|
||||
import weakref
|
||||
from collections import defaultdict
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, Sequence
|
||||
from functools import lru_cache, partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
|
||||
|
|
@ -50,6 +50,7 @@ from litellm.constants import (
|
|||
DEFAULT_HEALTH_CHECK_INTERVAL,
|
||||
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
|
||||
DEFAULT_MAX_LRU_CACHE_SIZE,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
)
|
||||
|
|
@ -116,6 +117,9 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
from litellm.router_utils.auto_router_model_naming import (
|
||||
AUTO_ROUTER_MODEL_PREFIX,
|
||||
classify_strategy_router_model,
|
||||
count_heuristic_v2_routers,
|
||||
heuristic_v2_limit_violation,
|
||||
uses_heuristic_v2_classifier,
|
||||
)
|
||||
from litellm.router_utils.batch_utils import (
|
||||
_get_router_metadata_variable_name,
|
||||
|
|
@ -210,6 +214,7 @@ from litellm.types.router import (
|
|||
DeploymentTypedDict,
|
||||
FallbackAccessCheck,
|
||||
GuardrailTypedDict,
|
||||
HeuristicV2RouterLimit,
|
||||
LiteLLM_Params,
|
||||
MockRouterTestingParams,
|
||||
ModelGroupInfo,
|
||||
|
|
@ -231,6 +236,7 @@ from litellm.types.router import (
|
|||
)
|
||||
from litellm.types.services import ServiceTypes
|
||||
from litellm.types.utils import (
|
||||
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
|
||||
PROMPT_QUOTING_ROUTING_DECISION_FIELDS,
|
||||
CustomPricingLiteLLMParams,
|
||||
GenericBudgetConfigType,
|
||||
|
|
@ -681,6 +687,7 @@ class Router:
|
|||
background_health_check_model_groups: Sequence[str] | None = None,
|
||||
enable_weighted_failover: bool = False,
|
||||
fallback_access_check: FallbackAccessCheck | None = None,
|
||||
heuristic_v2_router_limit: HeuristicV2RouterLimit | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
|
||||
|
|
@ -757,6 +764,7 @@ class Router:
|
|||
|
||||
self.set_verbose = set_verbose
|
||||
self.ignore_invalid_deployments = ignore_invalid_deployments
|
||||
self.heuristic_v2_router_limit = heuristic_v2_router_limit
|
||||
self.fallback_access_check: Final = fallback_access_check
|
||||
self.debug_level = debug_level
|
||||
self.enable_pre_call_checks = enable_pre_call_checks
|
||||
|
|
@ -2168,6 +2176,76 @@ class Router:
|
|||
verbose_router_logger.debug("Error occurred while printing deployment - %s", e)
|
||||
raise e
|
||||
|
||||
@staticmethod
|
||||
def _deployment_params_with_request_reasoning_override(
|
||||
deployment_params: Mapping[str, object], request_kwargs: Mapping[str, object]
|
||||
) -> dict[str, object]: # mutable-ok: litellm's request pipeline consumes a mutable kwargs mapping
|
||||
"""Return deployment params whose equivalent effort controls cannot outrank a request override.
|
||||
|
||||
Providers expose the same setting through several native carriers. A request-level
|
||||
``reasoning_effort`` is the portable override, so a deployment's ``thinking`` or nested
|
||||
``*.effort`` must not remain beside it and either win or trigger a conflicting-params 400.
|
||||
Every changed mapping is copied so the Router's shared deployment config stays immutable.
|
||||
"""
|
||||
sanitized: Final = dict(deployment_params) # mutable-ok: request-local copy protects shared Router state
|
||||
if request_kwargs.get("reasoning_effort") is None:
|
||||
return sanitized
|
||||
|
||||
sanitized.pop("thinking", None)
|
||||
Router._pop_effort_from_nested_carrier(sanitized, "output_config")
|
||||
Router._pop_effort_from_nested_carrier(sanitized, "reasoning")
|
||||
|
||||
extra_body: Final = sanitized.get("extra_body")
|
||||
if isinstance(extra_body, Mapping):
|
||||
sanitized_extra_body: Final = dict(extra_body) # mutable-ok: request-local nested copy
|
||||
sanitized_extra_body.pop("reasoning_effort", None)
|
||||
sanitized_extra_body.pop("thinking", None)
|
||||
Router._pop_effort_from_nested_carrier(sanitized_extra_body, "output_config")
|
||||
Router._pop_effort_from_nested_carrier(sanitized_extra_body, "reasoning")
|
||||
if sanitized_extra_body:
|
||||
sanitized["extra_body"] = sanitized_extra_body
|
||||
else:
|
||||
sanitized.pop("extra_body", None)
|
||||
return sanitized
|
||||
|
||||
@staticmethod
|
||||
def _is_classifier_internal_call(kwargs: Mapping[str, object]) -> bool:
|
||||
metadata: Final = kwargs.get("metadata")
|
||||
litellm_metadata: Final = kwargs.get("litellm_metadata")
|
||||
return any(
|
||||
isinstance(candidate, Mapping)
|
||||
and candidate.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
||||
for candidate in (metadata, litellm_metadata)
|
||||
)
|
||||
|
||||
def _drop_unsupported_classifier_reasoning_effort(
|
||||
self,
|
||||
deployment: DeploymentTypedDict,
|
||||
model: str,
|
||||
kwargs: dict[str, object], # mutable-ok: fallback must update the active request and its log body together
|
||||
) -> None:
|
||||
"""Let a classifier fallback without reasoning support remain a usable fallback.
|
||||
|
||||
The dashboard only offers explicitly advertised levels, but an existing config can outlive
|
||||
a model change and fallbacks can target a different group. Unknown capability fails open;
|
||||
only a provider that explicitly rejects the parameter has it removed.
|
||||
"""
|
||||
if kwargs.get("reasoning_effort") is None or not self._is_classifier_internal_call(kwargs):
|
||||
return
|
||||
if self._deployment_accepts_param(deployment, model, "reasoning_effort"):
|
||||
return
|
||||
verbose_router_logger.warning(
|
||||
"litellm.router.py: dropping classifier reasoning_effort for model=%s because the selected deployment does not support it",
|
||||
model,
|
||||
)
|
||||
kwargs.pop("reasoning_effort", None)
|
||||
proxy_server_request: Final = kwargs.get("proxy_server_request")
|
||||
if not isinstance(proxy_server_request, dict):
|
||||
return
|
||||
body: Final = proxy_server_request.get("body")
|
||||
if isinstance(body, dict):
|
||||
body.pop("reasoning_effort", None)
|
||||
|
||||
### COMPLETION, EMBEDDING, IMG GENERATION FUNCTIONS
|
||||
|
||||
def completion(self, model: str, messages: list[dict[str, str]], **kwargs) -> ModelResponse | CustomStreamWrapper:
|
||||
|
|
@ -2203,9 +2281,16 @@ class Router:
|
|||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
self._drop_unsupported_classifier_reasoning_effort(
|
||||
deployment=cast(DeploymentTypedDict, deployment), # cast-ok: selection returns a router deployment
|
||||
model=model,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
# Check for silent model experiment
|
||||
# Make a local copy of litellm_params to avoid mutating the Router's state
|
||||
litellm_params: Final = deployment["litellm_params"].copy()
|
||||
litellm_params: Final = self._deployment_params_with_request_reasoning_override(
|
||||
deployment["litellm_params"], kwargs
|
||||
)
|
||||
silent_model: Final = litellm_params.pop("silent_model", None)
|
||||
|
||||
if silent_model is not None:
|
||||
|
|
@ -3216,6 +3301,11 @@ class Router:
|
|||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
self._drop_unsupported_classifier_reasoning_effort(
|
||||
deployment=cast(DeploymentTypedDict, deployment), # cast-ok: selection returns a router deployment
|
||||
model=model,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
_timeout_debug_deployment_dict = deployment
|
||||
end_time: Final = time.time()
|
||||
|
|
@ -3237,7 +3327,9 @@ class Router:
|
|||
|
||||
# Check for silent model experiment
|
||||
# Make a local copy of litellm_params to avoid mutating the Router's state
|
||||
litellm_params: Final = deployment["litellm_params"].copy()
|
||||
litellm_params: Final = self._deployment_params_with_request_reasoning_override(
|
||||
deployment["litellm_params"], kwargs
|
||||
)
|
||||
silent_model: Final = litellm_params.pop("silent_model", None)
|
||||
|
||||
if silent_model is not None:
|
||||
|
|
@ -8710,6 +8802,30 @@ class Router:
|
|||
"""
|
||||
return classify_strategy_router_model(litellm_params.model) == "complexity"
|
||||
|
||||
def config_deployments(self) -> Iterator[Mapping[str, object]]:
|
||||
"""The model_list rows that came from config.yaml rather than the DB (``model_info.db_model`` unset)."""
|
||||
for deployment in self.model_list:
|
||||
if not isinstance(deployment, Mapping):
|
||||
continue
|
||||
model_info = deployment.get("model_info")
|
||||
if not (isinstance(model_info, Mapping) and model_info.get("db_model")):
|
||||
yield deployment
|
||||
|
||||
def heuristic_v2_router_limit_violation(self) -> str | None:
|
||||
"""
|
||||
Why one more heuristic_v2 router cannot join this router, or None when it can.
|
||||
|
||||
Judged against every deployment currently on the model_list; an upsert pops the row being
|
||||
edited first, so an edit of an existing heuristic_v2 router keeps its own slot. The limit is
|
||||
resolved on every call through ``heuristic_v2_router_limit``; unset means unlimited, which
|
||||
is the SDK default, and the proxy injects a resolver backed by its license.
|
||||
"""
|
||||
limit: Final = self.heuristic_v2_router_limit() if self.heuristic_v2_router_limit is not None else None
|
||||
others: Final = count_heuristic_v2_routers(
|
||||
deployment for deployment in self.model_list if isinstance(deployment, Mapping)
|
||||
)
|
||||
return heuristic_v2_limit_violation(held=others + 1, limit=limit)
|
||||
|
||||
def init_complexity_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
Initialize the complexity-router deployment.
|
||||
|
|
@ -8727,6 +8843,10 @@ class Router:
|
|||
)
|
||||
|
||||
complexity_router_config: Final[dict | None] = deployment.litellm_params.complexity_router_config
|
||||
if uses_heuristic_v2_classifier(complexity_router_config):
|
||||
limit_violation: Final = self.heuristic_v2_router_limit_violation()
|
||||
if limit_violation is not None:
|
||||
raise ValueError(limit_violation)
|
||||
|
||||
default_model: str | None = deployment.litellm_params.complexity_router_default_model
|
||||
|
||||
|
|
@ -9550,8 +9670,16 @@ class Router:
|
|||
raise e
|
||||
|
||||
def _restore_deployment_after_failed_upsert(self, previous_deployment: Deployment | None, model_id: str) -> None:
|
||||
"""Put a deployment back the way it was before a failed upsert popped it.
|
||||
|
||||
A rollback re-admits state that was already serving, so it does not go through the
|
||||
heuristic_v2 ceiling a newcomer gets: with the ceiling tightened since the deployment first
|
||||
registered, judging the rollback would drop a serving router over an unrelated failed edit.
|
||||
"""
|
||||
if previous_deployment is None or self.has_model_id(model_id):
|
||||
return
|
||||
limit_resolver: Final = self.heuristic_v2_router_limit
|
||||
self.heuristic_v2_router_limit = None
|
||||
try:
|
||||
self.add_deployment(deployment=previous_deployment)
|
||||
verbose_router_logger.info(
|
||||
|
|
@ -9566,6 +9694,8 @@ class Router:
|
|||
model_id,
|
||||
restore_error,
|
||||
)
|
||||
finally:
|
||||
self.heuristic_v2_router_limit = limit_resolver
|
||||
|
||||
@staticmethod
|
||||
def _backend_cost_map_keys(model: str, custom_llm_provider: str | None) -> tuple[str, ...]:
|
||||
|
|
@ -9895,6 +10025,17 @@ class Router:
|
|||
coerce_token_limit(model_info.get("max_output_tokens")),
|
||||
)
|
||||
|
||||
def get_configured_mode(self, model_name: str) -> "str | None":
|
||||
"""Return the mode explicitly configured for a concrete deployment."""
|
||||
deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name)
|
||||
if deployment is None:
|
||||
return None
|
||||
|
||||
mode: Final = deployment.model_info.get("mode")
|
||||
if isinstance(mode, str) and mode.strip():
|
||||
return mode
|
||||
return None
|
||||
|
||||
def get_configured_display_name(self, model_name: str) -> "str | None":
|
||||
"""
|
||||
Return the display_name explicitly configured in a concrete deployment's
|
||||
|
|
@ -10255,6 +10396,8 @@ class Router:
|
|||
total_itpm: int | None = None
|
||||
total_otpm: int | None = None
|
||||
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
|
||||
reasoning_efforts_initialized = False
|
||||
reasoning_efforts_unknown = False
|
||||
model_list: Final = self.get_model_list(model_name=model_group)
|
||||
if model_list is None:
|
||||
return None
|
||||
|
|
@ -10441,10 +10584,23 @@ class Router:
|
|||
if model_info.get("rpm", None) is not None and _deployment_rpm is None:
|
||||
_deployment_rpm = model_info.get("rpm")
|
||||
|
||||
model_group_info.supported_reasoning_efforts = intersect_supported_reasoning_efforts(
|
||||
model_group_info.supported_reasoning_efforts,
|
||||
resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=deployment_is_mapped),
|
||||
deployment_reasoning_efforts = (
|
||||
resolve_supported_reasoning_efforts( # rebind-ok: recalculated per deployment
|
||||
model_info, deployment_is_mapped=deployment_is_mapped
|
||||
)
|
||||
)
|
||||
if deployment_reasoning_efforts is None:
|
||||
reasoning_efforts_unknown = True
|
||||
model_group_info.supported_reasoning_efforts = None
|
||||
elif not reasoning_efforts_initialized:
|
||||
reasoning_efforts_initialized = True
|
||||
if not reasoning_efforts_unknown:
|
||||
model_group_info.supported_reasoning_efforts = deployment_reasoning_efforts
|
||||
elif not reasoning_efforts_unknown:
|
||||
model_group_info.supported_reasoning_efforts = intersect_supported_reasoning_efforts(
|
||||
model_group_info.supported_reasoning_efforts,
|
||||
deployment_reasoning_efforts,
|
||||
)
|
||||
|
||||
if _deployment_tpm is not None:
|
||||
if total_tpm is None:
|
||||
|
|
@ -12308,10 +12464,12 @@ class Router:
|
|||
@staticmethod
|
||||
def _pop_effort_from_nested_carrier(request_kwargs: dict[str, object], carrier: str) -> None:
|
||||
nested: Final = request_kwargs.get(carrier)
|
||||
if not isinstance(nested, dict):
|
||||
if not isinstance(nested, Mapping):
|
||||
return
|
||||
nested.pop("effort", None)
|
||||
if not nested:
|
||||
sanitized: Final = {key: value for key, value in nested.items() if key != "effort"}
|
||||
if sanitized:
|
||||
request_kwargs[carrier] = sanitized # rebind-ok: copy-on-write, so a shared nested carrier is never edited
|
||||
else:
|
||||
request_kwargs.pop(carrier, None)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -2,23 +2,41 @@
|
|||
Auto-Routing Strategy that works with a Semantic Router Config
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.internal_call_metadata import (
|
||||
effective_turn_off_message_logging,
|
||||
forwarded_internal_call_metadata,
|
||||
parent_session_kwargs,
|
||||
)
|
||||
from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from semantic_router.routers import SemanticRouter
|
||||
from semantic_router.routers.base import Route
|
||||
|
||||
from litellm.router import Router
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
else:
|
||||
Router = Any
|
||||
PreRoutingHookResponse = Any
|
||||
Route = Any
|
||||
SemanticRouter = Any
|
||||
LiteLLMRouterEncoder = Any
|
||||
|
||||
|
||||
class _CallerMetadata(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
metadata: Mapping[str, object] | None = None
|
||||
litellm_metadata: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
class AutoRouter(CustomLogger):
|
||||
|
|
@ -50,6 +68,8 @@ class AutoRouter(CustomLogger):
|
|||
"""
|
||||
from semantic_router.routers import SemanticRouter
|
||||
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import LiteLLMRouterEncoder
|
||||
|
||||
self.auto_router_config_path: str | None = auto_router_config_path
|
||||
self.auto_router_config: str | None = auto_router_config
|
||||
self.auto_sync_value = self.DEFAULT_AUTO_SYNC_VALUE
|
||||
|
|
@ -59,6 +79,11 @@ class AutoRouter(CustomLogger):
|
|||
self.embedding_model: str = embedding_model
|
||||
self.max_input_chars: int = max_input_chars
|
||||
self.litellm_router_instance: Router = litellm_router_instance
|
||||
self.encoder: LiteLLMRouterEncoder = LiteLLMRouterEncoder(
|
||||
litellm_router_instance=litellm_router_instance,
|
||||
model_name=embedding_model,
|
||||
max_input_chars=max_input_chars,
|
||||
)
|
||||
|
||||
def _load_semantic_routing_routes(self) -> list[Route]:
|
||||
from semantic_router.routers import SemanticRouter
|
||||
|
|
@ -129,9 +154,6 @@ class AutoRouter(CustomLogger):
|
|||
from semantic_router.routers import SemanticRouter
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
|
||||
from litellm.router_strategy.auto_router.litellm_encoder import (
|
||||
LiteLLMRouterEncoder,
|
||||
)
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
resolved_messages: Final = (
|
||||
|
|
@ -149,34 +171,47 @@ class AutoRouter(CustomLogger):
|
|||
#######################
|
||||
routelayer = SemanticRouter(
|
||||
routes=self.loaded_routes,
|
||||
encoder=LiteLLMRouterEncoder(
|
||||
litellm_router_instance=self.litellm_router_instance,
|
||||
model_name=self.embedding_model,
|
||||
max_input_chars=self.max_input_chars,
|
||||
),
|
||||
encoder=self.encoder,
|
||||
auto_sync=self.auto_sync_value,
|
||||
)
|
||||
self.routelayer = routelayer
|
||||
|
||||
message_content: Final = self._extract_text_from_messages(resolved_messages)
|
||||
route_name: Final = self._matched_route_name(routelayer, message_content)
|
||||
route_name: Final = await self._matched_route_name(routelayer, message_content, request_kwargs)
|
||||
|
||||
return PreRoutingHookResponse(
|
||||
model=route_name or self.default_model,
|
||||
messages=messages,
|
||||
)
|
||||
|
||||
def _matched_route_name(self, routelayer: "SemanticRouter", text: str) -> str | None:
|
||||
async def _matched_route_name(
|
||||
self, routelayer: "SemanticRouter", text: str, request_kwargs: Mapping[str, object]
|
||||
) -> str | None:
|
||||
"""Name of the route `text` matches, or None when nothing matched or the match failed.
|
||||
|
||||
The route layer embeds `text` to compare it against the routes, and that embedding call can
|
||||
`text` is embedded here rather than by `routelayer(text=...)` so the caller's metadata reaches
|
||||
`aembedding()` and the embedding's spend lands on the key/team that sent the request;
|
||||
SemanticRouter has no way to pass kwargs through to its encoder. That embedding call can
|
||||
fail (context limit, timeout, provider error). Choosing a model is a routing decision, so a
|
||||
failure here falls back to the default model rather than failing the user's request.
|
||||
"""
|
||||
from semantic_router.schema import RouteChoice
|
||||
|
||||
try:
|
||||
route_choice: Final = routelayer(text=text)
|
||||
caller: Final = _CallerMetadata.model_validate(request_kwargs)
|
||||
query_vector: Final = (
|
||||
await self.encoder.aencode_queries(
|
||||
[text],
|
||||
metadata=forwarded_internal_call_metadata(caller.metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN),
|
||||
litellm_metadata=forwarded_internal_call_metadata(
|
||||
caller.litellm_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN
|
||||
),
|
||||
proxy_server_request={"body": {"model": self.embedding_model, "input": [text]}},
|
||||
turn_off_message_logging=effective_turn_off_message_logging(request_kwargs),
|
||||
**parent_session_kwargs(request_kwargs),
|
||||
)
|
||||
)[0]
|
||||
route_choice: Final = await routelayer.acall(vector=query_vector)
|
||||
except Exception as e: # noqa: BLE001 -- the embedding call behind the route layer can fail many ways (context limit, timeout, provider/network error); none of them may fail the request
|
||||
verbose_router_logger.warning(
|
||||
"AutoRouter: semantic routing failed (%s), falling back to default model %s", e, self.default_model
|
||||
|
|
|
|||
|
|
@ -255,7 +255,8 @@ model_list:
|
|||
classifier_type: heuristic_first
|
||||
heuristic_first_max_tier: SIMPLE
|
||||
classifier_llm_config:
|
||||
model: gpt-4o-mini
|
||||
model: gpt-5-mini
|
||||
reasoning_effort: low
|
||||
tiers:
|
||||
SIMPLE: gpt-4o-mini
|
||||
MEDIUM: gpt-4o
|
||||
|
|
@ -263,6 +264,10 @@ model_list:
|
|||
REASONING: o1-preview
|
||||
```
|
||||
|
||||
`classifier_llm_config.reasoning_effort` applies only to the internal classifier call. Omit it to
|
||||
keep the classifier deployment or provider default, or set a supported value such as `none` or
|
||||
`low` to override that call.
|
||||
|
||||
A request short-circuits, meaning it routes on the scorer's own tier with no classifier call, when
|
||||
two things hold: the scorer landed at or below `heuristic_first_max_tier`, and it produced at least
|
||||
one signal. Everything else goes to the classifier, which then decides as it normally would.
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from pydantic import BaseModel, create_model
|
|||
from litellm._logging import verbose_router_logger
|
||||
from litellm.constants import (
|
||||
EMPTY_MAPPING,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
RETURN_RAW_MODEL_NAME_METADATA_KEY,
|
||||
SESSION_ID_GENERATED_METADATA_KEY,
|
||||
)
|
||||
|
|
@ -42,6 +43,7 @@ from litellm.router_strategy.complexity_router.tier_predictor import (
|
|||
TierSuccessPredictor,
|
||||
resolve_tier_artifact,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import (
|
||||
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
|
||||
ModelResponse,
|
||||
|
|
@ -1668,20 +1670,27 @@ class ComplexityRouter(CustomLogger):
|
|||
)
|
||||
|
||||
request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata")
|
||||
metadata: Final = forwarded_internal_call_metadata(request_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN)
|
||||
metadata: Final = { # mutable-ok: SDK metadata kwarg is enriched by the request pipeline
|
||||
**forwarded_internal_call_metadata(request_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN),
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY: AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
|
||||
}
|
||||
turn_off_message_logging: Final = _effective_turn_off_message_logging(request_kwargs)
|
||||
|
||||
messages_for_call: Final = [
|
||||
messages_for_call: Final[list[AllMessageValues]] = [ # mutable-ok: SDK request payload list is built once
|
||||
{"role": "system", "content": classifier_system_prompt},
|
||||
{"role": "user", "content": user_payload},
|
||||
]
|
||||
response_format: Final = classifier_response_format
|
||||
classifier_call_params: Mapping[str, str] = EMPTY_MAPPING
|
||||
if llm_config.reasoning_effort is not None:
|
||||
classifier_call_params = MappingProxyType({"reasoning_effort": llm_config.reasoning_effort})
|
||||
|
||||
proxy_server_request: Final = {
|
||||
"body": {
|
||||
"model": llm_config.model,
|
||||
"messages": messages_for_call,
|
||||
"response_format": response_format,
|
||||
**classifier_call_params,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1693,6 +1702,7 @@ class ComplexityRouter(CustomLogger):
|
|||
metadata=metadata,
|
||||
proxy_server_request=proxy_server_request,
|
||||
turn_off_message_logging=turn_off_message_logging,
|
||||
**classifier_call_params,
|
||||
**_parent_session_kwargs(request_kwargs),
|
||||
)
|
||||
content: Final = response.choices[0].message.content
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing import Annotated, Final, Literal
|
|||
|
||||
from pydantic import BaseModel, ConfigDict, Field, SkipValidation, field_serializer, field_validator, model_validator
|
||||
|
||||
from litellm.types.llms.openai import REASONING_EFFORT
|
||||
from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin
|
||||
|
||||
from .tier_predictor import TrainedTierArtifact
|
||||
|
|
@ -432,6 +433,13 @@ class ClassifierLLMConfig(BaseModel):
|
|||
model: str = Field(
|
||||
description="Model name (from the router's model_list) to call for classification",
|
||||
)
|
||||
reasoning_effort: REASONING_EFFORT | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Reasoning effort override for classifier calls. Leave unset to use "
|
||||
"the classifier deployment or provider default."
|
||||
),
|
||||
)
|
||||
timeout_ms: int = Field(
|
||||
default=3000,
|
||||
description="Timeout budget for the classification call, in milliseconds",
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ the router silently dropping the deployment at load time under
|
|||
``ignore_invalid_deployments``.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
|
@ -163,6 +163,38 @@ def strategy_router_dependencies(
|
|||
)
|
||||
|
||||
|
||||
def uses_heuristic_v2_classifier(complexity_router_config: object) -> bool:
|
||||
"""Whether this complexity config classifies with the bundled heuristic_v2 model."""
|
||||
return _mapping(complexity_router_config).get("classifier_type") == "heuristic_v2"
|
||||
|
||||
|
||||
def is_heuristic_v2_router(litellm_params: Mapping[str, object]) -> bool:
|
||||
"""Whether this deployment is a complexity router that classifies with heuristic_v2."""
|
||||
return classify_strategy_router_model(str(litellm_params.get("model") or "")) == "complexity" and (
|
||||
uses_heuristic_v2_classifier(litellm_params.get("complexity_router_config"))
|
||||
)
|
||||
|
||||
|
||||
def count_heuristic_v2_routers(deployments: Iterable[Mapping[str, object]]) -> int:
|
||||
"""How many of ``deployments`` (router model_list entries or config.yaml rows) are heuristic_v2 routers."""
|
||||
return sum(1 for deployment in deployments if is_heuristic_v2_router(_mapping(deployment.get("litellm_params"))))
|
||||
|
||||
|
||||
def heuristic_v2_limit_violation(*, held: int, limit: int | None) -> str | None:
|
||||
"""Why holding ``held`` heuristic_v2 routers exceeds ``limit``, or None when it fits.
|
||||
|
||||
``limit`` None means unlimited. The message is shared by every enforcement point (config
|
||||
load, model writes, router registration) and stays SDK-neutral: it names the cap and what
|
||||
the caller can change; the proxy appends how its license lifts the cap.
|
||||
"""
|
||||
if limit is None or held <= limit:
|
||||
return None
|
||||
return (
|
||||
f"At most {limit} auto-router(s) with classifier_type 'heuristic_v2' can be registered but this would make "
|
||||
f"{held}. Use classifier_type 'heuristic' for this router or remove an existing heuristic_v2 router."
|
||||
)
|
||||
|
||||
|
||||
def validate_complexity_router_config_write(complexity_router_config: Mapping[str, object] | None) -> str | None:
|
||||
"""Reject a complexity config the router would refuse to build a deployment from.
|
||||
|
||||
|
|
|
|||
|
|
@ -82,6 +82,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
"""
|
||||
|
||||
CACHE_KEY_PREFIX = "deployment_affinity:v1"
|
||||
USER_ID_AFFINITY_PREFIX: Final = "user_id:"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -253,15 +254,6 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
hashed_user_key: Final = cls._hash_user_key(user_key) if user_key is not None else "unscoped"
|
||||
return f"{cls.CACHE_KEY_PREFIX}:session:{model_group}:{hashed_user_key}:{session_id}"
|
||||
|
||||
@staticmethod
|
||||
def _get_user_key_from_metadata_dict(metadata: dict) -> str | None:
|
||||
# NOTE: affinity is keyed on the *API key hash* provided by the proxy (not the
|
||||
# OpenAI `user` parameter, which is an end-user identifier).
|
||||
user_key: Final = metadata.get("user_api_key_hash")
|
||||
if user_key is None:
|
||||
return None
|
||||
return str(user_key)
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id_from_metadata_dict(metadata: dict) -> str | None:
|
||||
session_id: Final = metadata.get("session_id")
|
||||
|
|
@ -285,22 +277,30 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return metadata_dicts
|
||||
|
||||
@staticmethod
|
||||
def _get_user_key_from_request_kwargs(request_kwargs: dict) -> str | None:
|
||||
def _first_metadata_value(metadata_dicts: Sequence[dict], key: str) -> str | None:
|
||||
value: Final = next((metadata[key] for metadata in metadata_dicts if metadata.get(key) is not None), None)
|
||||
return None if value is None else str(value)
|
||||
|
||||
@classmethod
|
||||
def _get_user_key_from_request_kwargs(cls, request_kwargs: dict) -> str | None:
|
||||
"""
|
||||
Extract a stable affinity key from request kwargs.
|
||||
|
||||
Source (proxy): `metadata.user_api_key_hash`
|
||||
Source (proxy): `metadata.user_api_key_hash` for virtual-key callers. JWT-authenticated
|
||||
callers carry no key hash, so their `metadata.user_api_key_user_id` stands in for it,
|
||||
namespaced under `USER_ID_AFFINITY_PREFIX` so a user id can never alias a key hash.
|
||||
|
||||
Note: the OpenAI `user` parameter is an end-user identifier and is intentionally
|
||||
not used for deployment affinity.
|
||||
"""
|
||||
# Check metadata dicts (Proxy usage)
|
||||
for metadata in DeploymentAffinityCheck._iter_metadata_dicts(request_kwargs):
|
||||
user_key = DeploymentAffinityCheck._get_user_key_from_metadata_dict(metadata=metadata)
|
||||
if user_key is not None:
|
||||
return user_key
|
||||
|
||||
return None
|
||||
metadata_dicts: Final = cls._iter_metadata_dicts(request_kwargs)
|
||||
user_api_key_hash: Final = cls._first_metadata_value(metadata_dicts, "user_api_key_hash")
|
||||
if user_api_key_hash is not None:
|
||||
return user_api_key_hash
|
||||
user_id: Final = cls._first_metadata_value(metadata_dicts, "user_api_key_user_id")
|
||||
if user_id is None:
|
||||
return None
|
||||
return f"{cls.USER_ID_AFFINITY_PREFIX}{user_id}"
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id_from_request_kwargs(request_kwargs: dict) -> str | None:
|
||||
|
|
@ -533,9 +533,9 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return typed_healthy_deployments
|
||||
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: api-key affinity hit -> deployment=%s user_key=%s",
|
||||
"DeploymentAffinityCheck: caller affinity hit -> deployment=%s user_key=%s",
|
||||
model_id,
|
||||
self._shorten_for_logs(user_key),
|
||||
self._shorten_for_logs(self._hash_user_key(user_key)),
|
||||
)
|
||||
return [deployment]
|
||||
|
||||
|
|
@ -626,7 +626,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
deployment_model_name,
|
||||
model_id,
|
||||
self.ttl_seconds,
|
||||
self._shorten_for_logs(user_key),
|
||||
self._shorten_for_logs(self._hash_user_key(user_key)),
|
||||
)
|
||||
else:
|
||||
verbose_router_logger.debug(
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
|
@ -29,6 +30,13 @@ def is_interception_internal_key(
|
|||
return any(key.startswith(prefix) for prefix in prefixes)
|
||||
|
||||
|
||||
CONVERTED_STREAM_KEYS: Final = frozenset(f"{prefix}_converted_stream" for prefix in INTERCEPTION_INTERNAL_PREFIXES)
|
||||
|
||||
|
||||
def converted_stream_requested(params: Mapping[str, object]) -> bool:
|
||||
return any(bool(params.get(key)) for key in CONVERTED_STREAM_KEYS)
|
||||
|
||||
|
||||
class AgenticLoopSafetyError(ValueError):
|
||||
"""
|
||||
Raised when an agentic-loop safety rail refuses a rerun.
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ from pydantic import (
|
|||
ConfigDict,
|
||||
Discriminator,
|
||||
Field,
|
||||
NonNegativeInt,
|
||||
PrivateAttr,
|
||||
SerializerFunctionWrapHandler,
|
||||
field_serializer,
|
||||
|
|
@ -1321,6 +1322,18 @@ class ResponseAPIUsage(BaseLiteLLMOpenAIResponseObject):
|
|||
model_config = {"extra": "allow"}
|
||||
|
||||
|
||||
class WebSearchToolUsage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
num_requests: NonNegativeInt
|
||||
|
||||
|
||||
class ResponsesToolUsage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
web_search: WebSearchToolUsage | None = None
|
||||
|
||||
|
||||
ResponsesAPIStatus = Literal["completed", "failed", "in_progress", "cancelled", "queued", "incomplete"]
|
||||
"""
|
||||
The status of the response generation.
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTableWithKeyCount,
|
||||
|
|
@ -9,6 +11,17 @@ from litellm.proxy._types import (
|
|||
)
|
||||
|
||||
|
||||
class InsensitiveContains(TypedDict):
|
||||
contains: ReadOnly[str]
|
||||
mode: ReadOnly[Literal["insensitive"]]
|
||||
|
||||
|
||||
class UserSearchWhere(TypedDict):
|
||||
"""Prisma filter behind `/user/list?search=`: user_id or user_email contains the term, case-insensitive."""
|
||||
|
||||
OR: ReadOnly[tuple[Mapping[Literal["user_id", "user_email"], InsensitiveContains], ...]]
|
||||
|
||||
|
||||
class UserListResponse(BaseModel):
|
||||
"""
|
||||
Response model for the user list endpoint
|
||||
|
|
|
|||
|
|
@ -11,8 +11,8 @@ class ModelInfoMetadata(TypedDict):
|
|||
|
||||
class ModelInfoResponse(TypedDict):
|
||||
"""OpenAI-compatible model object. `mode`, `max_input_tokens`, and
|
||||
`max_output_tokens` are attached when the cost map knows them; `metadata`
|
||||
is present only when the endpoint is called with include_metadata=true.
|
||||
`max_output_tokens` are attached when the cost map or deployment config
|
||||
knows them; `metadata` is present only with include_metadata=true.
|
||||
"""
|
||||
|
||||
id: str
|
||||
|
|
|
|||
|
|
@ -885,6 +885,18 @@ class FallbackAccessCheck(Protocol):
|
|||
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
|
||||
|
||||
|
||||
class HeuristicV2RouterLimit(Protocol):
|
||||
"""
|
||||
Resolves how many heuristic_v2 complexity routers the Router may hold right now; None means unlimited.
|
||||
|
||||
The Router calls it on every registration and limit query instead of caching the answer, so the
|
||||
proxy can keep the limit on its license object (re-verified on config load) rather than hand
|
||||
over a snapshot.
|
||||
"""
|
||||
|
||||
def __call__(self) -> int | None: ...
|
||||
|
||||
|
||||
class LiteLLM_RouterFileObject(TypedDict, total=False):
|
||||
"""
|
||||
Tracking the litellm params hash, used for mapping the file id to the right model
|
||||
|
|
|
|||
|
|
@ -2543,6 +2543,7 @@ class ImageResponse(OpenAIImageResponse, BaseLiteLLMOpenAIResponseObject):
|
|||
)
|
||||
super().__init__(created=created, data=_data, usage=_usage)
|
||||
|
||||
self.background = kwargs.get("background", None)
|
||||
self.quality = kwargs.get("quality", None)
|
||||
self.output_format = kwargs.get("output_format", None)
|
||||
self.size = kwargs.get("size", None)
|
||||
|
|
@ -3611,6 +3612,7 @@ all_litellm_params = (
|
|||
"litellm_system_prompt",
|
||||
"provider_specific_header",
|
||||
"prompt_version",
|
||||
"prompt_environment",
|
||||
"api_base",
|
||||
"force_timeout",
|
||||
"logger_fn",
|
||||
|
|
|
|||
|
|
@ -3338,6 +3338,9 @@ def get_optional_params_image_gen(
|
|||
continue
|
||||
passed_params[k] = v
|
||||
|
||||
provider_supported_params: Final[tuple[str, ...]] = (
|
||||
tuple(provider_config.get_supported_openai_params(model=model or "")) if provider_config is not None else ()
|
||||
)
|
||||
default_params: Final = {
|
||||
"n": None,
|
||||
"quality": None,
|
||||
|
|
@ -3348,6 +3351,7 @@ def get_optional_params_image_gen(
|
|||
"imageConfig": None,
|
||||
"tools": None,
|
||||
"web_search_options": None,
|
||||
**{k: None for k in provider_supported_params},
|
||||
}
|
||||
|
||||
non_default_params: Final = _get_non_default_params(
|
||||
|
|
@ -3407,10 +3411,9 @@ def get_optional_params_image_gen(
|
|||
if size is not None:
|
||||
optional_params["aspectRatio"] = _map_openai_size_to_vertex_ai_aspect_ratio(size)
|
||||
|
||||
openai_params: list[str] = list(default_params.keys())
|
||||
if provider_config is not None:
|
||||
supported_params = provider_config.get_supported_openai_params(model=model or "")
|
||||
openai_params = list(supported_params)
|
||||
openai_params: Final[list[str]] = (
|
||||
list(provider_supported_params) if provider_config is not None else list(default_params.keys())
|
||||
)
|
||||
|
||||
optional_params = add_provider_specific_params_to_optional_params(
|
||||
optional_params=optional_params,
|
||||
|
|
|
|||
|
|
@ -29277,6 +29277,75 @@
|
|||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": true
|
||||
},
|
||||
"gpt-6-astra": {
|
||||
"cache_creation_input_token_cost": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_priority": 5e-05,
|
||||
"cache_creation_input_token_cost_flex": 6.25e-06,
|
||||
"cache_creation_input_token_cost_priority": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": 4e-06,
|
||||
"cache_read_input_token_cost_flex": 5e-07,
|
||||
"cache_read_input_token_cost_priority": 2e-06,
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens_priority": 4e-05,
|
||||
"input_cost_per_token_batches": 5e-06,
|
||||
"input_cost_per_token_flex": 5e-06,
|
||||
"input_cost_per_token_priority": 2e-05,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 5e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 7.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 3.75e-05,
|
||||
"output_cost_per_token_above_272k_tokens_priority": 0.00015,
|
||||
"output_cost_per_token_batches": 2.5e-05,
|
||||
"output_cost_per_token_flex": 2.5e-05,
|
||||
"output_cost_per_token_priority": 0.0001,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
"regional_processing_uplift_multiplier_us": 1.1,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.01,
|
||||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_minimal_reasoning_effort": false,
|
||||
"supports_native_streaming": true,
|
||||
"supports_none_reasoning_effort": false,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_cache_breakpoint": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
|
|
@ -52911,6 +52980,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -52933,7 +53007,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-terra": {
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
|
|
@ -52944,6 +53019,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 4.4e-07,
|
||||
"output_cost_per_token": 1.32e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 1.98e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -52966,7 +53046,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-cyber": {
|
||||
"input_cost_per_token": 1.375e-05,
|
||||
|
|
@ -53005,6 +53086,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 4.4e-08,
|
||||
"output_cost_per_token": 1.32e-06,
|
||||
"output_cost_per_token_above_272k_tokens": 1.98e-06,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -53027,7 +53113,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"us.openai.gpt-5.6-sol": {
|
||||
"input_cost_per_token": 4.4e-06,
|
||||
|
|
@ -53192,6 +53279,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.95e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -53213,7 +53305,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.4": {
|
||||
"input_cost_per_token": 2.75e-06,
|
||||
|
|
@ -53222,6 +53315,11 @@
|
|||
"cache_read_input_token_cost_above_272k_tokens": 5.5e-07,
|
||||
"output_cost_per_token": 1.65e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 2.475e-05,
|
||||
"search_context_cost_per_query": {
|
||||
"search_context_size_high": 0.012,
|
||||
"search_context_size_low": 0.012,
|
||||
"search_context_size_medium": 0.012
|
||||
},
|
||||
"litellm_provider": "bedrock_mantle",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
|
|
@ -53243,7 +53341,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"bedrock_mantle/google.gemma-4-31b": {
|
||||
"input_cost_per_token": 1.4e-07,
|
||||
|
|
|
|||
|
|
@ -67,8 +67,8 @@ proxy = [
|
|||
"azure-identity>=1.25.2,<2.0",
|
||||
"azure-storage-blob>=12.28.0,<13.0",
|
||||
"mcp>=1.28.1,<2.0",
|
||||
"litellm-proxy-extras==0.4.92",
|
||||
"litellm-enterprise==0.1.63",
|
||||
"litellm-proxy-extras==0.4.93",
|
||||
"litellm-enterprise==0.1.64",
|
||||
"RestrictedPython>=8.5,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
"InquirerPy>=0.3.4,<1.0",
|
||||
|
|
@ -174,6 +174,8 @@ litellm-proxy = "litellm.proxy.client.cli:cli"
|
|||
[dependency-groups]
|
||||
dev = [
|
||||
"diff-cover==9.7.2",
|
||||
"hypothesis==6.165.10",
|
||||
"reportlab==5.0.1",
|
||||
"basedpyright==1.39.7",
|
||||
"keyring==25.7.0",
|
||||
"pytest==9.0.3",
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@
|
|||
"limit": 52
|
||||
},
|
||||
"B010": {
|
||||
"limit": 189
|
||||
"limit": 188
|
||||
},
|
||||
"B018": {
|
||||
"limit": 2
|
||||
|
|
|
|||
|
|
@ -207,13 +207,16 @@ async def completions(request: Request) -> Response:
|
|||
|
||||
async def embeddings(request: Request) -> Response:
|
||||
body = await _parse_body(request)
|
||||
model = _requested_model(body)
|
||||
if model == _SLOW_MODEL:
|
||||
await asyncio.sleep(_SLOW_RESPONSE_SECONDS)
|
||||
raw_input = body.get("input", "")
|
||||
count = len(raw_input) if isinstance(raw_input, list) else 1
|
||||
return JSONResponse(
|
||||
{
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": i, "embedding": [0.0] * 1536} for i in range(max(count, 1))],
|
||||
"model": _requested_model(body),
|
||||
"model": model,
|
||||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -172,6 +172,7 @@ pylint: >=3.3.9 # GPLv2 license
|
|||
langchain-mcp-adapters: >=0.2.1 # MIT License
|
||||
langgraph: >=1.0.10 # MIT License
|
||||
langgraph-prebuilt: >=1.0.8 # MIT License - https://github.com/langchain-ai/langgraph/blob/main/LICENSE
|
||||
hypothesis: >=6.165.10 # MPL 2.0 license
|
||||
pytest-rerunfailures: >=15.1 # MPL 2.0 license
|
||||
pytest-recording: >=0.13.4 # MIT license
|
||||
expression: >=5.6.0 # MIT License - https://github.com/cognitedata/Expression/blob/main/LICENSE
|
||||
|
|
|
|||
|
|
@ -23,10 +23,12 @@ test.describe("AI Hub (internal admin view)", () => {
|
|||
await expect(modal.getByText(/Select All \(\d+\)/)).toBeVisible({ timeout: 5_000 });
|
||||
|
||||
// Step 1: pick the seeded models via "Select All"
|
||||
await modal.getByText(/Select All/i).click();
|
||||
await modal.getByRole("checkbox", { name: /Select All/ }).check();
|
||||
|
||||
// Move to confirm step
|
||||
await modal.getByRole("button", { name: "Next" }).click();
|
||||
const next = modal.getByRole("button", { name: "Next" });
|
||||
await expect(next).toBeEnabled();
|
||||
await next.click();
|
||||
await expect(modal.getByText("Confirm Making Models Public")).toBeVisible({ timeout: 5_000 });
|
||||
|
||||
// Submit
|
||||
|
|
|
|||
|
|
@ -703,3 +703,170 @@ class TestSpendLogsPartitionDetectionMissingPsycopg:
|
|||
assert any(
|
||||
"psycopg is not installed" in record.message for record in caplog.records
|
||||
)
|
||||
|
||||
|
||||
_ATTEMPT_BUDGET = 4
|
||||
|
||||
_P3005_STDERR = """Error: P3005
|
||||
|
||||
The database schema is not empty. Read more about how to baseline an existing production database: https://pris.ly/d/migrate-baseline
|
||||
"""
|
||||
|
||||
|
||||
def _p3018_stderr(migration_name):
|
||||
return f"""Error: P3018
|
||||
|
||||
A migration failed to apply. New migrations cannot be applied before the error is recovered from.
|
||||
|
||||
Migration name: {migration_name}
|
||||
|
||||
Database error code: 42P07
|
||||
|
||||
Database error:
|
||||
ERROR: relation "SomeTable" already exists
|
||||
"""
|
||||
|
||||
|
||||
class _MigrateDeployHarness:
|
||||
"""Drives _setup_database_v2 with a scripted sequence of
|
||||
`prisma migrate deploy` outcomes, with every recovery command faked out so
|
||||
nothing touches a database or the packaged migrations directory."""
|
||||
|
||||
def __init__(self, monkeypatch, tmp_path, outcomes, repeat_last=False):
|
||||
import subprocess as subprocess_module
|
||||
|
||||
import litellm_proxy_extras.utils as utils_module
|
||||
|
||||
self.deploy_calls = []
|
||||
self.resolved = []
|
||||
self.baselines = 0
|
||||
self._outcomes = list(outcomes)
|
||||
self._repeat_last = repeat_last
|
||||
self._subprocess_module = subprocess_module
|
||||
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager, "_get_prisma_dir", staticmethod(lambda: str(tmp_path))
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_create_baseline_migration",
|
||||
staticmethod(self._fake_baseline),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_roll_back_migration",
|
||||
staticmethod(lambda name: None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ProxyExtrasDBManager,
|
||||
"_resolve_specific_migration",
|
||||
staticmethod(self.resolved.append),
|
||||
)
|
||||
monkeypatch.setattr(utils_module.subprocess, "run", self._fake_run)
|
||||
monkeypatch.setattr(utils_module.time, "sleep", lambda seconds: None)
|
||||
|
||||
self.baseline_succeeds = True
|
||||
|
||||
def _fake_baseline(self, *args, **kwargs):
|
||||
self.baselines += 1
|
||||
return self.baseline_succeeds
|
||||
|
||||
def _next_outcome(self):
|
||||
if self._outcomes:
|
||||
if self._repeat_last and len(self._outcomes) == 1:
|
||||
return self._outcomes[0]
|
||||
return self._outcomes.pop(0)
|
||||
raise AssertionError("prisma migrate deploy called more times than scripted")
|
||||
|
||||
def _fake_run(self, cmd, **kwargs):
|
||||
assert cmd[1:] == ["migrate", "deploy"], f"unexpected prisma command: {cmd}"
|
||||
self.deploy_calls.append(cmd)
|
||||
outcome = self._next_outcome()
|
||||
if outcome == "ok":
|
||||
return _FakeCompleted()
|
||||
if outcome == "timeout":
|
||||
raise self._subprocess_module.TimeoutExpired(cmd, 1)
|
||||
raise self._subprocess_module.CalledProcessError(1, cmd, stderr=outcome)
|
||||
|
||||
def run(self):
|
||||
return ProxyExtrasDBManager._setup_database_v2(use_migrate=True)
|
||||
|
||||
|
||||
class TestMigrateDeployAttemptAccounting:
|
||||
"""A `prisma db push` database has a full schema and no ledger, so the v2
|
||||
resolver baselines it and then works through every migration whose objects
|
||||
already exist. Those recoveries make progress, so they must not spend the
|
||||
retry budget, which is there to stop a run that is getting nowhere."""
|
||||
|
||||
def test_a_push_created_database_finishes_bootstrapping(
|
||||
self, monkeypatch, tmp_path
|
||||
):
|
||||
already_there = [
|
||||
"20250329084805_new_cron_job_table",
|
||||
"20250806095134_rename_alias_to_server_name_mcp_table",
|
||||
"20260224203854_add_agent_object_permissions_table",
|
||||
"20260301120000_fourth_table",
|
||||
"20260302120000_fifth_table",
|
||||
"20260303120000_sixth_table",
|
||||
]
|
||||
harness = _MigrateDeployHarness(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
[_P3005_STDERR]
|
||||
+ [_p3018_stderr(name) for name in already_there]
|
||||
+ ["ok"],
|
||||
)
|
||||
|
||||
assert harness.run() is True
|
||||
assert harness.baselines == 1
|
||||
assert harness.resolved == already_there
|
||||
assert len(harness.deploy_calls) == len(already_there) + 2
|
||||
|
||||
def test_repeated_recovery_of_one_migration_still_gives_up(
|
||||
self, monkeypatch, tmp_path
|
||||
):
|
||||
harness = _MigrateDeployHarness(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
[_p3018_stderr("20250329084805_new_cron_job_table")],
|
||||
repeat_last=True,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
harness.run()
|
||||
assert len(harness.deploy_calls) <= _ATTEMPT_BUDGET + 1
|
||||
|
||||
def test_timeouts_still_spend_the_budget(self, monkeypatch, tmp_path):
|
||||
harness = _MigrateDeployHarness(
|
||||
monkeypatch, tmp_path, ["timeout"], repeat_last=True
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
harness.run()
|
||||
assert len(harness.deploy_calls) == _ATTEMPT_BUDGET
|
||||
|
||||
def test_a_baseline_that_never_lands_stops_after_the_budget(
|
||||
self, monkeypatch, tmp_path
|
||||
):
|
||||
harness = _MigrateDeployHarness(
|
||||
monkeypatch, tmp_path, [_P3005_STDERR], repeat_last=True
|
||||
)
|
||||
harness.baseline_succeeds = False
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
harness.run()
|
||||
assert len(harness.deploy_calls) == _ATTEMPT_BUDGET
|
||||
|
||||
def test_an_unrecoverable_error_is_not_retried(self, monkeypatch, tmp_path):
|
||||
harness = _MigrateDeployHarness(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
["Error: P3018\n\nMigration name: 20260101000000_x\n\nERROR: syntax error at or near \"SLECT\"\n"],
|
||||
repeat_last=True,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
harness.run()
|
||||
assert len(harness.deploy_calls) == 1
|
||||
assert harness.resolved == []
|
||||
|
|
|
|||
|
|
@ -2879,7 +2879,6 @@ def response_format_tests(response: litellm.ModelResponse):
|
|||
"model",
|
||||
[
|
||||
"bedrock/mistral.mistral-large-2407-v1:0",
|
||||
"bedrock/cohere.command-r-plus-v1:0",
|
||||
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
|
||||
"mistral.mistral-7b-instruct-v0:2",
|
||||
"meta.llama3-8b-instruct-v1:0",
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import litellm
|
||||
from litellm import completion, completion_cost, embedding
|
||||
from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE
|
||||
|
||||
litellm.set_verbose = False
|
||||
|
||||
|
|
@ -269,11 +270,14 @@ def test_openai_azure_embedding_timeouts():
|
|||
def test_openai_embedding_timeouts():
|
||||
try:
|
||||
response = embedding(
|
||||
model="text-embedding-ada-002",
|
||||
model="openai/slow-endpoint",
|
||||
input=["good morning from litellm"],
|
||||
timeout=0.00001,
|
||||
api_base=FAKE_OPENAI_API_BASE,
|
||||
api_key="fake-key",
|
||||
timeout=0.5,
|
||||
)
|
||||
print(response)
|
||||
pytest.fail("Expected timeout error, the request returned instead")
|
||||
except openai.APITimeoutError:
|
||||
print("Good job got OpenAI timeout error!")
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -1552,8 +1552,9 @@ def test_router_timeout():
|
|||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "os.environ/OPENAI_API_KEY",
|
||||
"model": "openai/slow-endpoint",
|
||||
"api_base": FAKE_OPENAI_API_BASE,
|
||||
"api_key": "fake-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
|
|
@ -1562,7 +1563,7 @@ def test_router_timeout():
|
|||
start_time = time.time()
|
||||
try:
|
||||
res = router.completion(
|
||||
model="gpt-3.5-turbo", messages=messages, timeout=0.0001
|
||||
model="gpt-3.5-turbo", messages=messages, timeout=0.5
|
||||
)
|
||||
print(res)
|
||||
pytest.fail("this should have timed out")
|
||||
|
|
|
|||
|
|
@ -1168,7 +1168,6 @@ async def test_completion_replicate_llama3_streaming(sync_mode):
|
|||
"model, region",
|
||||
[
|
||||
# ["bedrock/ai21.jamba-instruct-v1:0", "us-east-1"],
|
||||
# ["bedrock/cohere.command-r-plus-v1:0", None],
|
||||
["us.anthropic.claude-sonnet-4-5-20250929-v1:0", None],
|
||||
# ["mistral.mistral-7b-instruct-v0:2", None],
|
||||
# ["meta.llama3-8b-instruct-v1:0", None],
|
||||
|
|
@ -1271,7 +1270,7 @@ def test_bedrock_claude_3_streaming():
|
|||
"model",
|
||||
[
|
||||
"claude-haiku-4-5-20251001",
|
||||
"cohere.command-r-plus-v1:0", # bedrock
|
||||
"bedrock/mistral.mistral-7b-instruct-v0:2",
|
||||
"gpt-3.5-turbo",
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import openai
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -216,13 +217,16 @@ def test_timeout_streaming():
|
|||
litellm.set_verbose = False
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="gpt-3.5-turbo",
|
||||
model="openai/slow-endpoint",
|
||||
messages=[{"role": "user", "content": "hello, write a 20 pg essay"}],
|
||||
timeout=0.0001,
|
||||
api_base=FAKE_OPENAI_API_BASE,
|
||||
api_key="fake-key",
|
||||
timeout=0.5,
|
||||
stream=True,
|
||||
)
|
||||
for chunk in response:
|
||||
print(chunk)
|
||||
pytest.fail("Did not raise error `openai.APITimeoutError`. The stream completed instead")
|
||||
except openai.APITimeoutError as e:
|
||||
print(
|
||||
"Passed: Raised correct exception. Got openai.APITimeoutError\nGood Job", e
|
||||
|
|
|
|||
|
|
@ -2815,6 +2815,7 @@ async def test_mcp_server_manager_with_access_groups_integration():
|
|||
"""Integration test for MCPServerManager with access group filtering"""
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
MCPServerAccess,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
|
@ -2848,11 +2849,11 @@ async def test_mcp_server_manager_with_access_groups_integration():
|
|||
)
|
||||
|
||||
# Mock the permission lookup to return staff access group
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers") as mock_get_allowed:
|
||||
mock_get_allowed.return_value = [
|
||||
"staff-server-id",
|
||||
"ops-server-id",
|
||||
] # User has access to staff and ops
|
||||
with patch.object(MCPRequestHandler, "get_mcp_server_access") as mock_get_allowed: # test-quality-ok: manager resolver seam
|
||||
mock_get_allowed.return_value = MCPServerAccess(
|
||||
server_ids=("staff-server-id", "ops-server-id"),
|
||||
scope="scoped",
|
||||
)
|
||||
|
||||
allowed_servers = await test_manager.get_allowed_mcp_servers(user_auth)
|
||||
|
||||
|
|
@ -2901,6 +2902,7 @@ async def test_get_allowed_mcp_servers_returns_empty_for_non_admin_without_permi
|
|||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
MCPServerAccess,
|
||||
)
|
||||
|
||||
test_manager = MCPServerManager()
|
||||
|
|
@ -2923,9 +2925,9 @@ async def test_get_allowed_mcp_servers_returns_empty_for_non_admin_without_permi
|
|||
)
|
||||
|
||||
with patch.object(
|
||||
MCPRequestHandler, "get_allowed_mcp_servers", new_callable=AsyncMock
|
||||
MCPRequestHandler, "get_mcp_server_access", new_callable=AsyncMock
|
||||
) as mock_permission_lookup:
|
||||
mock_permission_lookup.return_value = []
|
||||
mock_permission_lookup.return_value = MCPServerAccess(server_ids=())
|
||||
allowed_servers = await test_manager.get_allowed_mcp_servers(user_auth)
|
||||
|
||||
assert allowed_servers == []
|
||||
|
|
|
|||
|
|
@ -1,6 +1,13 @@
|
|||
import time
|
||||
from collections.abc import Iterator
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from openai import OpenAI, BadRequestError, NotFoundError, APIStatusError
|
||||
import pytest
|
||||
from openai import APIStatusError, BadRequestError, NotFoundError, OpenAI, Stream
|
||||
from openai.types.responses import ResponseStreamEvent
|
||||
|
||||
BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: Final = 90
|
||||
|
||||
|
||||
def generate_key():
|
||||
|
|
@ -153,43 +160,48 @@ def test_cancel_response():
|
|||
raise e
|
||||
|
||||
|
||||
def admitted_response_id(chunk: ResponseStreamEvent) -> str | None:
|
||||
response: Final = getattr(chunk, "response", None)
|
||||
return None if response is None else response.id
|
||||
|
||||
|
||||
def events_until_admission(stream: Stream[ResponseStreamEvent], started: float) -> Iterator[ResponseStreamEvent]:
|
||||
for chunk in stream:
|
||||
print("stream chunk=", chunk)
|
||||
yield chunk
|
||||
if admitted_response_id(chunk) is not None:
|
||||
return
|
||||
if time.monotonic() - started > BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS:
|
||||
return
|
||||
|
||||
|
||||
def test_cancel_streaming_response():
|
||||
try:
|
||||
client = get_test_client()
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
client: Final = get_test_client()
|
||||
started: Final = time.monotonic()
|
||||
stream: Final = client.responses.create(
|
||||
model="gpt-5.5",
|
||||
input="count from 1 to 500, one number per line",
|
||||
stream=True,
|
||||
background=True,
|
||||
timeout=BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS,
|
||||
)
|
||||
|
||||
stream = client.responses.create(
|
||||
model="gpt-5.5",
|
||||
input="just respond with the word 'ping'",
|
||||
stream=True,
|
||||
background=True,
|
||||
with stream:
|
||||
events: Final = tuple(events_until_admission(stream, started))
|
||||
|
||||
elapsed: Final = time.monotonic() - started
|
||||
keepalive_events: Final = sum(1 for chunk in events if chunk.type == "keepalive")
|
||||
response_id: Final = next((rid for rid in map(admitted_response_id, events) if rid is not None), None)
|
||||
if response_id is None and keepalive_events:
|
||||
pytest.skip(
|
||||
f"OpenAI held the background stream in keepalive for {elapsed:.0f}s "
|
||||
f"({keepalive_events} keepalive events) without creating the response"
|
||||
)
|
||||
assert response_id is not None, f"no response event within {elapsed:.0f}s of streaming a background response"
|
||||
|
||||
collected_chunks = []
|
||||
response_id = None
|
||||
for chunk in stream:
|
||||
print("stream chunk=", chunk)
|
||||
collected_chunks.append(chunk)
|
||||
# Extract response ID from the first chunk that has it
|
||||
if (
|
||||
response_id is None
|
||||
and hasattr(chunk, "response")
|
||||
and hasattr(chunk.response, "id")
|
||||
):
|
||||
response_id = chunk.response.id
|
||||
|
||||
assert len(collected_chunks) > 0
|
||||
|
||||
# cancel the response if we got a response ID
|
||||
if response_id:
|
||||
cancel_response = client.responses.cancel(response_id)
|
||||
print("CANCEL streaming response=", cancel_response)
|
||||
assert hasattr(cancel_response, "id")
|
||||
except Exception as e:
|
||||
if "Cannot cancel a completed response" in str(e):
|
||||
pass
|
||||
else:
|
||||
raise e
|
||||
cancel_response: Final = client.responses.cancel(response_id)
|
||||
print("CANCEL streaming response=", cancel_response)
|
||||
assert cancel_response.status == "cancelled"
|
||||
|
||||
|
||||
def test_cancel_invalid_response_id():
|
||||
|
|
|
|||
|
|
@ -1,148 +1,105 @@
|
|||
# Rust ↔ Python SDK parity harness
|
||||
# Rust/Python migration harness
|
||||
|
||||
This folder is the operator-facing harness for the Rust migration test plan. It runs pytest normally, listens to test events in-process, and redraws a live matrix grouped by testing strategy and SDK-level function.
|
||||
This local harness follows [the agreed structure](AGENTS.md). The root command selects strategies and combines their reports. Each strategy has an independent entry point
|
||||
|
||||
The matrix always has these SDK columns:
|
||||
|
||||
- `ocr / aocr`
|
||||
- `messages / amessages`
|
||||
- `responses / aresponses`
|
||||
- `count_tokens`
|
||||
- `chat_completions / acompletion`
|
||||
- `transcription / atranscription`
|
||||
|
||||
The harness has four deliberately broad test-strategy folders:
|
||||
|
||||
| Strategy | Folder |
|
||||
| --- | --- |
|
||||
| Public SDK parity over generated and recorded inputs | [`e2e_fuzz_tests/`](e2e_fuzz_tests/) |
|
||||
| Focused tests of Rust-owned behavior | [`unit_tests_rust/`](unit_tests_rust/) |
|
||||
| Isolated transform and Python-to-Rust helper coverage | [`validate_sub_methods/`](validate_sub_methods/) |
|
||||
| Already-existing live-API SDK tests | [`existing_e2e_test_sdk/`](existing_e2e_test_sdk/) |
|
||||
|
||||
## Run it
|
||||
|
||||
From the repository root:
|
||||
|
||||
```bash
|
||||
poetry run python -m tests.rust-python-harness
|
||||
```text
|
||||
strategies/
|
||||
e2e_parity/runner.py
|
||||
sdk/ocr/fixtures/
|
||||
sdk/messages/
|
||||
sdk/chat_completions/
|
||||
sdk/responses/
|
||||
gateway/
|
||||
existing_e2e_test_sdk/runner.py
|
||||
trace_parity/runner.py
|
||||
sdk/
|
||||
gateway/
|
||||
unit_tests/
|
||||
runner.py
|
||||
mapping_validator.py
|
||||
python_runner.py
|
||||
rust_runner.py
|
||||
shared/
|
||||
parity/
|
||||
tracing/
|
||||
reporting/
|
||||
```
|
||||
|
||||
The default runs every configured test once and updates all matching cells in real time. Narrow a run by strategy, SDK function, or both:
|
||||
## Run locally
|
||||
|
||||
```bash
|
||||
poetry run python -m tests.rust-python-harness --strategy e2e_fuzz_tests
|
||||
poetry run python -m tests.rust-python-harness --function messages
|
||||
poetry run python -m tests.rust-python-harness --strategy validate_sub_methods --function ocr
|
||||
uv run python -m tests.rust-python-harness --list
|
||||
uv run python -m tests.rust-python-harness --function ocr --plain
|
||||
uv run python -m tests.rust-python-harness --strategy e2e_parity --surface sdk --function ocr --plain
|
||||
uv run python -m tests.rust-python-harness.strategies.e2e_parity.runner --function ocr --plain
|
||||
uv run python -m tests.rust-python-harness.strategies.trace_parity.runner --plain
|
||||
uv run python -m tests.rust-python-harness.strategies.unit_tests.runner --plain
|
||||
uv run python -m tests.rust-python-harness.strategies.existing_e2e_test_sdk.runner --function transcription --plain
|
||||
```
|
||||
|
||||
For a guided run, use the interactive picker. It asks which strategy rows and SDK
|
||||
function columns to include, then hands the terminal to the live dashboard. It never
|
||||
captures keys while tests are running, so Ctrl-C and pytest debugging remain safe.
|
||||
Use `--interactive` for strategy and function selection, `--pytest-arg=-x` to stop pytest on its first failure, and `--coverage` to write Python coverage under `target/rust-python-harness/`. The harness enables pytest namespace-package discovery only for its own invocations
|
||||
|
||||
```bash
|
||||
poetry run python -m tests.rust-python-harness --interactive
|
||||
```
|
||||
This harness has no CI execution. A configured test that fails or disappears makes the command fail. An unconfigured strategy cell remains planned and contributes no passing evidence. Interruptions and collection errors stop execution; ordinary test failures remain in the combined report while later strategies run
|
||||
|
||||
Useful operator options:
|
||||
## Strategy responsibilities
|
||||
|
||||
```bash
|
||||
# Inspect coverage and pytest selectors without running anything.
|
||||
poetry run python -m tests.rust-python-harness --list
|
||||
E2E parity compares SDK objects, exceptions, callbacks, streams, and provider requests. Gateway tests compare HTTP responses. Both surfaces use the same strategy runner and keep execution details and fixtures in their own folders. OCR has recorded sync/async SDK coverage; the existing Messages and Responses bridge checks remain partial
|
||||
|
||||
# Stable line-oriented output for CI logs or redirected output.
|
||||
poetry run python -m tests.rust-python-harness --plain
|
||||
Trace parity compares operation names through an explicit Python/Rust mapping, call counts, and required completion-before-start ordering with `shared/tracing/compare.py`. Surface tests supply captured operation intervals. No production trace instrumentation or trace case is configured yet
|
||||
|
||||
# Measure Python reference lines exercised by this parity run and build an HTML heatmap.
|
||||
poetry run python -m tests.rust-python-harness --coverage
|
||||
Unit testing combines test mapping validation, separate Python processes with Rust disabled and enabled, backend verification, result comparison, and native Cargo tests. Native tests stay beside their Rust implementation. Existing Python tests stay at their original paths. No complete Python/native unit mapping is configured yet, so these cells remain planned
|
||||
|
||||
# Forward pytest options. Use the equals form when the value begins with a dash.
|
||||
poetry run python -m tests.rust-python-harness --pytest-arg=-x
|
||||
```
|
||||
The existing E2E SDK strategy retains the live provider tests configured upstream. It runs OCR, Chat Completions, and Transcription checks from their existing paths and reports them separately from parity tests. These tests require provider credentials
|
||||
|
||||
The process returns pytest's exit code. A configured selector that collects no test is also a failure. A planned cell has no selector yet and does not fail the run.
|
||||
## Configure cases
|
||||
|
||||
The dashboard adapts to narrow terminals, shows elapsed time and unique-test progress,
|
||||
and prints the three slowest tests when the run ends. Each failure includes a focused
|
||||
`poetry run pytest ... -q` command. Redirected output and CI automatically use the
|
||||
line-oriented plain renderer; `--plain` lets you opt into it locally.
|
||||
|
||||
The final screen includes a confidence score for every SDK section. It is the direct
|
||||
ratio of required strategy rows with passing evidence, such as `1/3 = 33%`; High means
|
||||
all required strategies passed, Medium means some passed, and Low means none passed.
|
||||
This behavioral score is intentionally shown separately from Python and Rust LOC.
|
||||
|
||||
Coverage reports are written outside the three strategy folders at
|
||||
`target/rust-python-harness/`. Open `python-html/index.html` to inspect executed and
|
||||
missing Python lines; `python.json` and `python.xml` are available for automation.
|
||||
Coverage is finalized after pytest exits, because worker processes must flush their
|
||||
data first.
|
||||
|
||||
## Port coverage and confidence
|
||||
|
||||
Treat these as separate signals instead of one ambiguous coverage percentage:
|
||||
|
||||
| Signal | Tool | What it proves |
|
||||
| --- | --- | --- |
|
||||
| Python reference LOC | `coverage.py` / `pytest-cov` via `--coverage` | The mapped Python behavior ran |
|
||||
| Rust port LOC | `cargo-llvm-cov` | The mapped Rust implementation ran |
|
||||
| Parity contracts | This harness matrix | Python and Rust had the same observable behavior |
|
||||
|
||||
`validate_sub_methods/` owns the future source-section inventory that maps a stable
|
||||
Python qualified symbol to its Rust symbol. That inventory is the denominator for
|
||||
per-function rollups; raw coverage for the entire LiteLLM repository would obscure
|
||||
the port's real gaps. `unit_tests_rust/` owns direct `cargo-llvm-cov` runs, while
|
||||
`e2e_fuzz_tests/` owns behavioral parity and fuzz-case counts. Keep Python, Rust, and
|
||||
parity percentages visible side by side and label section confidence High only when
|
||||
the mapped implementation exists, every required strategy passes, and both sides meet
|
||||
their LOC thresholds. Generated Rust LCOV/HTML and the combined index also belong in
|
||||
`target/rust-python-harness/`, not in a fourth strategy folder.
|
||||
|
||||
## Read the matrix
|
||||
|
||||
| Mark | Meaning |
|
||||
| --- | --- |
|
||||
| `✓` | All collected tests passed |
|
||||
| `✗` | At least one test failed |
|
||||
| `!` | Test setup or teardown failed |
|
||||
| `↷` | All collected tests skipped |
|
||||
| `?` | A configured selector did not collect a test |
|
||||
| `—` | Strategy is planned but has no test yet |
|
||||
| `n/a` | Strategy does not apply to this SDK function |
|
||||
| `◐` | The configured tests cover only part of the TDD's parity contract |
|
||||
|
||||
The initial end-to-end entries deliberately show `◐`: the repository has Rust bridge tests for OCR, Messages, and Responses websocket plumbing, but those are not yet frozen-Python-oracle comparisons. The remaining TDD cells stay visible as planned work instead of disappearing from a green summary.
|
||||
|
||||
## Attach parity tests
|
||||
|
||||
Each of the four folders contains a concise `README.md` and a `strategy.json`. Add a pytest file or node ID to the appropriate SDK function's `selectors` list:
|
||||
Each strategy has a `strategy.json`. Its `functions` object defines SDK cases for OCR, Messages, Responses, Count Tokens, Chat Completions, and Transcription. E2E and trace manifests also accept a `gateway` object keyed by API name. A case has `coverage`, `selectors`, and an optional `note`
|
||||
|
||||
```json
|
||||
{
|
||||
"coverage": "complete",
|
||||
"selectors": [
|
||||
"tests/rust-python-harness/validate_sub_methods/test_messages.py"
|
||||
]
|
||||
"coverage": "partial",
|
||||
"selectors": ["tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/test_sdk_parity.py"]
|
||||
}
|
||||
```
|
||||
|
||||
Selectors use the same syntax as pytest. A file selector aggregates every test in the file; a node selector can target one test or parametrized family; a selector ending in `/` aggregates every test in that folder, recursively. The runner deduplicates selectors, so one test may intentionally prove more than one cell without executing twice.
|
||||
Selectors use pytest file or node syntax. A selector ending in `/` includes tests recursively from that directory
|
||||
|
||||
Use these coverage values:
|
||||
Use `planned` with no selectors until an executable contract exists, `partial` for incomplete coverage, `complete` for the full contract, and `not_applicable` when a strategy does not apply. The dashboard shows passing evidence separately from coverage completeness and LOC coverage
|
||||
|
||||
- `complete`: implements the full strategy contract for that SDK function.
|
||||
- `partial`: useful coverage exists, but the TDD contract is not fully proven.
|
||||
- `planned`: no runnable parity test exists yet.
|
||||
- `not_applicable`: the strategy cannot apply, such as streaming for OCR.
|
||||
Unit cases use `unit_suite` instead of `selectors`, pointing to a repository-relative JSON file with this shape:
|
||||
|
||||
Keep comparison mechanics in shared harness modules and provider/function facts in the owning strategy folder. A Python/Rust mismatch is a test failure; do not normalize away observable return types, exception classes, private response fields, chunk ordering, or callback payload differences merely to make a cell green.
|
||||
```json
|
||||
{
|
||||
"python_selectors": ["tests/test_api.py::test_decode"],
|
||||
"cargo_manifest": "litellm-rust/Cargo.toml",
|
||||
"cargo_package": "litellm-core",
|
||||
"cargo_filter": "ocr::",
|
||||
"backend": {
|
||||
"environment_variable": "LITELLM_USE_RUST_OCR",
|
||||
"probe": "tests.rust-python-harness.strategies.unit_tests.python_runner:ocr_backend"
|
||||
},
|
||||
"mappings": [{"python": "tests/test_api.py::test_decode", "rust": "ocr::test_decode"}]
|
||||
}
|
||||
```
|
||||
|
||||
## Architecture
|
||||
Names match automatically when the collected Python and Rust test names agree. Explicit `mappings` handle different names, class names, and parametrized cases. Missing or ambiguous counterparts fail validation in either direction. The Cargo filter must select the same behavior as the Python selectors
|
||||
|
||||
- `catalog.py` validates and loads every strategy manifest.
|
||||
- `models.py` owns typed strategy, case, coverage, and run-state models.
|
||||
- `runner.py` maps live pytest events back to one or more matrix cells.
|
||||
- `ui.py` renders the interactive Rich dashboard and a dependency-free plain fallback.
|
||||
- `cli.py` handles filtering and preserves pytest exit semantics.
|
||||
The backend probe returns `python` or `rust` and runs at startup and before every test call, after fixtures have run. The OCR probe verifies the dispatch flag and native extension availability. Surface tests must also assert that calls reach their intended implementation to catch per-call fallback. Python outcomes must agree, and failed runs remain failures even if both backends fail identically
|
||||
|
||||
The harness is driven from Python, matching the SDK surface and existing test tooling. Rust remains responsible for the implementation under comparison; the harness does not move provider semantics into the PyO3 bridge.
|
||||
## OCR fixtures
|
||||
|
||||
Fixtures, provider configuration, input strategies, and recording commands live in [the OCR package](strategies/e2e_parity/sdk/ocr/fixtures/README.md). Record with provider credentials:
|
||||
|
||||
```bash
|
||||
uv run python -m tests.rust-python-harness.strategies.e2e_parity.sdk.ocr.fixtures.record --examples 1000
|
||||
```
|
||||
|
||||
`LITELLM_OCR_FIXTURE_DIR` and `--fixture-dir` override the default directory. Shared recording, replay, comparison, streaming, and cassette persistence live in `shared/parity/`
|
||||
|
||||
Run the harness's own checks locally:
|
||||
|
||||
```bash
|
||||
uv run pytest -o consider_namespace_packages=true tests/rust-python-harness/shared tests/rust-python-harness/strategies/unit_tests tests/test_rust_python_harness.py -q
|
||||
```
|
||||
|
||||
Existing OCR parity gaps remain visible: invalid-model provider errors differ, Reducto lacks a native contract, and the expanded Azure corpus exposes duplicate Content-Type headers. Moving the harness does not change provider responses or weaken assertions
|
||||
|
|
|
|||
|
|
@ -2,92 +2,74 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Final
|
||||
|
||||
from .models import Coverage, HarnessCase, SDK_FUNCTIONS, Strategy
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
STRATEGIES_ROOT = Path(__file__).parent
|
||||
from .shared.reporting.models import Coverage, HarnessCase, SDK_FUNCTIONS, Strategy
|
||||
|
||||
STRATEGIES_ROOT: Final = Path(__file__).parent / "strategies"
|
||||
|
||||
|
||||
def _require_string(value: Any, field: str, source: Path) -> str:
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError(f"{source}: {field} must be a non-empty string")
|
||||
return value
|
||||
class CaseSpec(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
coverage: Coverage
|
||||
selectors: tuple[str, ...] = ()
|
||||
note: str = ""
|
||||
unit_suite: str | None = None
|
||||
|
||||
|
||||
class StrategySpec(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
order: int
|
||||
id: str
|
||||
label: str
|
||||
description: str
|
||||
functions: dict[str, CaseSpec]
|
||||
gateway: dict[str, CaseSpec] = {}
|
||||
|
||||
|
||||
def _load_strategy(source: Path) -> Strategy:
|
||||
with source.open(encoding="utf-8") as stream:
|
||||
data = json.load(stream)
|
||||
|
||||
strategy_id = _require_string(data.get("id"), "id", source)
|
||||
label = _require_string(data.get("label"), "label", source)
|
||||
description = _require_string(data.get("description"), "description", source)
|
||||
order = data.get("order")
|
||||
if not isinstance(order, int):
|
||||
raise ValueError(f"{source}: order must be an integer")
|
||||
function_data = data.get("functions")
|
||||
if not isinstance(function_data, dict):
|
||||
raise ValueError(f"{source}: functions must be an object")
|
||||
|
||||
missing = set(SDK_FUNCTIONS) - set(function_data)
|
||||
extra = set(function_data) - set(SDK_FUNCTIONS)
|
||||
if missing or extra:
|
||||
raise ValueError(
|
||||
f"{source}: functions must exactly match {SDK_FUNCTIONS}; missing={missing}, extra={extra}"
|
||||
data: Final = StrategySpec.model_validate_json(source.read_text(encoding="utf-8"))
|
||||
if set(data.functions) != set(SDK_FUNCTIONS):
|
||||
raise ValueError(f"{source}: functions must exactly match {SDK_FUNCTIONS}")
|
||||
cases: Final = tuple(
|
||||
HarnessCase(
|
||||
strategy_id=data.id,
|
||||
strategy_label=data.label,
|
||||
sdk_function=name,
|
||||
coverage=case.coverage,
|
||||
selectors=case.selectors,
|
||||
note=case.note,
|
||||
surface=surface,
|
||||
unit_suite=case.unit_suite,
|
||||
)
|
||||
|
||||
cases: list[HarnessCase] = []
|
||||
for sdk_function in SDK_FUNCTIONS:
|
||||
case_data = function_data[sdk_function]
|
||||
if not isinstance(case_data, dict):
|
||||
raise ValueError(f"{source}: functions.{sdk_function} must be an object")
|
||||
try:
|
||||
coverage = Coverage(case_data.get("coverage"))
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"{source}: invalid coverage for {sdk_function}") from exc
|
||||
selectors = case_data.get("selectors", [])
|
||||
if not isinstance(selectors, list) or not all(
|
||||
isinstance(item, str) and item for item in selectors
|
||||
):
|
||||
raise ValueError(
|
||||
f"{source}: selectors for {sdk_function} must be a list of strings"
|
||||
)
|
||||
if coverage is Coverage.NOT_APPLICABLE and selectors:
|
||||
raise ValueError(
|
||||
f"{source}: not_applicable case {sdk_function} cannot have selectors"
|
||||
)
|
||||
cases.append(
|
||||
HarnessCase(
|
||||
strategy_id=strategy_id,
|
||||
strategy_label=label,
|
||||
sdk_function=sdk_function,
|
||||
coverage=coverage,
|
||||
selectors=tuple(selectors),
|
||||
note=str(case_data.get("note", "")),
|
||||
)
|
||||
)
|
||||
|
||||
return Strategy(
|
||||
order=order,
|
||||
id=strategy_id,
|
||||
label=label,
|
||||
description=description,
|
||||
directory=source.parent,
|
||||
cases=tuple(cases),
|
||||
for surface, functions in (("sdk", data.functions), ("gateway", data.gateway))
|
||||
for name in (SDK_FUNCTIONS if surface == "sdk" else functions)
|
||||
for case in (functions[name],)
|
||||
)
|
||||
for case in cases:
|
||||
if case.coverage in {Coverage.PLANNED, Coverage.NOT_APPLICABLE} and (case.selectors or case.unit_suite):
|
||||
raise ValueError(f"{source}: {case.coverage.value} case {case.key} cannot configure tests")
|
||||
if any(not selector.strip() for selector in case.selectors):
|
||||
raise ValueError(f"{source}: empty selector in {case.key}")
|
||||
if data.id == "unit_tests" and case.selectors:
|
||||
raise ValueError(f"{source}: unit_tests must configure unit_suite instead of pytest selectors")
|
||||
if data.id != "unit_tests" and case.unit_suite:
|
||||
raise ValueError(f"{source}: unit_suite is only valid for unit_tests")
|
||||
return Strategy(data.order, data.id, data.label, data.description, source.parent, cases)
|
||||
|
||||
|
||||
def load_catalog(root: Path = STRATEGIES_ROOT) -> tuple[Strategy, ...]:
|
||||
sources = sorted(root.glob("*/strategy.json"))
|
||||
sources: Final = tuple(sorted(root.glob("*/strategy.json")))
|
||||
if not sources:
|
||||
raise ValueError(f"No strategy manifests found below {root}")
|
||||
strategies = tuple(
|
||||
sorted(
|
||||
(_load_strategy(source) for source in sources),
|
||||
key=lambda strategy: strategy.order,
|
||||
)
|
||||
)
|
||||
ids = [strategy.id for strategy in strategies]
|
||||
if len(ids) != len(set(ids)):
|
||||
try:
|
||||
strategies: Final = tuple(sorted((_load_strategy(source) for source in sources), key=lambda item: item.order))
|
||||
except (ValidationError, json.JSONDecodeError) as error:
|
||||
raise ValueError(str(error)) from error
|
||||
if len({strategy.id for strategy in strategies}) != len(strategies):
|
||||
raise ValueError(f"Duplicate strategy id in {root}")
|
||||
return strategies
|
||||
|
|
|
|||
|
|
@ -6,10 +6,14 @@ from collections.abc import Sequence
|
|||
from pathlib import Path
|
||||
|
||||
from .catalog import load_catalog
|
||||
from .models import SDK_FUNCTIONS, HarnessCase, Strategy
|
||||
from .runner import run_pytest
|
||||
from .ui import make_dashboard
|
||||
from .shared.reporting.models import SDK_FUNCTIONS, HarnessCase, Strategy
|
||||
from .shared.reporting.orchestration import StrategyRunner, run_strategies
|
||||
from .shared.reporting.ui import make_dashboard
|
||||
from .strategies.e2e_parity.runner import run as run_e2e
|
||||
from .strategies.existing_e2e_test_sdk.runner import run as run_existing
|
||||
from .strategies.trace_parity.runner import run as run_trace
|
||||
from .strategies.unit_tests.mapping_validator import FunctionReport, build_function_report
|
||||
from .strategies.unit_tests.runner import run as run_units
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
COVERAGE_ROOT = REPO_ROOT / "target" / "rust-python-harness"
|
||||
|
|
@ -44,6 +48,7 @@ def _parser() -> argparse.ArgumentParser:
|
|||
choices=SDK_FUNCTIONS,
|
||||
help="run only this SDK function",
|
||||
)
|
||||
parser.add_argument("--surface", choices=("sdk", "gateway"), help="run only this API surface")
|
||||
parser.add_argument(
|
||||
"--validate-ledger",
|
||||
action="store_true",
|
||||
|
|
@ -135,9 +140,9 @@ def _print_catalog(strategies: Sequence[Strategy]) -> None:
|
|||
print(f"{strategy.id:20} {strategy.label}")
|
||||
for case in strategy.cases:
|
||||
selectors = (
|
||||
", ".join(case.selectors) if case.selectors else "no test configured"
|
||||
", ".join(case.selectors) if case.selectors else case.unit_suite or "no test configured"
|
||||
)
|
||||
print(f" {case.sdk_function:12} {case.coverage.value:14} {selectors}")
|
||||
print(f" {case.surface}/{case.sdk_function:12} {case.coverage.value:14} {selectors}")
|
||||
|
||||
|
||||
def _print_function_report(report: FunctionReport) -> None:
|
||||
|
|
@ -172,7 +177,21 @@ def _validate_ledger(sdk_functions: set[str]) -> int:
|
|||
return 0 if all(report.is_clean for report in reports) else 1
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
def _resolve_runner(strategy_id: str) -> StrategyRunner:
|
||||
match strategy_id:
|
||||
case "e2e_parity":
|
||||
return run_e2e
|
||||
case "trace_parity":
|
||||
return run_trace
|
||||
case "unit_tests":
|
||||
return run_units
|
||||
case "existing_e2e_test_sdk":
|
||||
return run_existing
|
||||
case _:
|
||||
raise ValueError(f"Unknown strategy: {strategy_id}")
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None, *, strategy_id: str | None = None) -> int:
|
||||
args = _parser().parse_args(argv)
|
||||
if args.coverage and importlib.util.find_spec("pytest_cov") is None:
|
||||
_parser().error(
|
||||
|
|
@ -181,7 +200,8 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||
)
|
||||
if args.validate_ledger:
|
||||
return _validate_ledger(set(args.sdk_functions))
|
||||
strategies = load_catalog()
|
||||
catalog = load_catalog()
|
||||
strategies = tuple(strategy for strategy in catalog if strategy_id is None or strategy.id == strategy_id)
|
||||
if args.list:
|
||||
_print_catalog(strategies)
|
||||
return 0
|
||||
|
|
@ -194,7 +214,8 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||
sdk_functions = sdk_functions or picked_functions
|
||||
|
||||
try:
|
||||
cases = _select(strategies, strategy_ids, sdk_functions)
|
||||
selected = _select(strategies, strategy_ids, sdk_functions)
|
||||
cases = tuple(case for case in selected if args.surface is None or case.surface == args.surface)
|
||||
except ValueError as exc:
|
||||
_parser().error(str(exc))
|
||||
selected_strategy_ids = {case.strategy_id for case in cases}
|
||||
|
|
@ -210,11 +231,12 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||
if args.coverage:
|
||||
pytest_args.extend(_coverage_pytest_args())
|
||||
with dashboard:
|
||||
exit_code, run = run_pytest(
|
||||
exit_code, run = run_strategies(
|
||||
cases=cases,
|
||||
repo_root=REPO_ROOT,
|
||||
on_update=dashboard.update,
|
||||
pytest_args=pytest_args,
|
||||
resolve_runner=_resolve_runner,
|
||||
)
|
||||
dashboard.finish(run, exit_code)
|
||||
if args.coverage and (COVERAGE_ROOT / "python.json").exists():
|
||||
|
|
|
|||
|
|
@ -1,3 +0,0 @@
|
|||
# End-to-end fuzz tests
|
||||
|
||||
Runs the same SDK call through the Python and Rust paths using generated inputs and recorded provider responses. It compares public results, streams, callbacks, and exceptions to catch behavior differences a unit test can miss.
|
||||
|
|
@ -1,14 +0,0 @@
|
|||
{
|
||||
"order": 10,
|
||||
"id": "e2e_fuzz_tests",
|
||||
"label": "End-to-end fuzz tests",
|
||||
"description": "Compare observable Python and Rust SDK behavior over generated and recorded inputs.",
|
||||
"functions": {
|
||||
"ocr": {"coverage": "partial", "selectors": ["tests/test_litellm/ocr/test_rust_bridge.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
|
||||
"messages": {"coverage": "partial", "selectors": ["tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
|
||||
"responses": {"coverage": "partial", "selectors": ["tests/test_litellm/responses/test_rust_bridge_websocket.py"], "note": "Covers the websocket bridge; full responses parity is still being added."},
|
||||
"count_tokens": {"coverage": "planned", "selectors": [], "note": "No Rust count_tokens parity test is present yet."},
|
||||
"chat_completions": {"coverage": "partial", "selectors": ["tests/test_litellm/rust_bridge/test_chat_completions.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
|
||||
"transcription": {"coverage": "partial", "selectors": ["tests/test_litellm/test_audio_transcription_rust_bridge.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."}
|
||||
}
|
||||
}
|
||||
91
tests/rust-python-harness/shared/parity/README.md
Normal file
91
tests/rust-python-harness/shared/parity/README.md
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
# Implementation parity testing through the SDK interface
|
||||
|
||||
> Given the same SDK call and identical provider behavior, do two implementations expose the same SDK contract?
|
||||
|
||||
## What the harness compares
|
||||
|
||||
- A fixture contains a LiteLLM SDK input and a recorded upstream provider response
|
||||
- The same LiteLLM input is transformed by isolated baseline and candidate implementations
|
||||
- The resulting provider requests must match in method, path, headers, and body, excluding runtime-specific HTTP metadata
|
||||
- The recorded provider response is then replayed unchanged to both workers
|
||||
- The harness compares the values returned through the Python SDK interface
|
||||
- Non-streaming responses are compared directly, including their concrete return type and public model fields
|
||||
- Streaming responses are consumed and compared chunk by chunk, including wrapper type, chunk type and order, termination, and public exception behavior
|
||||
- Failed SDK calls are compared by exception class, stable message, status, code, model, provider, and parameter fields
|
||||
- Traceback paths and line numbers are excluded because they are runtime-specific
|
||||
- Route-specific comparators and chunk normalizers handle differences in each public SDK contract
|
||||
|
||||
## Process isolation
|
||||
|
||||
- SDK object and stream parity runs both implementations sequentially in the same process so tests can retain returned objects
|
||||
- Every test saves and restores the original bridge state
|
||||
- A small subprocess smoke test verifies environment-based startup configuration and detects fallback to the Python HTTP implementation
|
||||
|
||||
## Streaming execution
|
||||
|
||||
The invocation callback passed to `run_in_process` must consume the stream before returning its `StreamOutcome`.
|
||||
Use `consume_sync_stream` inside that callback, or await `consume_async_stream` inside the callback passed to
|
||||
`run_in_process_async`. Provider requests are collected only after the callback completes. Streaming is explicit:
|
||||
an iterable return value alone does not select stream consumption
|
||||
|
||||
The consumers retain the wrapper type, iteration capabilities, chunk types and order, and any partial output before
|
||||
an error. Errors retain their creation or iteration phase and the full public `SDKError` fields, with traceback text
|
||||
removed. `capture_sync_stream` and `capture_async_stream` consume through the same helpers and then serialize the
|
||||
outcome for subprocess reports. A serialization failure raises as a harness failure rather than becoming an SDK error
|
||||
|
||||
Response models and stream chunks share a recursive comparator. It compares concrete model, container, and scalar
|
||||
types, public fields and extras, and exact values while ignoring Pydantic private attributes at every nesting level.
|
||||
An API may supply an explicit chunk normalizer for its public contract
|
||||
|
||||
Shared tests exercise a local SSE provider through recording, VCR cassette storage, replay, and typed event comparison
|
||||
in sync and async modes. They cover fragmented events, split UTF-8 characters, CRLF framing, coalesced events, and
|
||||
application errors within a normally completed HTTP stream. HTTP byte boundaries and decoded SDK event boundaries
|
||||
are checked separately
|
||||
|
||||
OCR remains the only integrated LiteLLM route. These tests validate shared streaming machinery, not another route's
|
||||
SDK parity. Connection interruption, early cancellation, and lifecycle timeout enforcement remain outside this coverage
|
||||
|
||||
## Hypothesis and property-based testing
|
||||
|
||||
- Hypothesis is a Python library for property-based testing
|
||||
- Example-based tests use inputs selected by the test author
|
||||
- Property-based tests define strategies for valid inputs and properties that must hold for every generated example
|
||||
- Hypothesis generates combinations from those strategies and normally shrinks a failing example to a smaller reproducible case
|
||||
- In this harness, Hypothesis is used only during fixture generation to expand the LiteLLM input corpus
|
||||
- Each API owns the strategies that vary its supported inputs
|
||||
- Fixture generation is deterministic, and each generated input is recorded with the raw provider response it received
|
||||
- The parity tests use committed fixtures and do not call the provider or generate new Hypothesis examples
|
||||
- Provider responses are replayed unchanged, so the parity test does not fuzz or validate provider behavior
|
||||
- Because Hypothesis does not run the parity assertion directly, parity failures are not automatically shrunk
|
||||
|
||||
## API-owned fixtures
|
||||
|
||||
The shared package owns recording, replay, persistence, execution, comparison, and route-neutral media constructors.
|
||||
Each API package owns its input models, explicit strategies, provider targets, route-specific assets, fixture directory,
|
||||
and regeneration command. See the API package documentation for its configured contracts and recording command
|
||||
|
||||
## VCR cassettes
|
||||
|
||||
Fixtures use VCR's YAML `version: 1` format with ordered request/response `interactions`. VCR handles text and binary
|
||||
body serialization. Each cassette also contains `recorded_at`, `ttl_seconds: 0` (committed fixtures never expire), and
|
||||
`x-litellm` metadata holding the SDK input and request provenance. Streaming responses carry
|
||||
`x-litellm-chunk-lengths` so local replay preserves the original byte boundaries
|
||||
|
||||
The recording server captures requests before forwarding their responses. Saved requests use the stable
|
||||
`http://parity-provider.invalid` origin and strip authentication headers and credential query parameters. The upstream
|
||||
request keeps its credentials. Provider response bytes and non-success statuses are preserved
|
||||
|
||||
Standard VCR can load these files and replay their interactions. Parity tests keep using the local HTTP server because
|
||||
Rust HTTP calls do not pass through VCR's Python patches. The harness still compares the two implementations' requests
|
||||
against each other; the saved request is available for inspection and VCR playback, not a new parity assertion
|
||||
|
||||
Refresh parity cassettes through the API's recording command. Generic VCR writers do not preserve the SDK metadata
|
||||
|
||||
Legacy JSON fixtures remain readable. Migrated cassettes mark reconstructed requests as `python_replay`; fresh
|
||||
recordings use `recorded`. The metadata extensions follow the filesystem cassette layout proposed in
|
||||
[PR #39338](https://github.com/BerriAI/litellm/pull/39338), without depending on its unmerged persistence backend
|
||||
|
||||
## References
|
||||
|
||||
- [Hypothesis documentation](https://hypothesis.readthedocs.io/en/latest/)
|
||||
- [Hypothesis documentation source](https://github.com/HypothesisWorks/hypothesis/tree/master/hypothesis/docs)
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
import pytest
|
||||
|
||||
pytest.register_assert_rewrite("tests.rust-python-harness.shared.parity.compare")
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue