mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_post_call_policy_pipeline
This commit is contained in:
commit
c64318cfbd
404 changed files with 23971 additions and 2031 deletions
|
|
@ -32,7 +32,7 @@ jobs:
|
|||
echo "An open sync PR already exists on branch $open_pr; skipping this run."
|
||||
fi
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GH_TOKEN }}
|
||||
GH_TOKEN: ${{ secrets.GH_TOKEN || github.token }}
|
||||
- name: Run the sync
|
||||
if: steps.existing.outputs.open_pr == ''
|
||||
run: |
|
||||
|
|
@ -65,4 +65,4 @@ jobs:
|
|||
--head "$branch" \
|
||||
--base litellm_internal_staging
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GH_TOKEN }}
|
||||
GH_TOKEN: ${{ secrets.GH_TOKEN || github.token }}
|
||||
|
|
|
|||
|
|
@ -114,4 +114,4 @@ jobs:
|
|||
|
||||
- name: Audit provider endpoints against the schema
|
||||
working-directory: terraform/provider
|
||||
run: go run ./tools/endpointaudit -provider-dir ./litellm -spec "${RUNNER_TEMP}/openapi.json"
|
||||
run: go run ./tools/endpointaudit -provider-dir ./litellm -spec "${RUNNER_TEMP}/openapi.json" -coverage-allowlist ./tools/endpointaudit/coverage_allowlist.txt
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 18483
|
||||
"limit": 17271
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2557
|
||||
"limit": 2539
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -12,25 +12,25 @@
|
|||
"limit": 480
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"limit": 113
|
||||
"limit": 112
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"limit": 40
|
||||
},
|
||||
"reportDeprecated": {
|
||||
"limit": 213
|
||||
"limit": 212
|
||||
},
|
||||
"reportDuplicateImport": {
|
||||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 5960
|
||||
"limit": 5486
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
},
|
||||
"reportGeneralTypeIssues": {
|
||||
"limit": 105
|
||||
"limit": 101
|
||||
},
|
||||
"reportIncompatibleMethodOverride": {
|
||||
"limit": 56
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5659
|
||||
"limit": 5658
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15482
|
||||
"limit": 15425
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -72,7 +72,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportOptionalMemberAccess": {
|
||||
"limit": 1058
|
||||
"limit": 1055
|
||||
},
|
||||
"reportOptionalOperand": {
|
||||
"limit": 0
|
||||
|
|
@ -93,7 +93,7 @@
|
|||
"limit": 213
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"limit": 26
|
||||
"limit": 25
|
||||
},
|
||||
"reportUndefinedVariable": {
|
||||
"limit": 0
|
||||
|
|
@ -105,13 +105,13 @@
|
|||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38779
|
||||
"limit": 38721
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19827
|
||||
"limit": 19778
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30348
|
||||
"limit": 30290
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 117
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 831
|
||||
"limit": 829
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
@ -138,9 +138,9 @@
|
|||
"limit": 138
|
||||
},
|
||||
"reportUnusedImport": {
|
||||
"limit": 544
|
||||
"limit": 543
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"limit": 145
|
||||
"limit": 139
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ GET - /audit/{id} - Get audit log by id
|
|||
GET - /audit - Get all audit logs
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Final, Optional
|
||||
|
||||
#### AUDIT LOGGING ####
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
|
@ -18,11 +18,16 @@ from litellm_enterprise.types.proxy.audit_logging_endpoints import (
|
|||
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.table_repositories import AuditLogRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _build_json_field_or_condition(json_key: str, value: str) -> Dict[str, Any]:
|
||||
def _build_json_field_or_condition(json_key: str, value: str) -> dict[str, object]:
|
||||
"""
|
||||
Build an OR condition that matches a value inside a JSON column at the
|
||||
given key, checking both before_value and updated_values.
|
||||
|
|
@ -101,46 +106,37 @@ async def get_audit_logs(
|
|||
detail={"message": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
# Build filter conditions
|
||||
where_conditions: Dict[str, Any] = {}
|
||||
if changed_by:
|
||||
where_conditions["changed_by"] = changed_by
|
||||
if changed_by_api_key:
|
||||
where_conditions["changed_by_api_key"] = changed_by_api_key
|
||||
if action:
|
||||
where_conditions["action"] = action
|
||||
if table_name:
|
||||
where_conditions["table_name"] = table_name
|
||||
if object_id:
|
||||
where_conditions["object_id"] = object_id
|
||||
if start_date or end_date:
|
||||
date_filter: Dict[str, Any] = {}
|
||||
if start_date:
|
||||
date_filter["gte"] = start_date
|
||||
if end_date:
|
||||
date_filter["lte"] = end_date
|
||||
where_conditions["updated_at"] = date_filter
|
||||
date_filter: Final[dict[str, str]] = {
|
||||
**({"gte": start_date} if start_date else {}),
|
||||
**({"lte": end_date} if end_date else {}),
|
||||
}
|
||||
|
||||
# JSON field filters (PostgreSQL only) — each filter is AND'd with the
|
||||
# others, but checks both before_value and updated_values internally (OR).
|
||||
if object_team_id:
|
||||
where_conditions["AND"] = where_conditions.get("AND", []) + [
|
||||
_build_json_field_or_condition("team_id", object_team_id)
|
||||
]
|
||||
if object_key_hash:
|
||||
where_conditions["AND"] = where_conditions.get("AND", []) + [
|
||||
_build_json_field_or_condition("token", object_key_hash)
|
||||
]
|
||||
json_field_conditions: Final[list[dict[str, object]]] = [
|
||||
*([_build_json_field_or_condition("team_id", object_team_id)] if object_team_id else []),
|
||||
*([_build_json_field_or_condition("token", object_key_hash)] if object_key_hash else []),
|
||||
]
|
||||
|
||||
# Build sort conditions
|
||||
order_by: Dict[str, Any] = {}
|
||||
if sort_by and isinstance(sort_by, str):
|
||||
order_by[sort_by] = sort_order
|
||||
else:
|
||||
order_by["updated_at"] = sort_order # Default sort by updated_at
|
||||
# Build filter conditions
|
||||
where_conditions: Final[dict[str, object]] = {
|
||||
**({"changed_by": changed_by} if changed_by else {}),
|
||||
**({"changed_by_api_key": changed_by_api_key} if changed_by_api_key else {}),
|
||||
**({"action": action} if action else {}),
|
||||
**({"table_name": table_name} if table_name else {}),
|
||||
**({"object_id": object_id} if object_id else {}),
|
||||
**({"updated_at": date_filter} if start_date or end_date else {}),
|
||||
**({"AND": json_field_conditions} if json_field_conditions else {}),
|
||||
}
|
||||
|
||||
order_by: Final[dict[str, str]] = (
|
||||
{sort_by: sort_order} if sort_by and isinstance(sort_by, str) else {"updated_at": sort_order}
|
||||
)
|
||||
|
||||
audit_log_table: Final[TableActions["prisma_models.LiteLLM_AuditLog"]] = AuditLogRepository(prisma_client).table
|
||||
|
||||
# Get paginated results
|
||||
audit_logs = await prisma_client.db.litellm_auditlog.find_many(
|
||||
audit_logs: Final = await audit_log_table.find_many(
|
||||
where=where_conditions,
|
||||
order=order_by,
|
||||
skip=(page - 1) * page_size,
|
||||
|
|
@ -148,13 +144,14 @@ async def get_audit_logs(
|
|||
)
|
||||
|
||||
# Get total count for pagination
|
||||
total_count = await prisma_client.db.litellm_auditlog.count(where=where_conditions)
|
||||
total_pages = -(-total_count // page_size) # Ceiling division
|
||||
total_count: Final = await audit_log_table.count(where=where_conditions)
|
||||
total_pages: Final = -(-total_count // page_size) # Ceiling division
|
||||
|
||||
# Return paginated response
|
||||
return PaginatedAuditLogResponse(
|
||||
audit_logs=[
|
||||
AuditLogResponse(**audit_log.model_dump()) for audit_log in audit_logs
|
||||
AuditLogResponse.model_validate(audit_log.model_dump())
|
||||
for audit_log in audit_logs
|
||||
]
|
||||
if audit_logs
|
||||
else [],
|
||||
|
|
@ -198,8 +195,10 @@ async def get_audit_log_by_id(
|
|||
detail={"message": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
audit_log_table: Final[TableActions["prisma_models.LiteLLM_AuditLog"]] = AuditLogRepository(prisma_client).table
|
||||
|
||||
# Get the audit log by ID
|
||||
audit_log = await prisma_client.db.litellm_auditlog.find_unique(where={"id": id})
|
||||
audit_log: Final = await audit_log_table.find_unique(where={"id": id})
|
||||
|
||||
if audit_log is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -207,4 +206,4 @@ async def get_audit_log_by_id(
|
|||
)
|
||||
|
||||
# Convert to response model
|
||||
return AuditLogResponse(**audit_log.model_dump())
|
||||
return AuditLogResponse.model_validate(audit_log.model_dump())
|
||||
|
|
|
|||
|
|
@ -473,19 +473,56 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
)
|
||||
|
||||
page_size: Final = min(limit or 20, 100)
|
||||
cursor_args: _CursorPageArgs = {"cursor": {"unified_object_id": after}, "skip": 1} if after else {}
|
||||
|
||||
batches = await _managed_object_table(self.prisma_client).find_many(
|
||||
where=where_clause,
|
||||
take=page_size + 1,
|
||||
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
||||
**cursor_args,
|
||||
matches: Final = await self._collect_listed_batches(
|
||||
where_clause=where_clause,
|
||||
after=after,
|
||||
wanted=page_size + 1,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return build_list_page(list(matches[:page_size]), has_more=len(matches) > page_size)
|
||||
|
||||
has_more = len(batches) > page_size
|
||||
async def _collect_listed_batches(
|
||||
self,
|
||||
where_clause: Mapping[str, object],
|
||||
after: Optional[str],
|
||||
wanted: int,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[LiteLLMBatch, ...]:
|
||||
"""Read chunks newest-first until ``wanted`` batches survive parsing and
|
||||
file-id resolution or the caller's rows run out, so a run of rows that will
|
||||
not parse refills the page instead of emptying it. The first chunk is
|
||||
``wanted`` rows, so a healthy page still costs one query; a scan that has to
|
||||
continue widens to ``FILE_LIST_CONTINUATION_CHUNK_SIZE`` like ``afile_list``,
|
||||
and every chunk advances the keyset cursor, so the walk ends once the
|
||||
caller's rows are exhausted."""
|
||||
matches: tuple[LiteLLMBatch, ...] = () # rebind-ok: accumulates survivors across chunks
|
||||
cursor_id: Optional[str] = after # rebind-ok: keyset cursor advances to each chunk's last row
|
||||
chunk_size: int = wanted # rebind-ok: widens once a scan has to continue past the first chunk
|
||||
while len(matches) < wanted:
|
||||
cursor_args: _CursorPageArgs = {"cursor": {"unified_object_id": cursor_id}, "skip": 1} if cursor_id else {}
|
||||
chunk = await _managed_object_table(self.prisma_client).find_many(
|
||||
where=where_clause,
|
||||
take=chunk_size,
|
||||
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
|
||||
**cursor_args,
|
||||
)
|
||||
matches = matches + await self._resolve_listed_rows(
|
||||
rows=chunk, wanted=wanted - len(matches), user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
if len(chunk) < chunk_size:
|
||||
break
|
||||
cursor_id = chunk[-1].unified_object_id
|
||||
chunk_size = max(chunk_size, FILE_LIST_CONTINUATION_CHUNK_SIZE)
|
||||
return matches
|
||||
|
||||
async def _resolve_listed_rows(
|
||||
self,
|
||||
rows: "Sequence[PrismaManagedObjectRow]",
|
||||
wanted: int,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[LiteLLMBatch, ...]:
|
||||
parsed_rows: Final = tuple(
|
||||
(row, batch_obj) for row in batches[:page_size] if (batch_obj := _parse_managed_batch_row(row)) is not None
|
||||
(row, batch_obj) for row in rows if (batch_obj := _parse_managed_batch_row(row)) is not None
|
||||
)
|
||||
unified_id_by_raw_id: Final = await map_raw_file_ids_to_unified(
|
||||
raw_file_ids=frozenset(
|
||||
|
|
@ -496,19 +533,19 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
),
|
||||
prisma_client=self.prisma_client,
|
||||
)
|
||||
resolved_batches: Final = [
|
||||
await self._resolve_listed_batch(
|
||||
resolved: Final[list[LiteLLMBatch]] = [] # mutable-ok: resolution stops as soon as the page is full
|
||||
for row, batch_obj in parsed_rows:
|
||||
if len(resolved) == wanted:
|
||||
break
|
||||
resolved_batch = await self._resolve_listed_batch(
|
||||
row=row,
|
||||
batch_obj=batch_obj,
|
||||
unified_id_by_raw_id=unified_id_by_raw_id,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
for row, batch_obj in parsed_rows
|
||||
]
|
||||
return build_list_page(
|
||||
[batch_obj for batch_obj in resolved_batches if batch_obj is not None],
|
||||
has_more=has_more,
|
||||
)
|
||||
if resolved_batch is not None:
|
||||
resolved.append(resolved_batch)
|
||||
return tuple(resolved)
|
||||
|
||||
async def _resolve_listed_batch(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from collections.abc import Sequence
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -26,39 +27,50 @@ from litellm.proxy.management_helpers.utils import (
|
|||
management_endpoint_wrapper,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.project_repository import ProjectRepository
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.repositories.verification_token_repository import VerificationTokenRepository
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
from prisma.actions import (
|
||||
LiteLLM_ProjectTableActions,
|
||||
LiteLLM_TeamTableActions,
|
||||
LiteLLM_VerificationTokenActions,
|
||||
)
|
||||
|
||||
from litellm import Router
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _team_table(prisma_client: PrismaClient) -> "LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable]":
|
||||
team_table: LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable] = prisma_client.db.litellm_teamtable
|
||||
return team_table
|
||||
_OBJECT_PERMISSION_PAYLOAD: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _project_table(prisma_client: PrismaClient) -> "LiteLLM_ProjectTableActions[prisma_models.LiteLLM_ProjectTable]":
|
||||
project_table: LiteLLM_ProjectTableActions[prisma_models.LiteLLM_ProjectTable] = (
|
||||
prisma_client.db.litellm_projecttable
|
||||
)
|
||||
return project_table
|
||||
def _team_table(prisma_client: PrismaClient) -> TableActions["prisma_models.LiteLLM_TeamTable"]:
|
||||
return TeamRepository(prisma_client).table
|
||||
|
||||
|
||||
def _project_table(prisma_client: PrismaClient) -> TableActions["prisma_models.LiteLLM_ProjectTable"]:
|
||||
return ProjectRepository(prisma_client).table
|
||||
|
||||
|
||||
def _verification_token_table(
|
||||
prisma_client: PrismaClient,
|
||||
) -> "LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken]":
|
||||
verification_token_table: LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken] = (
|
||||
prisma_client.db.litellm_verificationtoken
|
||||
)
|
||||
return verification_token_table
|
||||
) -> TableActions["prisma_models.LiteLLM_VerificationToken"]:
|
||||
return VerificationTokenRepository(prisma_client).table
|
||||
|
||||
|
||||
def _budget_table(prisma_client: PrismaClient) -> TableActions["prisma_models.LiteLLM_BudgetTable"]:
|
||||
return BudgetRepository(prisma_client).table
|
||||
|
||||
|
||||
def _object_permission_table(
|
||||
prisma_client: PrismaClient,
|
||||
) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]:
|
||||
return ObjectPermissionRepository(prisma_client).table
|
||||
|
||||
|
||||
def _user_table(prisma_client: PrismaClient) -> TableActions["prisma_models.LiteLLM_UserTable"]:
|
||||
return UserRepository(prisma_client).table
|
||||
|
||||
|
||||
def _jsonified(prisma_client: PrismaClient, payload: dict[str, object]) -> dict[str, object]:
|
||||
|
|
@ -329,7 +341,7 @@ async def _create_budget_for_project(
|
|||
|
||||
new_budget = _jsonified(prisma_client, budget_row.model_dump(exclude_none=True))
|
||||
|
||||
_budget: prisma_models.LiteLLM_BudgetTable = await prisma_client.db.litellm_budgettable.create(
|
||||
_budget: Final = await _budget_table(prisma_client).create(
|
||||
data={
|
||||
**new_budget,
|
||||
"created_by": user_id or litellm_proxy_admin_name,
|
||||
|
|
@ -352,10 +364,8 @@ async def _set_project_object_permission(
|
|||
return None
|
||||
|
||||
if data.object_permission is not None:
|
||||
created_object_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data.object_permission.model_dump(exclude_none=True),
|
||||
)
|
||||
created_object_permission: Final = await _object_permission_table(prisma_client).create(
|
||||
data=data.object_permission.model_dump(exclude_none=True),
|
||||
)
|
||||
del data.object_permission
|
||||
return created_object_permission.object_permission_id
|
||||
|
|
@ -586,10 +596,8 @@ async def new_project(
|
|||
new_project_row = _remove_budget_fields_from_project_data(new_project_row)
|
||||
|
||||
verbose_proxy_logger.info(f"new_project_row: {json.dumps(new_project_row, indent=2)}")
|
||||
response: prisma_models.LiteLLM_ProjectTable = await prisma_client.db.litellm_projecttable.create(
|
||||
data={
|
||||
**new_project_row, # type: ignore
|
||||
},
|
||||
response: Final = await _project_table(prisma_client).create(
|
||||
data={**new_project_row},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
||||
|
|
@ -776,7 +784,7 @@ async def update_project(
|
|||
|
||||
if budget_updates and existing_project.budget_id:
|
||||
# Update existing budget
|
||||
await prisma_client.db.litellm_budgettable.update(
|
||||
await _budget_table(prisma_client).update(
|
||||
where={"budget_id": existing_project.budget_id},
|
||||
data={
|
||||
**budget_updates,
|
||||
|
|
@ -791,18 +799,17 @@ async def update_project(
|
|||
if "object_permission" in update_data:
|
||||
object_permission_data = update_data.pop("object_permission")
|
||||
if object_permission_data:
|
||||
object_permission_payload: Final = _OBJECT_PERMISSION_PAYLOAD.validate_python(object_permission_data)
|
||||
if existing_project.object_permission_id:
|
||||
# Update existing permission
|
||||
await prisma_client.db.litellm_objectpermissiontable.update(
|
||||
await _object_permission_table(prisma_client).update(
|
||||
where={"object_permission_id": existing_project.object_permission_id},
|
||||
data=object_permission_data,
|
||||
data=object_permission_payload,
|
||||
)
|
||||
else:
|
||||
# Create new permission
|
||||
created_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=object_permission_data,
|
||||
)
|
||||
created_permission: Final = await _object_permission_table(prisma_client).create(
|
||||
data=object_permission_payload,
|
||||
)
|
||||
update_data["object_permission_id"] = created_permission.object_permission_id
|
||||
|
||||
|
|
@ -818,7 +825,7 @@ async def update_project(
|
|||
update_data = _remove_budget_fields_from_project_data(update_data)
|
||||
|
||||
# Update project
|
||||
updated_project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.update(
|
||||
updated_project: Final = await _project_table(prisma_client).update(
|
||||
where={"project_id": data.project_id},
|
||||
data=update_data,
|
||||
include={"litellm_budget_table": True, "object_permission": True},
|
||||
|
|
@ -1058,7 +1065,7 @@ async def list_projects(
|
|||
# Look up the user's team memberships via the reverse-index on
|
||||
# LiteLLM_UserTable.teams (maintained by team_member_add alongside
|
||||
# members_with_roles). This avoids a full scan of all team rows.
|
||||
user_record: prisma_models.LiteLLM_UserTable | None = await prisma_client.db.litellm_usertable.find_unique(
|
||||
user_record: Final = await _user_table(prisma_client).find_unique(
|
||||
where={"user_id": user_api_key_dict.user_id},
|
||||
)
|
||||
user_team_ids: list[str] = user_record.teams if user_record is not None and user_record.teams else []
|
||||
|
|
|
|||
340
litellm-rust/Cargo.lock
generated
340
litellm-rust/Cargo.lock
generated
|
|
@ -2,6 +2,36 @@
|
|||
# It is not intended for manual editing.
|
||||
version = 4
|
||||
|
||||
[[package]]
|
||||
name = "aho-corasick"
|
||||
version = "1.1.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c982642fa9e8606056828ee9a8505737230110bb1099153c79efe865c59d12ba"
|
||||
dependencies = [
|
||||
"memchr",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "alloca"
|
||||
version = "0.4.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e5a7d05ea6aea7e9e64d25b9156ba2fee3fdd659e34e41063cd2fc7cd020d7f4"
|
||||
dependencies = [
|
||||
"cc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "anes"
|
||||
version = "0.1.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299"
|
||||
|
||||
[[package]]
|
||||
name = "anstyle"
|
||||
version = "1.0.14"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000"
|
||||
|
||||
[[package]]
|
||||
name = "arc-swap"
|
||||
version = "1.9.2"
|
||||
|
|
@ -506,6 +536,12 @@ dependencies = [
|
|||
"either",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cast"
|
||||
version = "0.3.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5"
|
||||
|
||||
[[package]]
|
||||
name = "cc"
|
||||
version = "1.3.0"
|
||||
|
|
@ -541,6 +577,58 @@ dependencies = [
|
|||
"rand_core 0.10.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ciborium"
|
||||
version = "0.2.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "42e69ffd6f0917f5c029256a24d0161db17cea3997d185db0d35926308770f0e"
|
||||
dependencies = [
|
||||
"ciborium-io",
|
||||
"ciborium-ll",
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ciborium-io"
|
||||
version = "0.2.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "05afea1e0a06c9be33d539b876f1ce3692f4afea2cb41f740e7743225ed1c757"
|
||||
|
||||
[[package]]
|
||||
name = "ciborium-ll"
|
||||
version = "0.2.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "57663b653d948a338bfb3eeba9bb2fd5fcfaecb9e199e87e1eda4d9e8b240fd9"
|
||||
dependencies = [
|
||||
"ciborium-io",
|
||||
"half",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "clap"
|
||||
version = "4.6.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "473c7e07f409a8d772161724aa8db6a765a2532a70f9667eeb7b49d3d02fbdca"
|
||||
dependencies = [
|
||||
"clap_builder",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "clap_builder"
|
||||
version = "4.6.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7b48fea5a88e9ae728a2dcbedbfc0e730f7d60da42e1cb049a83c9fb8b789889"
|
||||
dependencies = [
|
||||
"anstyle",
|
||||
"clap_lex",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "clap_lex"
|
||||
version = "1.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9"
|
||||
|
||||
[[package]]
|
||||
name = "cmake"
|
||||
version = "0.1.58"
|
||||
|
|
@ -596,6 +684,72 @@ dependencies = [
|
|||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "criterion"
|
||||
version = "0.8.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "950046b2aa2492f9a536f5f4f9a3de7b9e2476e575e05bd6c333371add4d98f3"
|
||||
dependencies = [
|
||||
"alloca",
|
||||
"anes",
|
||||
"cast",
|
||||
"ciborium",
|
||||
"clap",
|
||||
"criterion-plot",
|
||||
"itertools",
|
||||
"num-traits",
|
||||
"oorandom",
|
||||
"page_size",
|
||||
"plotters",
|
||||
"rayon",
|
||||
"regex",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tinytemplate",
|
||||
"walkdir",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "criterion-plot"
|
||||
version = "0.8.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d8d80a2f4f5b554395e47b5d8305bc3d27813bacb73493eb1001e8f76dae29ea"
|
||||
dependencies = [
|
||||
"cast",
|
||||
"itertools",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "crossbeam-deque"
|
||||
version = "0.8.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5181e0de7b61eb03a81e347d6dd8797bae9da5146707b51077e2d71a54ec0ceb"
|
||||
dependencies = [
|
||||
"crossbeam-epoch",
|
||||
"crossbeam-utils",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "crossbeam-epoch"
|
||||
version = "0.9.20"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2d6914041f254d6e9176c01941b21115dcfb7089e55135a35411081bd106ef3f"
|
||||
dependencies = [
|
||||
"crossbeam-utils",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "crossbeam-utils"
|
||||
version = "0.8.22"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17"
|
||||
|
||||
[[package]]
|
||||
name = "crunchy"
|
||||
version = "0.2.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5"
|
||||
|
||||
[[package]]
|
||||
name = "crypto-common"
|
||||
version = "0.1.7"
|
||||
|
|
@ -856,6 +1010,17 @@ dependencies = [
|
|||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "half"
|
||||
version = "2.7.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"crunchy",
|
||||
"zerocopy",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hashbrown"
|
||||
version = "0.17.1"
|
||||
|
|
@ -1179,6 +1344,15 @@ version = "2.12.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2"
|
||||
|
||||
[[package]]
|
||||
name = "itertools"
|
||||
version = "0.13.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "413ee7dfc52ee1a4949ceeb7dbc8a33f2d6c088194d9f922fb8318faf1f01186"
|
||||
dependencies = [
|
||||
"either",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "itoa"
|
||||
version = "1.0.18"
|
||||
|
|
@ -1255,10 +1429,13 @@ dependencies = [
|
|||
name = "litellm-python-bridge"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"criterion",
|
||||
"litellm-ai-gateway",
|
||||
"litellm-core",
|
||||
"pyo3",
|
||||
"pyo3-async-runtimes",
|
||||
"pythonize",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
]
|
||||
|
|
@ -1340,6 +1517,12 @@ version = "1.21.4"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50"
|
||||
|
||||
[[package]]
|
||||
name = "oorandom"
|
||||
version = "11.1.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e"
|
||||
|
||||
[[package]]
|
||||
name = "openssl-probe"
|
||||
version = "0.2.1"
|
||||
|
|
@ -1352,6 +1535,16 @@ version = "0.5.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1a80800c0488c3a21695ea981a54918fbb37abf04f4d0720c453632255e2ff0e"
|
||||
|
||||
[[package]]
|
||||
name = "page_size"
|
||||
version = "0.6.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "30d5b2194ed13191c1999ae0704b7839fb18384fa22e49b57eeaa97d79ce40da"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"winapi",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "percent-encoding"
|
||||
version = "2.3.2"
|
||||
|
|
@ -1376,6 +1569,34 @@ version = "0.3.33"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e"
|
||||
|
||||
[[package]]
|
||||
name = "plotters"
|
||||
version = "0.3.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5aeb6f403d7a4911efb1e33402027fc44f29b5bf6def3effcc22d7bb75f2b747"
|
||||
dependencies = [
|
||||
"num-traits",
|
||||
"plotters-backend",
|
||||
"plotters-svg",
|
||||
"wasm-bindgen",
|
||||
"web-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "plotters-backend"
|
||||
version = "0.3.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "df42e13c12958a16b3f7f4386b9ab1f3e7933914ecea48da7139435263a4172a"
|
||||
|
||||
[[package]]
|
||||
name = "plotters-svg"
|
||||
version = "0.3.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "51bae2ac328883f7acdfea3d66a7c35751187f870bc81f94563733a154d7a670"
|
||||
dependencies = [
|
||||
"plotters-backend",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "portable-atomic"
|
||||
version = "1.14.0"
|
||||
|
|
@ -1486,6 +1707,16 @@ dependencies = [
|
|||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pythonize"
|
||||
version = "0.29.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6ec376e1216e0c929a74964ce2020012a1a39f32d80e78aa688721219ea7fb89"
|
||||
dependencies = [
|
||||
"pyo3",
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "quinn"
|
||||
version = "0.11.11"
|
||||
|
|
@ -1613,12 +1844,61 @@ dependencies = [
|
|||
"rand_core 0.10.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rayon"
|
||||
version = "1.12.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fb39b166781f92d482534ef4b4b1b2568f42613b53e5b6c160e24cfbfa30926d"
|
||||
dependencies = [
|
||||
"either",
|
||||
"rayon-core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rayon-core"
|
||||
version = "1.13.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "22e18b0f0062d30d4230b2e85ff77fdfe4326feb054b9783a3460d8435c8ab91"
|
||||
dependencies = [
|
||||
"crossbeam-deque",
|
||||
"crossbeam-utils",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "regex"
|
||||
version = "1.13.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f020237b6c8eed93db2e2cb53c00c60a8e1bc73da7d073199a1180401450218d"
|
||||
dependencies = [
|
||||
"aho-corasick",
|
||||
"memchr",
|
||||
"regex-automata",
|
||||
"regex-syntax",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "regex-automata"
|
||||
version = "0.4.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ad8553b9b26413251cbf30e620595c7a41b3887f03da04579c0e6b0d6a06b4b2"
|
||||
dependencies = [
|
||||
"aho-corasick",
|
||||
"memchr",
|
||||
"regex-syntax",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "regex-lite"
|
||||
version = "0.1.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cab834c73d247e67f4fae452806d17d3c7501756d98c8808d7c9c7aa7d18f973"
|
||||
|
||||
[[package]]
|
||||
name = "regex-syntax"
|
||||
version = "0.8.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4"
|
||||
|
||||
[[package]]
|
||||
name = "reqwest"
|
||||
version = "0.12.28"
|
||||
|
|
@ -1774,6 +2054,15 @@ version = "1.0.23"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f"
|
||||
|
||||
[[package]]
|
||||
name = "same-file"
|
||||
version = "1.0.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502"
|
||||
dependencies = [
|
||||
"winapi-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "schannel"
|
||||
version = "0.1.29"
|
||||
|
|
@ -2099,6 +2388,16 @@ dependencies = [
|
|||
"zerovec",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tinytemplate"
|
||||
version = "1.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "be4d6b5f19ff7664e8c98d03e2139cb510db9b0a60b55f8e8709b689d939b6bc"
|
||||
dependencies = [
|
||||
"serde",
|
||||
"serde_json",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tinyvec"
|
||||
version = "1.12.0"
|
||||
|
|
@ -2363,6 +2662,16 @@ version = "0.8.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5c3082ca00d5a5ef149bb8b555a72ae84c9c59f7250f013ac822ac2e49b19c64"
|
||||
|
||||
[[package]]
|
||||
name = "walkdir"
|
||||
version = "2.5.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b"
|
||||
dependencies = [
|
||||
"same-file",
|
||||
"winapi-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "want"
|
||||
version = "0.3.1"
|
||||
|
|
@ -2475,6 +2784,37 @@ dependencies = [
|
|||
"rustls-pki-types",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "winapi"
|
||||
version = "0.3.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419"
|
||||
dependencies = [
|
||||
"winapi-i686-pc-windows-gnu",
|
||||
"winapi-x86_64-pc-windows-gnu",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "winapi-i686-pc-windows-gnu"
|
||||
version = "0.4.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6"
|
||||
|
||||
[[package]]
|
||||
name = "winapi-util"
|
||||
version = "0.1.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
|
||||
dependencies = [
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "winapi-x86_64-pc-windows-gnu"
|
||||
version = "0.4.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f"
|
||||
|
||||
[[package]]
|
||||
name = "windows-link"
|
||||
version = "0.2.1"
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ litellm-ai-gateway = { path = "crates/ai-gateway", default-features = false }
|
|||
axum = "0.7"
|
||||
pyo3 = "0.29.0"
|
||||
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
|
||||
pythonize = "0.29.0"
|
||||
rand = "0.8"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls", "http2", "stream"] }
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
|
|
|
|||
|
|
@ -9,10 +9,23 @@ repository.workspace = true
|
|||
name = "_native"
|
||||
crate-type = ["cdylib"]
|
||||
|
||||
[features]
|
||||
default = ["extension-module"]
|
||||
extension-module = ["pyo3/extension-module"]
|
||||
|
||||
[dependencies]
|
||||
litellm-core = { workspace = true, features = ["bedrock-auth"] }
|
||||
litellm-ai-gateway = { workspace = true, default-features = false }
|
||||
pyo3 = { workspace = true, features = ["extension-module"] }
|
||||
pyo3.workspace = true
|
||||
pyo3-async-runtimes.workspace = true
|
||||
pythonize.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
criterion = "0.8.2"
|
||||
|
||||
[[bench]]
|
||||
name = "serialization"
|
||||
harness = false
|
||||
|
|
|
|||
103
litellm-rust/crates/python-bridge/benches/serialization.rs
Normal file
103
litellm-rust/crates/python-bridge/benches/serialization.rs
Normal file
|
|
@ -0,0 +1,103 @@
|
|||
use std::hint::black_box;
|
||||
use std::time::Duration;
|
||||
|
||||
use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
const PAYLOAD_SIZES: &[(&str, usize)] = &[
|
||||
("1_KiB", 1024),
|
||||
("64_KiB", 64 * 1024),
|
||||
("1_MiB", 1024 * 1024),
|
||||
("4_MiB", 4 * 1024 * 1024),
|
||||
("16_MiB", 16 * 1024 * 1024),
|
||||
];
|
||||
|
||||
fn former_json_roundtrip_from_py(py: Python<'_>, value: &Bound<'_, PyAny>) -> Value {
|
||||
let json = py.import("json").expect("Python json module should import");
|
||||
let encoded: String = json
|
||||
.call_method1("dumps", (value,))
|
||||
.expect("payload should serialize")
|
||||
.extract()
|
||||
.expect("json.dumps should return a string");
|
||||
serde_json::from_str(&encoded).expect("serialized JSON should parse")
|
||||
}
|
||||
|
||||
fn pythonize_from_py(value: &Bound<'_, PyAny>) -> Value {
|
||||
pythonize::depythonize(value).expect("payload should depythonize")
|
||||
}
|
||||
|
||||
fn former_json_roundtrip_to_py(py: Python<'_>, value: &Value) -> Py<PyAny> {
|
||||
let json = py.import("json").expect("Python json module should import");
|
||||
let encoded = serde_json::to_string(value).expect("response should serialize");
|
||||
json.call_method1("loads", (encoded,))
|
||||
.expect("serialized response should parse in Python")
|
||||
.unbind()
|
||||
}
|
||||
|
||||
fn pythonize_to_py(py: Python<'_>, value: &Value) -> Py<PyAny> {
|
||||
pythonize::pythonize(py, value)
|
||||
.expect("response should pythonize")
|
||||
.unbind()
|
||||
}
|
||||
|
||||
fn serialization(c: &mut Criterion) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
for &(label, payload_bytes) in PAYLOAD_SIZES {
|
||||
let data_uri = format!("data:image/png;base64,{}", "A".repeat(payload_bytes));
|
||||
let document = PyDict::new(py);
|
||||
document
|
||||
.set_item("type", "image_url")
|
||||
.expect("document type should be set");
|
||||
document
|
||||
.set_item("image_url", &data_uri)
|
||||
.expect("document URL should be set");
|
||||
let response = json!({
|
||||
"pages": [{
|
||||
"index": 0,
|
||||
"markdown": "OCR text",
|
||||
"images": [{"image_base64": data_uri}],
|
||||
}],
|
||||
"model": "mistral-ocr-latest",
|
||||
"document_annotation": null,
|
||||
"usage_info": {"pages_processed": 1},
|
||||
"object": "ocr",
|
||||
});
|
||||
|
||||
c.bench_with_input(
|
||||
BenchmarkId::new("python_to_rust_json", label),
|
||||
&document,
|
||||
|b, document| {
|
||||
b.iter(|| former_json_roundtrip_from_py(py, black_box(document.as_any())))
|
||||
},
|
||||
);
|
||||
c.bench_with_input(
|
||||
BenchmarkId::new("python_to_rust_pythonize", label),
|
||||
&document,
|
||||
|b, document| b.iter(|| pythonize_from_py(black_box(document.as_any()))),
|
||||
);
|
||||
c.bench_with_input(
|
||||
BenchmarkId::new("rust_to_python_json", label),
|
||||
&response,
|
||||
|b, response| b.iter(|| former_json_roundtrip_to_py(py, black_box(response))),
|
||||
);
|
||||
c.bench_with_input(
|
||||
BenchmarkId::new("rust_to_python_pythonize", label),
|
||||
&response,
|
||||
|b, response| b.iter(|| pythonize_to_py(py, black_box(response))),
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
criterion_group! {
|
||||
name = benches;
|
||||
config = Criterion::default()
|
||||
.sample_size(20)
|
||||
.warm_up_time(Duration::from_secs(1))
|
||||
.measurement_time(Duration::from_secs(4));
|
||||
targets = serialization
|
||||
}
|
||||
criterion_main!(benches);
|
||||
|
|
@ -19,6 +19,9 @@ use pyo3::types::{PyAny, PyDict};
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
mod gil;
|
||||
mod marshal;
|
||||
|
||||
use marshal::{from_py, to_py};
|
||||
|
||||
pyo3::create_exception!(
|
||||
_native,
|
||||
|
|
@ -41,35 +44,18 @@ type MarshaledOcrInputs = (
|
|||
Option<Duration>,
|
||||
);
|
||||
|
||||
fn py_to_json(py: Python<'_>, value: &Bound<'_, PyAny>) -> PyResult<Value> {
|
||||
let json = py.import("json")?;
|
||||
let encoded: String = json.call_method1("dumps", (value,))?.extract()?;
|
||||
serde_json::from_str(&encoded).map_err(|err| PyValueError::new_err(err.to_string()))
|
||||
}
|
||||
|
||||
fn json_to_py(py: Python<'_>, value: Value) -> PyResult<Py<PyAny>> {
|
||||
let json = py.import("json")?;
|
||||
let encoded =
|
||||
serde_json::to_string(&value).map_err(|err| PyValueError::new_err(err.to_string()))?;
|
||||
Ok(json.call_method1("loads", (encoded,))?.unbind())
|
||||
}
|
||||
|
||||
fn messages_response_to_py(
|
||||
py: Python<'_>,
|
||||
response: AnthropicMessagesResponse,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let value =
|
||||
serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?;
|
||||
json_to_py(py, value)
|
||||
to_py(py, &response)
|
||||
}
|
||||
|
||||
fn chat_completions_response_to_py(
|
||||
py: Python<'_>,
|
||||
response: ChatCompletionsResponse,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let value =
|
||||
serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?;
|
||||
json_to_py(py, value)
|
||||
to_py(py, &response)
|
||||
}
|
||||
|
||||
fn core_error_to_pyerr(err: CoreError) -> PyErr {
|
||||
|
|
@ -116,7 +102,7 @@ fn optional_object_to_map(
|
|||
value: Option<Py<PyAny>>,
|
||||
) -> PyResult<Map<String, Value>> {
|
||||
match value {
|
||||
Some(value) => match py_to_json(py, value.bind(py))? {
|
||||
Some(value) => match from_py(value.bind(py))? {
|
||||
Value::Object(map) => Ok(map),
|
||||
_ => Err(PyValueError::new_err(format!("{name} must be a dict"))),
|
||||
},
|
||||
|
|
@ -139,7 +125,7 @@ fn marshal_headers(
|
|||
headers: Option<Py<PyAny>>,
|
||||
) -> PyResult<HashMap<String, String>> {
|
||||
let value = match headers {
|
||||
Some(headers) => py_to_json(py, headers.bind(py))?,
|
||||
Some(headers) => from_py(headers.bind(py))?,
|
||||
None => Value::Object(Map::new()),
|
||||
};
|
||||
let Value::Object(headers) = value else {
|
||||
|
|
@ -211,7 +197,7 @@ fn marshal_inputs(
|
|||
optional_params: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<MarshaledOcrInputs> {
|
||||
let document = py_to_json(py, document.bind(py))?;
|
||||
let document = from_py(document.bind(py))?;
|
||||
let extra_headers = match extra_headers {
|
||||
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
|
||||
None => None,
|
||||
|
|
@ -262,7 +248,7 @@ fn ocr(
|
|||
});
|
||||
|
||||
match result {
|
||||
Ok(value) => json_to_py(py, value),
|
||||
Ok(value) => to_py(py, &value),
|
||||
Err(err) => Err(core_error_to_pyerr(err)),
|
||||
}
|
||||
}
|
||||
|
|
@ -307,7 +293,7 @@ fn aocr(
|
|||
.await
|
||||
.map_err(core_error_to_pyerr)?;
|
||||
|
||||
Python::attach(|py| json_to_py(py, value))
|
||||
Python::attach(|py| to_py(py, &value))
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -325,7 +311,7 @@ fn transcription(
|
|||
optional_params: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let audio = py_to_json(py, audio.bind(py))?;
|
||||
let audio = from_py(audio.bind(py))?;
|
||||
let extra_headers = match extra_headers {
|
||||
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
|
||||
None => None,
|
||||
|
|
@ -351,7 +337,7 @@ fn transcription(
|
|||
))
|
||||
});
|
||||
match result {
|
||||
Ok(value) => json_to_py(py, value),
|
||||
Ok(value) => to_py(py, &value),
|
||||
Err(err) => Err(core_error_to_pyerr(err)),
|
||||
}
|
||||
}
|
||||
|
|
@ -370,7 +356,7 @@ fn atranscription(
|
|||
optional_params: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let audio = py_to_json(py, audio.bind(py))?;
|
||||
let audio = from_py(audio.bind(py))?;
|
||||
let extra_headers = match extra_headers {
|
||||
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
|
||||
None => None,
|
||||
|
|
@ -394,7 +380,7 @@ fn atranscription(
|
|||
})
|
||||
.await
|
||||
.map_err(core_error_to_pyerr)?;
|
||||
Python::attach(|py| json_to_py(py, value))
|
||||
Python::attach(|py| to_py(py, &value))
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -406,7 +392,7 @@ fn marshal_messages_inputs(
|
|||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<MarshaledMessagesInputs> {
|
||||
let body = py_to_json(py, body.bind(py))?;
|
||||
let body: Value = from_py(body.bind(py))?;
|
||||
if !body.is_object() {
|
||||
return Err(PyValueError::new_err("body must be a dict"));
|
||||
}
|
||||
|
|
@ -498,7 +484,7 @@ fn marshal_chat_completions_inputs(
|
|||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<MarshaledChatCompletionsInputs> {
|
||||
let messages = py_to_json(py, messages.bind(py))?;
|
||||
let messages: Value = from_py(messages.bind(py))?;
|
||||
if !messages.is_array() {
|
||||
return Err(PyValueError::new_err("messages must be a list"));
|
||||
}
|
||||
|
|
@ -527,7 +513,7 @@ fn chat_completions_decline(
|
|||
optional_params: Option<Py<PyAny>>,
|
||||
custom_llm_provider: Option<String>,
|
||||
) -> PyResult<Option<String>> {
|
||||
let messages = py_to_json(py, messages.bind(py))?;
|
||||
let messages = from_py(messages.bind(py))?;
|
||||
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
|
||||
Ok(chat_completions_decline_reason(
|
||||
&model,
|
||||
|
|
|
|||
20
litellm-rust/crates/python-bridge/src/marshal.rs
Normal file
20
litellm-rust/crates/python-bridge/src/marshal.rs
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
use pyo3::exceptions::PyValueError;
|
||||
use pyo3::prelude::*;
|
||||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
pub fn from_py<T>(value: &Bound<'_, PyAny>) -> PyResult<T>
|
||||
where
|
||||
T: DeserializeOwned,
|
||||
{
|
||||
pythonize::depythonize(value).map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
|
||||
pub fn to_py<T>(py: Python<'_>, value: &T) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + ?Sized,
|
||||
{
|
||||
pythonize::pythonize(py, value)
|
||||
.map(Bound::unbind)
|
||||
.map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
52
litellm-rust/crates/python-bridge/tests/marshal_boundary.rs
Normal file
52
litellm-rust/crates/python-bridge/tests/marshal_boundary.rs
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
use std::fs;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
const DISALLOWED_OUTSIDE_MARSHAL: &[&str] = &[
|
||||
"py.import(\"json\")",
|
||||
"pythonize::",
|
||||
"serde_json::to_string",
|
||||
"serde_json::from_str",
|
||||
];
|
||||
|
||||
fn source_root() -> PathBuf {
|
||||
Path::new(env!("CARGO_MANIFEST_DIR")).join("src")
|
||||
}
|
||||
|
||||
fn rust_sources(directory: &Path) -> Vec<PathBuf> {
|
||||
fs::read_dir(directory)
|
||||
.expect("bridge source directory should be readable")
|
||||
.map(|entry| {
|
||||
entry
|
||||
.expect("bridge source entry should be readable")
|
||||
.path()
|
||||
})
|
||||
.flat_map(|path| {
|
||||
if path.is_dir() {
|
||||
rust_sources(&path)
|
||||
} else if path.extension().is_some_and(|extension| extension == "rs") {
|
||||
vec![path]
|
||||
} else {
|
||||
Vec::new()
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serialization_is_centralized_in_marshal_module() {
|
||||
let root = source_root();
|
||||
|
||||
for path in rust_sources(&root) {
|
||||
if path == root.join("marshal.rs") {
|
||||
continue;
|
||||
}
|
||||
let source = fs::read_to_string(&path).expect("bridge source should be readable");
|
||||
for disallowed in DISALLOWED_OUTSIDE_MARSHAL {
|
||||
assert!(
|
||||
!source.contains(disallowed),
|
||||
"{} bypasses the typed marshal module with `{disallowed}`",
|
||||
path.display()
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -18,7 +18,10 @@ until they're actually needed.
|
|||
import importlib
|
||||
import sys
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Final, cast
|
||||
from types import ModuleType
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
# Import all the data structures that define what can be lazy-loaded
|
||||
# These are just lists of names and maps of where to find them
|
||||
|
|
@ -53,6 +56,9 @@ from ._lazy_imports_registry import (
|
|||
UTILS_NAMES,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tiktoken import Encoding
|
||||
|
||||
|
||||
def get_litellm_globals() -> dict:
|
||||
"""
|
||||
|
|
@ -78,10 +84,10 @@ def _get_utils_globals() -> dict:
|
|||
# They're separate from the main lazy import system because they have specific use cases
|
||||
|
||||
# Lazy loader for default encoding - avoids importing heavy tiktoken library at startup
|
||||
_default_encoding: Any | None = None
|
||||
_default_encoding: "Encoding | None" = None
|
||||
|
||||
|
||||
def _get_default_encoding() -> Any:
|
||||
def _get_default_encoding() -> "Encoding":
|
||||
"""
|
||||
Lazily load and cache the default OpenAI encoding.
|
||||
|
||||
|
|
@ -100,10 +106,10 @@ def _get_default_encoding() -> Any:
|
|||
|
||||
|
||||
# Lazy loader for get_modified_max_tokens to avoid importing token_counter at module import time
|
||||
_get_modified_max_tokens_func: Any | None = None
|
||||
_get_modified_max_tokens_func: "Callable[..., int | None] | None" = None
|
||||
|
||||
|
||||
def _get_modified_max_tokens() -> Any:
|
||||
def _get_modified_max_tokens() -> "Callable[..., int | None]":
|
||||
"""
|
||||
Lazily load and cache the get_modified_max_tokens function.
|
||||
|
||||
|
|
@ -124,10 +130,10 @@ def _get_modified_max_tokens() -> Any:
|
|||
|
||||
|
||||
# Lazy loader for token_counter to avoid importing token_counter module at module import time
|
||||
_token_counter_new_func: Any | None = None
|
||||
_token_counter_new_func: "Callable[..., int] | None" = None
|
||||
|
||||
|
||||
def _get_token_counter_new() -> Any:
|
||||
def _get_token_counter_new() -> "Callable[..., int]":
|
||||
"""
|
||||
Lazily load and cache the token_counter function (aliased as token_counter_new).
|
||||
|
||||
|
|
@ -154,10 +160,10 @@ def _get_token_counter_new() -> Any:
|
|||
# This registry maps attribute names (like "ModelResponse") to handler functions
|
||||
# It's built once the first time someone accesses a lazy-loaded attribute
|
||||
# Example: {"ModelResponse": _lazy_import_utils, "Cache": _lazy_import_caching, ...}
|
||||
_LAZY_IMPORT_REGISTRY: dict[str, Callable[[str], Any]] | None = None
|
||||
_LAZY_IMPORT_REGISTRY: dict[str, Callable[[str], object]] | None = None
|
||||
|
||||
|
||||
def _get_lazy_import_registry() -> dict[str, Callable[[str], Any]]:
|
||||
def _get_lazy_import_registry() -> dict[str, Callable[[str], object]]:
|
||||
"""
|
||||
Build the registry that maps attribute names to their handler functions.
|
||||
|
||||
|
|
@ -206,7 +212,18 @@ def _get_lazy_import_registry() -> dict[str, Callable[[str], Any]]:
|
|||
return _LAZY_IMPORT_REGISTRY
|
||||
|
||||
|
||||
def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], category: str) -> Any:
|
||||
class _AttributeView(TypedDict):
|
||||
"""Holds one module attribute so the lazily fetched value is read back as ``object``."""
|
||||
|
||||
value: ReadOnly[object]
|
||||
|
||||
|
||||
def _module_attribute(module: ModuleType, attr_name: str) -> object:
|
||||
attribute: Final[_AttributeView] = {"value": getattr(module, attr_name)}
|
||||
return attribute["value"]
|
||||
|
||||
|
||||
def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], category: str) -> object:
|
||||
"""
|
||||
Generic function that handles lazy importing for most attributes.
|
||||
|
||||
|
|
@ -255,7 +272,7 @@ def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], cate
|
|||
|
||||
# Step 6: Get the actual attribute from the module
|
||||
# Example: getattr(utils_module, "ModelResponse") returns the ModelResponse class
|
||||
value: Final = getattr(module, attr_name)
|
||||
value: Final = _module_attribute(module, attr_name)
|
||||
|
||||
# Step 7: Cache it so we don't have to import again next time
|
||||
_globals[name] = value
|
||||
|
|
@ -272,62 +289,62 @@ def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], cate
|
|||
# The registry (above) maps attribute names to these handler functions.
|
||||
|
||||
|
||||
def _lazy_import_utils(name: str) -> Any:
|
||||
def _lazy_import_utils(name: str) -> object:
|
||||
"""Handler for utils module attributes (ModelResponse, token_counter, etc.)"""
|
||||
return _generic_lazy_import(name, _UTILS_IMPORT_MAP, "Utils")
|
||||
|
||||
|
||||
def _lazy_import_cost_calculator(name: str) -> Any:
|
||||
def _lazy_import_cost_calculator(name: str) -> object:
|
||||
"""Handler for cost calculator functions (completion_cost, cost_per_token, etc.)"""
|
||||
return _generic_lazy_import(name, _COST_CALCULATOR_IMPORT_MAP, "Cost calculator")
|
||||
|
||||
|
||||
def _lazy_import_token_counter(name: str) -> Any:
|
||||
def _lazy_import_token_counter(name: str) -> object:
|
||||
"""Handler for token counter utilities"""
|
||||
return _generic_lazy_import(name, _TOKEN_COUNTER_IMPORT_MAP, "Token counter")
|
||||
|
||||
|
||||
def _lazy_import_bedrock_types(name: str) -> Any:
|
||||
def _lazy_import_bedrock_types(name: str) -> object:
|
||||
"""Handler for Bedrock type aliases"""
|
||||
return _generic_lazy_import(name, _BEDROCK_TYPES_IMPORT_MAP, "Bedrock types")
|
||||
|
||||
|
||||
def _lazy_import_types_utils(name: str) -> Any:
|
||||
def _lazy_import_types_utils(name: str) -> object:
|
||||
"""Handler for types from litellm.types.utils (BudgetConfig, ImageObject, etc.)"""
|
||||
return _generic_lazy_import(name, _TYPES_UTILS_IMPORT_MAP, "Types utils")
|
||||
|
||||
|
||||
def _lazy_import_caching(name: str) -> Any:
|
||||
def _lazy_import_caching(name: str) -> object:
|
||||
"""Handler for caching classes (Cache, DualCache, RedisCache, etc.)"""
|
||||
return _generic_lazy_import(name, _CACHING_IMPORT_MAP, "Caching")
|
||||
|
||||
|
||||
def _lazy_import_dotprompt(name: str) -> Any:
|
||||
def _lazy_import_dotprompt(name: str) -> object:
|
||||
"""Handler for dotprompt integration globals"""
|
||||
return _generic_lazy_import(name, _DOTPROMPT_IMPORT_MAP, "Dotprompt")
|
||||
|
||||
|
||||
def _lazy_import_types(name: str) -> Any:
|
||||
def _lazy_import_types(name: str) -> object:
|
||||
"""Handler for type classes (GuardrailItem, etc.)"""
|
||||
return _generic_lazy_import(name, _TYPES_IMPORT_MAP, "Types")
|
||||
|
||||
|
||||
def _lazy_import_llm_configs(name: str) -> Any:
|
||||
def _lazy_import_llm_configs(name: str) -> object:
|
||||
"""Handler for LLM config classes (AnthropicConfig, OpenAILikeChatConfig, etc.)"""
|
||||
return _generic_lazy_import(name, _LLM_CONFIGS_IMPORT_MAP, "LLM config")
|
||||
|
||||
|
||||
def _lazy_import_litellm_logging(name: str) -> Any:
|
||||
def _lazy_import_litellm_logging(name: str) -> object:
|
||||
"""Handler for litellm_logging module (Logging, modify_integration)"""
|
||||
return _generic_lazy_import(name, _LITELLM_LOGGING_IMPORT_MAP, "Litellm logging")
|
||||
|
||||
|
||||
def _lazy_import_llm_provider_logic(name: str) -> Any:
|
||||
def _lazy_import_llm_provider_logic(name: str) -> object:
|
||||
"""Handler for LLM provider logic functions (get_llm_provider, etc.)"""
|
||||
return _generic_lazy_import(name, _LLM_PROVIDER_LOGIC_IMPORT_MAP, "LLM provider logic")
|
||||
|
||||
|
||||
def _lazy_import_utils_module(name: str) -> Any:
|
||||
def _lazy_import_utils_module(name: str) -> object:
|
||||
"""
|
||||
Handler for utils module lazy imports.
|
||||
|
||||
|
|
@ -355,7 +372,7 @@ def _lazy_import_utils_module(name: str) -> Any:
|
|||
module = importlib.import_module(module_path)
|
||||
|
||||
# Get the actual attribute from the module
|
||||
value: Final = getattr(module, attr_name)
|
||||
value: Final = _module_attribute(module, attr_name)
|
||||
|
||||
# Cache it so we don't have to import again next time
|
||||
_globals[name] = value
|
||||
|
|
@ -370,7 +387,7 @@ def _lazy_import_utils_module(name: str) -> Any:
|
|||
# These handlers have custom logic that doesn't fit the generic pattern
|
||||
|
||||
|
||||
def _lazy_import_llm_client_cache(name: str) -> Any:
|
||||
def _lazy_import_llm_client_cache(name: str) -> object:
|
||||
"""
|
||||
Handler for LLM client cache - has special logic for singleton instance.
|
||||
|
||||
|
|
@ -386,8 +403,7 @@ def _lazy_import_llm_client_cache(name: str) -> Any:
|
|||
return _globals[name]
|
||||
|
||||
# Import the class
|
||||
module: Final = importlib.import_module("litellm.caching.llm_caching_handler")
|
||||
LLMClientCache: Final = getattr(module, "LLMClientCache")
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
||||
# If they want the class itself, return it
|
||||
if name == "LLMClientCache":
|
||||
|
|
@ -403,7 +419,7 @@ def _lazy_import_llm_client_cache(name: str) -> Any:
|
|||
raise AttributeError(f"LLM client cache lazy import: unknown attribute {name!r}")
|
||||
|
||||
|
||||
def _lazy_import_http_handlers(name: str) -> Any:
|
||||
def _lazy_import_http_handlers(name: str) -> object:
|
||||
"""
|
||||
Handler for HTTP clients - has special logic for creating client instances.
|
||||
|
||||
|
|
|
|||
|
|
@ -17,11 +17,27 @@ A2A Streaming Events:
|
|||
- Artifact update (kind: "artifact-update") - Content/artifact delivery
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, MutableMapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from uuid import uuid4
|
||||
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
_STR_KEY_MAPPING_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _as_object_mapping(value: object) -> Mapping[str, object]:
|
||||
try:
|
||||
return _STR_KEY_MAPPING_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return {}
|
||||
|
||||
|
||||
class A2AStreamingContext:
|
||||
|
|
@ -30,7 +46,7 @@ class A2AStreamingContext:
|
|||
Tracks task_id, context_id, and message accumulation.
|
||||
"""
|
||||
|
||||
def __init__(self, request_id: str, input_message: dict[str, Any]):
|
||||
def __init__(self, request_id: str, input_message: Mapping[str, JsonValue]):
|
||||
self.request_id = request_id
|
||||
self.task_id = str(uuid4())
|
||||
self.context_id = str(uuid4())
|
||||
|
|
@ -46,44 +62,46 @@ class A2ACompletionBridgeTransformation:
|
|||
"""
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_a2a_parts(parts: list[dict[str, Any]]) -> str:
|
||||
def _text_from_a2a_part(part: JsonValue) -> str | None:
|
||||
if not isinstance(part, dict):
|
||||
return None
|
||||
text: Final = part.get("text")
|
||||
if text is None:
|
||||
return None
|
||||
if part.get("kind") not in (None, "", "text"):
|
||||
return None
|
||||
return str(text)
|
||||
|
||||
@staticmethod
|
||||
def _extract_text_from_a2a_parts(parts: Sequence[JsonValue]) -> str:
|
||||
"""Extract text from A2A parts (with or without explicit ``kind``)."""
|
||||
content_parts: Final[list[str]] = []
|
||||
for part in parts:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
kind = part.get("kind")
|
||||
text = part.get("text")
|
||||
if text is None:
|
||||
continue
|
||||
if kind in (None, "", "text"):
|
||||
content_parts.append(str(text))
|
||||
return "\n".join(content_parts)
|
||||
extracted: Final = (A2ACompletionBridgeTransformation._text_from_a2a_part(part) for part in parts)
|
||||
return "\n".join(text for text in extracted if text is not None)
|
||||
|
||||
@staticmethod
|
||||
def get_forward_metadata(
|
||||
a2a_message: dict[str, Any],
|
||||
params: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
a2a_message: Mapping[str, JsonValue],
|
||||
params: Mapping[str, JsonValue] | None = None,
|
||||
) -> Mapping[str, JsonValue] | None:
|
||||
"""
|
||||
Merge A2A metadata from MessageSendParams and the message for downstream providers.
|
||||
|
||||
Forwarded once on the LangGraph run payload (``metadata``), not duplicated on
|
||||
each input message — see ``apply_forward_metadata_to_completion_params``.
|
||||
"""
|
||||
merged: Final[dict[str, Any]] = {}
|
||||
if params and isinstance(params.get("metadata"), dict):
|
||||
merged.update(params["metadata"])
|
||||
params_metadata: Final = params.get("metadata") if params else None
|
||||
message_metadata: Final = a2a_message.get("metadata")
|
||||
if isinstance(message_metadata, dict):
|
||||
merged.update(message_metadata)
|
||||
merged: Final[dict[str, JsonValue]] = {
|
||||
**(params_metadata if isinstance(params_metadata, dict) else {}),
|
||||
**(message_metadata if isinstance(message_metadata, dict) else {}),
|
||||
}
|
||||
return merged or None
|
||||
|
||||
@staticmethod
|
||||
def apply_forward_metadata_to_completion_params(
|
||||
completion_params: dict[str, Any],
|
||||
a2a_message: dict[str, Any],
|
||||
params: dict[str, Any] | None = None,
|
||||
completion_params: MutableMapping[str, object],
|
||||
a2a_message: Mapping[str, JsonValue],
|
||||
params: Mapping[str, JsonValue] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Attach A2A metadata to completion kwargs for provider bridges (e.g. LangGraph).
|
||||
|
|
@ -97,24 +115,20 @@ class A2ACompletionBridgeTransformation:
|
|||
if not forward_metadata:
|
||||
return
|
||||
|
||||
extra_body = completion_params.get("extra_body")
|
||||
if not isinstance(extra_body, dict):
|
||||
extra_body = {}
|
||||
extra_body: Final = _as_object_mapping(completion_params.get("extra_body"))
|
||||
# Layer client-supplied A2A metadata under any agent-owner-configured
|
||||
# ``extra_body.metadata`` so the configured keys remain authoritative
|
||||
# and an A2A caller cannot overwrite server-set run metadata.
|
||||
existing_metadata: Final = extra_body.get("metadata")
|
||||
existing_dict: Final[dict[str, Any]] = existing_metadata if isinstance(existing_metadata, dict) else {}
|
||||
merged_metadata: Final[dict[str, Any]] = {**forward_metadata, **existing_dict}
|
||||
extra_body = {**extra_body, "metadata": merged_metadata}
|
||||
completion_params["extra_body"] = extra_body
|
||||
existing_dict: Final = _as_object_mapping(extra_body.get("metadata"))
|
||||
merged_metadata: Final[dict[str, object]] = {**forward_metadata, **existing_dict}
|
||||
completion_params["extra_body"] = {**extra_body, "metadata": merged_metadata}
|
||||
|
||||
verbose_logger.debug("A2A -> completion forward metadata keys=%s", list(forward_metadata.keys()))
|
||||
|
||||
@staticmethod
|
||||
def a2a_message_to_openai_messages(
|
||||
a2a_message: dict[str, Any],
|
||||
) -> list[dict[str, Any]]:
|
||||
a2a_message: Mapping[str, JsonValue],
|
||||
) -> list[dict[str, object]]:
|
||||
"""
|
||||
Transform an A2A message to OpenAI message format.
|
||||
|
||||
|
|
@ -125,25 +139,19 @@ class A2ACompletionBridgeTransformation:
|
|||
List of OpenAI-format messages
|
||||
"""
|
||||
role: Final = a2a_message.get("role", "user")
|
||||
parts = a2a_message.get("parts", [])
|
||||
raw_parts: Final = a2a_message.get("parts", [])
|
||||
|
||||
# Map A2A roles to OpenAI roles
|
||||
openai_role = role
|
||||
if role == "user":
|
||||
openai_role = "user"
|
||||
elif role == "assistant":
|
||||
openai_role = "assistant"
|
||||
elif role == "system":
|
||||
openai_role = "system"
|
||||
|
||||
if not isinstance(parts, list):
|
||||
parts = []
|
||||
openai_role: Final = (
|
||||
"user" if role == "user" else "assistant" if role == "assistant" else "system" if role == "system" else role
|
||||
)
|
||||
parts: Final = raw_parts if isinstance(raw_parts, list) else []
|
||||
|
||||
content: Final = A2ACompletionBridgeTransformation._extract_text_from_a2a_parts(parts)
|
||||
|
||||
# Do not attach A2A message.metadata here — the completion bridge forwards it
|
||||
# once at run level via extra_body.metadata (LangGraph POST /runs/wait shape).
|
||||
openai_message: Final[dict[str, Any]] = {"role": openai_role, "content": content}
|
||||
openai_message: Final[dict[str, object]] = {"role": openai_role, "content": content}
|
||||
|
||||
verbose_logger.debug(
|
||||
"A2A -> OpenAI transform: role=%s -> %s, content_length=%s", role, openai_role, len(content)
|
||||
|
|
@ -151,11 +159,20 @@ class A2ACompletionBridgeTransformation:
|
|||
|
||||
return [openai_message]
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_content(response: "ModelResponse | CustomStreamWrapper") -> str:
|
||||
if not isinstance(response, ModelResponse) or not response.choices:
|
||||
return ""
|
||||
choice: Final = response.choices[0]
|
||||
if not choice.message:
|
||||
return ""
|
||||
return choice.message.content or ""
|
||||
|
||||
@staticmethod
|
||||
def openai_response_to_a2a_response(
|
||||
response: Any,
|
||||
response: "ModelResponse | CustomStreamWrapper",
|
||||
request_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Transform a LiteLLM ModelResponse to A2A SendMessageResponse format.
|
||||
|
||||
|
|
@ -166,12 +183,7 @@ class A2ACompletionBridgeTransformation:
|
|||
Returns:
|
||||
A2A SendMessageResponse dict
|
||||
"""
|
||||
# Extract content from response
|
||||
content = ""
|
||||
if hasattr(response, "choices") and response.choices:
|
||||
choice: Final = response.choices[0]
|
||||
if hasattr(choice, "message") and choice.message:
|
||||
content = choice.message.content or ""
|
||||
content: Final = A2ACompletionBridgeTransformation._extract_response_content(response)
|
||||
|
||||
# Build A2A message
|
||||
a2a_message: Final = {
|
||||
|
|
@ -182,7 +194,7 @@ class A2ACompletionBridgeTransformation:
|
|||
}
|
||||
|
||||
# Build A2A response
|
||||
a2a_response: Final = {
|
||||
a2a_response: Final[dict[str, object]] = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": a2a_message,
|
||||
|
|
@ -200,7 +212,7 @@ class A2ACompletionBridgeTransformation:
|
|||
@staticmethod
|
||||
def create_task_event(
|
||||
ctx: A2AStreamingContext,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Create the initial task event with status 'submitted'.
|
||||
|
||||
|
|
@ -235,7 +247,7 @@ class A2ACompletionBridgeTransformation:
|
|||
state: str,
|
||||
final: bool = False,
|
||||
message_text: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Create a status update event.
|
||||
|
||||
|
|
@ -245,7 +257,7 @@ class A2ACompletionBridgeTransformation:
|
|||
final: Whether this is the final event
|
||||
message_text: Optional message text for 'working' status
|
||||
"""
|
||||
status: Final[dict[str, Any]] = {
|
||||
status: Final[dict[str, object]] = {
|
||||
"state": state,
|
||||
"timestamp": A2ACompletionBridgeTransformation._get_timestamp(),
|
||||
}
|
||||
|
|
@ -277,7 +289,7 @@ class A2ACompletionBridgeTransformation:
|
|||
def create_artifact_update_event(
|
||||
ctx: A2AStreamingContext,
|
||||
text: str,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Create an artifact update event with content.
|
||||
|
||||
|
|
|
|||
|
|
@ -86,7 +86,7 @@ A2ACardResolver: Final = LiteLLMA2ACardResolver
|
|||
|
||||
|
||||
def _set_usage_on_logging_obj(
|
||||
kwargs: dict[str, Any],
|
||||
kwargs: Mapping[str, object],
|
||||
prompt_tokens: int,
|
||||
completion_tokens: int,
|
||||
) -> None:
|
||||
|
|
@ -99,7 +99,7 @@ def _set_usage_on_logging_obj(
|
|||
completion_tokens: Number of output tokens
|
||||
"""
|
||||
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
if isinstance(litellm_logging_obj, Logging):
|
||||
usage: Final = litellm.Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
|
|
@ -109,7 +109,7 @@ def _set_usage_on_logging_obj(
|
|||
|
||||
|
||||
def _set_agent_id_on_logging_obj(
|
||||
kwargs: dict[str, Any],
|
||||
kwargs: Mapping[str, object],
|
||||
agent_id: str | None,
|
||||
) -> None:
|
||||
"""
|
||||
|
|
@ -123,7 +123,7 @@ def _set_agent_id_on_logging_obj(
|
|||
return
|
||||
|
||||
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
if isinstance(litellm_logging_obj, Logging):
|
||||
# Set agent_id directly on model_call_details (same pattern as custom_llm_provider)
|
||||
litellm_logging_obj.model_call_details["agent_id"] = agent_id
|
||||
|
||||
|
|
@ -132,7 +132,7 @@ _A2A_COST_PARAM_KEYS: Final = ("cost_per_query", "input_cost_per_token", "output
|
|||
|
||||
|
||||
def _set_litellm_params_on_logging_obj(
|
||||
kwargs: dict[str, Any],
|
||||
kwargs: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> None:
|
||||
"""
|
||||
|
|
@ -144,18 +144,22 @@ def _set_litellm_params_on_logging_obj(
|
|||
context, so merge the pricing keys in rather than replacing the dict.
|
||||
"""
|
||||
logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if logging_obj is None:
|
||||
if not isinstance(logging_obj, Logging):
|
||||
return
|
||||
|
||||
cost_params = {key: litellm_params[key] for key in _A2A_COST_PARAM_KEYS if litellm_params.get(key) is not None}
|
||||
cost_params: Final = {
|
||||
key: litellm_params[key] for key in _A2A_COST_PARAM_KEYS if litellm_params.get(key) is not None
|
||||
}
|
||||
if not cost_params:
|
||||
return
|
||||
|
||||
existing: Final = logging_obj.model_call_details.get("litellm_params") or {}
|
||||
logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params}
|
||||
logging_obj.model_call_details["litellm_params"] = {
|
||||
**(logging_obj.model_call_details.get("litellm_params") or {}),
|
||||
**cost_params,
|
||||
}
|
||||
|
||||
|
||||
def _get_a2a_model_info(a2a_client: "A2AClientType", kwargs: dict[str, Any]) -> str:
|
||||
def _get_a2a_model_info(a2a_client: "A2AClientType", kwargs: Mapping[str, object]) -> str:
|
||||
"""
|
||||
Extract agent info and set model/custom_llm_provider for cost tracking.
|
||||
|
||||
|
|
@ -175,7 +179,7 @@ def _get_a2a_model_info(a2a_client: "A2AClientType", kwargs: dict[str, Any]) ->
|
|||
|
||||
# Set on litellm_logging_obj if available (for standard logging payload)
|
||||
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
if isinstance(litellm_logging_obj, Logging):
|
||||
litellm_logging_obj.model = model
|
||||
litellm_logging_obj.custom_llm_provider = custom_llm_provider
|
||||
litellm_logging_obj.model_call_details["model"] = model
|
||||
|
|
@ -498,7 +502,7 @@ async def asend_message(
|
|||
response: Final = LiteLLMSendMessageResponse.from_a2a_response(a2a_response, request_id=str(request.id))
|
||||
|
||||
# Calculate token usage from request and response
|
||||
response_dict: Final[dict[str, object]] = a2a_response.model_dump(mode="json", exclude_none=True)
|
||||
response_dict: Final[dict[str, object]] = a2a_response.root.model_dump(mode="json", exclude_none=True)
|
||||
(
|
||||
prompt_tokens,
|
||||
completion_tokens,
|
||||
|
|
|
|||
|
|
@ -390,7 +390,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
|
||||
] = "openai",
|
||||
logging_obj: Any | None = None,
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
):
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
|
|
|
|||
|
|
@ -4,11 +4,38 @@ BitBucket API client for fetching .prompt files from BitBucket repositories.
|
|||
|
||||
import base64
|
||||
import urllib.parse
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, TypedDict
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
|
||||
class BitBucketSrcEntry(TypedDict):
|
||||
path: ReadOnly[NotRequired[str]]
|
||||
type: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
class BitBucketSrcListing(TypedDict):
|
||||
values: ReadOnly[NotRequired[list[BitBucketSrcEntry]]]
|
||||
|
||||
|
||||
class BitBucketBranch(TypedDict):
|
||||
name: ReadOnly[NotRequired[str]]
|
||||
type: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
class BitBucketBranchListing(TypedDict):
|
||||
values: ReadOnly[NotRequired[list[BitBucketBranch]]]
|
||||
|
||||
|
||||
class BitBucketFileMetadata(TypedDict):
|
||||
content_type: ReadOnly[str | None]
|
||||
content_length: ReadOnly[str | None]
|
||||
last_modified: ReadOnly[str | None]
|
||||
|
||||
|
||||
def _sanitize_file_path(file_path: str) -> str:
|
||||
"""Reject path traversal and URL-encode each path segment."""
|
||||
if "#" in file_path or "?" in file_path:
|
||||
|
|
@ -31,7 +58,7 @@ class BitBucketClient:
|
|||
- Branch-specific file fetching
|
||||
"""
|
||||
|
||||
def __init__(self, config: dict[str, Any]):
|
||||
def __init__(self, config: Mapping[str, object]):
|
||||
"""
|
||||
Initialize the BitBucket client.
|
||||
|
||||
|
|
@ -135,8 +162,8 @@ class BitBucketClient:
|
|||
response: Final = self.http_handler.get(url, headers=self.headers)
|
||||
response.raise_for_status()
|
||||
|
||||
data: Final = response.json()
|
||||
files: Final = []
|
||||
data: Final[BitBucketSrcListing] = response.json()
|
||||
files: Final[list[str]] = []
|
||||
|
||||
for item in data.get("values", []):
|
||||
if item.get("type") == "commit_file":
|
||||
|
|
@ -162,7 +189,7 @@ class BitBucketClient:
|
|||
else:
|
||||
raise Exception(f"Error listing files in '{directory_path}': {e}")
|
||||
|
||||
def get_repository_info(self) -> dict[str, Any]:
|
||||
def get_repository_info(self) -> Mapping[str, object]:
|
||||
"""
|
||||
Get information about the repository.
|
||||
|
||||
|
|
@ -191,7 +218,7 @@ class BitBucketClient:
|
|||
except Exception:
|
||||
return False
|
||||
|
||||
def get_branches(self) -> list[dict[str, Any]]:
|
||||
def get_branches(self) -> list[BitBucketBranch]:
|
||||
"""
|
||||
Get list of branches in the repository.
|
||||
|
||||
|
|
@ -204,12 +231,12 @@ class BitBucketClient:
|
|||
response: Final = self.http_handler.get(url, headers=self.headers)
|
||||
response.raise_for_status()
|
||||
|
||||
data: Final = response.json()
|
||||
data: Final[BitBucketBranchListing] = response.json()
|
||||
return data.get("values", [])
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to get branches: {e}")
|
||||
|
||||
def get_file_metadata(self, file_path: str) -> dict[str, Any] | None:
|
||||
def get_file_metadata(self, file_path: str) -> BitBucketFileMetadata | None:
|
||||
"""
|
||||
Get metadata about a file (size, last modified, etc.).
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ litellm_content_retrieve tool calls server-side via the typed agentic loop plan.
|
|||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, ClassVar, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.compression import compress
|
||||
|
|
@ -22,6 +22,9 @@ from litellm.types.integrations.custom_logger import (
|
|||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
LITELLM_CONTENT_RETRIEVE_TOOL_NAME: Final = "litellm_content_retrieve"
|
||||
_CACHE_TTL_SECONDS: Final = 15 * 60
|
||||
|
||||
|
|
@ -222,7 +225,7 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
response: Any,
|
||||
anthropic_messages_provider_config: Any,
|
||||
anthropic_messages_optional_request_params: dict,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj | None",
|
||||
stream: bool,
|
||||
kwargs: dict,
|
||||
) -> AgenticLoopPlan:
|
||||
|
|
|
|||
|
|
@ -9,9 +9,12 @@ Flow:
|
|||
from __future__ import annotations
|
||||
|
||||
import gzip
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, Protocol
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
|
|
@ -28,6 +31,34 @@ _MAVVRIK_ALLOWED_SUFFIXES: Final = (".mavvrik.dev", ".mavvrik.ai", ".mavvrik.app
|
|||
_GCS_CHUNK_SIZE: Final = 8 * 1024 * 1024 # 8 MB
|
||||
|
||||
|
||||
class MavvrikRegisterBody(TypedDict):
|
||||
metricsMarker: ReadOnly[NotRequired[int | str]]
|
||||
|
||||
|
||||
class MavvrikUploadUrlBody(TypedDict):
|
||||
url: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
class _RegisterResponse(Protocol):
|
||||
def json(self) -> MavvrikRegisterBody: ...
|
||||
|
||||
|
||||
class _UploadUrlResponse(Protocol):
|
||||
def json(self) -> MavvrikUploadUrlBody: ...
|
||||
|
||||
|
||||
def _register_body(response: _RegisterResponse) -> MavvrikRegisterBody:
|
||||
return response.json()
|
||||
|
||||
|
||||
def _upload_url_body(response: _UploadUrlResponse) -> MavvrikUploadUrlBody:
|
||||
return response.json()
|
||||
|
||||
|
||||
def _header_value(headers: Mapping[str, str], name: str) -> str | None:
|
||||
return headers.get(name)
|
||||
|
||||
|
||||
def _validate_api_endpoint(api_endpoint: str) -> None:
|
||||
if not api_endpoint.startswith("https://"):
|
||||
raise ValueError("MAVVRIK_API_ENDPOINT must be an HTTPS URL")
|
||||
|
|
@ -56,12 +87,12 @@ class FocusMavvrikDestination(FocusDestination):
|
|||
self,
|
||||
*,
|
||||
prefix: str,
|
||||
config: dict[str, Any] | None = None,
|
||||
config: Mapping[str, str] | None = None,
|
||||
) -> None:
|
||||
config = config or {}
|
||||
api_key: Final = config.get("api_key")
|
||||
api_endpoint: Final = config.get("api_endpoint")
|
||||
connection_id: Final = config.get("connection_id")
|
||||
resolved_config: Final[Mapping[str, str]] = config or {}
|
||||
api_key: Final = resolved_config.get("api_key")
|
||||
api_endpoint: Final = resolved_config.get("api_endpoint")
|
||||
connection_id: Final = resolved_config.get("connection_id")
|
||||
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
|
|
@ -100,7 +131,7 @@ class FocusMavvrikDestination(FocusDestination):
|
|||
def _auth_headers(self) -> dict[str, str]:
|
||||
return {"Content-Type": "application/json", "x-api-key": self.api_key}
|
||||
|
||||
async def _ensure_registered(self) -> int | None:
|
||||
async def _ensure_registered(self) -> int | str | None:
|
||||
"""POST agent endpoint to register/initialize the connector (once per instance).
|
||||
|
||||
Returns metricsMarker from the Mavvrik response — the last date index
|
||||
|
|
@ -127,7 +158,7 @@ class FocusMavvrikDestination(FocusDestination):
|
|||
if resp.status_code >= 400:
|
||||
raise RuntimeError(f"Mavvrik FOCUS destination: register failed ({resp.status_code}): {resp.text[:200]}")
|
||||
self._registered = True
|
||||
metrics_marker: Final = resp.json().get("metricsMarker", 0)
|
||||
metrics_marker: Final = _register_body(resp).get("metricsMarker", 0)
|
||||
verbose_logger.debug(
|
||||
"Mavvrik FOCUS destination: connector registered (metricsMarker=%s)",
|
||||
metrics_marker,
|
||||
|
|
@ -148,7 +179,7 @@ class FocusMavvrikDestination(FocusDestination):
|
|||
raise RuntimeError(
|
||||
f"Mavvrik FOCUS destination: failed to get signed URL ({resp.status_code}): {resp.text[:200]}"
|
||||
)
|
||||
signed_url: Final = resp.json().get("url")
|
||||
signed_url: Final = _upload_url_body(resp).get("url")
|
||||
if not signed_url:
|
||||
raise RuntimeError(f"Mavvrik FOCUS destination: response missing 'url' field: {resp.json()}")
|
||||
_validate_gcs_url(signed_url, "signed URL")
|
||||
|
|
@ -190,7 +221,7 @@ class FocusMavvrikDestination(FocusDestination):
|
|||
f"Mavvrik FOCUS destination: GCS session init failed ({init_resp.status_code}): {init_resp.text[:400]}"
|
||||
)
|
||||
|
||||
session_uri: Final = init_resp.headers.get("Location")
|
||||
session_uri: Final = _header_value(init_resp.headers, "Location")
|
||||
if not session_uri:
|
||||
raise RuntimeError("Mavvrik FOCUS destination: GCS session init missing Location header")
|
||||
_validate_gcs_url(session_uri, "session URI")
|
||||
|
|
@ -264,7 +295,7 @@ class FocusMavvrikDestination(FocusDestination):
|
|||
)
|
||||
verbose_logger.debug("Mavvrik FOCUS destination: metricsMarker advanced to %s", date_epoch)
|
||||
|
||||
async def get_metrics_marker(self) -> int | None:
|
||||
async def get_metrics_marker(self) -> int | str | None:
|
||||
"""Register with Mavvrik and return the current metricsMarker.
|
||||
|
||||
Always calls the Mavvrik register API — unlike deliver() which skips
|
||||
|
|
@ -287,7 +318,7 @@ class FocusMavvrikDestination(FocusDestination):
|
|||
if resp.status_code >= 400:
|
||||
raise RuntimeError(f"Mavvrik FOCUS destination: register failed ({resp.status_code}): {resp.text[:200]}")
|
||||
self._registered = True
|
||||
metrics_marker: Final = resp.json().get("metricsMarker", 0)
|
||||
metrics_marker: Final = _register_body(resp).get("metricsMarker", 0)
|
||||
verbose_logger.debug("Mavvrik FOCUS destination: got metricsMarker=%s", metrics_marker)
|
||||
return metrics_marker
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,9 @@ import time
|
|||
from datetime import datetime, timedelta
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm import get_secret
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -18,10 +21,32 @@ PROMETHEUS_URL: Final[str | None] = get_secret("PROMETHEUS_URL")
|
|||
PROMETHEUS_SELECTED_INSTANCE: Final[str | None] = get_secret("PROMETHEUS_SELECTED_INSTANCE")
|
||||
async_http_handler: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
|
||||
|
||||
_RAW_JSON_PAYLOAD: Final = TypeAdapter(object)
|
||||
|
||||
|
||||
class PrometheusRangeSample(BaseModel):
|
||||
"""One ``matrix`` series of the Prometheus HTTP query API."""
|
||||
|
||||
metric: dict[str, object]
|
||||
values: list[tuple[float, str]]
|
||||
|
||||
|
||||
class PrometheusQueryData(BaseModel):
|
||||
result: list[PrometheusRangeSample]
|
||||
|
||||
|
||||
class PrometheusQueryResponse(BaseModel):
|
||||
data: PrometheusQueryData
|
||||
|
||||
|
||||
class PrometheusDailySpend(TypedDict):
|
||||
date: ReadOnly[str]
|
||||
spend: ReadOnly[float]
|
||||
|
||||
|
||||
async def get_metric_from_prometheus(
|
||||
metric_name: str,
|
||||
):
|
||||
) -> list[PrometheusRangeSample]:
|
||||
# Get the start of the current day in Unix timestamp
|
||||
if PROMETHEUS_URL is None:
|
||||
raise ValueError("PROMETHEUS_URL not set please set 'PROMETHEUS_URL=<>' in .env")
|
||||
|
|
@ -31,13 +56,13 @@ async def get_metric_from_prometheus(
|
|||
response: Final = await async_http_handler.get(
|
||||
f"{PROMETHEUS_URL}/api/v1/query", params={"query": query, "time": now}
|
||||
) # End of the day
|
||||
_json_response: Final = response.json()
|
||||
_json_response: Final = _RAW_JSON_PAYLOAD.validate_python(response.json())
|
||||
verbose_logger.debug("json response from prometheus /query api %s", _json_response)
|
||||
results: Final = response.json()["data"]["result"]
|
||||
results: Final = PrometheusQueryResponse.model_validate(_json_response).data.result
|
||||
return results
|
||||
|
||||
|
||||
async def get_fallback_metric_from_prometheus():
|
||||
async def get_fallback_metric_from_prometheus() -> str:
|
||||
"""
|
||||
Gets fallback metrics from prometheus for the last 24 hours
|
||||
"""
|
||||
|
|
@ -55,17 +80,17 @@ async def get_fallback_metric_from_prometheus():
|
|||
verbose_logger.debug("response json %s", response_json)
|
||||
for result in response_json:
|
||||
verbose_logger.debug("result= %s", result)
|
||||
metric = result["metric"]
|
||||
metric_values = result["values"]
|
||||
metric_labels = result.metric
|
||||
metric_values = result.values
|
||||
most_recent_value = metric_values[0]
|
||||
|
||||
if PROMETHEUS_SELECTED_INSTANCE is not None:
|
||||
if metric.get("instance") != PROMETHEUS_SELECTED_INSTANCE:
|
||||
if metric_labels.get("instance") != PROMETHEUS_SELECTED_INSTANCE:
|
||||
continue
|
||||
|
||||
value = int(float(most_recent_value[1])) # Convert value to integer
|
||||
primary_model = metric.get("primary_model", "Unknown")
|
||||
fallback_model = metric.get("fallback_model", "Unknown")
|
||||
primary_model = metric_labels.get("primary_model", "Unknown")
|
||||
fallback_model = metric_labels.get("fallback_model", "Unknown")
|
||||
response_message += f"`{value} successful fallback requests` with primary model=`{primary_model}` -> fallback model=`{fallback_model}`"
|
||||
response_message += "\n"
|
||||
verbose_logger.debug("response message %s", response_message)
|
||||
|
|
@ -96,7 +121,7 @@ def _quote_promql_string_literal(value: str) -> str:
|
|||
return json.dumps(value, ensure_ascii=False)
|
||||
|
||||
|
||||
async def get_daily_spend_from_prometheus(api_key: str | None):
|
||||
async def get_daily_spend_from_prometheus(api_key: str | None) -> list[PrometheusDailySpend]:
|
||||
"""
|
||||
Expected Response Format:
|
||||
[
|
||||
|
|
@ -133,17 +158,16 @@ async def get_daily_spend_from_prometheus(api_key: str | None):
|
|||
}
|
||||
|
||||
response: Final = await async_http_handler.get(url, params=params)
|
||||
_json_response: Final = response.json()
|
||||
_json_response: Final = _RAW_JSON_PAYLOAD.validate_python(response.json())
|
||||
verbose_logger.debug("json response from prometheus /query api %s", _json_response)
|
||||
results: Final = response.json()["data"]["result"]
|
||||
formatted_results: Final = []
|
||||
|
||||
for result in results:
|
||||
metric_data = result["values"]
|
||||
for timestamp, value in metric_data:
|
||||
# Convert timestamp to ISO 8601 string with UTC offset
|
||||
date = datetime.fromtimestamp(float(timestamp)).isoformat() + "+00:00"
|
||||
spend = float(value)
|
||||
formatted_results.append({"date": date, "spend": spend})
|
||||
results: Final = PrometheusQueryResponse.model_validate(_json_response).data.result
|
||||
formatted_results: Final[list[PrometheusDailySpend]] = [
|
||||
{
|
||||
"date": datetime.fromtimestamp(float(timestamp)).isoformat() + "+00:00",
|
||||
"spend": float(value),
|
||||
}
|
||||
for result in results
|
||||
for timestamp, value in result.values
|
||||
]
|
||||
|
||||
return formatted_results
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import litellm
|
|||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
|
||||
from litellm.integrations.s3 import get_s3_object_key, resolve_sse_params
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
|
@ -222,7 +223,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
protocol: Final = "https://" if self.s3_endpoint_url.startswith("https://") else "http://"
|
||||
return f"{protocol}{self.s3_bucket_name}.{endpoint_host}/{encoded_key}"
|
||||
return f"{self.s3_endpoint_url}/{self.s3_bucket_name}/{encoded_key}"
|
||||
return f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{encoded_key}"
|
||||
return (
|
||||
f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}."
|
||||
f"{get_aws_dns_suffix(self.s3_region_name)}/{encoded_key}"
|
||||
)
|
||||
|
||||
def _sse_headers(self) -> Mapping[str, str]:
|
||||
candidates: Final = {
|
||||
|
|
|
|||
55
litellm/litellm_core_utils/aws_partition.py
Normal file
55
litellm/litellm_core_utils/aws_partition.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
import re
|
||||
from types import MappingProxyType
|
||||
from typing import Final, NamedTuple
|
||||
|
||||
|
||||
class AwsPartition(NamedTuple):
|
||||
partition: str
|
||||
dns_suffix: str
|
||||
|
||||
|
||||
_COMMERCIAL_PARTITION: Final = AwsPartition(partition="aws", dns_suffix="amazonaws.com")
|
||||
|
||||
_PARTITIONS_BY_REGION_PREFIX: Final = MappingProxyType(
|
||||
{
|
||||
"cn-": AwsPartition(partition="aws-cn", dns_suffix="amazonaws.com.cn"),
|
||||
"us-gov-": AwsPartition(partition="aws-us-gov", dns_suffix="amazonaws.com"),
|
||||
"us-isob-": AwsPartition(partition="aws-iso-b", dns_suffix="sc2s.sgov.gov"),
|
||||
"us-isof-": AwsPartition(partition="aws-iso-f", dns_suffix="csp.hci.ic.gov"),
|
||||
"us-iso-": AwsPartition(partition="aws-iso", dns_suffix="c2s.ic.gov"),
|
||||
"eu-isoe-": AwsPartition(partition="aws-iso-e", dns_suffix="cloud.adc-e.uk"),
|
||||
}
|
||||
)
|
||||
|
||||
_BEDROCK_ARN_PATTERN: Final = re.compile(r"arn:aws(?:-[a-z0-9-]+)?:bedrock")
|
||||
_BEDROCK_ARN_PREFIX_PATTERN: Final = re.compile(r"\Aarn:aws(?:-[a-z0-9-]+)?:bedrock:")
|
||||
_AWS_ARN_PATTERN: Final = re.compile(r"arn:aws(?:-[a-z0-9-]+)?:")
|
||||
|
||||
|
||||
def get_aws_partition(aws_region_name: str | None) -> AwsPartition:
|
||||
if not aws_region_name:
|
||||
return _COMMERCIAL_PARTITION
|
||||
return next(
|
||||
(partition for prefix, partition in _PARTITIONS_BY_REGION_PREFIX.items() if aws_region_name.startswith(prefix)),
|
||||
_COMMERCIAL_PARTITION,
|
||||
)
|
||||
|
||||
|
||||
def get_aws_dns_suffix(aws_region_name: str | None) -> str:
|
||||
return get_aws_partition(aws_region_name).dns_suffix
|
||||
|
||||
|
||||
def get_aws_arn_prefix(aws_region_name: str | None) -> str:
|
||||
return f"arn:{get_aws_partition(aws_region_name).partition}:"
|
||||
|
||||
|
||||
def contains_bedrock_arn(value: str) -> bool:
|
||||
return _BEDROCK_ARN_PATTERN.search(value) is not None
|
||||
|
||||
|
||||
def is_bedrock_arn(value: str) -> bool:
|
||||
return _BEDROCK_ARN_PREFIX_PATTERN.match(value) is not None
|
||||
|
||||
|
||||
def contains_aws_arn(value: str) -> bool:
|
||||
return _AWS_ARN_PATTERN.search(value) is not None
|
||||
|
|
@ -2,6 +2,7 @@ from collections.abc import Mapping
|
|||
from typing import Final
|
||||
|
||||
import litellm
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
|
||||
|
||||
def _form_field_value(value: object) -> str:
|
||||
|
|
@ -13,18 +14,31 @@ def _form_field_value(value: object) -> str:
|
|||
|
||||
|
||||
def _flatten_form_field(key: str, value: object) -> tuple[tuple[str, str], ...]:
|
||||
if isinstance(value, Mapping):
|
||||
return tuple(
|
||||
item for subkey, subvalue in value.items() for item in _flatten_form_field(f"{key}[{subkey}]", subvalue)
|
||||
)
|
||||
if isinstance(value, (list, tuple)):
|
||||
return tuple(item for entry in value for item in _flatten_form_field(f"{key}[]", entry))
|
||||
if value is None:
|
||||
return ()
|
||||
serialized: Final = _form_field_value(value)
|
||||
if not serialized:
|
||||
return ()
|
||||
return ((key, serialized),)
|
||||
pending_fields: Final[ # mutable-ok: depth-capped stack walks nested JSON into multipart names
|
||||
list[tuple[str, object, int]]
|
||||
] = [ # mutable-ok: depth-capped stack walks nested JSON into multipart names
|
||||
(key, value, 0)
|
||||
]
|
||||
flat_fields: Final[list[tuple[str, str]]] = [] # mutable-ok: local accumulator
|
||||
while pending_fields:
|
||||
current_key, current_value, depth = pending_fields.pop()
|
||||
if depth > DEFAULT_MAX_RECURSE_DEPTH:
|
||||
raise ValueError("form field nesting exceeds max depth")
|
||||
if isinstance(current_value, Mapping):
|
||||
pending_fields.extend(
|
||||
(f"{current_key}[{subkey}]", subvalue, depth + 1)
|
||||
for subkey, subvalue in reversed(tuple(current_value.items()))
|
||||
)
|
||||
continue
|
||||
if isinstance(current_value, (list, tuple)):
|
||||
pending_fields.extend((f"{current_key}[]", entry, depth + 1) for entry in reversed(tuple(current_value)))
|
||||
continue
|
||||
if current_value is None:
|
||||
continue
|
||||
serialized = _form_field_value(current_value)
|
||||
if serialized:
|
||||
flat_fields.append((current_key, serialized))
|
||||
return tuple(flat_fields)
|
||||
|
||||
|
||||
def _is_form_scalar(value: object) -> bool:
|
||||
|
|
@ -32,23 +46,36 @@ def _is_form_scalar(value: object) -> bool:
|
|||
|
||||
|
||||
def _flatten_form_data_field(key: str, value: object) -> tuple[tuple[str, str | tuple[str, ...]], ...]:
|
||||
if isinstance(value, Mapping):
|
||||
return tuple(
|
||||
item
|
||||
for subkey, subvalue in value.items()
|
||||
for item in _flatten_form_data_field(f"{key}[{subkey}]", subvalue)
|
||||
)
|
||||
if isinstance(value, (list, tuple)):
|
||||
if all(_is_form_scalar(entry) for entry in value):
|
||||
serialized_fields: Final = tuple(field for entry in value if (field := _form_field_value(entry)))
|
||||
return ((key, serialized_fields),) if serialized_fields else ()
|
||||
return tuple(item for entry in value for item in _flatten_form_data_field(f"{key}[]", entry))
|
||||
if value is None:
|
||||
return ()
|
||||
serialized: Final = _form_field_value(value)
|
||||
if not serialized:
|
||||
return ()
|
||||
return ((key, serialized),)
|
||||
pending_fields: Final[ # mutable-ok: depth-capped stack walks nested JSON into multipart names
|
||||
list[tuple[str, object, int]]
|
||||
] = [ # mutable-ok: depth-capped stack walks nested JSON into multipart names
|
||||
(key, value, 0)
|
||||
]
|
||||
flat_fields: Final[list[tuple[str, str | tuple[str, ...]]]] = [] # mutable-ok: local accumulator
|
||||
while pending_fields:
|
||||
current_key, current_value, depth = pending_fields.pop()
|
||||
if depth > DEFAULT_MAX_RECURSE_DEPTH:
|
||||
raise ValueError("form field nesting exceeds max depth")
|
||||
if isinstance(current_value, Mapping):
|
||||
pending_fields.extend(
|
||||
(f"{current_key}[{subkey}]", subvalue, depth + 1)
|
||||
for subkey, subvalue in reversed(tuple(current_value.items()))
|
||||
)
|
||||
continue
|
||||
if isinstance(current_value, (list, tuple)):
|
||||
if all(_is_form_scalar(entry) for entry in current_value):
|
||||
serialized_fields = tuple(field for entry in current_value if (field := _form_field_value(entry)))
|
||||
if serialized_fields:
|
||||
flat_fields.append((current_key, serialized_fields))
|
||||
continue
|
||||
pending_fields.extend((f"{current_key}[]", entry, depth + 1) for entry in reversed(tuple(current_value)))
|
||||
continue
|
||||
if current_value is None:
|
||||
continue
|
||||
serialized = _form_field_value(current_value)
|
||||
if serialized:
|
||||
flat_fields.append((current_key, serialized))
|
||||
return tuple(flat_fields)
|
||||
|
||||
|
||||
def flatten_form_field_values(*sources: Mapping[str, object] | None) -> tuple[tuple[str, str | tuple[str, ...]], ...]:
|
||||
|
|
|
|||
|
|
@ -2,9 +2,21 @@
|
|||
Utility functions for ModelResponse and ModelResponseStream objects.
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from litellm.types.utils import Delta, ModelResponseBase, ModelResponseStream
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.types.utils import Delta, ModelResponseBase, ModelResponseStream, StreamingChoices
|
||||
|
||||
|
||||
class _AttributeView(TypedDict):
|
||||
value: ReadOnly[object]
|
||||
|
||||
|
||||
def _attribute_of(source: object, name: str) -> object:
|
||||
attribute: Final[_AttributeView] = {"value": getattr(source, name)}
|
||||
return attribute["value"]
|
||||
|
||||
|
||||
def is_model_response_stream_empty(model_response: ModelResponseStream) -> bool:
|
||||
|
|
@ -40,10 +52,10 @@ def is_model_response_stream_empty(model_response: ModelResponseStream) -> bool:
|
|||
return False
|
||||
|
||||
# Check model_extra for dynamically added fields (this is where Pydantic stores them)
|
||||
if hasattr(model_response, "model_extra") and model_response.model_extra:
|
||||
for extra_field_name, extra_field_value in model_response.model_extra.items():
|
||||
if _has_meaningful_content(extra_field_value):
|
||||
return False
|
||||
stream_extra_fields: Final[Mapping[str, object]] = model_response.model_extra or {}
|
||||
for extra_field_value in stream_extra_fields.values():
|
||||
if _has_meaningful_content(extra_field_value):
|
||||
return False
|
||||
|
||||
# Check for any non-base fields that are set
|
||||
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
|
||||
|
|
@ -57,7 +69,7 @@ def is_model_response_stream_empty(model_response: ModelResponseStream) -> bool:
|
|||
continue
|
||||
|
||||
# Check if any other field has meaningful content
|
||||
model_response_value = getattr(model_response, model_response_field, None)
|
||||
model_response_value: object = getattr(model_response, model_response_field, None)
|
||||
if _has_meaningful_content(model_response_value):
|
||||
return False
|
||||
|
||||
|
|
@ -71,7 +83,7 @@ def is_model_response_stream_empty(model_response: ModelResponseStream) -> bool:
|
|||
return True
|
||||
|
||||
|
||||
def _has_meaningful_content(value: Any) -> bool:
|
||||
def _has_meaningful_content(value: object) -> bool:
|
||||
"""
|
||||
Check if a value contains meaningful content.
|
||||
|
||||
|
|
@ -102,7 +114,7 @@ def _has_meaningful_content(value: Any) -> bool:
|
|||
return True
|
||||
|
||||
|
||||
def _is_choice_non_empty(choice: Any) -> bool:
|
||||
def _is_choice_non_empty(choice: StreamingChoices) -> bool:
|
||||
"""
|
||||
Deep check if a choice contains any meaningful content.
|
||||
|
||||
|
|
@ -113,41 +125,41 @@ def _is_choice_non_empty(choice: Any) -> bool:
|
|||
bool: True if the choice has meaningful content, False otherwise
|
||||
"""
|
||||
# Check finish_reason
|
||||
if hasattr(choice, "finish_reason") and choice.finish_reason is not None:
|
||||
if getattr(choice, "finish_reason", None) is not None:
|
||||
return True
|
||||
|
||||
# Check logprobs
|
||||
if hasattr(choice, "logprobs") and choice.logprobs is not None:
|
||||
if getattr(choice, "logprobs", None) is not None:
|
||||
return True
|
||||
|
||||
# Check enhancements (if present)
|
||||
if hasattr(choice, "enhancements") and choice.enhancements is not None:
|
||||
if getattr(choice, "enhancements", None) is not None:
|
||||
return True
|
||||
|
||||
# Deep check delta object
|
||||
if hasattr(choice, "delta") and choice.delta is not None:
|
||||
if _is_delta_non_empty(choice.delta):
|
||||
return True
|
||||
choice_delta: Final[Delta | None] = getattr(choice, "delta", None)
|
||||
if choice_delta is not None and _is_delta_non_empty(choice_delta):
|
||||
return True
|
||||
|
||||
# Check model_extra for dynamically added fields on the choice
|
||||
if hasattr(choice, "model_extra") and choice.model_extra:
|
||||
for extra_field_name, extra_field_value in choice.model_extra.items():
|
||||
# Skip certain structural fields that are just default/None placeholders
|
||||
if extra_field_name == "index" and extra_field_value == 0:
|
||||
continue
|
||||
if extra_field_name in {"finish_reason", "logprobs"} and extra_field_value is None:
|
||||
continue
|
||||
if extra_field_name == "delta":
|
||||
continue
|
||||
if _has_meaningful_content(extra_field_value):
|
||||
return True
|
||||
choice_extra_fields: Final[Mapping[str, object]] = choice.model_extra or {}
|
||||
for extra_field_name, extra_field_value in choice_extra_fields.items():
|
||||
# Skip certain structural fields that are just default/None placeholders
|
||||
if extra_field_name == "index" and extra_field_value == 0:
|
||||
continue
|
||||
if extra_field_name in {"finish_reason", "logprobs"} and extra_field_value is None:
|
||||
continue
|
||||
if extra_field_name == "delta":
|
||||
continue
|
||||
if _has_meaningful_content(extra_field_value):
|
||||
return True
|
||||
|
||||
# Check for any other non-standard fields on the choice
|
||||
for attr_name in dir(choice):
|
||||
# Skip private attributes, methods, and known empty fields
|
||||
if (
|
||||
attr_name.startswith("_")
|
||||
or callable(getattr(choice, attr_name))
|
||||
or callable(_attribute_of(choice, attr_name))
|
||||
or attr_name.startswith("model_")
|
||||
or attr_name
|
||||
in {
|
||||
|
|
@ -160,8 +172,8 @@ def _is_choice_non_empty(choice: Any) -> bool:
|
|||
):
|
||||
continue
|
||||
|
||||
attr_value = getattr(choice, attr_name, None)
|
||||
if _has_meaningful_content(attr_value):
|
||||
choice_attr_value: object = getattr(choice, attr_name, None)
|
||||
if _has_meaningful_content(choice_attr_value):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
@ -178,20 +190,20 @@ def _is_delta_non_empty(delta: Delta) -> bool:
|
|||
bool: True if the delta has meaningful content, False otherwise
|
||||
"""
|
||||
# Check model_extra for dynamically added fields (this is where Pydantic stores them)
|
||||
if hasattr(delta, "model_extra") and delta.model_extra:
|
||||
for extra_field_name, extra_field_value in delta.model_extra.items():
|
||||
# Even structural fields are meaningful if they have actual content
|
||||
if _has_meaningful_content(extra_field_value):
|
||||
return True
|
||||
delta_extra_fields: Final[Mapping[str, object]] = delta.model_extra or {}
|
||||
for extra_field_value in delta_extra_fields.values():
|
||||
# Even structural fields are meaningful if they have actual content
|
||||
if _has_meaningful_content(extra_field_value):
|
||||
return True
|
||||
|
||||
# Check all regular attributes of the delta object
|
||||
for attr_name in dir(delta):
|
||||
# Skip private attributes, methods, and Pydantic-specific fields
|
||||
if attr_name.startswith("_") or callable(getattr(delta, attr_name)) or attr_name.startswith("model_"):
|
||||
if attr_name.startswith("_") or callable(_attribute_of(delta, attr_name)) or attr_name.startswith("model_"):
|
||||
continue
|
||||
|
||||
attr_value = getattr(delta, attr_name, None)
|
||||
if _has_meaningful_content(attr_value):
|
||||
delta_attr_value: object = getattr(delta, attr_name, None)
|
||||
if _has_meaningful_content(delta_attr_value):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -21,13 +21,61 @@ Admins can opt out via two ``litellm`` globals (wired from proxy config):
|
|||
|
||||
import socket
|
||||
from ipaddress import ip_address, ip_network
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, Protocol
|
||||
from urllib.parse import quote, urlparse, urlunparse
|
||||
|
||||
import httpx
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
|
||||
_SockAddr = tuple[str, int] | tuple[str, int, int, int] | tuple[int, bytes]
|
||||
|
||||
|
||||
class _LocationHeaderView(TypedDict):
|
||||
location: ReadOnly[object]
|
||||
|
||||
|
||||
class _ResponseView(TypedDict):
|
||||
response: ReadOnly[httpx.Response]
|
||||
|
||||
|
||||
class _UrlFetcher(Protocol):
|
||||
"""The slice of ``httpx.Client`` / ``HTTPHandler`` that ``safe_get`` drives."""
|
||||
|
||||
def get(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
headers: dict[str, str] | None = None,
|
||||
follow_redirects: bool = False,
|
||||
) -> httpx.Response: ...
|
||||
|
||||
|
||||
class _AsyncUrlFetcher(Protocol):
|
||||
"""The slice of ``httpx.AsyncClient`` / ``AsyncHTTPHandler`` that ``async_safe_get`` drives."""
|
||||
|
||||
async def get(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
headers: dict[str, str] | None = None,
|
||||
follow_redirects: bool = False,
|
||||
) -> httpx.Response: ...
|
||||
|
||||
|
||||
class _FetcherView(TypedDict):
|
||||
fetcher: ReadOnly[_UrlFetcher]
|
||||
|
||||
|
||||
class _AsyncFetcherView(TypedDict):
|
||||
fetcher: ReadOnly[_AsyncUrlFetcher]
|
||||
|
||||
|
||||
class _CallerHeadersView(TypedDict):
|
||||
headers: ReadOnly[dict[str, str]]
|
||||
|
||||
|
||||
# Globally-routable IPs that are cloud-internal. Everything else
|
||||
# non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented by
|
||||
# Python's ``ipaddress`` module). This list only holds IPs that are
|
||||
|
|
@ -44,7 +92,7 @@ class SSRFError(ValueError):
|
|||
"""Raised when a URL targets a blocked network."""
|
||||
|
||||
|
||||
def encode_url_path_segment(value: Any, *, field_name: str = "path parameter") -> str:
|
||||
def encode_url_path_segment(value: object, *, field_name: str = "path parameter") -> str:
|
||||
"""Percent-encode one user-controlled URL path segment.
|
||||
|
||||
``urllib.parse.quote(..., safe="")`` intentionally leaves RFC 3986
|
||||
|
|
@ -64,7 +112,7 @@ def encode_url_path_segment(value: Any, *, field_name: str = "path parameter") -
|
|||
return quote(value_str, safe="")
|
||||
|
||||
|
||||
def encode_url_path_segments(value: Any, *, field_name: str = "path") -> str:
|
||||
def encode_url_path_segments(value: object, *, field_name: str = "path") -> str:
|
||||
"""Percent-encode a user-controlled URL path made of multiple segments.
|
||||
|
||||
Empty segments are rejected, so leading, trailing, or consecutive slashes
|
||||
|
|
@ -77,11 +125,7 @@ def encode_url_path_segments(value: Any, *, field_name: str = "path") -> str:
|
|||
if value_str == "":
|
||||
raise ValueError(f"{field_name} is required")
|
||||
|
||||
encoded_segments: Final = []
|
||||
for segment in value_str.split("/"):
|
||||
encoded_segments.append(encode_url_path_segment(segment, field_name=field_name))
|
||||
|
||||
return "/".join(encoded_segments)
|
||||
return "/".join(encode_url_path_segment(segment, field_name=field_name) for segment in value_str.split("/"))
|
||||
|
||||
|
||||
def _is_blocked_ip(addr: str) -> bool:
|
||||
|
|
@ -202,7 +246,7 @@ def _format_host_header(hostname: str, port: int, default_port: int) -> str:
|
|||
return f"{bracketed}:{port}"
|
||||
|
||||
|
||||
def _sockaddr_host(sockaddr: Any) -> str:
|
||||
def _sockaddr_host(sockaddr: _SockAddr) -> str:
|
||||
"""Return the host element of a ``getaddrinfo`` sockaddr as ``str``.
|
||||
|
||||
``getaddrinfo`` with ``IPPROTO_TCP`` returns AF_INET / AF_INET6 sockaddrs
|
||||
|
|
@ -285,8 +329,8 @@ def validate_url(url: str) -> tuple[str, str]:
|
|||
raise SSRFError(f"No addresses found for '{hostname}'")
|
||||
|
||||
if not is_allowlisted:
|
||||
for family, type_, proto, canonname, sockaddr in addrinfo:
|
||||
resolved_ip = _sockaddr_host(sockaddr)
|
||||
for addrinfo_entry in addrinfo:
|
||||
resolved_ip = _sockaddr_host(addrinfo_entry[4])
|
||||
if _is_blocked_ip(resolved_ip):
|
||||
raise SSRFError(
|
||||
f"URL targets a blocked address ({resolved_ip}). "
|
||||
|
|
@ -363,9 +407,10 @@ def assert_same_origin(candidate_url: str, expected_url: str) -> None:
|
|||
_MAX_REDIRECTS: Final = 10
|
||||
|
||||
|
||||
def _extract_redirect_url(response: Any, request_url: str) -> str:
|
||||
def _extract_redirect_url(response: httpx.Response, request_url: str) -> str:
|
||||
"""Extract and resolve the redirect target from a response's Location header."""
|
||||
location: Final = response.headers.get("location")
|
||||
header_view: Final[_LocationHeaderView] = {"location": response.headers.get("location")}
|
||||
location: Final = header_view["location"]
|
||||
if not isinstance(location, str) or not location:
|
||||
raise SSRFError("Redirect response has no Location header")
|
||||
# Resolve relative URLs against the request URL
|
||||
|
|
@ -393,14 +438,17 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
|||
"""
|
||||
if not getattr(litellm, "user_url_validation", True):
|
||||
kwargs.setdefault("follow_redirects", True)
|
||||
return client.get(url, **kwargs)
|
||||
unvalidated: Final[_ResponseView] = {"response": client.get(url, **kwargs)}
|
||||
return unvalidated["response"]
|
||||
fetcher_view: Final[_FetcherView] = {"fetcher": client}
|
||||
fetcher: Final = fetcher_view["fetcher"]
|
||||
kwargs.pop("follow_redirects", None)
|
||||
caller_headers: Final = kwargs.pop("headers", {})
|
||||
headers_view: Final[_CallerHeadersView] = {"headers": kwargs.pop("headers", {})}
|
||||
for _ in range(_MAX_REDIRECTS):
|
||||
validated_url, original_host = validate_url(url)
|
||||
response = client.get(
|
||||
response = fetcher.get(
|
||||
validated_url,
|
||||
headers={**caller_headers, "Host": original_host},
|
||||
headers={**headers_view["headers"], "Host": original_host},
|
||||
follow_redirects=False,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -416,14 +464,17 @@ async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
|||
"""Async version of safe_get."""
|
||||
if not getattr(litellm, "user_url_validation", True):
|
||||
kwargs.setdefault("follow_redirects", True)
|
||||
return await client.get(url, **kwargs)
|
||||
unvalidated: Final[_ResponseView] = {"response": await client.get(url, **kwargs)}
|
||||
return unvalidated["response"]
|
||||
fetcher_view: Final[_AsyncFetcherView] = {"fetcher": client}
|
||||
fetcher: Final = fetcher_view["fetcher"]
|
||||
kwargs.pop("follow_redirects", None)
|
||||
caller_headers: Final = kwargs.pop("headers", {})
|
||||
headers_view: Final[_CallerHeadersView] = {"headers": kwargs.pop("headers", {})}
|
||||
for _ in range(_MAX_REDIRECTS):
|
||||
validated_url, original_host = validate_url(url)
|
||||
response = await client.get(
|
||||
response = await fetcher.get(
|
||||
validated_url,
|
||||
headers={**caller_headers, "Host": original_host},
|
||||
headers={**headers_view["headers"], "Host": original_host},
|
||||
follow_redirects=False,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ A2A Protocol Transformation for LiteLLM
|
|||
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -20,6 +20,11 @@ from ..common_utils import (
|
|||
)
|
||||
from .streaming_iterator import A2AModelResponseIterator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class A2AConfig(BaseConfig):
|
||||
"""
|
||||
|
|
@ -246,12 +251,12 @@ class A2AConfig(BaseConfig):
|
|||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
request_data: dict,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -169,7 +171,7 @@ class AimlImageGenerationConfig(BaseImageGenerationConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -16,6 +16,8 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.utils import Choices, ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -66,7 +68,7 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Translate from OpenAI's `/v1/chat/completions` to Amazon Nova's `/v1/chat/completions`
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -16,6 +16,9 @@ from litellm.types.utils import ModelResponse
|
|||
|
||||
from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AmazonNovaChatConfig(OpenAILikeChatConfig):
|
||||
max_completion_tokens: int | None = None
|
||||
|
|
@ -83,7 +86,7 @@ class AmazonNovaChatConfig(OpenAILikeChatConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest
|
|||
from litellm.types.utils import LiteLLMBatch, LlmProviders, ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
LoggingClass = LiteLLMLoggingObj
|
||||
|
|
@ -261,7 +263,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -92,6 +92,8 @@ from ..common_utils import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
LoggingClass = LiteLLMLoggingObj
|
||||
|
|
@ -2575,7 +2577,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ Litellm provider slug: `anthropic_text/<model_name>`
|
|||
import json
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -32,6 +32,9 @@ from litellm.types.utils import (
|
|||
Usage,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AnthropicTextError(BaseLLMException):
|
||||
def __init__(self, status_code, message):
|
||||
|
|
@ -182,7 +185,7 @@ class AnthropicTextConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: str,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
@ -202,9 +205,10 @@ class AnthropicTextConfig(BaseConfig):
|
|||
model_response.choices[0].finish_reason = completion_response["stop_reason"]
|
||||
|
||||
## CALCULATING USAGE
|
||||
prompt_tokens: Final = len(encoding.encode(prompt)) ##[TODO] use the anthropic tokenizer here
|
||||
tokenizer: Final = encoding if encoding is not None else litellm.encoding
|
||||
prompt_tokens: Final = len(tokenizer.encode(prompt)) ##[TODO] use the anthropic tokenizer here
|
||||
completion_tokens: Final = len(
|
||||
encoding.encode(model_response["choices"][0]["message"].get("content", ""))
|
||||
tokenizer.encode(model_response["choices"][0]["message"].get("content", ""))
|
||||
) ##[TODO] use the anthropic tokenizer here
|
||||
|
||||
model_response.created = int(time.time())
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
import inspect
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.anthropic import AppliedEdit
|
||||
|
|
@ -11,7 +11,13 @@ from .constants import CLEAR_TOOL_USES_EDIT_TYPE, COMPACT_EDIT_TYPE
|
|||
from .editors import apply_clear_tool_uses_20250919, apply_compact_20260112
|
||||
from .result import PolyfillResult
|
||||
|
||||
EditorFn = Callable[..., Any]
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.router import Router
|
||||
|
||||
EditorResult: TypeAlias = "PolyfillResult | tuple[list[dict[str, object]], AppliedEdit | None]"
|
||||
|
||||
EditorFn: TypeAlias = "Callable[..., EditorResult | Awaitable[EditorResult]]"
|
||||
|
||||
_EDITOR_REGISTRY: Final[dict[str, EditorFn]] = {
|
||||
CLEAR_TOOL_USES_EDIT_TYPE: apply_clear_tool_uses_20250919,
|
||||
|
|
@ -19,23 +25,31 @@ _EDITOR_REGISTRY: Final[dict[str, EditorFn]] = {
|
|||
}
|
||||
|
||||
|
||||
def _normalize_spec(
|
||||
spec: dict[str, Any] | list[dict[str, Any]] | None,
|
||||
) -> list[dict[str, Any]] | None:
|
||||
"""Accept Anthropic-native dict form or OpenAI list form; return edits list."""
|
||||
if isinstance(spec, list):
|
||||
# Local import to avoid an import cycle at module load.
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
spec = AnthropicConfig.map_openai_context_management_to_anthropic(spec)
|
||||
|
||||
edits: Final = spec.get("edits") if isinstance(spec, dict) else None
|
||||
def _edits_from(normalized: dict[str, object] | None) -> list[dict[str, object]] | None:
|
||||
edits: Final = normalized.get("edits") if isinstance(normalized, dict) else None
|
||||
if not edits or not isinstance(edits, list):
|
||||
return None
|
||||
return [edit for edit in edits if isinstance(edit, dict)]
|
||||
|
||||
|
||||
def _wrap_editor_return(raw: Any, *, fallback_system: Any) -> PolyfillResult:
|
||||
def _normalize_spec(
|
||||
spec: dict[str, object] | list[dict[str, object]] | None,
|
||||
) -> list[dict[str, object]] | None:
|
||||
"""Accept Anthropic-native dict form or OpenAI list form; return edits list."""
|
||||
if isinstance(spec, list):
|
||||
# Local import to avoid an import cycle at module load.
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
return _edits_from(AnthropicConfig.map_openai_context_management_to_anthropic(spec))
|
||||
|
||||
return _edits_from(spec)
|
||||
|
||||
|
||||
def _wrap_editor_return(
|
||||
raw: EditorResult,
|
||||
*,
|
||||
fallback_system: str | list[dict[str, object]] | None,
|
||||
) -> PolyfillResult:
|
||||
"""Coerce an editor's native return shape into a ``PolyfillResult``.
|
||||
|
||||
v0 sync editors (e.g. ``clear_tool_uses_20250919``) return a 2-tuple
|
||||
|
|
@ -46,7 +60,7 @@ def _wrap_editor_return(raw: Any, *, fallback_system: Any) -> PolyfillResult:
|
|||
return raw
|
||||
# Legacy 2-tuple return — sync editors don't mutate ``system``, so
|
||||
# carry the caller's value forward.
|
||||
messages, applied = cast(tuple[list[dict[str, Any]], Any], raw)
|
||||
messages, applied = raw
|
||||
return PolyfillResult(
|
||||
messages=messages,
|
||||
system=fallback_system,
|
||||
|
|
@ -57,13 +71,13 @@ def _wrap_editor_return(raw: Any, *, fallback_system: Any) -> PolyfillResult:
|
|||
async def apply_context_management(
|
||||
*,
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]] | None,
|
||||
system: Any,
|
||||
context_management_spec: dict[str, Any] | list[dict[str, Any]] | None,
|
||||
litellm_metadata: dict[str, Any] | None = None,
|
||||
llm_router: Any = None,
|
||||
user_api_key_auth: Any = None,
|
||||
messages: list[dict[str, object]],
|
||||
tools: list[dict[str, object]] | None,
|
||||
system: str | list[dict[str, object]] | None,
|
||||
context_management_spec: dict[str, object] | list[dict[str, object]] | None,
|
||||
litellm_metadata: dict[str, object] | None = None,
|
||||
llm_router: "Router | None" = None,
|
||||
user_api_key_auth: "UserAPIKeyAuth | None" = None,
|
||||
) -> PolyfillResult:
|
||||
"""Run edits in order; return a single ``PolyfillResult``.
|
||||
|
||||
|
|
@ -92,22 +106,30 @@ async def apply_context_management(
|
|||
)
|
||||
continue
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": current_messages,
|
||||
"tools": tools,
|
||||
"system": current_system,
|
||||
"edit_spec": edit_spec,
|
||||
}
|
||||
# Only async editors accept these — passing them to sync v0 editors
|
||||
# would break their signature.
|
||||
if inspect.iscoroutinefunction(editor):
|
||||
kwargs["litellm_metadata"] = litellm_metadata
|
||||
kwargs["llm_router"] = llm_router
|
||||
kwargs["user_api_key_auth"] = user_api_key_auth
|
||||
raw_result = await cast(Callable[..., Awaitable[Any]], editor)(**kwargs)
|
||||
else:
|
||||
raw_result = editor(**kwargs)
|
||||
editor_is_async = inspect.iscoroutinefunction(editor)
|
||||
editor_return = (
|
||||
editor(
|
||||
model=model,
|
||||
messages=current_messages,
|
||||
tools=tools,
|
||||
system=current_system,
|
||||
edit_spec=edit_spec,
|
||||
litellm_metadata=litellm_metadata,
|
||||
llm_router=llm_router,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if editor_is_async
|
||||
else editor(
|
||||
model=model,
|
||||
messages=current_messages,
|
||||
tools=tools,
|
||||
system=current_system,
|
||||
edit_spec=edit_spec,
|
||||
)
|
||||
)
|
||||
raw_result = editor_return if isinstance(editor_return, (PolyfillResult, tuple)) else await editor_return
|
||||
|
||||
result = _wrap_editor_return(raw_result, fallback_system=current_system)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.anthropic import AppliedEdit
|
||||
|
|
@ -14,7 +16,18 @@ from ..constants import (
|
|||
from ..placeholders import build_cleared_tool_result_content
|
||||
|
||||
|
||||
def _count_tool_uses(messages: list[dict[str, Any]]) -> int:
|
||||
class ClearToolUsesEditSpec(TypedDict, total=False):
|
||||
"""The ``clear_tool_uses_20250919`` entry of a ``context_management`` spec."""
|
||||
|
||||
type: ReadOnly[str]
|
||||
trigger: ReadOnly[dict[str, object]]
|
||||
keep: ReadOnly[dict[str, object]]
|
||||
clear_at_least: ReadOnly[object]
|
||||
exclude_tools: ReadOnly[object]
|
||||
clear_tool_inputs: ReadOnly[object]
|
||||
|
||||
|
||||
def _count_tool_uses(messages: list[dict[str, object]]) -> int:
|
||||
"""Return the number of tool_use content blocks across all messages.
|
||||
|
||||
Only counts blocks with a string ``id`` to stay consistent with
|
||||
|
|
@ -32,7 +45,7 @@ def _count_tool_uses(messages: list[dict[str, Any]]) -> int:
|
|||
return count
|
||||
|
||||
|
||||
def _collect_tool_use_ids_in_order(messages: list[dict[str, Any]]) -> list[str]:
|
||||
def _collect_tool_use_ids_in_order(messages: list[dict[str, object]]) -> list[str]:
|
||||
"""Return tool_use ids in the chronological order they appear in messages."""
|
||||
ids: Final[list[str]] = []
|
||||
for msg in messages:
|
||||
|
|
@ -47,10 +60,10 @@ def _collect_tool_use_ids_in_order(messages: list[dict[str, Any]]) -> list[str]:
|
|||
|
||||
|
||||
def _trigger_met(
|
||||
trigger: dict[str, Any],
|
||||
trigger: dict[str, object],
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]] | None,
|
||||
messages: list[dict[str, object]],
|
||||
tools: list[dict[str, object]] | None,
|
||||
) -> tuple[bool, int | None]:
|
||||
"""Return (trigger_met, input_tokens if counted for reuse)."""
|
||||
trigger_type: Final = trigger.get("type", "input_tokens")
|
||||
|
|
@ -73,7 +86,7 @@ def _trigger_met(
|
|||
return current_tokens > threshold, current_tokens
|
||||
|
||||
|
||||
def _resolve_keep_count(keep: dict[str, Any]) -> int:
|
||||
def _resolve_keep_count(keep: dict[str, object]) -> int:
|
||||
keep_type: Final = keep.get("type", "tool_uses")
|
||||
if keep_type != "tool_uses":
|
||||
return DEFAULT_KEEP_TOOL_USES
|
||||
|
|
@ -84,7 +97,7 @@ def _resolve_keep_count(keep: dict[str, Any]) -> int:
|
|||
|
||||
|
||||
def _last_completed_tool_use_id(
|
||||
messages: list[dict[str, Any]],
|
||||
messages: list[dict[str, object]],
|
||||
) -> str | None:
|
||||
"""Latest completed tool_result id; never cleared."""
|
||||
last_id: str | None = None
|
||||
|
|
@ -99,17 +112,19 @@ def _last_completed_tool_use_id(
|
|||
return last_id
|
||||
|
||||
|
||||
def _clear_tool_results(messages: list[dict[str, Any]], ids_to_clear: set) -> tuple[list[dict[str, Any]], int]:
|
||||
def _clear_tool_results(
|
||||
messages: list[dict[str, object]], ids_to_clear: set[str]
|
||||
) -> tuple[list[dict[str, object]], int]:
|
||||
"""Clear matching tool_result content; return (messages, cleared_count)."""
|
||||
cleared = 0
|
||||
new_messages: Final[list[dict[str, Any]]] = []
|
||||
new_messages: Final[list[dict[str, object]]] = []
|
||||
for msg in messages:
|
||||
content = msg.get("content")
|
||||
if not isinstance(content, list):
|
||||
new_messages.append(msg)
|
||||
continue
|
||||
|
||||
new_blocks: list[Any] = []
|
||||
new_blocks: list[object] = []
|
||||
mutated = False
|
||||
for block in content:
|
||||
if (
|
||||
|
|
@ -138,11 +153,11 @@ def _clear_tool_results(messages: list[dict[str, Any]], ids_to_clear: set) -> tu
|
|||
def apply_clear_tool_uses_20250919(
|
||||
*,
|
||||
model: str,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]] | None,
|
||||
system: Any,
|
||||
edit_spec: dict[str, Any],
|
||||
) -> tuple[list[dict[str, Any]], AppliedEdit | None]:
|
||||
messages: list[dict[str, object]],
|
||||
tools: list[dict[str, object]] | None,
|
||||
system: str | list[dict[str, object]] | None,
|
||||
edit_spec: ClearToolUsesEditSpec,
|
||||
) -> tuple[list[dict[str, object]], AppliedEdit | None]:
|
||||
"""Apply clear_tool_uses; return (messages, AppliedEdit or None)."""
|
||||
ignored_knobs = [knob for knob in ("clear_at_least", "exclude_tools", "clear_tool_inputs") if knob in edit_spec]
|
||||
for ignored_knob in ignored_knobs:
|
||||
|
|
@ -153,11 +168,11 @@ def apply_clear_tool_uses_20250919(
|
|||
CLEAR_TOOL_USES_EDIT_TYPE,
|
||||
)
|
||||
|
||||
trigger: Final = edit_spec.get("trigger") or {
|
||||
trigger: Final[dict[str, object]] = edit_spec.get("trigger") or {
|
||||
"type": "input_tokens",
|
||||
"value": DEFAULT_INPUT_TOKENS_TRIGGER,
|
||||
}
|
||||
keep: Final = edit_spec.get("keep") or {
|
||||
keep: Final[dict[str, object]] = edit_spec.get("keep") or {
|
||||
"type": "tool_uses",
|
||||
"value": DEFAULT_KEEP_TOOL_USES,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -18,11 +18,14 @@ import asyncio
|
|||
import contextlib
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0
|
||||
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = (
|
||||
b"event: error\n"
|
||||
|
|
@ -181,7 +184,7 @@ class AgenticAnthropicStreamingIterator:
|
|||
messages: list[dict],
|
||||
anthropic_messages_provider_config: Any,
|
||||
anthropic_messages_optional_request_params: dict,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str,
|
||||
kwargs: dict,
|
||||
hold_back: bool = False,
|
||||
|
|
|
|||
|
|
@ -571,7 +571,34 @@ def anthropic_messages_handler(
|
|||
anthropic_messages_provider_config = OpenAILikeAnthropicMessagesConfig()
|
||||
if anthropic_messages_provider_config is None:
|
||||
# Route to Responses API for OpenAI / Azure, chat/completions for everything else.
|
||||
_shared_kwargs: Final = dict(
|
||||
if _should_route_to_responses_api(custom_llm_provider, original_model, model):
|
||||
return LiteLLMMessagesToResponsesAPIHandler.anthropic_messages_handler(
|
||||
max_tokens=max_tokens,
|
||||
messages=messages,
|
||||
model=original_model,
|
||||
metadata=metadata,
|
||||
stop_sequences=stop_sequences,
|
||||
stream=stream,
|
||||
system=system,
|
||||
temperature=temperature,
|
||||
thinking=thinking,
|
||||
tool_choice=tool_choice,
|
||||
tools=tools,
|
||||
top_k=top_k,
|
||||
top_p=top_p,
|
||||
_is_async=is_async,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# The in-gateway context_management polyfill runs inside
|
||||
# ``async_anthropic_messages_handler`` so it can ``await`` the
|
||||
# summarization model for ``compact_20260112``. ``context_management``
|
||||
# is passed through as a regular kwarg.
|
||||
return LiteLLMMessagesToCompletionTransformationHandler.anthropic_messages_handler(
|
||||
max_tokens=max_tokens,
|
||||
messages=messages,
|
||||
model=original_model,
|
||||
|
|
@ -592,16 +619,6 @@ def anthropic_messages_handler(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
if _should_route_to_responses_api(custom_llm_provider, original_model, model):
|
||||
return LiteLLMMessagesToResponsesAPIHandler.anthropic_messages_handler(**_shared_kwargs)
|
||||
|
||||
# The in-gateway context_management polyfill runs inside
|
||||
# ``async_anthropic_messages_handler`` so it can ``await`` the
|
||||
# summarization model for ``compact_20260112``. ``context_management``
|
||||
# is passed through as a regular kwarg.
|
||||
return LiteLLMMessagesToCompletionTransformationHandler.anthropic_messages_handler(
|
||||
**_shared_kwargs,
|
||||
)
|
||||
|
||||
if custom_llm_provider is None:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -342,7 +342,7 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
self.start_time = datetime.now()
|
||||
self.completion_start_time: datetime | None = None
|
||||
|
||||
async def _handle_streaming_logging(self, collected_chunks: list[bytes]):
|
||||
async def _handle_streaming_logging(self, collected_chunks: list[bytes], *, stream_teardown: bool = False):
|
||||
"""Handle the logging after all chunks have been collected."""
|
||||
from litellm.proxy.pass_through_endpoints.streaming_handler import (
|
||||
PassThroughStreamingHandler,
|
||||
|
|
@ -354,21 +354,26 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
if self.completion_start_time is not None:
|
||||
self.litellm_logging_obj.completion_start_time = self.completion_start_time
|
||||
self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time
|
||||
logging_coroutine: Final = PassThroughStreamingHandler._route_streaming_logging_to_handler(
|
||||
litellm_logging_obj=self.litellm_logging_obj,
|
||||
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
|
||||
url_route="/v1/messages",
|
||||
request_body=self.request_body or {},
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=self.start_time,
|
||||
raw_bytes=collected_chunks,
|
||||
end_time=end_time,
|
||||
)
|
||||
deferred_dispatch_armed: Final = (
|
||||
getattr(self.litellm_logging_obj, "_on_deferred_stream_complete", None) is not None
|
||||
)
|
||||
if deferred_dispatch_armed and not stream_teardown:
|
||||
self.litellm_logging_obj._deferred_stream_complete_args = (logging_coroutine,)
|
||||
return
|
||||
# Enqueue on the rooted logging worker rather than asyncio.create_task:
|
||||
# this also runs during generator teardown after a client disconnect,
|
||||
# where an unrooted task could be garbage-collected before it bills.
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine=PassThroughStreamingHandler._route_streaming_logging_to_handler(
|
||||
litellm_logging_obj=self.litellm_logging_obj,
|
||||
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
|
||||
url_route="/v1/messages",
|
||||
request_body=self.request_body or {},
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=self.start_time,
|
||||
raw_bytes=collected_chunks,
|
||||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=logging_coroutine)
|
||||
|
||||
def get_async_streaming_response_iterator(
|
||||
self,
|
||||
|
|
@ -433,7 +438,7 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
# post-loop logging below never runs and the tokens already streamed
|
||||
# (and billed by the provider) would never reach spend tracking. See LIT-5839.
|
||||
if collected_chunks:
|
||||
await self._handle_streaming_logging(collected_chunks)
|
||||
await self._handle_streaming_logging(collected_chunks, stream_teardown=True)
|
||||
raise
|
||||
|
||||
if not saw_terminal_event:
|
||||
|
|
|
|||
|
|
@ -5,10 +5,11 @@ Used when the target model is an OpenAI or Azure model.
|
|||
"""
|
||||
|
||||
from collections.abc import AsyncIterator, Coroutine, Mapping
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, TypeAlias
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.anthropic import (
|
||||
AllAnthropicMessageValues,
|
||||
AllAnthropicToolsValues,
|
||||
AnthropicMessagesRequest,
|
||||
AnthropicOutputConfig,
|
||||
|
|
@ -23,6 +24,8 @@ from ..utils import local_model_name
|
|||
from .streaming_iterator import AnthropicResponsesStreamWrapper
|
||||
from .transformation import LiteLLMAnthropicToResponsesAPIAdapter
|
||||
|
||||
AnthropicRequestMessages: TypeAlias = list[AllAnthropicMessageValues] | list[dict[str, object]]
|
||||
|
||||
_ADAPTER: Final = LiteLLMAnthropicToResponsesAPIAdapter()
|
||||
|
||||
|
||||
|
|
@ -34,22 +37,22 @@ def _forwarded_kwargs(extra_kwargs: Mapping[str, object] | None) -> Mapping[str,
|
|||
def _build_responses_kwargs(
|
||||
*,
|
||||
max_tokens: int,
|
||||
messages: list[dict],
|
||||
messages: AnthropicRequestMessages,
|
||||
model: str,
|
||||
context_management: dict | None = None,
|
||||
metadata: dict | None = None,
|
||||
context_management: dict[str, object] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
output_config: AnthropicOutputConfig | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict] | None = None,
|
||||
thinking: dict[str, object] | None = None,
|
||||
tool_choice: dict[str, object] | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict[str, object]] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
extra_kwargs: dict[str, Any] | None = None,
|
||||
extra_kwargs: Mapping[str, object] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Build the kwargs dict to pass directly to litellm.responses() / litellm.aresponses().
|
||||
|
|
@ -83,30 +86,32 @@ def _build_responses_kwargs(
|
|||
|
||||
anthropic_request: Final = AnthropicMessagesRequest(**request_data)
|
||||
responses_kwargs: Final = _ADAPTER.translate_request(anthropic_request)
|
||||
forwarded_kwargs: Final = _forwarded_kwargs(extra_kwargs)
|
||||
|
||||
# Normalize reasoning effort based on model capabilities
|
||||
# (e.g. "max" → "xhigh"/"high", "minimal" → "low" if unsupported)
|
||||
reasoning: Final = responses_kwargs.get("reasoning")
|
||||
if isinstance(reasoning, dict) and "effort" in reasoning:
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import (
|
||||
normalize_reasoning_effort_value,
|
||||
)
|
||||
if isinstance(reasoning, dict):
|
||||
effort: Final[object] = reasoning.get("effort")
|
||||
if isinstance(effort, str):
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import (
|
||||
normalize_reasoning_effort_value,
|
||||
)
|
||||
|
||||
effort: Final = reasoning["effort"]
|
||||
normalized: Final = normalize_reasoning_effort_value(
|
||||
effort,
|
||||
model=model,
|
||||
custom_llm_provider=(extra_kwargs or {}).get("custom_llm_provider"),
|
||||
)
|
||||
if normalized != effort:
|
||||
responses_kwargs["reasoning"] = {**reasoning, "effort": normalized}
|
||||
provider_hint: Final = forwarded_kwargs.get("custom_llm_provider")
|
||||
normalized: Final = normalize_reasoning_effort_value(
|
||||
effort,
|
||||
model=model,
|
||||
custom_llm_provider=provider_hint if isinstance(provider_hint, str) else None,
|
||||
)
|
||||
if normalized != effort:
|
||||
responses_kwargs["reasoning"] = {**reasoning, "effort": normalized}
|
||||
|
||||
if stream:
|
||||
responses_kwargs["stream"] = True
|
||||
|
||||
# Forward litellm-specific kwargs (api_key, api_base, logging obj, etc.)
|
||||
excluded: Final = {"anthropic_messages"}
|
||||
forwarded_kwargs: Final = _forwarded_kwargs(extra_kwargs)
|
||||
for key, value in forwarded_kwargs.items():
|
||||
if key == "litellm_logging_obj" and value is not None:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
|
|
@ -140,18 +145,18 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
@staticmethod
|
||||
async def async_anthropic_messages_handler(
|
||||
max_tokens: int,
|
||||
messages: list[dict],
|
||||
messages: AnthropicRequestMessages,
|
||||
model: str,
|
||||
context_management: dict | None = None,
|
||||
metadata: dict | None = None,
|
||||
context_management: dict[str, object] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
output_config: AnthropicOutputConfig | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict] | None = None,
|
||||
thinking: dict[str, object] | None = None,
|
||||
tool_choice: dict[str, object] | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict[str, object]] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
|
|
@ -193,18 +198,18 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
@staticmethod
|
||||
def anthropic_messages_handler(
|
||||
max_tokens: int,
|
||||
messages: list[dict],
|
||||
messages: AnthropicRequestMessages,
|
||||
model: str,
|
||||
context_management: dict | None = None,
|
||||
metadata: dict | None = None,
|
||||
context_management: dict[str, object] | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
output_config: AnthropicOutputConfig | None = None,
|
||||
stop_sequences: list[str] | None = None,
|
||||
stream: bool | None = False,
|
||||
system: str | None = None,
|
||||
temperature: float | None = None,
|
||||
thinking: dict | None = None,
|
||||
tool_choice: dict | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict] | None = None,
|
||||
thinking: dict[str, object] | None = None,
|
||||
tool_choice: dict[str, object] | None = None,
|
||||
tools: list[AllAnthropicToolsValues | dict[str, object]] | None = None,
|
||||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
|
|
|
|||
|
|
@ -2,9 +2,10 @@
|
|||
Anthropic Skills API configuration and transformations
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
|
|
@ -22,6 +23,8 @@ from litellm.types.llms.anthropic_skills import (
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
_RAW_JSON_PAYLOAD: Final = TypeAdapter(object)
|
||||
|
||||
|
||||
class AnthropicSkillsConfig(BaseSkillsAPIConfig):
|
||||
"""Anthropic-specific Skills API configuration"""
|
||||
|
|
@ -104,10 +107,10 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Skill:
|
||||
"""Transform Anthropic response to Skill object"""
|
||||
response_json: Final = raw_response.json()
|
||||
response_json: Final = _RAW_JSON_PAYLOAD.validate_python(raw_response.json())
|
||||
verbose_logger.debug("Transforming create skill response: %s", response_json)
|
||||
|
||||
return Skill(**response_json)
|
||||
return Skill.model_validate(response_json)
|
||||
|
||||
def transform_list_skills_request(
|
||||
self,
|
||||
|
|
@ -122,13 +125,12 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
|
|||
url: Final = self.get_complete_url(api_base=api_base, endpoint="skills")
|
||||
|
||||
# Build query parameters
|
||||
query_params: Final[dict[str, Any]] = {}
|
||||
if "limit" in list_params and list_params["limit"]:
|
||||
query_params["limit"] = list_params["limit"]
|
||||
if "page" in list_params and list_params["page"]:
|
||||
query_params["page"] = list_params["page"]
|
||||
if "source" in list_params and list_params["source"]:
|
||||
query_params["source"] = list_params["source"]
|
||||
limit: Final = list_params.get("limit")
|
||||
page: Final = list_params.get("page")
|
||||
source: Final = list_params.get("source")
|
||||
query_params: Final[dict[str, int | str]] = {
|
||||
key: value for key, value in (("limit", limit), ("page", page), ("source", source)) if value
|
||||
}
|
||||
|
||||
verbose_logger.debug(
|
||||
"List skills request made to Anthropic Skills endpoint with params: %s",
|
||||
|
|
@ -143,10 +145,10 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ListSkillsResponse:
|
||||
"""Transform Anthropic response to ListSkillsResponse"""
|
||||
response_json: Final = raw_response.json()
|
||||
response_json: Final = _RAW_JSON_PAYLOAD.validate_python(raw_response.json())
|
||||
verbose_logger.debug("Transforming list skills response: %s", response_json)
|
||||
|
||||
return ListSkillsResponse(**response_json)
|
||||
return ListSkillsResponse.model_validate(response_json)
|
||||
|
||||
def transform_get_skill_request(
|
||||
self,
|
||||
|
|
@ -168,10 +170,10 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> Skill:
|
||||
"""Transform Anthropic response to Skill object"""
|
||||
response_json: Final = raw_response.json()
|
||||
response_json: Final = _RAW_JSON_PAYLOAD.validate_python(raw_response.json())
|
||||
verbose_logger.debug("Transforming get skill response: %s", response_json)
|
||||
|
||||
return Skill(**response_json)
|
||||
return Skill.model_validate(response_json)
|
||||
|
||||
def transform_delete_skill_request(
|
||||
self,
|
||||
|
|
@ -193,7 +195,7 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> DeleteSkillResponse:
|
||||
"""Transform Anthropic response to DeleteSkillResponse"""
|
||||
response_json: Final = raw_response.json()
|
||||
response_json: Final = _RAW_JSON_PAYLOAD.validate_python(raw_response.json())
|
||||
verbose_logger.debug("Transforming delete skill response: %s", response_json)
|
||||
|
||||
return DeleteSkillResponse(**response_json)
|
||||
return DeleteSkillResponse.model_validate(response_json)
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Final, Union
|
|||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.llms.base_llm.text_to_speech.transformation import (
|
||||
BaseTextToSpeechConfig,
|
||||
TextToSpeechRequestData,
|
||||
|
|
@ -238,7 +239,7 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM):
|
|||
return api_base.rstrip("/") + "/v1/speech"
|
||||
|
||||
aws_region_name: Final = litellm_params.get("aws_region_name", self.DEFAULT_REGION)
|
||||
return f"https://polly.{aws_region_name}.amazonaws.com/v1/speech"
|
||||
return f"https://polly.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}/v1/speech"
|
||||
|
||||
def is_ssml_input(self, input: str) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from collections.abc import Coroutine
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from openai import AsyncAzureOpenAI, AzureOpenAI
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -16,6 +16,9 @@ from litellm.utils import (
|
|||
from .azure import AzureChatCompletion
|
||||
from .common_utils import AzureOpenAIError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class AzureAudioTranscription(AzureChatCompletion):
|
||||
def audio_transcriptions(
|
||||
|
|
@ -23,7 +26,7 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
model: str,
|
||||
audio_file: FileTypes,
|
||||
optional_params: dict,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
model_response: TranscriptionResponse,
|
||||
timeout: float,
|
||||
max_retries: int,
|
||||
|
|
@ -112,7 +115,7 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
data: dict,
|
||||
model_response: TranscriptionResponse,
|
||||
timeout: float,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
api_version: str | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@ from ...base_llm.chat.transformation import BaseConfig
|
|||
from ..common_utils import AzureOpenAIError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
LoggingClass = LiteLLMLoggingObj
|
||||
|
|
@ -271,7 +273,7 @@ class AzureOpenAIConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -193,7 +193,7 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
data: dict,
|
||||
timeout: Any,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: Any,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
max_retries: int,
|
||||
azure_ad_token: str | None = None,
|
||||
client=None, # this is the AsyncAzureOpenAI
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
|
|||
verbose_logger.debug("create_file_data=%s", create_file_data)
|
||||
response = await openai_client.files.create(**self._prepare_create_file_data(create_file_data))
|
||||
verbose_logger.debug("create_file_response=%s", response)
|
||||
return OpenAIFileObject(**response.model_dump())
|
||||
return OpenAIFileObject.model_validate(response.model_dump())
|
||||
|
||||
def create_file(
|
||||
self,
|
||||
|
|
@ -60,8 +60,8 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> OpenAIFileObject | Coroutine[Any, Any, OpenAIFileObject]:
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
) -> OpenAIFileObject | Coroutine[object, object, OpenAIFileObject]:
|
||||
openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
|
|
@ -84,7 +84,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
|
|||
response: Final = cast(AzureOpenAI | OpenAI, openai_client).files.create(
|
||||
**self._prepare_create_file_data(create_file_data)
|
||||
)
|
||||
return OpenAIFileObject(**response.model_dump())
|
||||
return OpenAIFileObject.model_validate(response.model_dump())
|
||||
|
||||
async def afile_content(
|
||||
self,
|
||||
|
|
@ -104,8 +104,8 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
|
|||
max_retries: int | None,
|
||||
api_version: str | None = None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
) -> HttpxBinaryResponseContent | Coroutine[Any, Any, HttpxBinaryResponseContent]:
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
) -> HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent]:
|
||||
openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
|
|
@ -150,7 +150,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
|
|||
max_retries: int | None,
|
||||
api_version: str | None = None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
):
|
||||
openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
litellm_params=litellm_params or {},
|
||||
|
|
@ -200,7 +200,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
|
|||
organization: str | None = None,
|
||||
api_version: str | None = None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
):
|
||||
openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
litellm_params=litellm_params or {},
|
||||
|
|
@ -252,7 +252,7 @@ class AzureOpenAIFilesAPI(BaseAzureLLM):
|
|||
purpose: str | None = None,
|
||||
api_version: str | None = None,
|
||||
client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = None,
|
||||
litellm_params: dict | None = None,
|
||||
litellm_params: dict[str, object] | None = None,
|
||||
):
|
||||
openai_client: AzureOpenAI | AsyncAzureOpenAI | OpenAI | AsyncOpenAI | None = self.get_azure_openai_client(
|
||||
litellm_params=litellm_params or {},
|
||||
|
|
|
|||
|
|
@ -34,6 +34,8 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.utils import ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
|
|
@ -295,7 +297,7 @@ class AzureAIAgentsConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ The Model Router is a special Azure AI deployment that automatically routes requ
|
|||
to the best available model. It has specific cost tracking requirements.
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from httpx import Response
|
||||
|
||||
|
|
@ -14,6 +14,9 @@ from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
|
|||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AzureModelRouterConfig(AzureAIStudioConfig):
|
||||
"""
|
||||
|
|
@ -56,7 +59,7 @@ class AzureModelRouterConfig(AzureAIStudioConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import copy
|
||||
import enum
|
||||
import re
|
||||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
|
@ -25,6 +25,9 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.types.utils import ModelResponse, ProviderField
|
||||
from litellm.utils import _add_path_to_api_base, supports_tool_choice
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AzureFoundryErrorStrings(str, enum.Enum):
|
||||
SET_EXTRA_PARAMETERS_TO_PASS_THROUGH = "Set extra-parameters to 'pass-through'"
|
||||
|
|
@ -258,7 +261,7 @@ class AzureAIStudioConfig(OpenAIConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from litellm.types.utils import ImageResponse
|
|||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
|
|
@ -199,7 +200,7 @@ class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import asyncio
|
|||
import re
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
|
|
@ -41,6 +41,9 @@ from litellm.llms.base_llm.ocr.transformation import (
|
|||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR: Final = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
|
||||
|
||||
|
||||
|
|
@ -676,7 +679,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
|
|
@ -751,7 +754,7 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
**kwargs,
|
||||
) -> OCRResponse:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import httpx
|
|||
import litellm
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.types.utils import ModelResponse, TextCompletionResponse
|
||||
|
||||
|
|
@ -19,7 +20,7 @@ class BaseLLM:
|
|||
response: httpx.Response,
|
||||
model_response: "ModelResponse",
|
||||
stream: bool,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
optional_params: dict,
|
||||
api_key: str,
|
||||
data: dict | str,
|
||||
|
|
@ -38,7 +39,7 @@ class BaseLLM:
|
|||
response: httpx.Response,
|
||||
model_response: "TextCompletionResponse",
|
||||
stream: bool,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
optional_params: dict,
|
||||
api_key: str,
|
||||
data: dict | str,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.utils import FileTypes, ModelResponse, TranscriptionResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -110,7 +112,7 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -4,9 +4,10 @@ Bridge for transforming API requests to another API requests
|
|||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm import LiteLLMLoggingObj, ModelResponse
|
||||
|
|
@ -38,7 +39,7 @@ class CompletionTransformationBridge(ABC):
|
|||
messages: list["AllMessageValues"],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> "ModelResponse":
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@ from litellm.types.llms.openai import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -342,7 +344,7 @@ class BaseConfig(ABC):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> "ModelResponse":
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUser
|
|||
from litellm.types.utils import ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -66,7 +68,7 @@ class BaseTextCompletionConfig(BaseConfig, ABC):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
|
|||
from litellm.types.utils import EmbeddingResponse, ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -78,7 +80,7 @@ class BaseEmbeddingConfig(BaseConfig, ABC):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ from litellm.types.utils import LlmProviders, ModelResponse
|
|||
from ..chat.transformation import BaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
from litellm.router import Router as _Router
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
|
@ -207,7 +209,7 @@ class BaseFilesConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -11,6 +11,8 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.utils import ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -91,7 +93,7 @@ class BaseImageGenerationConfig(ABC):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -17,6 +17,8 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -80,7 +82,7 @@ class BaseImageVariationConfig(BaseConfig, ABC):
|
|||
image: FileTypes,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
) -> ImageResponse:
|
||||
pass
|
||||
|
|
@ -96,7 +98,7 @@ class BaseImageVariationConfig(BaseConfig, ABC):
|
|||
image: FileTypes,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
) -> ImageResponse:
|
||||
pass
|
||||
|
|
@ -123,7 +125,7 @@ class BaseImageVariationConfig(BaseConfig, ABC):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from litellm.constants import (
|
|||
BEDROCK_MAX_POLICY_SIZE,
|
||||
STS_CREDENTIAL_EXPIRY_SAFETY_MARGIN_SECONDS,
|
||||
)
|
||||
from litellm.litellm_core_utils.aws_partition import contains_bedrock_arn, get_aws_dns_suffix
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.secret_managers.main import get_secret, get_secret_str
|
||||
|
||||
|
|
@ -348,7 +349,7 @@ class BaseAWSLLM:
|
|||
def _get_aws_region_from_model_arn(self, model: str | None) -> str | None:
|
||||
try:
|
||||
# First check if the string contains the expected prefix
|
||||
if not isinstance(model, str) or "arn:aws:bedrock" not in model:
|
||||
if not isinstance(model, str) or not contains_bedrock_arn(model):
|
||||
return None
|
||||
|
||||
# Split the ARN and check if we have enough parts
|
||||
|
|
@ -625,24 +626,29 @@ class BaseAWSLLM:
|
|||
return match.group(1) if match else None
|
||||
|
||||
@staticmethod
|
||||
def _resolve_sts_region(aws_sts_endpoint: str | None = None) -> str | None:
|
||||
"""STS signing region: parsed from aws_sts_endpoint else AWS_REGION / AWS_DEFAULT_REGION."""
|
||||
def _resolve_sts_region(
|
||||
aws_sts_endpoint: str | None = None,
|
||||
aws_region_name: str | None = None,
|
||||
) -> str | None:
|
||||
"""STS signing region: parsed from aws_sts_endpoint, else AWS_REGION / AWS_DEFAULT_REGION, else the configured aws_region_name."""
|
||||
return (
|
||||
BaseAWSLLM._parse_sts_region_from_endpoint(aws_sts_endpoint)
|
||||
or os.getenv("AWS_REGION")
|
||||
or os.getenv("AWS_DEFAULT_REGION")
|
||||
or aws_region_name
|
||||
)
|
||||
|
||||
def _build_sts_client_kwargs(
|
||||
self,
|
||||
aws_sts_endpoint: str | None = None,
|
||||
ssl_verify: bool | str | None = None,
|
||||
aws_region_name: str | None = None,
|
||||
) -> dict:
|
||||
"""STS client kwargs with aligned endpoint_url and region_name (SigV4)."""
|
||||
kwargs: Final[dict] = {"verify": self._get_ssl_verify(ssl_verify)}
|
||||
if aws_sts_endpoint is not None:
|
||||
kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
sts_region: Final = self._resolve_sts_region(aws_sts_endpoint)
|
||||
sts_region: Final = self._resolve_sts_region(aws_sts_endpoint, aws_region_name)
|
||||
if sts_region is not None:
|
||||
kwargs["region_name"] = sts_region
|
||||
return kwargs
|
||||
|
|
@ -837,6 +843,7 @@ class BaseAWSLLM:
|
|||
sts_client_kwargs: Final = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
|
|
@ -948,6 +955,7 @@ class BaseAWSLLM:
|
|||
aws_external_id: str | None = None,
|
||||
aws_sts_endpoint: str | None = None,
|
||||
ssl_verify: bool | str | None = None,
|
||||
aws_region_name: str | None = None,
|
||||
) -> dict:
|
||||
"""Handle cross-account role assumption for IRSA."""
|
||||
import boto3
|
||||
|
|
@ -961,6 +969,7 @@ class BaseAWSLLM:
|
|||
irsa_sts_kwargs: Final = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
# Create an STS client without credentials
|
||||
|
|
@ -1017,6 +1026,7 @@ class BaseAWSLLM:
|
|||
aws_external_id: str | None = None,
|
||||
aws_sts_endpoint: str | None = None,
|
||||
ssl_verify: bool | str | None = None,
|
||||
aws_region_name: str | None = None,
|
||||
) -> dict:
|
||||
"""Handle same-account role assumption for IRSA."""
|
||||
import boto3
|
||||
|
|
@ -1024,6 +1034,7 @@ class BaseAWSLLM:
|
|||
irsa_sts_kwargs: Final = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
verbose_logger.debug("Same account role assumption, using automatic IRSA")
|
||||
|
|
@ -1153,6 +1164,7 @@ class BaseAWSLLM:
|
|||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
else:
|
||||
sts_response = self._handle_irsa_same_account(
|
||||
|
|
@ -1161,6 +1173,7 @@ class BaseAWSLLM:
|
|||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
return self._extract_credentials_and_ttl(sts_response)
|
||||
|
|
@ -1182,6 +1195,7 @@ class BaseAWSLLM:
|
|||
sts_client_kwargs: Final = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
if aws_access_key_id is None and aws_secret_access_key is None:
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
|
|
@ -1363,14 +1377,15 @@ class BaseAWSLLM:
|
|||
"""
|
||||
Select the default endpoint url based on the endpoint type
|
||||
|
||||
Default endpoint url is https://bedrock-runtime.{aws_region_name}.amazonaws.com
|
||||
Default endpoint url is https://bedrock-runtime.{aws_region_name}.{partition dns suffix}
|
||||
"""
|
||||
dns_suffix: Final = get_aws_dns_suffix(aws_region_name)
|
||||
if endpoint_type == "agent":
|
||||
return f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com"
|
||||
return f"https://bedrock-agent-runtime.{aws_region_name}.{dns_suffix}"
|
||||
elif endpoint_type == "agentcore":
|
||||
return f"https://bedrock-agentcore.{aws_region_name}.amazonaws.com"
|
||||
return f"https://bedrock-agentcore.{aws_region_name}.{dns_suffix}"
|
||||
else:
|
||||
return f"https://bedrock-runtime.{aws_region_name}.amazonaws.com"
|
||||
return f"https://bedrock-runtime.{aws_region_name}.{dns_suffix}"
|
||||
|
||||
def _get_boto_credentials_from_optional_params(
|
||||
self, optional_params: dict, model: str | None = None
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from openai.types.batch import BatchRequestCounts
|
||||
from openai.types.batch import Metadata as OpenAIBatchMetadata
|
||||
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -68,6 +70,19 @@ def _predict_output_file_uri(output_prefix: str, input_uri: str, job_id: str | N
|
|||
return f"{output_prefix}{job_id}/{input_basename}.out"
|
||||
|
||||
|
||||
def _record_counts_from_response(response: Mapping[str, object]) -> BatchRequestCounts | None:
|
||||
total_records: Final = response.get("totalRecordCount")
|
||||
success_records: Final = response.get("successRecordCount")
|
||||
if not isinstance(total_records, int) or not isinstance(success_records, int):
|
||||
return None
|
||||
error_records: Final = response.get("errorRecordCount")
|
||||
return BatchRequestCounts(
|
||||
total=total_records,
|
||||
completed=success_records,
|
||||
failed=error_records if isinstance(error_records, int) else 0,
|
||||
)
|
||||
|
||||
|
||||
def _to_epoch(value: Any) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
|
|
@ -271,11 +286,11 @@ class BedrockBatchesHandler:
|
|||
``aws_external_id``). Unknown keys are ignored.
|
||||
|
||||
Returns:
|
||||
``LiteLLMBatch`` shaped like an OpenAI Batch resource. Note that
|
||||
``request_counts`` is always ``(0, 0, 0)`` because
|
||||
``GetModelInvocationJob`` does not surface per-record counts;
|
||||
callers that need accurate counts should parse
|
||||
``manifest.json.out`` from the output S3 prefix.
|
||||
``LiteLLMBatch`` shaped like an OpenAI Batch resource.
|
||||
``request_counts`` maps ``GetModelInvocationJob``'s
|
||||
``totalRecordCount`` / ``successRecordCount`` / ``errorRecordCount``
|
||||
when the provider reports them, and is ``None`` when it does not
|
||||
(older botocore, or a status that omits counts).
|
||||
"""
|
||||
try:
|
||||
import boto3
|
||||
|
|
@ -323,7 +338,9 @@ class BedrockBatchesHandler:
|
|||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": {"jobIdentifier": batch_id},
|
||||
"api_base": (f"https://bedrock.{region}.amazonaws.com/model-invocation-job/{url_path_id}"),
|
||||
"api_base": (
|
||||
f"https://bedrock.{region}.{get_aws_dns_suffix(region)}/model-invocation-job/{url_path_id}"
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -386,7 +403,7 @@ class BedrockBatchesHandler:
|
|||
failed_at=completed_at if openai_status == "failed" else None,
|
||||
cancelled_at=completed_at if openai_status == "cancelled" else None,
|
||||
expired_at=completed_at if openai_status == "expired" else None,
|
||||
request_counts=BatchRequestCounts(total=0, completed=0, failed=0),
|
||||
request_counts=_record_counts_from_response(response),
|
||||
metadata=openai_batch_metadata,
|
||||
completion_window="24h",
|
||||
endpoint="/v1/chat/completions",
|
||||
|
|
|
|||
|
|
@ -1,11 +1,12 @@
|
|||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Final, Literal, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
|
||||
from httpx import Headers, Response
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix, is_bedrock_arn
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
BEDROCK_MANAGED_S3_BATCH_PREFIX,
|
||||
)
|
||||
|
|
@ -34,6 +35,9 @@ from ..common_utils import (
|
|||
resolve_s3_encryption_key_id,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
# Bedrock batch input files are uploaded as
|
||||
# s3://bucket/litellm-bedrock-files-{model, ":" -> "-"}-{uuid4}.jsonl (see
|
||||
# BedrockFilesTransformation._get_s3_object_name). A uuid4 is always 36 hex/dash
|
||||
|
|
@ -138,8 +142,10 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
aws_region_name: Final = self._get_aws_region_name(request_params, model)
|
||||
|
||||
# Bedrock model invocation job endpoint
|
||||
# Format: https://bedrock.{region}.amazonaws.com/model-invocation-job
|
||||
bedrock_endpoint: Final = f"https://bedrock.{aws_region_name}.amazonaws.com/model-invocation-job"
|
||||
# Format: https://bedrock.{region}.{partition dns suffix}/model-invocation-job
|
||||
bedrock_endpoint: Final = (
|
||||
f"https://bedrock.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}/model-invocation-job"
|
||||
)
|
||||
|
||||
return bedrock_endpoint
|
||||
|
||||
|
|
@ -238,8 +244,9 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
# For Bedrock, we need to return a pre-signed request with AWS auth headers
|
||||
# Use common utility for AWS signing
|
||||
request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params)
|
||||
aws_region_name: Final = self._get_aws_region_name(request_params, model)
|
||||
endpoint_url: Final = (
|
||||
f"https://bedrock.{self._get_aws_region_name(request_params, model)}.amazonaws.com/model-invocation-job"
|
||||
f"https://bedrock.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}/model-invocation-job"
|
||||
)
|
||||
signed_headers, signed_data = self.common_utils.sign_aws_request(
|
||||
service_name="bedrock",
|
||||
|
|
@ -261,7 +268,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
self,
|
||||
model: str | None,
|
||||
raw_response: Response,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: dict,
|
||||
) -> LiteLLMBatch:
|
||||
"""
|
||||
|
|
@ -371,7 +378,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
"""
|
||||
# For Bedrock, batch_id should be the full job ARN
|
||||
# The GetModelInvocationJob API expects the full ARN as the identifier
|
||||
if not batch_id.startswith("arn:aws:bedrock:"):
|
||||
if not is_bedrock_arn(batch_id):
|
||||
raise ValueError(f"Invalid batch_id format. Expected ARN, got: {batch_id}")
|
||||
|
||||
# Extract the job identifier from the ARN - use the full ARN path part
|
||||
|
|
@ -390,7 +397,9 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
import urllib.parse as _ul
|
||||
|
||||
encoded_arn: Final = _ul.quote(batch_id, safe="")
|
||||
endpoint_url: Final = f"https://bedrock.{region}.amazonaws.com/model-invocation-job/{encoded_arn}"
|
||||
endpoint_url: Final = (
|
||||
f"https://bedrock.{region}.{get_aws_dns_suffix(region)}/model-invocation-job/{encoded_arn}"
|
||||
)
|
||||
|
||||
# Use common utility for AWS signing
|
||||
request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params)
|
||||
|
|
@ -527,7 +536,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
self,
|
||||
model: str | None,
|
||||
raw_response: Response,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: dict,
|
||||
) -> LiteLLMBatch:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import httpx
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
|
@ -38,6 +39,8 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
|
|
@ -97,7 +100,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
if aws_bedrock_runtime_endpoint:
|
||||
base_url = aws_bedrock_runtime_endpoint
|
||||
else:
|
||||
base_url = f"https://bedrock-agentcore.{region}.amazonaws.com"
|
||||
base_url = f"https://bedrock-agentcore.{region}.{get_aws_dns_suffix(region)}"
|
||||
|
||||
# Based on boto3 client.invoke_agent_runtime, the path is:
|
||||
# /runtimes/{URL-ENCODED-ARN}/invocations?qualifier=<value>
|
||||
|
|
@ -974,7 +977,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import json
|
|||
import time
|
||||
import types
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, Literal, cast, overload
|
||||
from typing import TYPE_CHECKING, Final, Literal, cast, overload
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -94,6 +94,9 @@ from ..common_utils import (
|
|||
normalize_bedrock_opus_output_config_effort,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
# Computer use tool prefixes supported by Bedrock
|
||||
BEDROCK_COMPUTER_USE_TOOLS: Final = [
|
||||
"computer_use_preview",
|
||||
|
|
@ -1770,7 +1773,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -37,6 +37,8 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -436,7 +438,7 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from httpx import Response
|
||||
|
||||
|
|
@ -24,6 +24,9 @@ from litellm.types.utils import (
|
|||
|
||||
from .amazon_llama_transformation import AmazonLlamaConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AmazonDeepSeekR1Config(AmazonLlamaConfig):
|
||||
def transform_response(
|
||||
|
|
@ -36,7 +39,7 @@ class AmazonDeepSeekR1Config(AmazonLlamaConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.utils import Choices
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -200,7 +202,7 @@ class AmazonMoonshotConfig(AmazonInvokeConfig, MoonshotChatConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> "ModelResponse":
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ Inherits from `AmazonConverseConfig`
|
|||
Nova + Invoke API Tutorial: https://docs.aws.amazon.com/nova/latest/userguide/using-invoke-api.html
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -18,6 +18,9 @@ from litellm.types.utils import ModelResponse
|
|||
from ..converse_transformation import AmazonConverseConfig
|
||||
from .base_invoke_transformation import AmazonInvokeConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig):
|
||||
"""
|
||||
|
|
@ -70,7 +73,7 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ The main difference is in the response format: Qwen2 uses "text" field while Qwe
|
|||
Qwen2 + Invoke API Tutorial: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -20,6 +20,9 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
|
|||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AmazonQwen2Config(AmazonQwen3Config):
|
||||
"""
|
||||
|
|
@ -41,7 +44,7 @@ class AmazonQwen2Config(AmazonQwen3Config):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ Inherits from `AmazonInvokeConfig`
|
|||
Qwen3 + Invoke API Tutorial: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -18,6 +18,9 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
|
|||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
|
||||
"""
|
||||
|
|
@ -167,7 +170,7 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -25,6 +25,8 @@ from litellm.types.utils import ModelResponse, Usage
|
|||
from litellm.utils import get_base64_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -188,7 +190,7 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -29,6 +29,8 @@ from litellm.types.utils import ModelResponse
|
|||
from litellm.utils import _supports_factory
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -397,7 +399,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -34,6 +34,8 @@ from litellm.types.utils import ModelResponse, Usage
|
|||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -286,7 +288,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ import httpx
|
|||
|
||||
import litellm
|
||||
from litellm import verbose_logger
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.llms.base_llm.anthropic_messages.transformation import (
|
||||
BaseAnthropicMessagesConfig,
|
||||
)
|
||||
|
|
@ -434,15 +435,15 @@ def init_bedrock_client(
|
|||
ssl_verify: Final = _get_bedrock_client_ssl_verify()
|
||||
|
||||
### SET REGION NAME
|
||||
if region_name:
|
||||
pass
|
||||
elif aws_region_name:
|
||||
region_name = aws_region_name
|
||||
elif litellm_aws_region_name:
|
||||
region_name = litellm_aws_region_name
|
||||
elif standard_aws_region_name:
|
||||
region_name = standard_aws_region_name
|
||||
else:
|
||||
resolved_region_name: Final = next(
|
||||
(
|
||||
candidate
|
||||
for candidate in (region_name, aws_region_name, litellm_aws_region_name, standard_aws_region_name)
|
||||
if isinstance(candidate, str) and candidate
|
||||
),
|
||||
None,
|
||||
)
|
||||
if resolved_region_name is None:
|
||||
raise BedrockError(
|
||||
message="AWS region not set: set AWS_REGION_NAME or AWS_REGION env variable or in .env file",
|
||||
status_code=401,
|
||||
|
|
@ -455,7 +456,7 @@ def init_bedrock_client(
|
|||
elif env_aws_bedrock_runtime_endpoint:
|
||||
endpoint_url = env_aws_bedrock_runtime_endpoint
|
||||
else:
|
||||
endpoint_url = f"https://bedrock-runtime.{region_name}.amazonaws.com"
|
||||
endpoint_url = f"https://bedrock-runtime.{resolved_region_name}.{get_aws_dns_suffix(resolved_region_name)}"
|
||||
|
||||
import boto3
|
||||
|
||||
|
|
@ -492,7 +493,7 @@ def init_bedrock_client(
|
|||
aws_access_key_id=sts_response["Credentials"]["AccessKeyId"],
|
||||
aws_secret_access_key=sts_response["Credentials"]["SecretAccessKey"],
|
||||
aws_session_token=sts_response["Credentials"]["SessionToken"],
|
||||
region_name=region_name,
|
||||
region_name=resolved_region_name,
|
||||
endpoint_url=endpoint_url,
|
||||
config=config,
|
||||
verify=ssl_verify,
|
||||
|
|
@ -513,7 +514,7 @@ def init_bedrock_client(
|
|||
aws_access_key_id=sts_response["Credentials"]["AccessKeyId"],
|
||||
aws_secret_access_key=sts_response["Credentials"]["SecretAccessKey"],
|
||||
aws_session_token=sts_response["Credentials"]["SessionToken"],
|
||||
region_name=region_name,
|
||||
region_name=resolved_region_name,
|
||||
endpoint_url=endpoint_url,
|
||||
config=config,
|
||||
verify=ssl_verify,
|
||||
|
|
@ -526,7 +527,7 @@ def init_bedrock_client(
|
|||
service_name="bedrock-runtime",
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
region_name=region_name,
|
||||
region_name=resolved_region_name,
|
||||
endpoint_url=endpoint_url,
|
||||
config=config,
|
||||
verify=ssl_verify,
|
||||
|
|
@ -536,7 +537,7 @@ def init_bedrock_client(
|
|||
|
||||
client = boto3.Session(profile_name=aws_profile_name).client(
|
||||
service_name="bedrock-runtime",
|
||||
region_name=region_name,
|
||||
region_name=resolved_region_name,
|
||||
endpoint_url=endpoint_url,
|
||||
config=config,
|
||||
verify=ssl_verify,
|
||||
|
|
@ -547,7 +548,7 @@ def init_bedrock_client(
|
|||
|
||||
client = boto3.client(
|
||||
service_name="bedrock-runtime",
|
||||
region_name=region_name,
|
||||
region_name=resolved_region_name,
|
||||
endpoint_url=endpoint_url,
|
||||
config=config,
|
||||
verify=ssl_verify,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,10 @@ to AWS Bedrock's CountTokens API format and vice versa.
|
|||
"""
|
||||
|
||||
import re
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, Literal
|
||||
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.bedrock.common_utils import get_bedrock_base_model
|
||||
|
|
@ -17,6 +20,48 @@ from litellm.llms.bedrock.common_utils import get_bedrock_base_model
|
|||
DEFAULT_ANTHROPIC_INVOKE_MODEL_MAX_TOKENS: Final = 1024
|
||||
|
||||
|
||||
def _json_dict(value: JsonValue) -> dict[str, JsonValue]:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _json_list(value: JsonValue) -> list[JsonValue]:
|
||||
return value if isinstance(value, list) else []
|
||||
|
||||
|
||||
def _to_converse_content(content: JsonValue) -> list[JsonValue]:
|
||||
if isinstance(content, str):
|
||||
return [{"text": content}]
|
||||
if isinstance(content, list):
|
||||
return content
|
||||
return []
|
||||
|
||||
|
||||
def _to_converse_message(message: JsonValue) -> dict[str, JsonValue]:
|
||||
fields: Final = _json_dict(message)
|
||||
return {
|
||||
"role": fields.get("role"),
|
||||
"content": _to_converse_content(fields.get("content", "")),
|
||||
}
|
||||
|
||||
|
||||
def _sanitized_bedrock_tool_name(raw_name: JsonValue) -> str:
|
||||
name: Final = re.sub(r"[^a-zA-Z0-9_]", "_", raw_name if isinstance(raw_name, str) else "")
|
||||
prefixed: Final = name if not name or name[0].isalpha() else f"t_{name}"
|
||||
return prefixed[:64]
|
||||
|
||||
|
||||
def _to_bedrock_tool_spec(tool: JsonValue) -> dict[str, JsonValue]:
|
||||
fields: Final = _json_dict(tool)
|
||||
name: Final = _sanitized_bedrock_tool_name(fields.get("name", ""))
|
||||
return {
|
||||
"toolSpec": {
|
||||
"name": name,
|
||||
"description": fields.get("description") or name,
|
||||
"inputSchema": {"json": fields.get("input_schema", {"type": "object", "properties": {}})},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
class BedrockCountTokensConfig(BaseAWSLLM):
|
||||
"""
|
||||
Configuration and transformation logic for AWS Bedrock CountTokens API.
|
||||
|
|
@ -27,7 +72,7 @@ class BedrockCountTokensConfig(BaseAWSLLM):
|
|||
- Response: {"inputTokens": <number>}
|
||||
"""
|
||||
|
||||
def _detect_input_type(self, request_data: dict[str, Any]) -> str:
|
||||
def _detect_input_type(self, request_data: Mapping[str, JsonValue]) -> Literal["converse", "invokeModel"]:
|
||||
"""
|
||||
Detect whether to use 'converse' or 'invokeModel' input format.
|
||||
|
||||
|
|
@ -57,8 +102,8 @@ class BedrockCountTokensConfig(BaseAWSLLM):
|
|||
|
||||
def transform_anthropic_to_bedrock_count_tokens(
|
||||
self,
|
||||
request_data: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
request_data: Mapping[str, JsonValue],
|
||||
) -> dict[str, JsonValue]:
|
||||
"""
|
||||
Transform request to Bedrock CountTokens format.
|
||||
Supports both Converse and InvokeModel input types.
|
||||
|
|
@ -95,27 +140,16 @@ class BedrockCountTokensConfig(BaseAWSLLM):
|
|||
else:
|
||||
return self._transform_to_invoke_model_format(request_data)
|
||||
|
||||
def _transform_to_converse_format(self, request_data: dict[str, Any]) -> dict[str, Any]:
|
||||
def _transform_to_converse_format(self, request_data: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
"""Transform to Converse input format, including system and tools."""
|
||||
messages: Final = request_data.get("messages", [])
|
||||
messages: Final = _json_list(request_data.get("messages"))
|
||||
system: Final = request_data.get("system")
|
||||
tools: Final = request_data.get("tools")
|
||||
|
||||
# Transform messages
|
||||
user_messages: Final = []
|
||||
for message in messages:
|
||||
transformed_message: dict[str, Any] = {
|
||||
"role": message.get("role"),
|
||||
"content": [],
|
||||
}
|
||||
content = message.get("content", "")
|
||||
if isinstance(content, str):
|
||||
transformed_message["content"].append({"text": content})
|
||||
elif isinstance(content, list):
|
||||
transformed_message["content"] = content
|
||||
user_messages.append(transformed_message)
|
||||
user_messages: Final[list[JsonValue]] = [_to_converse_message(message) for message in messages]
|
||||
|
||||
converse_input: Final[dict[str, Any]] = {"messages": user_messages}
|
||||
converse_input: Final[dict[str, JsonValue]] = {"messages": user_messages}
|
||||
|
||||
# Transform system prompt (string or list of blocks → Bedrock format)
|
||||
system_blocks: Final = self._transform_system(system)
|
||||
|
|
@ -129,7 +163,7 @@ class BedrockCountTokensConfig(BaseAWSLLM):
|
|||
|
||||
return {"input": {"converse": converse_input}}
|
||||
|
||||
def _transform_system(self, system: Any | None) -> list[dict[str, Any]]:
|
||||
def _transform_system(self, system: JsonValue) -> list[JsonValue]:
|
||||
"""Transform Anthropic system prompt to Bedrock system blocks."""
|
||||
if system is None:
|
||||
return []
|
||||
|
|
@ -140,36 +174,16 @@ class BedrockCountTokensConfig(BaseAWSLLM):
|
|||
return [{"text": block.get("text", "")} for block in system if isinstance(block, dict)]
|
||||
return []
|
||||
|
||||
def _transform_tools(self, tools: list[dict[str, Any]] | None) -> dict[str, Any] | None:
|
||||
def _transform_tools(self, tools: JsonValue) -> dict[str, JsonValue] | None:
|
||||
"""Transform Anthropic tools to Bedrock toolConfig format."""
|
||||
if not tools:
|
||||
return None
|
||||
|
||||
bedrock_tools: Final = []
|
||||
for tool in tools:
|
||||
name = tool.get("name", "")
|
||||
# Bedrock tool names must match [a-zA-Z][a-zA-Z0-9_]* and max 64 chars
|
||||
name = re.sub(r"[^a-zA-Z0-9_]", "_", name)
|
||||
if name and not name[0].isalpha():
|
||||
name = "t_" + name
|
||||
name = name[:64]
|
||||
|
||||
description = tool.get("description") or name
|
||||
input_schema = tool.get("input_schema", {"type": "object", "properties": {}})
|
||||
|
||||
bedrock_tools.append(
|
||||
{
|
||||
"toolSpec": {
|
||||
"name": name,
|
||||
"description": description,
|
||||
"inputSchema": {"json": input_schema},
|
||||
}
|
||||
}
|
||||
)
|
||||
bedrock_tools: Final[list[JsonValue]] = [_to_bedrock_tool_spec(tool) for tool in _json_list(tools)]
|
||||
|
||||
return {"tools": bedrock_tools}
|
||||
|
||||
def _transform_to_invoke_model_format(self, request_data: dict[str, Any]) -> dict[str, Any]:
|
||||
def _transform_to_invoke_model_format(self, request_data: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
"""Transform to InvokeModel input format."""
|
||||
import base64
|
||||
import json
|
||||
|
|
@ -223,7 +237,9 @@ class BedrockCountTokensConfig(BaseAWSLLM):
|
|||
|
||||
return endpoint
|
||||
|
||||
def transform_bedrock_response_to_anthropic(self, bedrock_response: dict[str, Any]) -> dict[str, Any]:
|
||||
def transform_bedrock_response_to_anthropic(
|
||||
self, bedrock_response: Mapping[str, JsonValue]
|
||||
) -> dict[str, JsonValue]:
|
||||
"""
|
||||
Transform Bedrock CountTokens response to Anthropic format.
|
||||
|
||||
|
|
@ -241,7 +257,7 @@ class BedrockCountTokensConfig(BaseAWSLLM):
|
|||
|
||||
return {"input_tokens": input_tokens}
|
||||
|
||||
def validate_count_tokens_request(self, request_data: dict[str, Any]) -> None:
|
||||
def validate_count_tokens_request(self, request_data: Mapping[str, JsonValue]) -> None:
|
||||
"""
|
||||
Validate the incoming count tokens request.
|
||||
Supports both Converse and InvokeModel input formats.
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ import copy
|
|||
import json
|
||||
import urllib.parse
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Final, get_args
|
||||
from typing import TYPE_CHECKING, Any, Final, get_args
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -37,6 +37,9 @@ from .amazon_titan_v2_transformation import AmazonTitanV2Config
|
|||
from .cohere_transformation import BedrockCohereEmbeddingConfig
|
||||
from .twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class BedrockEmbedding(BaseAWSLLM):
|
||||
def _load_credentials(
|
||||
|
|
@ -58,6 +61,7 @@ class BedrockEmbedding(BaseAWSLLM):
|
|||
aws_profile_name: Final = optional_params.pop("aws_profile_name", None)
|
||||
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
|
||||
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
|
||||
aws_external_id: Final = optional_params.pop("aws_external_id", None)
|
||||
|
||||
### SET REGION NAME ###
|
||||
if aws_region_name is None:
|
||||
|
|
@ -84,6 +88,7 @@ class BedrockEmbedding(BaseAWSLLM):
|
|||
aws_role_name=aws_role_name,
|
||||
aws_web_identity_token=aws_web_identity_token,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
)
|
||||
return credentials, aws_region_name
|
||||
|
||||
|
|
@ -233,7 +238,7 @@ class BedrockEmbedding(BaseAWSLLM):
|
|||
endpoint_url: str,
|
||||
aws_region_name: str,
|
||||
model: str,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
|
||||
api_key: str | None = None,
|
||||
is_async_invoke: bool | None = False,
|
||||
|
|
@ -301,7 +306,7 @@ class BedrockEmbedding(BaseAWSLLM):
|
|||
endpoint_url: str,
|
||||
aws_region_name: str,
|
||||
model: str,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
|
||||
api_key: str | None = None,
|
||||
is_async_invoke: bool | None = False,
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm._logging import verbose_logger
|
|||
from litellm._uuid import uuid
|
||||
from litellm.constants import BEDROCK_INVOKE_PROVIDERS_LITERAL
|
||||
from litellm.files.utils import FilesAPIUtils
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
BEDROCK_MANAGED_S3_BATCH_PREFIX,
|
||||
BEDROCK_MANAGED_S3_PREFIXES,
|
||||
|
|
@ -413,7 +414,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
|
||||
# S3 endpoint URL format
|
||||
s3_endpoint_url: Final = (
|
||||
request_params.get("s3_endpoint_url") or f"https://s3.{aws_region_name}.amazonaws.com"
|
||||
request_params.get("s3_endpoint_url")
|
||||
or f"https://s3.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}"
|
||||
).rstrip("/")
|
||||
|
||||
return f"{s3_endpoint_url}/{bucket_name}/{encoded_object_name}"
|
||||
|
|
@ -1249,7 +1251,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
region_params: Final[dict[str, str | None]] = {"aws_region_name": region_preference}
|
||||
aws_region_name: Final = self._get_aws_region_name(optional_params=region_params, model="")
|
||||
|
||||
s3_endpoint_url = (request_params.s3_endpoint_url or f"https://s3.{aws_region_name}.amazonaws.com").rstrip("/")
|
||||
s3_endpoint_url = (
|
||||
request_params.s3_endpoint_url or f"https://s3.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}"
|
||||
).rstrip("/")
|
||||
url: Final = f"{s3_endpoint_url}/{bucket_name}/{encode_s3_object_key_for_url(object_key)}"
|
||||
|
||||
litellm_params[S3_SIGNED_GET_HEADERS_PARAM] = self._sign_s3_get_request(
|
||||
|
|
|
|||
|
|
@ -7,18 +7,66 @@ This uses aws_sdk_bedrock_runtime for bidirectional streaming with Nova Sonic.
|
|||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
from typing import Any, Final
|
||||
from typing import Final, Protocol
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.types.realtime import RealtimeResponseTransformInput
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import BedrockError
|
||||
from .transformation import BedrockRealtimeConfig
|
||||
|
||||
_CLIENT_MODALITIES_ADAPTER: Final[TypeAdapter["list[str] | None"]] = TypeAdapter(list[str] | None)
|
||||
_CLIENT_MESSAGE_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
|
||||
|
||||
def _json_dict(value: JsonValue) -> dict[str, JsonValue]:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _json_str(value: JsonValue) -> str | None:
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
class RealtimeClientWebSocket(Protocol):
|
||||
"""The client-facing websocket surface the realtime bridge talks to."""
|
||||
|
||||
async def receive_text(self) -> str: ...
|
||||
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
||||
async def close(self, code: int = 1000, reason: str | None = None) -> None: ...
|
||||
|
||||
|
||||
class BedrockInputStream(Protocol):
|
||||
async def send(self, event: object) -> None: ...
|
||||
|
||||
async def close(self) -> None: ...
|
||||
|
||||
|
||||
class BedrockPayloadPart(Protocol):
|
||||
@property
|
||||
def bytes_(self) -> bytes | None: ...
|
||||
|
||||
|
||||
class BedrockOutputChunk(Protocol):
|
||||
@property
|
||||
def value(self) -> BedrockPayloadPart | None: ...
|
||||
|
||||
|
||||
class BedrockOutputStream(Protocol):
|
||||
async def receive(self) -> BedrockOutputChunk | None: ...
|
||||
|
||||
|
||||
class BedrockBidirectionalStream(Protocol):
|
||||
@property
|
||||
def input_stream(self) -> BedrockInputStream: ...
|
||||
|
||||
async def await_output(self) -> tuple[object, BedrockOutputStream]: ...
|
||||
|
||||
|
||||
class BedrockRealtime(BaseAWSLLM):
|
||||
|
|
@ -30,7 +78,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
async def async_realtime(
|
||||
self,
|
||||
model: str,
|
||||
websocket: Any,
|
||||
websocket: RealtimeClientWebSocket,
|
||||
logging_obj: LiteLLMLogging,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
|
|
@ -81,7 +129,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
elif aws_bedrock_runtime_endpoint is not None:
|
||||
endpoint_uri = aws_bedrock_runtime_endpoint
|
||||
else:
|
||||
endpoint_uri = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com"
|
||||
endpoint_uri = f"https://bedrock-runtime.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}"
|
||||
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Connecting to %s with model %s", endpoint_uri, model)
|
||||
|
||||
|
|
@ -132,7 +180,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
verbose_proxy_logger.debug("Bedrock Realtime: sent session.created to client on connect")
|
||||
|
||||
# Track state for transformation
|
||||
session_state: Final = {
|
||||
session_state: Final[RealtimeResponseTransformInput] = {
|
||||
"current_output_item_id": None,
|
||||
"current_response_id": None,
|
||||
"current_conversation_id": None,
|
||||
|
|
@ -182,11 +230,11 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
|
||||
async def _forward_client_to_bedrock(
|
||||
self,
|
||||
client_ws: Any,
|
||||
bedrock_stream: Any,
|
||||
client_ws: RealtimeClientWebSocket,
|
||||
bedrock_stream: BedrockBidirectionalStream,
|
||||
transformation_config: BedrockRealtimeConfig,
|
||||
model: str,
|
||||
session_state: dict,
|
||||
session_state: RealtimeResponseTransformInput,
|
||||
logging_obj: LiteLLMLogging | None = None,
|
||||
):
|
||||
"""Forward messages from client WebSocket to Bedrock stream."""
|
||||
|
|
@ -223,11 +271,11 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
client_message_type: str | None = None
|
||||
requested_modalities: list[str] | None = None
|
||||
with contextlib.suppress(Exception):
|
||||
parsed_client_message = json.loads(message)
|
||||
client_message_type = parsed_client_message.get("type")
|
||||
parsed_client_message = _json_dict(_CLIENT_MESSAGE_ADAPTER.validate_json(message))
|
||||
client_message_type = _json_str(parsed_client_message.get("type"))
|
||||
if client_message_type == "session.update":
|
||||
requested_modalities = _CLIENT_MODALITIES_ADAPTER.validate_python(
|
||||
parsed_client_message.get("session", {}).get("modalities")
|
||||
_json_dict(parsed_client_message.get("session")).get("modalities")
|
||||
)
|
||||
if client_message_type == "session.update":
|
||||
await client_ws.send_text(
|
||||
|
|
@ -246,12 +294,12 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
|
||||
async def _forward_bedrock_to_client(
|
||||
self,
|
||||
bedrock_stream: Any,
|
||||
client_ws: Any,
|
||||
bedrock_stream: BedrockBidirectionalStream,
|
||||
client_ws: RealtimeClientWebSocket,
|
||||
transformation_config: BedrockRealtimeConfig,
|
||||
model: str,
|
||||
logging_obj: LiteLLMLogging,
|
||||
session_state: dict,
|
||||
session_state: RealtimeResponseTransformInput,
|
||||
):
|
||||
"""Forward messages from Bedrock stream to client WebSocket."""
|
||||
try:
|
||||
|
|
@ -264,13 +312,12 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
verbose_proxy_logger.debug("Bedrock Realtime: Bedrock stream ended")
|
||||
break
|
||||
|
||||
if result.value and result.value.bytes_:
|
||||
bedrock_response = result.value.bytes_.decode("utf-8")
|
||||
payload_bytes = result.value.bytes_ if result.value else None
|
||||
if payload_bytes:
|
||||
bedrock_response = payload_bytes.decode("utf-8")
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Received from Bedrock: %s", bedrock_response[:200])
|
||||
|
||||
# Transform Bedrock format to OpenAI format
|
||||
from litellm.types.realtime import RealtimeResponseTransformInput
|
||||
|
||||
realtime_response_transform_input: RealtimeResponseTransformInput = {
|
||||
"current_output_item_id": session_state.get("current_output_item_id"),
|
||||
"current_response_id": session_state.get("current_response_id"),
|
||||
|
|
|
|||
|
|
@ -29,6 +29,8 @@ from ..common_utils import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -256,7 +258,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@ from litellm.utils import CustomStreamWrapper, ModelResponse, Usage
|
|||
from ..common_utils import API_BASE, BytezError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -185,7 +187,7 @@ class BytezChatConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -2,9 +2,11 @@ import base64
|
|||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
import httpx
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
|
|
@ -27,6 +29,16 @@ DEVICE_CODE_TIMEOUT_SECONDS: Final = 15 * 60
|
|||
DEVICE_CODE_COOLDOWN_SECONDS: Final = 5 * 60
|
||||
DEVICE_CODE_POLL_SLEEP_SECONDS: Final = 5
|
||||
|
||||
OPENAI_AUTH_CLAIM_KEY: Final = "https://api.openai.com/auth"
|
||||
|
||||
JsonObject: TypeAlias = Mapping[str, JsonValue]
|
||||
|
||||
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(JsonObject)
|
||||
|
||||
|
||||
def _optional_str(value: JsonValue | None) -> str | None:
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
class Authenticator:
|
||||
def __init__(self) -> None:
|
||||
|
|
@ -43,10 +55,10 @@ class Authenticator:
|
|||
def get_access_token(self) -> str:
|
||||
auth_data: Final = self._read_auth_file()
|
||||
if auth_data:
|
||||
access_token: Final = auth_data.get("access_token")
|
||||
access_token: Final = _optional_str(auth_data.get("access_token"))
|
||||
if access_token and not self._is_token_expired(auth_data, access_token):
|
||||
return access_token
|
||||
refresh_token: Final = auth_data.get("refresh_token")
|
||||
refresh_token: Final = _optional_str(auth_data.get("refresh_token"))
|
||||
if refresh_token:
|
||||
try:
|
||||
refreshed: Final = self._refresh_tokens(refresh_token)
|
||||
|
|
@ -67,48 +79,47 @@ class Authenticator:
|
|||
auth_data: Final = self._read_auth_file()
|
||||
if not auth_data:
|
||||
return None
|
||||
account_id: Final = auth_data.get("account_id")
|
||||
account_id: Final = _optional_str(auth_data.get("account_id"))
|
||||
if account_id:
|
||||
return account_id
|
||||
id_token: Final = auth_data.get("id_token")
|
||||
access_token: Final = auth_data.get("access_token")
|
||||
derived: Final = self._extract_account_id(id_token or access_token)
|
||||
derived: Final = self._extract_account_id(_optional_str(id_token or access_token))
|
||||
if derived:
|
||||
auth_data["account_id"] = derived
|
||||
self._write_auth_file(auth_data)
|
||||
self._write_auth_file({**auth_data, "account_id": derived})
|
||||
return derived
|
||||
|
||||
def _ensure_token_dir(self) -> None:
|
||||
if not os.path.exists(self.token_dir):
|
||||
os.makedirs(self.token_dir, exist_ok=True)
|
||||
|
||||
def _read_auth_file(self) -> dict[str, Any] | None:
|
||||
def _read_auth_file(self) -> JsonObject | None:
|
||||
try:
|
||||
with open(self.auth_file, "r") as f:
|
||||
return json.load(f)
|
||||
return _JSON_OBJECT_ADAPTER.validate_python(json.load(f))
|
||||
except OSError:
|
||||
return None
|
||||
except json.JSONDecodeError as exc:
|
||||
except (json.JSONDecodeError, ValidationError) as exc:
|
||||
verbose_logger.warning("Invalid ChatGPT auth file: %s", exc)
|
||||
return None
|
||||
|
||||
def _write_auth_file(self, data: dict[str, Any]) -> None:
|
||||
def _write_auth_file(self, data: JsonObject) -> None:
|
||||
try:
|
||||
with open(self.auth_file, "w") as f:
|
||||
json.dump(data, f)
|
||||
except OSError as exc:
|
||||
verbose_logger.error("Failed to write ChatGPT auth file: %s", exc)
|
||||
|
||||
def _is_token_expired(self, auth_data: dict[str, Any], access_token: str) -> bool:
|
||||
expires_at = auth_data.get("expires_at")
|
||||
if expires_at is None:
|
||||
expires_at = self._get_expires_at(access_token)
|
||||
if expires_at:
|
||||
auth_data["expires_at"] = expires_at
|
||||
self._write_auth_file(auth_data)
|
||||
if expires_at is None:
|
||||
def _is_token_expired(self, auth_data: JsonObject, access_token: str) -> bool:
|
||||
stored_expires_at: Final = auth_data.get("expires_at")
|
||||
if isinstance(stored_expires_at, (int, float)):
|
||||
return time.time() >= float(stored_expires_at) - TOKEN_EXPIRY_SKEW_SECONDS
|
||||
derived_expires_at: Final = self._get_expires_at(access_token)
|
||||
if derived_expires_at:
|
||||
self._write_auth_file({**auth_data, "expires_at": derived_expires_at})
|
||||
if derived_expires_at is None:
|
||||
return True
|
||||
return time.time() >= float(expires_at) - TOKEN_EXPIRY_SKEW_SECONDS
|
||||
return time.time() >= float(derived_expires_at) - TOKEN_EXPIRY_SKEW_SECONDS
|
||||
|
||||
def _get_expires_at(self, token: str) -> int | None:
|
||||
claims: Final = self._decode_jwt_claims(token)
|
||||
|
|
@ -117,15 +128,14 @@ class Authenticator:
|
|||
return int(exp)
|
||||
return None
|
||||
|
||||
def _decode_jwt_claims(self, token: str) -> dict[str, Any]:
|
||||
def _decode_jwt_claims(self, token: str) -> JsonObject:
|
||||
try:
|
||||
parts: Final = token.split(".")
|
||||
if len(parts) < 2:
|
||||
return {}
|
||||
payload_b64 = parts[1]
|
||||
payload_b64 += "=" * (-len(payload_b64) % 4)
|
||||
payload_b64: Final = parts[1] + "=" * (-len(parts[1]) % 4)
|
||||
payload_bytes: Final = base64.urlsafe_b64decode(payload_b64)
|
||||
return json.loads(payload_bytes.decode("utf-8"))
|
||||
return _JSON_OBJECT_ADAPTER.validate_python(json.loads(payload_bytes.decode("utf-8")))
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
|
@ -133,7 +143,7 @@ class Authenticator:
|
|||
if not token:
|
||||
return None
|
||||
claims: Final = self._decode_jwt_claims(token)
|
||||
auth_claims: Final = claims.get("https://api.openai.com/auth")
|
||||
auth_claims: Final = claims.get(OPENAI_AUTH_CLAIM_KEY)
|
||||
if isinstance(auth_claims, dict):
|
||||
account_id: Final = auth_claims.get("chatgpt_account_id")
|
||||
if isinstance(account_id, str) and account_id:
|
||||
|
|
@ -170,7 +180,7 @@ class Authenticator:
|
|||
json={"client_id": CHATGPT_CLIENT_ID},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data: Final = resp.json()
|
||||
data: Final = _JSON_OBJECT_ADAPTER.validate_python(resp.json())
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise GetDeviceCodeError(
|
||||
message=f"Failed to request device code: {exc}",
|
||||
|
|
@ -182,8 +192,8 @@ class Authenticator:
|
|||
status_code=400,
|
||||
)
|
||||
|
||||
device_auth_id: Final = data.get("device_auth_id")
|
||||
user_code: Final = data.get("user_code") or data.get("usercode")
|
||||
device_auth_id: Final = _optional_str(data.get("device_auth_id"))
|
||||
user_code: Final = _optional_str(data.get("user_code") or data.get("usercode"))
|
||||
interval: Final = data.get("interval")
|
||||
if not device_auth_id or not user_code:
|
||||
raise GetDeviceCodeError(
|
||||
|
|
@ -210,16 +220,16 @@ class Authenticator:
|
|||
},
|
||||
)
|
||||
if resp.status_code == 200:
|
||||
data = resp.json()
|
||||
if all(
|
||||
key in data
|
||||
for key in (
|
||||
"authorization_code",
|
||||
"code_challenge",
|
||||
"code_verifier",
|
||||
)
|
||||
):
|
||||
return data
|
||||
data = _JSON_OBJECT_ADAPTER.validate_python(resp.json())
|
||||
authorization_code = _optional_str(data.get("authorization_code"))
|
||||
code_challenge = _optional_str(data.get("code_challenge"))
|
||||
code_verifier = _optional_str(data.get("code_verifier"))
|
||||
if authorization_code and code_challenge and code_verifier:
|
||||
return {
|
||||
"authorization_code": authorization_code,
|
||||
"code_challenge": code_challenge,
|
||||
"code_verifier": code_verifier,
|
||||
}
|
||||
if resp.status_code in (403, 404):
|
||||
time.sleep(max(interval, DEVICE_CODE_POLL_SLEEP_SECONDS))
|
||||
continue
|
||||
|
|
@ -262,7 +272,7 @@ class Authenticator:
|
|||
content=body,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data: Final = resp.json()
|
||||
data: Final = _JSON_OBJECT_ADAPTER.validate_python(resp.json())
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise GetAccessTokenError(
|
||||
message=f"Token exchange failed: {exc}",
|
||||
|
|
@ -274,15 +284,18 @@ class Authenticator:
|
|||
status_code=400,
|
||||
)
|
||||
|
||||
if not all(key in data for key in ("access_token", "refresh_token", "id_token")):
|
||||
access_token: Final = _optional_str(data.get("access_token"))
|
||||
refresh_token: Final = _optional_str(data.get("refresh_token"))
|
||||
id_token: Final = _optional_str(data.get("id_token"))
|
||||
if not access_token or not refresh_token or not id_token:
|
||||
raise GetAccessTokenError(
|
||||
message=f"Token exchange response missing fields: {data}",
|
||||
status_code=400,
|
||||
)
|
||||
return {
|
||||
"access_token": data["access_token"],
|
||||
"refresh_token": data["refresh_token"],
|
||||
"id_token": data["id_token"],
|
||||
"access_token": access_token,
|
||||
"refresh_token": refresh_token,
|
||||
"id_token": id_token,
|
||||
}
|
||||
|
||||
def _refresh_tokens(self, refresh_token: str) -> dict[str, str]:
|
||||
|
|
@ -298,7 +311,7 @@ class Authenticator:
|
|||
},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data: Final = resp.json()
|
||||
data: Final = _JSON_OBJECT_ADAPTER.validate_python(resp.json())
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise RefreshAccessTokenError(
|
||||
message=f"Refresh token failed: {exc}",
|
||||
|
|
@ -310,8 +323,8 @@ class Authenticator:
|
|||
status_code=400,
|
||||
)
|
||||
|
||||
access_token: Final = data.get("access_token")
|
||||
id_token: Final = data.get("id_token")
|
||||
access_token: Final = _optional_str(data.get("access_token"))
|
||||
id_token: Final = _optional_str(data.get("id_token"))
|
||||
if not access_token or not id_token:
|
||||
raise RefreshAccessTokenError(
|
||||
message=f"Refresh response missing fields: {data}",
|
||||
|
|
@ -320,14 +333,14 @@ class Authenticator:
|
|||
|
||||
refreshed: Final = {
|
||||
"access_token": access_token,
|
||||
"refresh_token": data.get("refresh_token", refresh_token),
|
||||
"refresh_token": _optional_str(data.get("refresh_token")) or refresh_token,
|
||||
"id_token": id_token,
|
||||
}
|
||||
auth_data: Final = self._build_auth_record(refreshed)
|
||||
self._write_auth_file(auth_data)
|
||||
return refreshed
|
||||
|
||||
def _build_auth_record(self, tokens: dict[str, str]) -> dict[str, Any]:
|
||||
def _build_auth_record(self, tokens: dict[str, str]) -> JsonObject:
|
||||
access_token: Final = tokens.get("access_token")
|
||||
id_token: Final = tokens.get("id_token")
|
||||
expires_at: Final = self._get_expires_at(access_token) if access_token else None
|
||||
|
|
@ -340,31 +353,30 @@ class Authenticator:
|
|||
"account_id": account_id,
|
||||
}
|
||||
|
||||
def _get_device_code_cooldown_remaining(self, auth_data: dict[str, Any] | None) -> float:
|
||||
def _get_device_code_cooldown_remaining(self, auth_data: JsonObject | None) -> float:
|
||||
if not auth_data:
|
||||
return 0.0
|
||||
requested_at = auth_data.get("device_code_requested_at")
|
||||
requested_at: Final = auth_data.get("device_code_requested_at")
|
||||
if not isinstance(requested_at, (int, float, str)):
|
||||
return 0.0
|
||||
try:
|
||||
requested_at = float(requested_at)
|
||||
requested_seconds: Final = float(requested_at)
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
elapsed: Final = time.time() - requested_at
|
||||
elapsed: Final = time.time() - requested_seconds
|
||||
remaining: Final = DEVICE_CODE_COOLDOWN_SECONDS - elapsed
|
||||
return max(0.0, remaining)
|
||||
|
||||
def _record_device_code_request(self) -> None:
|
||||
auth_data: Final = self._read_auth_file() or {}
|
||||
auth_data["device_code_requested_at"] = time.time()
|
||||
self._write_auth_file(auth_data)
|
||||
self._write_auth_file({**auth_data, "device_code_requested_at": time.time()})
|
||||
|
||||
def _wait_for_access_token(self, timeout_seconds: float) -> str | None:
|
||||
deadline: Final = time.time() + timeout_seconds
|
||||
while time.time() < deadline:
|
||||
auth_data = self._read_auth_file()
|
||||
if auth_data:
|
||||
access_token = auth_data.get("access_token")
|
||||
access_token = _optional_str(auth_data.get("access_token"))
|
||||
if access_token and not self._is_token_expired(auth_data, access_token):
|
||||
return access_token
|
||||
sleep_for = min(DEVICE_CODE_POLL_SLEEP_SECONDS, max(0.0, deadline - time.time()))
|
||||
|
|
|
|||
|
|
@ -4,7 +4,27 @@ Streaming utilities for ChatGPT provider.
|
|||
Normalizes non-spec-compliant tool_call chunks from the ChatGPT backend API.
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from collections.abc import Awaitable
|
||||
from typing import Final, Protocol
|
||||
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionDeltaCustomToolCall,
|
||||
ChatCompletionDeltaToolCall,
|
||||
Delta,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
||||
|
||||
class ChatGPTChunkStream(Protocol):
|
||||
"""A ChatGPT chunk source driven either synchronously or asynchronously."""
|
||||
|
||||
def __next__(self) -> ModelResponseStream: ...
|
||||
|
||||
def __anext__(self) -> Awaitable[ModelResponseStream]: ...
|
||||
|
||||
|
||||
def _first_choice_delta(chunk: ModelResponseStream) -> Delta | None:
|
||||
return chunk.choices[0].delta
|
||||
|
||||
|
||||
class ChatGPTToolCallNormalizer:
|
||||
|
|
@ -20,13 +40,13 @@ class ChatGPTToolCallNormalizer:
|
|||
chunks to the consumer.
|
||||
"""
|
||||
|
||||
def __init__(self, stream: Any):
|
||||
self._stream = stream
|
||||
def __init__(self, stream: ChatGPTChunkStream):
|
||||
self._stream: Final = stream
|
||||
self._seen_ids: dict[str, int] = {} # tool_call_id -> assigned_index
|
||||
self._next_index: int = 0
|
||||
self._last_id: str | None = None # tracks which tool call the next delta belongs to
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
def __getattr__(self, name: str) -> object:
|
||||
return getattr(self._stream, name)
|
||||
|
||||
def __iter__(self):
|
||||
|
|
@ -35,30 +55,30 @@ class ChatGPTToolCallNormalizer:
|
|||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
def __next__(self) -> ModelResponseStream:
|
||||
while True:
|
||||
chunk = next(self._stream)
|
||||
result = self._normalize(chunk)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
async def __anext__(self):
|
||||
async def __anext__(self) -> ModelResponseStream:
|
||||
while True:
|
||||
chunk = await self._stream.__anext__()
|
||||
result = self._normalize(chunk)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
def _normalize(self, chunk: Any) -> Any:
|
||||
def _normalize(self, chunk: ModelResponseStream) -> ModelResponseStream | None:
|
||||
"""Fix tool_calls in the chunk. Returns None to skip duplicate chunks."""
|
||||
if not chunk.choices:
|
||||
return chunk
|
||||
|
||||
delta: Final = chunk.choices[0].delta
|
||||
delta: Final = _first_choice_delta(chunk)
|
||||
if delta is None or not delta.tool_calls:
|
||||
return chunk
|
||||
|
||||
normalized: Final = []
|
||||
normalized: Final[list[ChatCompletionDeltaToolCall | ChatCompletionDeltaCustomToolCall]] = []
|
||||
for tc in delta.tool_calls:
|
||||
if tc.id and tc.id not in self._seen_ids:
|
||||
# New tool call — assign correct index
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
|
|
@ -28,6 +28,9 @@ from ..common_utils import (
|
|||
get_chatgpt_default_instructions,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
def __init__(self) -> None:
|
||||
|
|
@ -107,7 +110,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
self,
|
||||
model: str,
|
||||
raw_response: Any,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
):
|
||||
body_text: Final = raw_response.text or ""
|
||||
if not self._should_parse_as_sse(raw_response=raw_response, body_text=body_text):
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ from litellm.types.utils import ModelResponse
|
|||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -85,7 +87,7 @@ class ClarifaiConfig(OpenAIGPTConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -15,6 +15,8 @@ from ..common_utils import ModelResponseIterator as CohereModelResponseIterator
|
|||
from ..common_utils import validate_environment as cohere_validate_environment
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -225,7 +227,7 @@ class CohereChatConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ from ..common_utils import CohereError, CohereV2ModelResponseIterator
|
|||
from ..common_utils import validate_environment as cohere_validate_environment
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -189,7 +191,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -3,8 +3,7 @@ Legacy /v1/embedding handler for Bedrock Cohere.
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -20,6 +19,9 @@ from litellm.types.utils import EmbeddingResponse
|
|||
|
||||
from .v1_transformation import CohereEmbeddingConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
|
||||
def validate_environment(api_key, headers: dict):
|
||||
# Create a lowercase key lookup to avoid duplicate headers with different cases
|
||||
|
|
@ -58,7 +60,7 @@ async def async_embedding(
|
|||
api_base: str,
|
||||
api_key: str | None,
|
||||
headers: dict,
|
||||
encoding: Callable,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
):
|
||||
## LOGGING
|
||||
|
|
@ -120,7 +122,7 @@ def embedding(
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
data: dict | CohereEmbeddingRequest | None = None,
|
||||
complete_api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.types.utils import GenericGuardrailAPIInputs
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.rerank import RerankResponse
|
||||
|
||||
|
||||
|
|
@ -42,7 +43,7 @@ class CohereRerankHandler(BaseTranslation):
|
|||
self,
|
||||
data: dict,
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Any | None = None,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> Any:
|
||||
"""
|
||||
Process input text fields ('query' and 'instruction') by applying
|
||||
|
|
@ -94,7 +95,7 @@ class CohereRerankHandler(BaseTranslation):
|
|||
self,
|
||||
response: "RerankResponse",
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Any | None = None,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
user_api_key_dict: Any | None = None,
|
||||
request_data: dict | None = None,
|
||||
) -> Any:
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -130,7 +132,7 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from litellm.types.utils import ModelResponse
|
|||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -49,7 +51,7 @@ class CompactifAIChatConfig(OpenAIGPTConfig):
|
|||
messages: list,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -26,6 +26,8 @@ from litellm.types.utils import HttpHandlerRequestFields, ImageResponse, LlmProv
|
|||
from litellm.utils import CustomStreamWrapper, ModelResponse, ProviderConfigManager
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -266,7 +268,7 @@ class BaseLLMAIOHTTPHandler:
|
|||
messages: list,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
client: ClientSession | None = None,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -5,20 +5,38 @@ import os
|
|||
import ssl
|
||||
import typing
|
||||
import urllib.request
|
||||
from collections.abc import Callable
|
||||
from typing import Any, ClassVar, Final
|
||||
from collections.abc import Callable, Generator
|
||||
from typing import ClassVar, Final
|
||||
|
||||
import aiohttp
|
||||
import aiohttp.client_exceptions
|
||||
import aiohttp.http_exceptions
|
||||
import httpx
|
||||
from aiohttp.client import ClientResponse, ClientSession
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
||||
AIOHTTP_EXC_MAP: Final[dict] = {
|
||||
|
||||
class HttpxTimeoutExtension(BaseModel):
|
||||
connect: float | None = None
|
||||
read: float | None = None
|
||||
write: float | None = None
|
||||
pool: float | None = None
|
||||
|
||||
|
||||
class AiohttpSslRequestOption(TypedDict, total=False):
|
||||
ssl: ReadOnly[bool | ssl.SSLContext]
|
||||
|
||||
|
||||
_TIMEOUT_EXTENSION: Final = TypeAdapter(HttpxTimeoutExtension)
|
||||
_EMPTY_TIMEOUT: Final[HttpxTimeoutExtension] = HttpxTimeoutExtension()
|
||||
_NO_SSL_OVERRIDE: Final[AiohttpSslRequestOption] = {}
|
||||
|
||||
AIOHTTP_EXC_MAP: Final[dict[type[BaseException], type[Exception]]] = {
|
||||
# Order matters here, most specific exception first
|
||||
# Timeout related exceptions
|
||||
asyncio.TimeoutError: httpx.TimeoutException,
|
||||
|
|
@ -58,11 +76,11 @@ except ImportError:
|
|||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def map_aiohttp_exceptions() -> typing.Iterator[None]:
|
||||
def map_aiohttp_exceptions() -> Generator[None, None, None]:
|
||||
try:
|
||||
yield
|
||||
except Exception as exc:
|
||||
mapped_exc = None
|
||||
mapped_exc: type[Exception] | None = None
|
||||
|
||||
for from_exc, to_exc in AIOHTTP_EXC_MAP.items():
|
||||
if not isinstance(exc, from_exc):
|
||||
|
|
@ -222,7 +240,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
if session.closed:
|
||||
return
|
||||
|
||||
session_loop: Final = getattr(session, "_loop", None)
|
||||
session_loop: Final[asyncio.AbstractEventLoop | None] = getattr(session, "_loop", None)
|
||||
try:
|
||||
current_loop: asyncio.AbstractEventLoop | None = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
|
|
@ -278,7 +296,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
|
||||
# Check if the existing session is still valid for the current event loop
|
||||
try:
|
||||
session_loop: Final = getattr(self.client, "_loop", None)
|
||||
session_loop: Final[asyncio.AbstractEventLoop | None] = getattr(self.client, "_loop", None)
|
||||
current_loop: Final = asyncio.get_running_loop()
|
||||
|
||||
# If session is from a different or closed loop, recreate it
|
||||
|
|
@ -312,7 +330,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
self,
|
||||
client_session: ClientSession,
|
||||
request: httpx.Request,
|
||||
timeout: dict,
|
||||
timeout: HttpxTimeoutExtension,
|
||||
proxy: str | None,
|
||||
sni_hostname: str | None,
|
||||
ssl_verify: bool | ssl.SSLContext | None = None,
|
||||
|
|
@ -323,7 +341,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
Args:
|
||||
client_session: The aiohttp ClientSession to use
|
||||
request: The httpx Request to send
|
||||
timeout: Timeout settings dict with 'connect', 'read', 'pool' keys
|
||||
timeout: Timeout settings with 'connect', 'read', 'pool' fields
|
||||
proxy: Optional proxy URL
|
||||
sni_hostname: Optional SNI hostname for SSL
|
||||
ssl_verify: Optional SSL verification setting (False to disable, SSLContext for custom)
|
||||
|
|
@ -346,25 +364,24 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
# Only pass ssl kwarg when explicitly configured, to avoid
|
||||
# overriding the session/connector defaults with None (which is
|
||||
# not a valid value for aiohttp's ssl parameter).
|
||||
request_kwargs: Final[dict[str, Any]] = {
|
||||
"method": request.method,
|
||||
"url": YarlURL(str(request.url), encoded=True),
|
||||
"headers": request.headers,
|
||||
"data": data,
|
||||
"allow_redirects": False,
|
||||
"auto_decompress": False,
|
||||
"timeout": ClientTimeout(
|
||||
sock_connect=timeout.get("connect"),
|
||||
sock_read=timeout.get("read"),
|
||||
connect=timeout.get("pool"),
|
||||
),
|
||||
"proxy": proxy,
|
||||
"server_hostname": sni_hostname,
|
||||
}
|
||||
if ssl_verify is not None:
|
||||
request_kwargs["ssl"] = ssl_verify
|
||||
ssl_option: Final[AiohttpSslRequestOption] = _NO_SSL_OVERRIDE if ssl_verify is None else {"ssl": ssl_verify}
|
||||
|
||||
response: Final = await client_session.request(**request_kwargs).__aenter__()
|
||||
response: Final = await client_session.request(
|
||||
method=request.method,
|
||||
url=YarlURL(str(request.url), encoded=True),
|
||||
headers=request.headers,
|
||||
data=data,
|
||||
allow_redirects=False,
|
||||
auto_decompress=False,
|
||||
timeout=ClientTimeout(
|
||||
sock_connect=timeout.connect,
|
||||
sock_read=timeout.read,
|
||||
connect=timeout.pool,
|
||||
),
|
||||
proxy=proxy,
|
||||
server_hostname=sni_hostname,
|
||||
**ssl_option,
|
||||
).__aenter__()
|
||||
|
||||
return response
|
||||
|
||||
|
|
@ -372,8 +389,8 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
|
|||
self,
|
||||
request: httpx.Request,
|
||||
) -> httpx.Response:
|
||||
timeout: Final = request.extensions.get("timeout", {})
|
||||
sni_hostname: Final = request.extensions.get("sni_hostname")
|
||||
timeout: Final = _TIMEOUT_EXTENSION.validate_python(request.extensions.get("timeout", _EMPTY_TIMEOUT))
|
||||
sni_hostname: Final[str | None] = request.extensions.get("sni_hostname")
|
||||
|
||||
# Use helper to ensure we have a valid session for the current event loop
|
||||
client_session = self._get_valid_client_session()
|
||||
|
|
|
|||
|
|
@ -165,6 +165,7 @@ def _rust_responses_websocket_enabled(
|
|||
from .http_handler import get_shared_realtime_ssl_context
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
from aiohttp import ClientSession
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
|
|
@ -405,7 +406,7 @@ class BaseLLMHTTPHandler:
|
|||
messages: list,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: object,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
json_mode: bool = False,
|
||||
|
|
@ -471,7 +472,7 @@ class BaseLLMHTTPHandler:
|
|||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
model_response: ModelResponse,
|
||||
encoding: object,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
optional_params: dict,
|
||||
timeout: float | httpx.Timeout,
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from .base import BaseLLM
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm import CustomStreamWrapper
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class CustomLLMError(Exception): # use this for all your exceptions
|
||||
|
|
@ -134,7 +135,7 @@ class CustomLLM(BaseLLM):
|
|||
api_base: str | None,
|
||||
model_response: ImageResponse,
|
||||
optional_params: dict,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
client: HTTPHandler | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
@ -148,7 +149,7 @@ class CustomLLM(BaseLLM):
|
|||
api_key: str | None, # dynamically set api_key - https://docs.litellm.ai/docs/set_keys#api_key
|
||||
api_base: str | None, # dynamically set api_base - https://docs.litellm.ai/docs/set_keys#api_base
|
||||
optional_params: dict,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
@ -160,7 +161,7 @@ class CustomLLM(BaseLLM):
|
|||
input: list,
|
||||
model_response: EmbeddingResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
optional_params: dict,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
|
|
@ -175,7 +176,7 @@ class CustomLLM(BaseLLM):
|
|||
input: list,
|
||||
model_response: EmbeddingResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
optional_params: dict,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
|
|
@ -193,7 +194,7 @@ class CustomLLM(BaseLLM):
|
|||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
optional_params: dict,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
client: HTTPHandler | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
@ -208,7 +209,7 @@ class CustomLLM(BaseLLM):
|
|||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
optional_params: dict,
|
||||
logging_obj: Any,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -38,6 +38,8 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -157,7 +159,7 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -136,6 +136,8 @@ def _split_parallel_tool_calls(messages: list[AllMessageValues]) -> list[AllMess
|
|||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -603,7 +605,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ Authentication priority:
|
|||
import os
|
||||
import re
|
||||
from typing import Any, Final, Literal
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
|
|
@ -224,11 +225,8 @@ class DatabricksBase:
|
|||
"""
|
||||
import requests
|
||||
|
||||
# Extract workspace URL from api_base
|
||||
workspace_url = api_base.rstrip("/")
|
||||
if "/serving-endpoints" in workspace_url:
|
||||
workspace_url = workspace_url.replace("/serving-endpoints", "")
|
||||
|
||||
api_base_parts: Final = urlsplit(api_base)
|
||||
workspace_url: Final = urlunsplit((api_base_parts.scheme, api_base_parts.netloc, "", "", ""))
|
||||
token_url: Final = f"{workspace_url}/oidc/v1/token"
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from litellm.types.utils import ImageObject, ImageResponse
|
|||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -185,7 +187,7 @@ class FalAIBriaConfig(FalAIBaseConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from litellm.types.utils import ImageObject, ImageResponse
|
|||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -192,7 +194,7 @@ class FalAIFluxProV11UltraConfig(FalAIBaseConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from litellm.types.utils import ImageObject, ImageResponse
|
|||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
|
|
@ -148,7 +150,7 @@ class FalAIIdeogramV3Config(FalAIBaseConfig):
|
|||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue