mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
chore(lint): strip inert type: ignore comments and zero LIT009, LIT010, LIT011 headroom
This commit is contained in:
parent
cdfefd7f41
commit
338e411103
443 changed files with 2449 additions and 2715 deletions
|
|
@ -1269,8 +1269,8 @@ from .llms.xai.common_utils import XAIModelInfo
|
|||
from litellm.types.utils import LlmProviders
|
||||
|
||||
## Lazy loading this is not straightforward, will leave it here for now.
|
||||
from .main import * # type: ignore
|
||||
from .compression import compress # type: ignore[no-redef]
|
||||
from .main import *
|
||||
from .compression import compress
|
||||
|
||||
# Skills API
|
||||
from .skills.main import (
|
||||
|
|
@ -1341,7 +1341,7 @@ from .assistants.main import *
|
|||
from .batches.main import *
|
||||
from .images.main import *
|
||||
from .videos.main import *
|
||||
from .batch_completion.main import * # type: ignore
|
||||
from .batch_completion.main import *
|
||||
from .rerank_api.main import *
|
||||
from .llms.anthropic.experimental_pass_through.messages.handler import *
|
||||
from .responses.main import *
|
||||
|
|
@ -2054,7 +2054,7 @@ if TYPE_CHECKING:
|
|||
supports_reasoning: Callable[..., bool]
|
||||
acreate: Callable[..., Any]
|
||||
get_max_tokens: Callable[..., int]
|
||||
get_model_info: Callable[..., _ModelInfoType] # type: ignore[no-redef]
|
||||
get_model_info: Callable[..., _ModelInfoType]
|
||||
register_prompt_template: Callable[..., None]
|
||||
validate_environment: Callable[..., dict]
|
||||
check_valid_key: Callable[..., bool]
|
||||
|
|
|
|||
|
|
@ -15,8 +15,8 @@ import os
|
|||
from collections.abc import Callable
|
||||
from typing import Final
|
||||
|
||||
import redis # type: ignore
|
||||
import redis.asyncio as async_redis # type: ignore
|
||||
import redis
|
||||
import redis.asyncio as async_redis
|
||||
|
||||
from litellm import get_secret, get_secret_str
|
||||
from litellm._redis_credential_provider import (
|
||||
|
|
@ -153,7 +153,7 @@ def _redis_kwargs_from_environment():
|
|||
|
||||
return_dict: Final = {}
|
||||
for k, v in mapping.items():
|
||||
value = get_secret(k, default_value=None) # type: ignore
|
||||
value = get_secret(k, default_value=None)
|
||||
if value is not None:
|
||||
return_dict[v] = value
|
||||
return return_dict
|
||||
|
|
@ -317,7 +317,7 @@ def create_azure_ad_redis_connect_func(
|
|||
# AzureADCredentialProvider for refresh-aware token retrieval. The raw
|
||||
# client_id/tenant_id/secret are intentionally NOT exposed here — the
|
||||
# credential closure already holds them.
|
||||
ad_connect._azure_credential = credential # type: ignore[attr-defined]
|
||||
ad_connect._azure_credential = credential
|
||||
return ad_connect
|
||||
|
||||
|
||||
|
|
@ -351,7 +351,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
for k, v in env_overrides.items():
|
||||
if isinstance(v, str) and v.startswith("os.environ/"):
|
||||
v = v.replace("os.environ/", "")
|
||||
value = get_secret(v) # type: ignore
|
||||
value = get_secret(v)
|
||||
env_overrides[k] = value
|
||||
|
||||
environment_kwargs: Final = _redis_kwargs_from_environment()
|
||||
|
|
@ -370,7 +370,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
**env_overrides,
|
||||
}
|
||||
|
||||
_startup_nodes: Final[str | list | None] = redis_kwargs.get("startup_nodes", None) or get_secret( # type: ignore
|
||||
_startup_nodes: Final[str | list | None] = redis_kwargs.get("startup_nodes", None) or get_secret(
|
||||
"REDIS_CLUSTER_NODES"
|
||||
)
|
||||
|
||||
|
|
@ -381,7 +381,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
elif _startup_nodes is None:
|
||||
redis_kwargs.pop("startup_nodes", None)
|
||||
|
||||
_sentinel_nodes: Final[str | list | None] = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore
|
||||
_sentinel_nodes: Final[str | list | None] = redis_kwargs.get("sentinel_nodes", None) or get_secret(
|
||||
"REDIS_SENTINEL_NODES"
|
||||
)
|
||||
|
||||
|
|
@ -395,7 +395,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
if _sentinel_password is not None:
|
||||
redis_kwargs["sentinel_password"] = _sentinel_password
|
||||
|
||||
_service_name: Final[str | None] = redis_kwargs.get("service_name", None) or get_secret( # type: ignore
|
||||
_service_name: Final[str | None] = redis_kwargs.get("service_name", None) or get_secret(
|
||||
"REDIS_SERVICE_NAME"
|
||||
)
|
||||
|
||||
|
|
@ -412,7 +412,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs
|
||||
)
|
||||
# Store GCP service account in redis_connect_func for async cluster access
|
||||
redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account # type: ignore[attr-defined]
|
||||
redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account
|
||||
|
||||
# Remove GCP-specific kwargs that shouldn't be passed to Redis client
|
||||
redis_kwargs.pop("gcp_service_account", None)
|
||||
|
|
@ -449,7 +449,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
# `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret
|
||||
# are intentionally NOT exposed on the function to avoid leaking
|
||||
# credentials via inspection or logging.
|
||||
redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True # type: ignore[attr-defined]
|
||||
redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True
|
||||
|
||||
# Always remove Azure-specific kwargs that shouldn't be passed to Redis client
|
||||
redis_kwargs.pop("azure_redis_ad_token", None)
|
||||
|
|
@ -481,7 +481,7 @@ def _get_redis_client_logic(**env_overrides):
|
|||
|
||||
|
||||
def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
|
||||
_redis_cluster_nodes_in_env: Final[str | None] = get_secret("REDIS_CLUSTER_NODES") # type: ignore
|
||||
_redis_cluster_nodes_in_env: Final[str | None] = get_secret("REDIS_CLUSTER_NODES")
|
||||
if _redis_cluster_nodes_in_env is not None:
|
||||
try:
|
||||
redis_kwargs["startup_nodes"] = json.loads(_redis_cluster_nodes_in_env)
|
||||
|
|
@ -505,7 +505,7 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster:
|
|||
new_startup_nodes.append(ClusterNode(**item))
|
||||
|
||||
cluster_kwargs.pop("startup_nodes", None)
|
||||
return redis.RedisCluster(startup_nodes=new_startup_nodes, **cluster_kwargs) # type: ignore
|
||||
return redis.RedisCluster(startup_nodes=new_startup_nodes, **cluster_kwargs)
|
||||
|
||||
|
||||
def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict:
|
||||
|
|
@ -638,7 +638,7 @@ def get_redis_async_client(
|
|||
# Create async RedisCluster with IAM token as password if available
|
||||
cluster_client: Final = async_redis.RedisCluster(
|
||||
startup_nodes=new_startup_nodes,
|
||||
**cluster_kwargs, # type: ignore
|
||||
**cluster_kwargs,
|
||||
)
|
||||
|
||||
return cluster_client
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import threading
|
|||
import time
|
||||
from typing import Any, Final
|
||||
|
||||
from redis.credentials import CredentialProvider # type: ignore[attr-defined]
|
||||
from redis.credentials import CredentialProvider
|
||||
|
||||
# Azure AD scope for Redis Cache for Azure.
|
||||
AZURE_REDIS_SCOPE: Final = "https://redis.azure.com/.default"
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Internal unified UUID helper.
|
|||
Always uses fastuuid for performance.
|
||||
"""
|
||||
|
||||
import fastuuid as _uuid # type: ignore
|
||||
import fastuuid as _uuid
|
||||
|
||||
# Expose a module-like alias so callers can use: uuid.uuid4()
|
||||
uuid = _uuid
|
||||
|
|
|
|||
|
|
@ -18,8 +18,8 @@ AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent-card.json"
|
|||
PREV_AGENT_CARD_WELL_KNOWN_PATH: str = "/.well-known/agent.json"
|
||||
|
||||
try:
|
||||
from a2a.client import A2ACardResolver as _A2ACardResolver # type: ignore[no-redef]
|
||||
from a2a.utils.constants import ( # type: ignore[no-redef]
|
||||
from a2a.client import A2ACardResolver as _A2ACardResolver
|
||||
from a2a.utils.constants import (
|
||||
AGENT_CARD_WELL_KNOWN_PATH,
|
||||
PREV_AGENT_CARD_WELL_KNOWN_PATH,
|
||||
)
|
||||
|
|
@ -102,7 +102,7 @@ def fix_agent_card_url(agent_card: "AgentCard", base_url: str) -> "AgentCard":
|
|||
return agent_card
|
||||
|
||||
|
||||
class LiteLLMA2ACardResolver(_A2ACardResolver): # type: ignore[misc]
|
||||
class LiteLLMA2ACardResolver(_A2ACardResolver):
|
||||
"""
|
||||
Custom A2A card resolver that supports multiple well-known paths.
|
||||
|
||||
|
|
|
|||
|
|
@ -29,9 +29,9 @@ try:
|
|||
A2A_SDK_AVAILABLE = True
|
||||
except ImportError:
|
||||
A2A_SDK_AVAILABLE = False
|
||||
Client = None # type: ignore[misc, assignment]
|
||||
ClientConfig = None # type: ignore[misc, assignment]
|
||||
create_client = None # type: ignore[misc, assignment]
|
||||
Client = None
|
||||
ClientConfig = None
|
||||
create_client = None
|
||||
|
||||
|
||||
class A2AExceptionCheckers:
|
||||
|
|
@ -219,6 +219,6 @@ async def handle_a2a_localhost_retry(
|
|||
streaming=is_streaming,
|
||||
),
|
||||
)
|
||||
new_client._litellm_httpx_client = httpx_client # type: ignore[attr-defined]
|
||||
new_client._litellm_agent_card = agent_card # type: ignore[attr-defined]
|
||||
new_client._litellm_httpx_client = httpx_client
|
||||
new_client._litellm_agent_card = agent_card
|
||||
return new_client
|
||||
|
|
|
|||
|
|
@ -271,7 +271,7 @@ class A2ACompletionBridgeHandler:
|
|||
# 3. Accumulate content and emit artifact update
|
||||
accumulated_text = ""
|
||||
chunk_count = 0
|
||||
async for chunk in response: # type: ignore[union-attr]
|
||||
async for chunk in response:
|
||||
chunk_count += 1
|
||||
|
||||
# Extract delta content
|
||||
|
|
|
|||
|
|
@ -59,9 +59,9 @@ try:
|
|||
|
||||
A2A_SDK_AVAILABLE = True
|
||||
except ImportError:
|
||||
Client = None # type: ignore[misc, assignment]
|
||||
ClientConfig = None # type: ignore[misc, assignment]
|
||||
create_client = None # type: ignore[misc, assignment]
|
||||
Client = None
|
||||
ClientConfig = None
|
||||
create_client = None
|
||||
|
||||
# Import our custom card resolver that supports multiple well-known paths
|
||||
from litellm.a2a_protocol.card_resolver import (
|
||||
|
|
@ -788,10 +788,10 @@ async def create_a2a_client(
|
|||
# Stash LiteLLM-owned handles on the client so the localhost-retry path can reuse
|
||||
# the configured httpx client (with this agent's trace-id/auth headers) without
|
||||
# excavating a2a-sdk private internals.
|
||||
a2a_client._litellm_httpx_client = httpx_client # type: ignore[attr-defined]
|
||||
a2a_client._litellm_httpx_client = httpx_client
|
||||
agent_card: Final = getattr(a2a_client, "_card", None)
|
||||
if agent_card is not None:
|
||||
a2a_client._litellm_agent_card = agent_card # type: ignore[attr-defined]
|
||||
a2a_client._litellm_agent_card = agent_card
|
||||
|
||||
verbose_logger.info("A2A client created for %s", base_url)
|
||||
|
||||
|
|
|
|||
|
|
@ -153,7 +153,7 @@ class AnthropicExceptionMapping:
|
|||
# Optionally add request_id if provided and not present
|
||||
if request_id and "request_id" not in parsed:
|
||||
parsed["request_id"] = request_id
|
||||
return parsed # type: ignore
|
||||
return parsed
|
||||
|
||||
# Extract message - use parsed dict if available, otherwise raw string
|
||||
if parsed is not None:
|
||||
|
|
|
|||
|
|
@ -51,9 +51,7 @@ async def aget_assistants(
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model="", custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -61,7 +59,7 @@ async def aget_assistants(
|
|||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model="",
|
||||
|
|
@ -98,7 +96,7 @@ def get_assistants(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -132,12 +130,12 @@ def get_assistants(
|
|||
max_retries=optional_params.max_retries,
|
||||
organization=organization,
|
||||
client=client,
|
||||
aget_assistants=aget_assistants, # type: ignore
|
||||
) # type: ignore
|
||||
aget_assistants=aget_assistants,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -145,14 +143,14 @@ def get_assistants(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_assistants_api.get_assistants(
|
||||
api_base=api_base,
|
||||
|
|
@ -162,7 +160,7 @@ def get_assistants(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
client=client,
|
||||
aget_assistants=aget_assistants, # type: ignore
|
||||
aget_assistants=aget_assistants,
|
||||
litellm_params=litellm_params_dict,
|
||||
)
|
||||
else:
|
||||
|
|
@ -173,7 +171,7 @@ def get_assistants(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -185,7 +183,7 @@ def get_assistants(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -210,9 +208,7 @@ async def acreate_assistants(
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -220,7 +216,7 @@ async def acreate_assistants(
|
|||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model=model,
|
||||
|
|
@ -267,7 +263,7 @@ def create_assistants(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -318,12 +314,12 @@ def create_assistants(
|
|||
organization=organization,
|
||||
create_assistant_data=create_assistant_data,
|
||||
client=client,
|
||||
async_create_assistants=async_create_assistants, # type: ignore
|
||||
) # type: ignore
|
||||
async_create_assistants=async_create_assistants,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -331,14 +327,14 @@ def create_assistants(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
if isinstance(client, OpenAI):
|
||||
client = None # only pass client if it's AzureOpenAI
|
||||
|
|
@ -363,7 +359,7 @@ def create_assistants(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
if response is None:
|
||||
|
|
@ -392,9 +388,7 @@ async def adelete_assistant(
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model="", custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -402,7 +396,7 @@ async def adelete_assistant(
|
|||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model="",
|
||||
|
|
@ -442,7 +436,7 @@ def delete_assistant(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -472,9 +466,9 @@ def delete_assistant(
|
|||
async_delete_assistants=async_delete_assistants,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -482,14 +476,14 @@ def delete_assistant(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
if isinstance(client, OpenAI):
|
||||
client = None # only pass client if it's AzureOpenAI
|
||||
|
|
@ -541,9 +535,7 @@ async def acreate_thread(custom_llm_provider: Literal["openai", "azure"], **kwar
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model="", custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -551,7 +543,7 @@ async def acreate_thread(custom_llm_provider: Literal["openai", "azure"], **kwar
|
|||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model="",
|
||||
|
|
@ -608,7 +600,7 @@ def create_thread(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -649,7 +641,7 @@ def create_thread(
|
|||
acreate_thread=acreate_thread,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -657,16 +649,16 @@ def create_thread(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
if isinstance(client, OpenAI):
|
||||
client = None # only pass client if it's AzureOpenAI
|
||||
|
|
@ -692,10 +684,10 @@ def create_thread(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response # type: ignore
|
||||
return response
|
||||
|
||||
|
||||
async def aget_thread(
|
||||
|
|
@ -715,9 +707,7 @@ async def aget_thread(
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model="", custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -725,7 +715,7 @@ async def aget_thread(
|
|||
response = await init_response
|
||||
else:
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model="",
|
||||
|
|
@ -758,7 +748,7 @@ def get_thread(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
api_base: str | None = None
|
||||
|
|
@ -797,9 +787,9 @@ def get_thread(
|
|||
aget_thread=aget_thread,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -807,14 +797,14 @@ def get_thread(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
if isinstance(client, OpenAI):
|
||||
client = None # only pass client if it's AzureOpenAI
|
||||
|
|
@ -839,10 +829,10 @@ def get_thread(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response # type: ignore
|
||||
return response
|
||||
|
||||
|
||||
### MESSAGES ###
|
||||
|
|
@ -879,9 +869,7 @@ async def a_add_message(
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model="", custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -890,7 +878,7 @@ async def a_add_message(
|
|||
else:
|
||||
# Call the synchronous function using run_in_executor
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model="",
|
||||
|
|
@ -937,7 +925,7 @@ def add_message(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
api_key: str | None = None
|
||||
|
|
@ -976,9 +964,9 @@ def add_message(
|
|||
a_add_message=a_add_message,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -986,14 +974,14 @@ def add_message(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_assistants_api.add_message(
|
||||
thread_id=thread_id,
|
||||
|
|
@ -1016,11 +1004,11 @@ def add_message(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
return response # type: ignore
|
||||
return response
|
||||
|
||||
|
||||
async def aget_messages(
|
||||
|
|
@ -1046,9 +1034,7 @@ async def aget_messages(
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model="", custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -1057,7 +1043,7 @@ async def aget_messages(
|
|||
else:
|
||||
# Call the synchronous function using run_in_executor
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model="",
|
||||
|
|
@ -1090,7 +1076,7 @@ def get_messages(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -1129,9 +1115,9 @@ def get_messages(
|
|||
aget_messages=aget_messages,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version: str | None = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -1139,14 +1125,14 @@ def get_messages(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token: str | None = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_assistants_api.get_messages(
|
||||
thread_id=thread_id,
|
||||
|
|
@ -1168,11 +1154,11 @@ def get_messages(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
return response # type: ignore
|
||||
return response
|
||||
|
||||
|
||||
### RUNS ###
|
||||
|
|
@ -1182,7 +1168,7 @@ def arun_thread_stream(
|
|||
**kwargs,
|
||||
) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]:
|
||||
kwargs["arun_thread"] = True
|
||||
return run_thread(stream=True, event_handler=event_handler, **kwargs) # type: ignore
|
||||
return run_thread(stream=True, event_handler=event_handler, **kwargs)
|
||||
|
||||
|
||||
async def arun_thread(
|
||||
|
|
@ -1222,9 +1208,7 @@ async def arun_thread(
|
|||
ctx: Final = contextvars.copy_context()
|
||||
func_with_context: Final = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider( # type: ignore
|
||||
model="", custom_llm_provider=custom_llm_provider
|
||||
) # type: ignore
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model="", custom_llm_provider=custom_llm_provider)
|
||||
|
||||
# Await normally
|
||||
init_response: Final = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -1233,7 +1217,7 @@ async def arun_thread(
|
|||
else:
|
||||
# Call the synchronous function using run_in_executor
|
||||
response = init_response
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise exception_type(
|
||||
model="",
|
||||
|
|
@ -1249,7 +1233,7 @@ def run_thread_stream(
|
|||
event_handler: AssistantEventHandler | None = None,
|
||||
**kwargs,
|
||||
) -> AssistantStreamManager[AssistantEventHandler]:
|
||||
return run_thread(stream=True, event_handler=event_handler, **kwargs) # type: ignore
|
||||
return run_thread(stream=True, event_handler=event_handler, **kwargs)
|
||||
|
||||
|
||||
def run_thread(
|
||||
|
|
@ -1283,7 +1267,7 @@ def run_thread(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -1329,9 +1313,9 @@ def run_thread(
|
|||
event_handler=event_handler,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -1339,14 +1323,14 @@ def run_thread(
|
|||
or litellm.azure_key
|
||||
or get_secret("AZURE_OPENAI_API_KEY")
|
||||
or get_secret("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body: Final = optional_params.get("extra_body", {})
|
||||
azure_ad_token = None
|
||||
if extra_body is not None:
|
||||
azure_ad_token = extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN") # type: ignore
|
||||
azure_ad_token = get_secret("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_assistants_api.run_thread(
|
||||
thread_id=thread_id,
|
||||
|
|
@ -1366,7 +1350,7 @@ def run_thread(
|
|||
client=client,
|
||||
arun_thread=arun_thread,
|
||||
litellm_params=litellm_params_dict,
|
||||
) # type: ignore
|
||||
)
|
||||
else:
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=f"LiteLLM doesn't support {custom_llm_provider} for 'run_thread'. Only 'openai' is supported.",
|
||||
|
|
@ -1375,7 +1359,7 @@ def run_thread(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response # type: ignore
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -274,7 +274,7 @@ async def _fetch_batch_output_file_content(
|
|||
credentials: Final = _extract_file_access_credentials(litellm_params)
|
||||
file_content_kwargs.update(credentials)
|
||||
|
||||
_file_content: Final = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType]
|
||||
_file_content: Final = await afile_content(**file_content_kwargs)
|
||||
return _file_content.content
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -287,7 +287,7 @@ def create_batch(
|
|||
if extra_body is not None:
|
||||
extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
get_secret_str("AZURE_AD_TOKEN") # type: ignore
|
||||
get_secret_str("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_batches_instance.create_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -327,7 +327,7 @@ def create_batch(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_batch", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_batch", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -370,7 +370,7 @@ async def aretrieve_batch(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
|
|
@ -436,7 +436,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
if extra_body is not None:
|
||||
extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
get_secret_str("AZURE_AD_TOKEN") # type: ignore
|
||||
get_secret_str("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_batches_instance.retrieve_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -498,7 +498,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="retrieve_batch", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="retrieve_batch", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -545,7 +545,7 @@ def retrieve_batch(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -677,7 +677,7 @@ async def alist_batches(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
|
|
@ -723,7 +723,7 @@ def list_batches(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -755,7 +755,7 @@ def list_batches(
|
|||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
|
|
@ -770,7 +770,7 @@ def list_batches(
|
|||
if extra_body is not None:
|
||||
extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
get_secret_str("AZURE_AD_TOKEN") # type: ignore
|
||||
get_secret_str("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_batches_instance.list_batches(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -813,7 +813,7 @@ def list_batches(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -909,7 +909,7 @@ def cancel_batch(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -959,7 +959,7 @@ def cancel_batch(
|
|||
if extra_body is not None:
|
||||
extra_body.pop("azure_ad_token", None)
|
||||
else:
|
||||
get_secret_str("AZURE_AD_TOKEN") # type: ignore
|
||||
get_secret_str("AZURE_AD_TOKEN")
|
||||
|
||||
response = azure_batches_instance.cancel_batch(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -999,7 +999,7 @@ def cancel_batch(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="cancel_batch", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="cancel_batch", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -534,11 +534,9 @@ class Cache:
|
|||
if isinstance(cached_response, dict):
|
||||
pass
|
||||
else:
|
||||
cached_response = json.loads(
|
||||
cached_response # type: ignore
|
||||
) # Convert string to dictionary
|
||||
cached_response = json.loads(cached_response) # Convert string to dictionary
|
||||
except Exception:
|
||||
cached_response = ast.literal_eval(cached_response) # type: ignore
|
||||
cached_response = ast.literal_eval(cached_response)
|
||||
return cached_response
|
||||
return cached_result
|
||||
|
||||
|
|
|
|||
|
|
@ -242,7 +242,7 @@ class LLMCachingHandler:
|
|||
or litellm.cache.get_cache_key(**self.request_kwargs)
|
||||
)
|
||||
if hasattr(cached_result, "_hidden_params"):
|
||||
cached_result._hidden_params["cache_key"] = cache_key # type: ignore
|
||||
cached_result._hidden_params["cache_key"] = cache_key
|
||||
return CachingHandlerResponse(cached_result=cached_result)
|
||||
elif (
|
||||
call_type == CallTypes.aembedding.value
|
||||
|
|
@ -356,7 +356,7 @@ class LLMCachingHandler:
|
|||
or litellm.cache.get_cache_key(**self.request_kwargs)
|
||||
)
|
||||
if hasattr(cached_result, "_hidden_params"):
|
||||
cached_result._hidden_params["cache_key"] = cache_key # type: ignore
|
||||
cached_result._hidden_params["cache_key"] = cache_key
|
||||
return CachingHandlerResponse(cached_result=cached_result)
|
||||
return CachingHandlerResponse(cached_result=cached_result)
|
||||
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ class DiskCache(BaseCache):
|
|||
original_cached_response: Final = self.disk_cache.get(key)
|
||||
if original_cached_response:
|
||||
try:
|
||||
cached_response = json.loads(original_cached_response) # type: ignore
|
||||
cached_response = json.loads(original_cached_response)
|
||||
except Exception:
|
||||
cached_response = original_cached_response
|
||||
return cached_response
|
||||
|
|
|
|||
|
|
@ -242,7 +242,7 @@ async def _run_under_circuit_breaker(
|
|||
return result
|
||||
|
||||
|
||||
def _redis_circuit_breaker_guard(method): # type: ignore
|
||||
def _redis_circuit_breaker_guard(method):
|
||||
"""
|
||||
Decorator for RedisCache async methods.
|
||||
Checks the circuit breaker before each call; records success/failure after.
|
||||
|
|
@ -256,7 +256,7 @@ def _redis_circuit_breaker_guard(method): # type: ignore
|
|||
"""
|
||||
|
||||
@functools.wraps(method)
|
||||
async def wrapper(self, *args, **kwargs): # type: ignore
|
||||
async def wrapper(self, *args, **kwargs):
|
||||
return await _run_under_circuit_breaker(
|
||||
self._circuit_breaker, method.__name__, lambda: method(self, *args, **kwargs)
|
||||
)
|
||||
|
|
@ -319,7 +319,7 @@ class RedisCache(BaseCache):
|
|||
self.redis_version = "Unknown"
|
||||
try:
|
||||
if not coroutine_checker.is_async_callable(self.redis_client):
|
||||
self.redis_version = self.redis_client.info()["redis_version"] # type: ignore
|
||||
self.redis_version = self.redis_client.info()["redis_version"]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -355,7 +355,7 @@ class RedisCache(BaseCache):
|
|||
# SYNC HEALTH PING
|
||||
try:
|
||||
if hasattr(self.redis_client, "ping"):
|
||||
self.redis_client.ping() # type: ignore
|
||||
self.redis_client.ping()
|
||||
except Exception as e:
|
||||
verbose_logger.error("Error connecting to Sync Redis client", extra={"error": str(e)})
|
||||
self._handle_sync_ping_error(e)
|
||||
|
|
@ -423,7 +423,7 @@ class RedisCache(BaseCache):
|
|||
redis_async_client = get_redis_async_client(connection_pool=self.async_redis_conn_pool, **self.redis_kwargs)
|
||||
in_memory_llm_clients_cache.set_cache(key=cache_key, value=redis_async_client)
|
||||
|
||||
self.redis_async_client = redis_async_client # type: ignore
|
||||
self.redis_async_client = redis_async_client
|
||||
return redis_async_client
|
||||
|
||||
def check_and_fix_namespace(self, key: str) -> str:
|
||||
|
|
@ -431,7 +431,7 @@ class RedisCache(BaseCache):
|
|||
Make sure each key starts with the given namespace
|
||||
"""
|
||||
if key is None:
|
||||
return key # type: ignore[return-value]
|
||||
return key
|
||||
if self.namespace is not None and not key.startswith(self.namespace):
|
||||
key = self.namespace + ":" + key
|
||||
|
||||
|
|
@ -493,7 +493,7 @@ class RedisCache(BaseCache):
|
|||
key = self.check_and_fix_namespace(key=key)
|
||||
try:
|
||||
start_time = time.time()
|
||||
result: Final[int] = _redis_client.incr(name=key, amount=value) # type: ignore
|
||||
result: Final[int] = _redis_client.incr(name=key, amount=value)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
|
|
@ -520,7 +520,7 @@ class RedisCache(BaseCache):
|
|||
if current_ttl == -1:
|
||||
# Key has no expiration
|
||||
start_time = time.time()
|
||||
_redis_client.expire(key, set_ttl) # type: ignore
|
||||
_redis_client.expire(key, set_ttl)
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
self.service_logger_obj.service_success_hook(
|
||||
|
|
@ -555,7 +555,7 @@ class RedisCache(BaseCache):
|
|||
return []
|
||||
|
||||
pattern = self.check_and_fix_namespace(key=pattern)
|
||||
async for key in _redis_client.scan_iter(match=pattern + "*", count=count): # type: ignore
|
||||
async for key in _redis_client.scan_iter(match=pattern + "*", count=count):
|
||||
keys.append(key)
|
||||
if len(keys) >= count:
|
||||
break
|
||||
|
|
@ -680,7 +680,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client()
|
||||
except Exception as e:
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -773,7 +773,7 @@ class RedisCache(BaseCache):
|
|||
_td: timedelta | None = None
|
||||
if ttl is not None:
|
||||
_td = timedelta(seconds=ttl)
|
||||
pipe.set( # type: ignore
|
||||
pipe.set(
|
||||
name=cache_key,
|
||||
value=json_cache_value,
|
||||
ex=_td,
|
||||
|
|
@ -849,7 +849,7 @@ class RedisCache(BaseCache):
|
|||
"""Helper function for async_set_cache_sadd. Separated for testing."""
|
||||
ttl = self.get_ttl(ttl=ttl)
|
||||
try:
|
||||
await redis_client.sadd(key, *value) # type: ignore
|
||||
await redis_client.sadd(key, *value)
|
||||
if ttl is not None:
|
||||
_td: Final = timedelta(seconds=ttl)
|
||||
await redis_client.expire(key, _td)
|
||||
|
|
@ -862,7 +862,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client()
|
||||
except Exception as e:
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
|
|
@ -945,7 +945,7 @@ class RedisCache(BaseCache):
|
|||
) -> float:
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
_used_ttl: Final = self.get_ttl(ttl=ttl)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
|
|
@ -1080,7 +1080,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
We use a wrapper so RedisCluster can override this method
|
||||
"""
|
||||
return self.redis_client.mget(keys=keys) # type: ignore
|
||||
return self.redis_client.mget(keys=keys)
|
||||
|
||||
async def _async_run_redis_mget_operation(self, keys: list[str]) -> list[Any]:
|
||||
"""
|
||||
|
|
@ -1089,7 +1089,7 @@ class RedisCache(BaseCache):
|
|||
We use a wrapper so RedisCluster can override this method
|
||||
"""
|
||||
async_redis_client: Final = self.init_async_client()
|
||||
return await async_redis_client.mget(keys=keys) # type: ignore
|
||||
return await async_redis_client.mget(keys=keys)
|
||||
|
||||
def batch_get_cache(
|
||||
self,
|
||||
|
|
@ -1147,7 +1147,7 @@ class RedisCache(BaseCache):
|
|||
async def async_get_cache(self, key, parent_otel_span: Span | None = None, **kwargs):
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client()
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
start_time: Final = time.time()
|
||||
|
||||
|
|
@ -1269,7 +1269,7 @@ class RedisCache(BaseCache):
|
|||
print_verbose("Pinging Sync Redis Cache")
|
||||
start_time: Final = time.time()
|
||||
try:
|
||||
response: Final[bool] = self.redis_client.ping() # type: ignore
|
||||
response: Final[bool] = self.redis_client.ping()
|
||||
print_verbose(f"Redis Cache PING: {response}")
|
||||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
|
|
@ -1339,7 +1339,7 @@ class RedisCache(BaseCache):
|
|||
await _redis_client.delete(*keys)
|
||||
|
||||
def client_list(self) -> list:
|
||||
client_list: Final[list] = self.redis_client.client_list() # type: ignore
|
||||
client_list: Final[list] = self.redis_client.client_list()
|
||||
return client_list
|
||||
|
||||
def info(self):
|
||||
|
|
@ -1376,10 +1376,10 @@ class RedisCache(BaseCache):
|
|||
redis_client: Final = redis_async.Redis(**self.redis_kwargs)
|
||||
|
||||
# Test the connection
|
||||
ping_result: Final = await redis_client.ping() # type: ignore[misc]
|
||||
ping_result: Final = await redis_client.ping()
|
||||
|
||||
# Close the connection
|
||||
await redis_client.aclose() # type: ignore[attr-defined]
|
||||
await redis_client.aclose()
|
||||
|
||||
if ping_result:
|
||||
return {
|
||||
|
|
@ -1448,7 +1448,7 @@ class RedisCache(BaseCache):
|
|||
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_redis_client: Final[Redis] = self.init_async_client() # type: ignore
|
||||
_redis_client: Final[Redis] = self.init_async_client()
|
||||
start_time: Final = time.time()
|
||||
|
||||
print_verbose(f"Increment Async Redis Cache Pipeline: increment list: {increment_list}")
|
||||
|
|
@ -1769,7 +1769,7 @@ class RedisCache(BaseCache):
|
|||
or None
|
||||
)
|
||||
except Exception:
|
||||
decoded_results.append(r) # type: ignore
|
||||
decoded_results.append(r)
|
||||
else:
|
||||
decoded_results.append(None)
|
||||
return decoded_results
|
||||
|
|
|
|||
|
|
@ -47,14 +47,14 @@ class RedisClusterCache(RedisCache):
|
|||
"""
|
||||
Overrides `_run_redis_mget_operation` in redis_cache.py
|
||||
"""
|
||||
return self.redis_client.mget_nonatomic(keys=keys) # type: ignore
|
||||
return self.redis_client.mget_nonatomic(keys=keys)
|
||||
|
||||
async def _async_run_redis_mget_operation(self, keys: list[str]) -> list[Any]:
|
||||
"""
|
||||
Overrides `_async_run_redis_mget_operation` in redis_cache.py
|
||||
"""
|
||||
async_redis_cluster_client: Final = self.init_async_client()
|
||||
return await async_redis_cluster_client.mget_nonatomic(keys=keys) # type: ignore
|
||||
return await async_redis_cluster_client.mget_nonatomic(keys=keys)
|
||||
|
||||
async def test_connection(self) -> dict:
|
||||
"""
|
||||
|
|
@ -78,14 +78,14 @@ class RedisClusterCache(RedisCache):
|
|||
# Create a fresh Redis Cluster client with current settings
|
||||
redis_client: Final = redis_async.RedisCluster(
|
||||
startup_nodes=new_startup_nodes,
|
||||
**cluster_kwargs, # type: ignore
|
||||
**cluster_kwargs,
|
||||
)
|
||||
|
||||
# Test the connection
|
||||
ping_result: Final = await redis_client.ping() # type: ignore[attr-defined, misc]
|
||||
ping_result: Final = await redis_client.ping()
|
||||
|
||||
# Close the connection
|
||||
await redis_client.aclose() # type: ignore[attr-defined]
|
||||
await redis_client.aclose()
|
||||
|
||||
if ping_result:
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -126,8 +126,8 @@ class RedisSemanticCache(BaseCache):
|
|||
# CustomTextVectorizer probes its embedding dimension at construction by
|
||||
# embedding "dimension test", so the first cache request issues one extra
|
||||
# billable embedding on top of the request's own.
|
||||
from redisvl.extensions.llmcache import SemanticCache # type: ignore[import-not-found, import-untyped]
|
||||
from redisvl.utils.vectorize import CustomTextVectorizer # type: ignore[import-not-found, import-untyped]
|
||||
from redisvl.extensions.llmcache import SemanticCache
|
||||
from redisvl.utils.vectorize import CustomTextVectorizer
|
||||
|
||||
try:
|
||||
cache_vectorizer: Final = CustomTextVectorizer(self._get_embedding)
|
||||
|
|
@ -207,7 +207,7 @@ class RedisSemanticCache(BaseCache):
|
|||
return {self.CACHE_KEY_FIELD_NAME: str(key)}
|
||||
|
||||
def _get_cache_key_filter_expression(self, key: str) -> Any:
|
||||
from redisvl.query.filter import Tag # type: ignore[import-not-found, import-untyped]
|
||||
from redisvl.query.filter import Tag
|
||||
|
||||
return Tag(self.CACHE_KEY_FIELD_NAME) == str(key)
|
||||
|
||||
|
|
|
|||
|
|
@ -146,7 +146,7 @@ class S3Cache(BaseCache):
|
|||
)
|
||||
|
||||
return cached_response
|
||||
except botocore.exceptions.ClientError as e: # type: ignore
|
||||
except botocore.exceptions.ClientError as e:
|
||||
if e.response["Error"]["Code"] == "NoSuchKey":
|
||||
verbose_logger.debug("S3 Cache: The specified key '%s' does not exist in the S3 bucket.", key)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -85,12 +85,8 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
resolved_url = None
|
||||
if sync_client is None or async_client is None:
|
||||
resolved_url = redis_url or self._build_valkey_url(host, port, password, ssl)
|
||||
self.sync_client = (
|
||||
sync_client if sync_client is not None else Redis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
self.async_client = (
|
||||
async_client if async_client is not None else AsyncRedis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
self.sync_client = sync_client if sync_client is not None else Redis.from_url(resolved_url)
|
||||
self.async_client = async_client if async_client is not None else AsyncRedis.from_url(resolved_url)
|
||||
|
||||
print_verbose(f"Valkey semantic-cache initializing index - {self.index_name}")
|
||||
|
||||
|
|
|
|||
|
|
@ -238,7 +238,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(result, model, custom_llm_provider)
|
||||
completion_stream: Final = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
streaming_response=result,
|
||||
sync_stream=True,
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
|
|
@ -336,7 +336,7 @@ class ResponsesToCompletionBridgeHandler:
|
|||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(result, model, custom_llm_provider)
|
||||
completion_stream: Final = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
streaming_response=result,
|
||||
sync_stream=False,
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -63,9 +63,9 @@ def _get_reasoning_items(
|
|||
msg: "AllMessageValues",
|
||||
) -> list[ChatCompletionReasoningItem]:
|
||||
"""Extract reasoning_items from a message dict with proper typing."""
|
||||
items: Final = msg.get("reasoning_items") # type: ignore[union-attr]
|
||||
items: Final = msg.get("reasoning_items")
|
||||
if items:
|
||||
return items # type: ignore[return-value]
|
||||
return items
|
||||
return []
|
||||
|
||||
|
||||
|
|
@ -261,8 +261,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
"type": "message",
|
||||
"role": role,
|
||||
"content": self._convert_content_to_responses_format(
|
||||
content, # type: ignore[arg-type]
|
||||
role, # type: ignore
|
||||
content,
|
||||
role,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
|
@ -336,7 +336,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
{
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"content": self._convert_content_to_responses_format(content, cast(str, role)), # type: ignore[arg-type]
|
||||
"content": self._convert_content_to_responses_format(content, cast(str, role)),
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -360,17 +360,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
elif key == "response_format":
|
||||
text_format = self._transform_response_format_to_text_format(value)
|
||||
if text_format:
|
||||
responses_api_request["text"] = text_format # type: ignore
|
||||
responses_api_request["text"] = text_format
|
||||
elif key == "tool_choice":
|
||||
responses_api_request["tool_choice"] = ( # type: ignore[assignment]
|
||||
self._normalize_tool_choice_for_responses_api(value)
|
||||
)
|
||||
responses_api_request["tool_choice"] = self._normalize_tool_choice_for_responses_api(value)
|
||||
elif key == "stream_options":
|
||||
stream_options = normalize_responses_api_stream_options(value)
|
||||
if stream_options is not None:
|
||||
responses_api_request["stream_options"] = stream_options
|
||||
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
|
||||
responses_api_request[key] = value # type: ignore
|
||||
responses_api_request[key] = value
|
||||
elif key == "previous_response_id":
|
||||
responses_api_request["previous_response_id"] = value
|
||||
elif key == "reasoning_effort":
|
||||
|
|
@ -524,7 +522,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
ResponseApplyPatchToolCall,
|
||||
)
|
||||
except ImportError:
|
||||
ResponseApplyPatchToolCall = None # type: ignore[assignment,misc]
|
||||
ResponseApplyPatchToolCall = None
|
||||
|
||||
from litellm.types.utils import Choices, Message
|
||||
|
||||
|
|
@ -942,7 +940,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
flat_custom["format"] = convert_custom_tool_format_to_responses_shape(custom_payload["format"])
|
||||
responses_tools.append(flat_custom)
|
||||
else:
|
||||
responses_tools.append(tool) # type: ignore
|
||||
responses_tools.append(tool)
|
||||
|
||||
return cast(list["ALL_RESPONSES_API_TOOL_PARAMS"], responses_tools)
|
||||
|
||||
|
|
@ -978,7 +976,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
def _map_reasoning_effort(self, reasoning_effort: str | dict[str, Any]) -> Reasoning | None:
|
||||
# If dict is passed, convert it directly to Reasoning object
|
||||
if isinstance(reasoning_effort, dict):
|
||||
return Reasoning(**reasoning_effort) # type: ignore[typeddict-item]
|
||||
return Reasoning(**reasoning_effort)
|
||||
|
||||
# Check if auto-summary is enabled via flag or environment variable
|
||||
# Priority: litellm.reasoning_auto_summary flag > LITELLM_REASONING_AUTO_SUMMARY env var
|
||||
|
|
@ -988,11 +986,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
# If string is passed, map with optional summary based on flag/env var
|
||||
if reasoning_effort == "none":
|
||||
return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none") # type: ignore
|
||||
return Reasoning(effort="none", summary="detailed") if auto_summary_enabled else Reasoning(effort="none")
|
||||
elif reasoning_effort == "high":
|
||||
return Reasoning(effort="high", summary="detailed") if auto_summary_enabled else Reasoning(effort="high")
|
||||
elif reasoning_effort == "xhigh":
|
||||
return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") # type: ignore[typeddict-item]
|
||||
return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh")
|
||||
elif reasoning_effort == "medium":
|
||||
return (
|
||||
Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium")
|
||||
|
|
@ -1108,7 +1106,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
verbose_logger.debug("Skipping unsupported annotation type: %s", type(annotation))
|
||||
continue
|
||||
|
||||
result.append(annotation_dict) # type: ignore
|
||||
result.append(annotation_dict)
|
||||
except Exception as e:
|
||||
# Skip malformed annotations
|
||||
verbose_logger.debug("Skipping malformed annotation: %s, error: %s", annotation, e)
|
||||
|
|
@ -1254,7 +1252,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
function=function_chunk,
|
||||
)
|
||||
if provider_specific_fields:
|
||||
tool_call_chunk.provider_specific_fields = provider_specific_fields # type: ignore
|
||||
tool_call_chunk.provider_specific_fields = provider_specific_fields
|
||||
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
|
|
|
|||
|
|
@ -84,7 +84,7 @@ def bm25_score_messages(
|
|||
# document tokens that start with that term (min 4 chars match). This lets
|
||||
# "cook" match "cooking" and "auth" match "authentication" without a full
|
||||
# stemmer dependency.
|
||||
def _expand_tf(query_term: str, tf_counts: Counter) -> int: # type: ignore[type-arg]
|
||||
def _expand_tf(query_term: str, tf_counts: Counter) -> int:
|
||||
"""Sum TF across all doc tokens that are prefixed by query_term."""
|
||||
exact: Final = tf_counts.get(query_term, 0)
|
||||
if exact:
|
||||
|
|
|
|||
|
|
@ -187,7 +187,7 @@ def create_container(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
|
|
@ -405,7 +405,7 @@ def list_containers(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
|
|
@ -596,7 +596,7 @@ def retrieve_container(
|
|||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
|
|
@ -811,7 +811,7 @@ def delete_container(
|
|||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
|
|
@ -1040,7 +1040,7 @@ def list_container_files(
|
|||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
|
|
@ -1291,7 +1291,7 @@ def upload_container_file(
|
|||
local_vars: Final = locals()
|
||||
try:
|
||||
resolved_custom_llm_provider: str = custom_llm_provider
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id")
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
|
|
|
|||
|
|
@ -53,7 +53,7 @@ class ContainerRequestUtils:
|
|||
|
||||
for param in valid_params:
|
||||
if param in passed_params and passed_params[param] is not None:
|
||||
container_create_optional_params[param] = passed_params[param] # type: ignore
|
||||
container_create_optional_params[param] = passed_params[param]
|
||||
|
||||
return container_create_optional_params
|
||||
|
||||
|
|
@ -69,7 +69,7 @@ class ContainerRequestUtils:
|
|||
filtered_params: Final = {k: v for k, v in container_create_optional_params.items() if k in supported_params}
|
||||
|
||||
return container_provider_config.map_openai_params(
|
||||
container_create_optional_params=filtered_params, # type: ignore
|
||||
container_create_optional_params=filtered_params,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
|
|
@ -90,7 +90,7 @@ class ContainerRequestUtils:
|
|||
|
||||
for param in valid_params:
|
||||
if param in passed_params and passed_params[param] is not None:
|
||||
container_list_optional_params[param] = passed_params[param] # type: ignore
|
||||
container_list_optional_params[param] = passed_params[param]
|
||||
|
||||
return container_list_optional_params
|
||||
|
||||
|
|
|
|||
|
|
@ -329,7 +329,7 @@ def cost_per_token(
|
|||
response: Any | None = None,
|
||||
### REQUEST MODEL ###
|
||||
request_model: str | None = None, # original request model for router detection
|
||||
) -> tuple[float, float]: # type: ignore
|
||||
) -> tuple[float, float]:
|
||||
"""
|
||||
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
|
||||
|
||||
|
|
@ -1514,7 +1514,7 @@ def completion_cost(
|
|||
# see https://replicate.com/pricing
|
||||
elif (model in litellm.replicate_models or "replicate" in model) and model not in litellm.model_cost:
|
||||
# for unmapped replicate model, default to replicate's time tracking logic
|
||||
return get_replicate_completion_pricing(completion_response, total_time) # type: ignore
|
||||
return get_replicate_completion_pricing(completion_response, total_time)
|
||||
|
||||
if model is None:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -141,7 +141,7 @@ def create_eval(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acreate_eval", False) is True
|
||||
|
||||
|
|
@ -153,7 +153,7 @@ def create_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -162,15 +162,15 @@ def create_eval(
|
|||
|
||||
# Build create request
|
||||
create_request: Final[CreateEvalRequest] = {
|
||||
"data_source_config": data_source_config, # type: ignore
|
||||
"testing_criteria": testing_criteria, # type: ignore
|
||||
"data_source_config": data_source_config,
|
||||
"testing_criteria": testing_criteria,
|
||||
}
|
||||
if name is not None:
|
||||
create_request["name"] = name
|
||||
|
||||
# Merge extra_body if provided
|
||||
if extra_body:
|
||||
create_request.update(extra_body) # type: ignore
|
||||
create_request.update(extra_body)
|
||||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
|
|
@ -199,7 +199,7 @@ def create_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.create_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.create_eval_handler(
|
||||
url=url,
|
||||
request_body=request_body,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -326,7 +326,7 @@ def list_evals(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("alist_evals", False) is True
|
||||
|
||||
|
|
@ -338,7 +338,7 @@ def list_evals(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -354,13 +354,13 @@ def list_evals(
|
|||
if before is not None:
|
||||
list_params["before"] = before
|
||||
if order is not None:
|
||||
list_params["order"] = order # type: ignore
|
||||
list_params["order"] = order
|
||||
if order_by is not None:
|
||||
list_params["order_by"] = order_by # type: ignore
|
||||
list_params["order_by"] = order_by
|
||||
|
||||
# Merge extra_query if provided
|
||||
if extra_query:
|
||||
list_params.update(extra_query) # type: ignore
|
||||
list_params.update(extra_query)
|
||||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
|
|
@ -385,7 +385,7 @@ def list_evals(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.list_evals_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.list_evals_handler(
|
||||
url=url,
|
||||
query_params=query_params,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -492,7 +492,7 @@ def get_eval(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("aget_eval", False) is True
|
||||
|
||||
|
|
@ -504,7 +504,7 @@ def get_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -536,7 +536,7 @@ def get_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.get_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.get_eval_handler(
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -657,7 +657,7 @@ def update_eval(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("aupdate_eval", False) is True
|
||||
|
||||
|
|
@ -669,7 +669,7 @@ def update_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -723,7 +723,7 @@ def update_eval(
|
|||
|
||||
# Merge extra_body if provided
|
||||
if extra_body:
|
||||
update_request.update(extra_body) # type: ignore
|
||||
update_request.update(extra_body)
|
||||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
|
|
@ -755,7 +755,7 @@ def update_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.update_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.update_eval_handler(
|
||||
url=url,
|
||||
request_body=request_body,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -862,7 +862,7 @@ def delete_eval(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("adelete_eval", False) is True
|
||||
|
||||
|
|
@ -874,7 +874,7 @@ def delete_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -906,7 +906,7 @@ def delete_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.delete_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.delete_eval_handler(
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1012,7 +1012,7 @@ def cancel_eval(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acancel_eval", False) is True
|
||||
|
||||
|
|
@ -1024,7 +1024,7 @@ def cancel_eval(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -1060,7 +1060,7 @@ def cancel_eval(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.cancel_eval_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.cancel_eval_handler(
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1191,7 +1191,7 @@ def create_run(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acreate_run", False) is True
|
||||
|
||||
|
|
@ -1203,7 +1203,7 @@ def create_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -1212,7 +1212,7 @@ def create_run(
|
|||
|
||||
# Build create request
|
||||
create_request: Final[CreateRunRequest] = {
|
||||
"data_source": data_source, # type: ignore
|
||||
"data_source": data_source,
|
||||
}
|
||||
if name is not None:
|
||||
create_request["name"] = name
|
||||
|
|
@ -1221,7 +1221,7 @@ def create_run(
|
|||
|
||||
# Merge extra_body if provided
|
||||
if extra_body:
|
||||
create_request.update(extra_body) # type: ignore
|
||||
create_request.update(extra_body)
|
||||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
|
|
@ -1248,7 +1248,7 @@ def create_run(
|
|||
)
|
||||
|
||||
# Make HTTP request (default 600s timeout for long-running operations)
|
||||
response: Final = base_llm_http_handler.create_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.create_run_handler(
|
||||
url=url,
|
||||
request_body=request_body,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -1375,7 +1375,7 @@ def list_runs(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("alist_runs", False) is True
|
||||
|
||||
|
|
@ -1387,7 +1387,7 @@ def list_runs(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -1403,11 +1403,11 @@ def list_runs(
|
|||
if before is not None:
|
||||
list_params["before"] = before
|
||||
if order is not None:
|
||||
list_params["order"] = order # type: ignore
|
||||
list_params["order"] = order
|
||||
|
||||
# Merge extra_query if provided
|
||||
if extra_query:
|
||||
list_params.update(extra_query) # type: ignore
|
||||
list_params.update(extra_query)
|
||||
|
||||
# Validate environment and get headers
|
||||
headers = extra_headers or {}
|
||||
|
|
@ -1433,7 +1433,7 @@ def list_runs(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.list_runs_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.list_runs_handler(
|
||||
url=url,
|
||||
query_params=query_params,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
|
|
@ -1545,7 +1545,7 @@ def get_run(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("aget_run", False) is True
|
||||
|
||||
|
|
@ -1557,7 +1557,7 @@ def get_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -1590,7 +1590,7 @@ def get_run(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.get_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.get_run_handler(
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1701,7 +1701,7 @@ def cancel_run(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acancel_run", False) is True
|
||||
|
||||
|
|
@ -1713,7 +1713,7 @@ def cancel_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -1750,7 +1750,7 @@ def cancel_run(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.cancel_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.cancel_run_handler(
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1866,7 +1866,7 @@ def delete_run(
|
|||
"""
|
||||
local_vars: Final = locals()
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("adelete_run", False) is True
|
||||
|
||||
|
|
@ -1878,7 +1878,7 @@ def delete_run(
|
|||
custom_llm_provider = "openai"
|
||||
|
||||
# Get provider config
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config( # type: ignore
|
||||
evals_api_provider_config: BaseEvalsAPIConfig | None = ProviderConfigManager.get_provider_evals_api_config(
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
|
|
@ -1915,7 +1915,7 @@ def delete_run(
|
|||
)
|
||||
|
||||
# Make HTTP request
|
||||
response: Final = base_llm_http_handler.delete_run_handler( # type: ignore
|
||||
response: Final = base_llm_http_handler.delete_run_handler(
|
||||
url=url,
|
||||
evals_api_provider_config=evals_api_provider_config,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -126,7 +126,7 @@ def _get_minimal_error_response() -> httpx.Response:
|
|||
return _MINIMAL_ERROR_RESPONSE
|
||||
|
||||
|
||||
class AuthenticationError(openai.AuthenticationError): # type: ignore
|
||||
class AuthenticationError(openai.AuthenticationError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -170,7 +170,7 @@ class AuthenticationError(openai.AuthenticationError): # type: ignore
|
|||
|
||||
|
||||
# raise when invalid models passed, example gpt-8
|
||||
class NotFoundError(openai.NotFoundError): # type: ignore
|
||||
class NotFoundError(openai.NotFoundError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -213,7 +213,7 @@ class NotFoundError(openai.NotFoundError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class BadRequestError(openai.BadRequestError): # type: ignore
|
||||
class BadRequestError(openai.BadRequestError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -288,7 +288,7 @@ class ImageFetchError(BadRequestError):
|
|||
)
|
||||
|
||||
|
||||
class UnprocessableEntityError(openai.UnprocessableEntityError): # type: ignore
|
||||
class UnprocessableEntityError(openai.UnprocessableEntityError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -327,7 +327,7 @@ class UnprocessableEntityError(openai.UnprocessableEntityError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class Timeout(openai.APITimeoutError): # type: ignore
|
||||
class Timeout(openai.APITimeoutError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -371,7 +371,7 @@ class Timeout(openai.APITimeoutError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class PermissionDeniedError(openai.PermissionDeniedError): # type: ignore
|
||||
class PermissionDeniedError(openai.PermissionDeniedError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -410,7 +410,7 @@ class PermissionDeniedError(openai.PermissionDeniedError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class RateLimitError(openai.RateLimitError): # type: ignore
|
||||
class RateLimitError(openai.RateLimitError):
|
||||
"""
|
||||
Unified rate-limit error.
|
||||
|
||||
|
|
@ -501,7 +501,7 @@ class RateLimitError(openai.RateLimitError): # type: ignore
|
|||
|
||||
|
||||
# sub class of rate limit error - meant to give more granularity for error handling context window exceeded errors
|
||||
class ContextWindowExceededError(BadRequestError): # type: ignore
|
||||
class ContextWindowExceededError(BadRequestError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -516,8 +516,8 @@ class ContextWindowExceededError(BadRequestError): # type: ignore
|
|||
self.litellm_debug_info = litellm_debug_info
|
||||
super().__init__(
|
||||
message=message,
|
||||
model=self.model, # type: ignore
|
||||
llm_provider=self.llm_provider, # type: ignore
|
||||
model=self.model,
|
||||
llm_provider=self.llm_provider,
|
||||
response=response,
|
||||
litellm_debug_info=self.litellm_debug_info,
|
||||
) # Call the base class constructor with the parameters it needs
|
||||
|
|
@ -543,7 +543,7 @@ class ContextWindowExceededError(BadRequestError): # type: ignore
|
|||
|
||||
|
||||
# sub class of bad request error - meant to help us catch guardrails-related errors on proxy.
|
||||
class RejectedRequestError(BadRequestError): # type: ignore
|
||||
class RejectedRequestError(BadRequestError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -562,8 +562,8 @@ class RejectedRequestError(BadRequestError): # type: ignore
|
|||
response: Final = httpx.Response(status_code=400, request=request)
|
||||
super().__init__(
|
||||
message=self.message,
|
||||
model=self.model, # type: ignore
|
||||
llm_provider=self.llm_provider, # type: ignore
|
||||
model=self.model,
|
||||
llm_provider=self.llm_provider,
|
||||
response=response,
|
||||
litellm_debug_info=self.litellm_debug_info,
|
||||
) # Call the base class constructor with the parameters it needs
|
||||
|
|
@ -585,7 +585,7 @@ class RejectedRequestError(BadRequestError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class ContentPolicyViolationError(BadRequestError): # type: ignore
|
||||
class ContentPolicyViolationError(BadRequestError):
|
||||
# Error code: 400 - {'error': {'code': 'content_policy_violation', 'message': 'Your request was rejected as a result of our safety system. Image descriptions generated from your prompt may contain text that is not allowed by our safety system. If you believe this was done in error, your request may succeed if retried, or by adjusting your prompt.', 'param': None, 'type': 'invalid_request_error'}}
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -605,8 +605,8 @@ class ContentPolicyViolationError(BadRequestError): # type: ignore
|
|||
self.provider_specific_fields = provider_specific_fields
|
||||
super().__init__(
|
||||
message=self.message,
|
||||
model=self.model, # type: ignore
|
||||
llm_provider=self.llm_provider, # type: ignore
|
||||
model=self.model,
|
||||
llm_provider=self.llm_provider,
|
||||
response=response,
|
||||
litellm_debug_info=self.litellm_debug_info,
|
||||
body=body,
|
||||
|
|
@ -630,7 +630,7 @@ class ContentPolicyViolationError(BadRequestError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class ServiceUnavailableError(openai.APIStatusError): # type: ignore
|
||||
class ServiceUnavailableError(openai.APIStatusError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -678,7 +678,7 @@ class ServiceUnavailableError(openai.APIStatusError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class BadGatewayError(openai.APIStatusError): # type: ignore
|
||||
class BadGatewayError(openai.APIStatusError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -726,7 +726,7 @@ class BadGatewayError(openai.APIStatusError): # type: ignore
|
|||
return _message
|
||||
|
||||
|
||||
class InternalServerError(openai.InternalServerError): # type: ignore
|
||||
class InternalServerError(openai.InternalServerError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -775,7 +775,7 @@ class InternalServerError(openai.InternalServerError): # type: ignore
|
|||
|
||||
|
||||
# raise this when the API returns an invalid response object - https://github.com/openai/openai-python/blob/1be14ee34a0f8e42d3f9aa5451aa4cb161f1781f/openai/api_requestor.py#L401
|
||||
class APIError(openai.APIError): # type: ignore
|
||||
class APIError(openai.APIError):
|
||||
def __init__(
|
||||
self,
|
||||
status_code: int,
|
||||
|
|
@ -796,7 +796,7 @@ class APIError(openai.APIError): # type: ignore
|
|||
self.num_retries = num_retries
|
||||
if request is None:
|
||||
request = httpx.Request(method="POST", url="https://api.openai.com/v1")
|
||||
super().__init__(self.message, request=request, body=None) # type: ignore
|
||||
super().__init__(self.message, request=request, body=None)
|
||||
|
||||
def __str__(self):
|
||||
_message = self.message
|
||||
|
|
@ -816,7 +816,7 @@ class APIError(openai.APIError): # type: ignore
|
|||
|
||||
|
||||
# raised if an invalid request (not get, delete, put, post) is made
|
||||
class APIConnectionError(openai.APIConnectionError): # type: ignore
|
||||
class APIConnectionError(openai.APIConnectionError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -855,7 +855,7 @@ class APIConnectionError(openai.APIConnectionError): # type: ignore
|
|||
|
||||
|
||||
# raised if an invalid request (not get, delete, put, post) is made
|
||||
class APIResponseValidationError(openai.APIResponseValidationError): # type: ignore
|
||||
class APIResponseValidationError(openai.APIResponseValidationError):
|
||||
def __init__(
|
||||
self,
|
||||
message,
|
||||
|
|
@ -902,7 +902,7 @@ class JSONSchemaValidationError(APIResponseValidationError):
|
|||
super().__init__(model=model, message=message, llm_provider=llm_provider)
|
||||
|
||||
|
||||
class OpenAIError(openai.OpenAIError): # type: ignore
|
||||
class OpenAIError(openai.OpenAIError):
|
||||
def __init__(self, original_exception=None):
|
||||
super().__init__()
|
||||
self.llm_provider = "openai"
|
||||
|
|
@ -987,7 +987,7 @@ class BudgetExceededError(Exception):
|
|||
|
||||
|
||||
## DEPRECATED ##
|
||||
class InvalidRequestError(openai.BadRequestError): # type: ignore
|
||||
class InvalidRequestError(openai.BadRequestError):
|
||||
def __init__(self, message, model, llm_provider):
|
||||
self.status_code = 400
|
||||
self.message = message
|
||||
|
|
@ -1024,7 +1024,7 @@ class MockException(openai.APIError):
|
|||
self.num_retries = num_retries
|
||||
if request is None:
|
||||
request = httpx.Request(method="POST", url="https://api.openai.com/v1")
|
||||
super().__init__(self.message, request=request, body=None) # type: ignore
|
||||
super().__init__(self.message, request=request, body=None)
|
||||
|
||||
|
||||
class LiteLLMUnknownProvider(BadRequestError):
|
||||
|
|
@ -1070,7 +1070,7 @@ class BlockedPiiEntityError(Exception):
|
|||
super().__init__(self.message)
|
||||
|
||||
|
||||
class MidStreamFallbackError(ServiceUnavailableError): # type: ignore
|
||||
class MidStreamFallbackError(ServiceUnavailableError):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ from mcp.client.stdio import stdio_client
|
|||
|
||||
streamable_http_client: Any | None = None
|
||||
try:
|
||||
import mcp.client.streamable_http as streamable_http_module # type: ignore
|
||||
import mcp.client.streamable_http as streamable_http_module
|
||||
|
||||
streamable_http_client = getattr(streamable_http_module, "streamable_http_client", None)
|
||||
except ImportError:
|
||||
|
|
|
|||
|
|
@ -131,7 +131,7 @@ async def acreate_file(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
|
|
@ -176,7 +176,7 @@ def create_file(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -252,7 +252,7 @@ def create_file(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -328,7 +328,7 @@ def file_retrieve(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -419,7 +419,7 @@ def file_retrieve(
|
|||
request=httpx.Request(
|
||||
method="create_thread",
|
||||
url="https://github.com/BerriAI/litellm",
|
||||
), # type: ignore
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -465,9 +465,9 @@ async def afile_delete(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
|
||||
return cast(FileDeleted, response) # type: ignore
|
||||
return cast(FileDeleted, response)
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
|
@ -511,7 +511,7 @@ def file_delete(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
_is_async: Final = kwargs.pop("is_async", False) is True
|
||||
|
|
@ -596,7 +596,7 @@ def file_delete(
|
|||
request=httpx.Request(
|
||||
method="create_thread",
|
||||
url="https://github.com/BerriAI/litellm",
|
||||
), # type: ignore
|
||||
),
|
||||
),
|
||||
)
|
||||
return cast(FileDeleted, response)
|
||||
|
|
@ -639,7 +639,7 @@ async def afile_list(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
|
|
@ -673,7 +673,7 @@ def file_list(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -755,7 +755,7 @@ def file_list(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="file_list", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="file_list", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -803,7 +803,7 @@ async def afile_content(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
|
|
@ -857,7 +857,7 @@ def file_content(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -987,7 +987,7 @@ def file_content(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -1065,7 +1065,7 @@ def file_content_streaming(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -93,9 +93,9 @@ class FileContentStreamingResponse:
|
|||
# are released promptly on client disconnects.
|
||||
with anyio.CancelScope(shield=True):
|
||||
if hasattr(stream_to_close, "aclose"):
|
||||
await cast(AsyncIterator[bytes], stream_to_close).aclose() # type: ignore[attr-defined]
|
||||
await cast(AsyncIterator[bytes], stream_to_close).aclose()
|
||||
elif hasattr(stream_to_close, "close"):
|
||||
result: Final = cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined]
|
||||
result: Final = cast(Iterator[bytes], stream_to_close).close()
|
||||
if result is not None:
|
||||
await result
|
||||
|
||||
|
|
@ -109,7 +109,7 @@ class FileContentStreamingResponse:
|
|||
self.stream_iterator = cast(Iterator[bytes] | AsyncIterator[bytes], iter(()))
|
||||
|
||||
if hasattr(stream_to_close, "close"):
|
||||
cast(Iterator[bytes], stream_to_close).close() # type: ignore[attr-defined]
|
||||
cast(Iterator[bytes], stream_to_close).close()
|
||||
|
||||
def _build_logging_response(self) -> dict[str, str]:
|
||||
response: Final = {
|
||||
|
|
|
|||
|
|
@ -119,7 +119,7 @@ async def acreate_fine_tuning_job(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
return response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -242,9 +242,9 @@ def create_fine_tuning_job(
|
|||
)
|
||||
# Azure OpenAI
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -252,7 +252,7 @@ def create_fine_tuning_job(
|
|||
or litellm.azure_key
|
||||
or get_secret_str("AZURE_OPENAI_API_KEY")
|
||||
or get_secret_str("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
if extra_body is not None:
|
||||
|
|
@ -321,7 +321,7 @@ def create_fine_tuning_job(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -362,7 +362,7 @@ async def acancel_fine_tuning_job(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
return response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -396,7 +396,7 @@ def cancel_fine_tuning_job(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -441,7 +441,7 @@ def cancel_fine_tuning_job(
|
|||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -449,7 +449,7 @@ def cancel_fine_tuning_job(
|
|||
or litellm.azure_key
|
||||
or get_secret_str("AZURE_OPENAI_API_KEY")
|
||||
or get_secret_str("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
if extra_body is not None:
|
||||
|
|
@ -473,7 +473,7 @@ def cancel_fine_tuning_job(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -514,7 +514,7 @@ async def alist_fine_tuning_jobs(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
return response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -550,7 +550,7 @@ def list_fine_tuning_jobs(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -594,9 +594,9 @@ def list_fine_tuning_jobs(
|
|||
)
|
||||
# Azure OpenAI
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -604,7 +604,7 @@ def list_fine_tuning_jobs(
|
|||
or litellm.azure_key
|
||||
or get_secret_str("AZURE_OPENAI_API_KEY")
|
||||
or get_secret_str("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
if extra_body is not None:
|
||||
|
|
@ -629,7 +629,7 @@ def list_fine_tuning_jobs(
|
|||
response=httpx.Response(
|
||||
status_code=400,
|
||||
content="Unsupported provider",
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="create_thread", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
@ -669,7 +669,7 @@ async def aretrieve_fine_tuning_job(
|
|||
if asyncio.iscoroutine(init_response):
|
||||
response = await init_response
|
||||
else:
|
||||
response = init_response # type: ignore
|
||||
response = init_response
|
||||
return response
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -700,7 +700,7 @@ def retrieve_fine_tuning_job(
|
|||
read_timeout: Final = timeout.read or 600
|
||||
timeout = read_timeout # default 10 min timeout
|
||||
elif timeout is not None and not isinstance(timeout, httpx.Timeout):
|
||||
timeout = float(timeout) # type: ignore
|
||||
timeout = float(timeout)
|
||||
elif timeout is None:
|
||||
timeout = 600.0
|
||||
|
||||
|
|
@ -733,9 +733,9 @@ def retrieve_fine_tuning_job(
|
|||
)
|
||||
# Azure OpenAI
|
||||
elif custom_llm_provider == "azure":
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") # type: ignore
|
||||
api_base = optional_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION") # type: ignore
|
||||
api_version = optional_params.api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -743,7 +743,7 @@ def retrieve_fine_tuning_job(
|
|||
or litellm.azure_key
|
||||
or get_secret_str("AZURE_OPENAI_API_KEY")
|
||||
or get_secret_str("AZURE_API_KEY")
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
extra_body = optional_params.get("extra_body", {})
|
||||
if extra_body is not None:
|
||||
|
|
@ -770,7 +770,7 @@ def retrieve_fine_tuning_job(
|
|||
request=httpx.Request(
|
||||
method="retrieve_fine_tuning_job",
|
||||
url="https://github.com/BerriAI/litellm",
|
||||
), # type: ignore
|
||||
),
|
||||
),
|
||||
)
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -156,7 +156,7 @@ class GenerateContentHelper:
|
|||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_body={}, # Will be handled by adapter
|
||||
generate_content_provider_config=None, # type: ignore
|
||||
generate_content_provider_config=None,
|
||||
generate_content_config_dict=dict(config or {}),
|
||||
native_request_fields={},
|
||||
litellm_params=litellm_params,
|
||||
|
|
@ -350,7 +350,7 @@ def generate_content(
|
|||
# Use the adapter to convert to completion format
|
||||
return GenerateContentToCompletionHandler.generate_content_handler(
|
||||
model=model,
|
||||
contents=contents, # type: ignore
|
||||
contents=contents,
|
||||
config=setup_result.generate_content_config_dict,
|
||||
tools=tools,
|
||||
_is_async=_is_async,
|
||||
|
|
@ -444,7 +444,7 @@ async def agenerate_content_stream(
|
|||
# Use the adapter to convert to completion format
|
||||
return await GenerateContentToCompletionHandler.async_generate_content_handler(
|
||||
model=model,
|
||||
contents=contents, # type: ignore
|
||||
contents=contents,
|
||||
config=setup_result.generate_content_config_dict,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
tools=tools,
|
||||
|
|
@ -534,7 +534,7 @@ def generate_content_stream(
|
|||
# Use the adapter to convert to completion format
|
||||
return GenerateContentToCompletionHandler.generate_content_handler(
|
||||
model=model,
|
||||
contents=contents, # type: ignore
|
||||
contents=contents,
|
||||
config=setup_result.generate_content_config_dict,
|
||||
_is_async=_is_async,
|
||||
litellm_params=setup_result.litellm_params,
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from litellm.utils import exception_type, get_litellm_params
|
|||
|
||||
#################### Initialize provider clients ####################
|
||||
llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler()
|
||||
from openai.types.audio.transcription_create_params import FileTypes # type: ignore
|
||||
from openai.types.audio.transcription_create_params import FileTypes
|
||||
|
||||
# BFL handlers
|
||||
from litellm.llms.black_forest_labs.image_edit.handler import bfl_image_edit
|
||||
|
|
@ -112,7 +112,7 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse:
|
|||
elif isinstance(init_response, ImageResponse): ## CACHING SCENARIO
|
||||
response = init_response
|
||||
elif asyncio.iscoroutine(init_response):
|
||||
response = await init_response # type: ignore
|
||||
response = await init_response
|
||||
|
||||
if response is None:
|
||||
raise ValueError("Unable to get Image Response. Please pass a valid llm_provider.")
|
||||
|
|
@ -207,12 +207,12 @@ def image_generation(
|
|||
aimg_generation: Final = kwargs.get("aimg_generation", False)
|
||||
litellm_call_id: Final = kwargs.get("litellm_call_id", None)
|
||||
logger_fn: Final = kwargs.get("logger_fn", None)
|
||||
mock_response: Final[str | None] = kwargs.get("mock_response", None) # type: ignore
|
||||
mock_response: Final[str | None] = kwargs.get("mock_response", None)
|
||||
proxy_server_request: Final = kwargs.get("proxy_server_request", None)
|
||||
azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None)
|
||||
model_info: Final = kwargs.get("model_info", None)
|
||||
metadata: Final = kwargs.get("metadata", {})
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
client: Final = kwargs.get("client", None)
|
||||
extra_headers: Final = kwargs.get("extra_headers", None)
|
||||
headers: Final[dict] = kwargs.get("headers", None) or {}
|
||||
|
|
@ -223,7 +223,7 @@ def image_generation(
|
|||
dynamic_api_key: str | None = None
|
||||
if model is not None or custom_llm_provider is not None:
|
||||
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(
|
||||
model=model, # type: ignore
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
|
@ -479,7 +479,7 @@ def image_generation(
|
|||
elif custom_llm_provider == "bedrock":
|
||||
if model is None:
|
||||
raise Exception("Model needs to be set for bedrock")
|
||||
model_response = bedrock_image_generation.image_generation( # type: ignore
|
||||
model_response = bedrock_image_generation.image_generation(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
timeout=timeout,
|
||||
|
|
@ -508,7 +508,7 @@ def image_generation(
|
|||
async_custom_client = client
|
||||
|
||||
## CALL FUNCTION
|
||||
model_response = custom_handler.aimage_generation( # type: ignore
|
||||
model_response = custom_handler.aimage_generation(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
api_key=api_key,
|
||||
|
|
@ -584,7 +584,7 @@ async def aimage_variation(*args, **kwargs) -> ImageResponse:
|
|||
init_response = ImageResponse(**init_response)
|
||||
response = init_response
|
||||
elif asyncio.iscoroutine(init_response):
|
||||
response = await init_response # type: ignore
|
||||
response = await init_response
|
||||
else:
|
||||
# Call the synchronous function using run_in_executor
|
||||
response = await loop.run_in_executor(None, func_with_context)
|
||||
|
|
@ -745,7 +745,7 @@ def image_edit(
|
|||
non_default_params: Final = {
|
||||
k: v for k, v in kwargs.items() if k not in default_params
|
||||
} # model-specific params - pass them straight to the model/provider
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
model_info: Final = kwargs.get("model_info", None)
|
||||
metadata: Final = kwargs.get("metadata", {})
|
||||
|
|
@ -860,7 +860,7 @@ def image_edit(
|
|||
if model is None:
|
||||
raise Exception("Model needs to be set for bedrock")
|
||||
image_edit_request_params.update(non_default_params)
|
||||
return bedrock_image_edit.image_edit( # type: ignore
|
||||
return bedrock_image_edit.image_edit(
|
||||
model=model,
|
||||
image=images,
|
||||
prompt=prompt,
|
||||
|
|
|
|||
|
|
@ -709,7 +709,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
"""Format an alert message for slack"""
|
||||
headers: Final = {f"{key} Name": key_val, "Provider": provider}
|
||||
if api_base is not None:
|
||||
headers["API Base"] = api_base # type: ignore
|
||||
headers["API Base"] = api_base
|
||||
|
||||
headers_str = "\n"
|
||||
for k, v in headers.items():
|
||||
|
|
@ -767,14 +767,11 @@ class SlackAlerting(CustomBatchLogger):
|
|||
|
||||
# Convert deployment_ids back to set if it was stored as a list
|
||||
if outage_value is not None:
|
||||
outage_value = self._restore_outage_value_from_cache(outage_value) # type: ignore
|
||||
outage_value = self._restore_outage_value_from_cache(outage_value)
|
||||
|
||||
if (
|
||||
getattr(exception, "status_code", None) is None
|
||||
or (
|
||||
exception.status_code != 408 # type: ignore
|
||||
and exception.status_code < 500 # type: ignore
|
||||
)
|
||||
or (exception.status_code != 408 and exception.status_code < 500)
|
||||
or self.llm_router is None
|
||||
):
|
||||
return
|
||||
|
|
@ -784,7 +781,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
_deployment_set.add(deployment_id)
|
||||
outage_value = ProviderRegionOutageModel(
|
||||
provider_region_id=cache_key,
|
||||
alerts=[exception.status_code], # type: ignore
|
||||
alerts=[exception.status_code],
|
||||
minor_alert_sent=False,
|
||||
major_alert_sent=False,
|
||||
last_updated_at=time.time(),
|
||||
|
|
@ -802,7 +799,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
return
|
||||
|
||||
if len(outage_value["alerts"]) < self.alerting_args.max_outage_alert_list_size:
|
||||
outage_value["alerts"].append(exception.status_code) # type: ignore
|
||||
outage_value["alerts"].append(exception.status_code)
|
||||
else: # prevent memory leaks
|
||||
pass
|
||||
_deployment_set = outage_value["deployment_ids"]
|
||||
|
|
@ -884,13 +881,10 @@ class SlackAlerting(CustomBatchLogger):
|
|||
max_alerts_size = 10
|
||||
"""
|
||||
try:
|
||||
outage_value: OutageModel | None = await self.internal_usage_cache.async_get_cache(key=deployment_id) # type: ignore
|
||||
outage_value: OutageModel | None = await self.internal_usage_cache.async_get_cache(key=deployment_id)
|
||||
if (
|
||||
getattr(exception, "status_code", None) is None
|
||||
or (
|
||||
exception.status_code != 408 # type: ignore
|
||||
and exception.status_code < 500 # type: ignore
|
||||
)
|
||||
or (exception.status_code != 408 and exception.status_code < 500)
|
||||
or self.llm_router is None
|
||||
):
|
||||
return
|
||||
|
|
@ -912,7 +906,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
if outage_value is None:
|
||||
outage_value = OutageModel(
|
||||
model_id=deployment_id,
|
||||
alerts=[exception.status_code], # type: ignore
|
||||
alerts=[exception.status_code],
|
||||
minor_alert_sent=False,
|
||||
major_alert_sent=False,
|
||||
last_updated_at=time.time(),
|
||||
|
|
@ -927,7 +921,7 @@ class SlackAlerting(CustomBatchLogger):
|
|||
return
|
||||
|
||||
if len(outage_value["alerts"]) < self.alerting_args.max_outage_alert_list_size:
|
||||
outage_value["alerts"].append(exception.status_code) # type: ignore
|
||||
outage_value["alerts"].append(exception.status_code)
|
||||
else: # prevent memory leaks
|
||||
pass
|
||||
|
||||
|
|
@ -1483,10 +1477,10 @@ Model Info:
|
|||
|
||||
if isinstance(response_obj, litellm.ModelResponse) and (
|
||||
hasattr(response_obj, "usage")
|
||||
and response_obj.usage is not None # type: ignore
|
||||
and hasattr(response_obj.usage, "completion_tokens") # type: ignore
|
||||
and response_obj.usage is not None
|
||||
and hasattr(response_obj.usage, "completion_tokens")
|
||||
):
|
||||
completion_tokens: Final = response_obj.usage.completion_tokens # type: ignore
|
||||
completion_tokens: Final = response_obj.usage.completion_tokens
|
||||
if completion_tokens is not None and completion_tokens > 0:
|
||||
final_value = float(response_s.total_seconds() / completion_tokens)
|
||||
if isinstance(final_value, timedelta):
|
||||
|
|
|
|||
|
|
@ -225,11 +225,11 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
# 1. if string, insert cache control in the message
|
||||
if isinstance(message_content, str):
|
||||
message["cache_control"] = control # type: ignore
|
||||
message["cache_control"] = control
|
||||
# 2. list of objects - only apply to last item per Anthropic spec
|
||||
elif isinstance(message_content, list):
|
||||
if len(message_content) > 0 and isinstance(message_content[-1], dict):
|
||||
message_content[-1]["cache_control"] = control # type: ignore
|
||||
message_content[-1]["cache_control"] = control
|
||||
return message
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import types
|
|||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel # type: ignore
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -56,8 +56,8 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
argilla_base_url=argilla_base_url,
|
||||
)
|
||||
self.sampling_rate: float = (
|
||||
float(os.getenv("ARGILLA_SAMPLING_RATE")) # type: ignore
|
||||
if os.getenv("ARGILLA_SAMPLING_RATE") is not None and os.getenv("ARGILLA_SAMPLING_RATE").strip().isdigit() # type: ignore
|
||||
float(os.getenv("ARGILLA_SAMPLING_RATE"))
|
||||
if os.getenv("ARGILLA_SAMPLING_RATE") is not None and os.getenv("ARGILLA_SAMPLING_RATE").strip().isdigit()
|
||||
else 1.0
|
||||
)
|
||||
|
||||
|
|
@ -196,9 +196,9 @@ class ArgillaLogger(CustomBatchLogger):
|
|||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
sampling_rate: Final = (
|
||||
float(os.getenv("LANGSMITH_SAMPLING_RATE")) # type: ignore
|
||||
float(os.getenv("LANGSMITH_SAMPLING_RATE"))
|
||||
if os.getenv("LANGSMITH_SAMPLING_RATE") is not None
|
||||
and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit() # type: ignore
|
||||
and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit()
|
||||
else 1.0
|
||||
)
|
||||
random_sample: Final = random.random()
|
||||
|
|
|
|||
|
|
@ -40,14 +40,14 @@ else:
|
|||
)
|
||||
except ImportError:
|
||||
LITELLM_TRACER_NAME = "litellm"
|
||||
OpenTelemetry = None # type: ignore
|
||||
OpenTelemetry = None
|
||||
|
||||
|
||||
ARIZE_HOSTED_PHOENIX_ENDPOINT: Final = "https://otlp.arize.com/v1/traces"
|
||||
_MAX_PROJECT_PROVIDERS: Final = 64
|
||||
|
||||
|
||||
class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
||||
class ArizePhoenixLogger(OpenTelemetry):
|
||||
"""
|
||||
Arize Phoenix logger that sends traces to a Phoenix endpoint.
|
||||
|
||||
|
|
@ -139,7 +139,7 @@ class ArizePhoenixLogger(OpenTelemetry): # type: ignore
|
|||
project_attributes["deployment.environment"] = deployment_environment
|
||||
|
||||
env_resource: Final = OTELResourceDetector().detect()
|
||||
project_resource: Final = Resource.create(project_attributes) # type: ignore[arg-type]
|
||||
project_resource: Final = Resource.create(project_attributes)
|
||||
return env_resource.merge(project_resource)
|
||||
|
||||
def _build_tracer_provider_for_project(self, project_name: str) -> TracerProvider:
|
||||
|
|
|
|||
|
|
@ -174,9 +174,7 @@ class ArizePhoenixTemplateManager:
|
|||
# Combine rendered content
|
||||
final_content = " ".join(rendered_content_parts)
|
||||
|
||||
rendered_messages.append(
|
||||
{"role": role, "content": final_content} # type: ignore
|
||||
)
|
||||
rendered_messages.append({"role": role, "content": final_content})
|
||||
|
||||
return rendered_messages
|
||||
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ def set_global_bitbucket_config(config: dict) -> None:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
litellm.global_bitbucket_config = config # type: ignore
|
||||
litellm.global_bitbucket_config = config
|
||||
|
||||
|
||||
def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement":
|
||||
|
|
|
|||
|
|
@ -292,9 +292,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
final_messages: list[AllMessageValues] = parsed_messages
|
||||
else:
|
||||
# If no messages were parsed, prepend the prompt to existing messages
|
||||
final_messages = [
|
||||
{"role": "user", "content": rendered_prompt} # type: ignore
|
||||
] + messages
|
||||
final_messages = [{"role": "user", "content": rendered_prompt}] + messages
|
||||
|
||||
# Update litellm_params with prompt metadata
|
||||
if litellm_params is None:
|
||||
|
|
@ -345,7 +343,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
{
|
||||
"role": current_role,
|
||||
"content": "\n".join(current_content).strip(),
|
||||
} # type: ignore
|
||||
}
|
||||
)
|
||||
current_role = "system"
|
||||
current_content = [line[7:].strip()] # Remove "System:" prefix
|
||||
|
|
@ -355,7 +353,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
{
|
||||
"role": current_role,
|
||||
"content": "\n".join(current_content).strip(),
|
||||
} # type: ignore
|
||||
}
|
||||
)
|
||||
current_role = "user"
|
||||
current_content = [line[5:].strip()] # Remove "User:" prefix
|
||||
|
|
@ -365,7 +363,7 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
{
|
||||
"role": current_role,
|
||||
"content": "\n".join(current_content).strip(),
|
||||
} # type: ignore
|
||||
}
|
||||
)
|
||||
current_role = "assistant"
|
||||
current_content = [line[10:].strip()] # Remove "Assistant:" prefix
|
||||
|
|
@ -379,9 +377,9 @@ class BitBucketPromptManager(CustomPromptManagement):
|
|||
|
||||
# If no role indicators found, treat as a single user message
|
||||
if not messages and prompt_content.strip():
|
||||
messages = [{"role": "user", "content": prompt_content.strip()}] # type: ignore
|
||||
messages = [{"role": "user", "content": prompt_content.strip()}]
|
||||
|
||||
return messages # type: ignore
|
||||
return messages
|
||||
|
||||
def post_call_hook(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -28,9 +28,9 @@ def get_utc_datetime():
|
|||
import datetime as dt
|
||||
|
||||
if hasattr(dt, "UTC"):
|
||||
return datetime.now(dt.UTC) # type: ignore
|
||||
return datetime.now(dt.UTC)
|
||||
else:
|
||||
return datetime.utcnow() # type: ignore
|
||||
return datetime.utcnow()
|
||||
|
||||
|
||||
class BraintrustLogger(CustomLogger):
|
||||
|
|
@ -43,7 +43,7 @@ class BraintrustLogger(CustomLogger):
|
|||
self.validate_environment(api_key=api_key)
|
||||
self.api_base = api_base or os.getenv("BRAINTRUST_API_BASE") or API_BASE
|
||||
self.default_project_id = None
|
||||
self.api_key: str = api_key or os.getenv("BRAINTRUST_API_KEY") # type: ignore
|
||||
self.api_key: str = api_key or os.getenv("BRAINTRUST_API_KEY")
|
||||
self.headers = {
|
||||
"Authorization": "Bearer " + self.api_key,
|
||||
"Content-Type": "application/json",
|
||||
|
|
|
|||
|
|
@ -150,7 +150,7 @@ def create_mock_braintrust_client():
|
|||
|
||||
if _original_http_handler_post is None:
|
||||
_original_http_handler_post = HTTPHandler.post
|
||||
HTTPHandler.post = _mock_http_handler_post # type: ignore
|
||||
HTTPHandler.post = _mock_http_handler_post
|
||||
verbose_logger.debug("[BRAINTRUST MOCK] Patched HTTPHandler.post")
|
||||
|
||||
# CRITICAL: Call the factory's initialization function to patch AsyncHTTPHandler.post
|
||||
|
|
|
|||
|
|
@ -133,7 +133,7 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
|
||||
self._prune_expired_cache()
|
||||
|
||||
compressed: Final = compress( # type: ignore
|
||||
compressed: Final = compress(
|
||||
messages=messages,
|
||||
model=model,
|
||||
call_type=CallTypes.anthropic_messages,
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ from litellm.types.utils import (
|
|||
try:
|
||||
from fastapi.exceptions import HTTPException
|
||||
except ImportError:
|
||||
HTTPException = None # type: ignore
|
||||
HTTPException = None
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -410,7 +410,7 @@ class CustomGuardrail(CustomLogger):
|
|||
if self.should_route_on_sensitive_data():
|
||||
try:
|
||||
self.raise_sensitive_data_route_exception(
|
||||
route_to_model=self.sensitive_data_route_to_model, # type: ignore
|
||||
route_to_model=self.sensitive_data_route_to_model,
|
||||
request_data=request_data,
|
||||
detection_info=detection_info,
|
||||
)
|
||||
|
|
@ -892,9 +892,9 @@ class CustomGuardrail(CustomLogger):
|
|||
if event_type is not None:
|
||||
guardrail_mode = event_type
|
||||
elif isinstance(self.event_hook, Mode):
|
||||
guardrail_mode = GuardrailMode(**dict(self.event_hook.model_dump())) # type: ignore[typeddict-item]
|
||||
guardrail_mode = GuardrailMode(**dict(self.event_hook.model_dump()))
|
||||
else:
|
||||
guardrail_mode = self.event_hook # type: ignore[assignment]
|
||||
guardrail_mode = self.event_hook
|
||||
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
filter_exceptions_from_params,
|
||||
|
|
|
|||
|
|
@ -783,13 +783,11 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
- Converting to string and then truncating the logged content catches this
|
||||
2. We want to avoid modifying the original `messages`, `response`, and `error_str` in the logging payload since these are in kwargs and could be returned to the user
|
||||
"""
|
||||
field_value: Final = standard_logging_object.get(field_name) # type: ignore
|
||||
field_value: Final = standard_logging_object.get(field_name)
|
||||
if field_value:
|
||||
str_value: Final = str(field_value)
|
||||
if len(str_value) > max_length:
|
||||
standard_logging_object[field_name] = self._truncate_text( # type: ignore
|
||||
text=str_value, max_length=max_length
|
||||
)
|
||||
standard_logging_object[field_name] = self._truncate_text(text=str_value, max_length=max_length)
|
||||
|
||||
def _truncate_text(self, text: str, max_length: int) -> str:
|
||||
"""Truncate text if it exceeds max_length"""
|
||||
|
|
@ -911,7 +909,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
for callback_obj in all_callbacks:
|
||||
if hasattr(callback_obj, "increment_callback_logging_failure"):
|
||||
verbose_logger.debug("Incrementing callback failure metric for %s", callback_name)
|
||||
callback_obj.increment_callback_logging_failure(callback_name=callback_name) # type: ignore
|
||||
callback_obj.increment_callback_logging_failure(callback_name=callback_name)
|
||||
return
|
||||
|
||||
verbose_logger.debug(
|
||||
|
|
|
|||
|
|
@ -500,7 +500,7 @@ class DataDogLogger(
|
|||
|
||||
response: Final = self.sync_client.post(
|
||||
url=self.intake_url,
|
||||
json=dd_payload, # type: ignore
|
||||
json=dd_payload,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
|
@ -616,7 +616,7 @@ class DataDogLogger(
|
|||
|
||||
response: Final = await self.async_client.post(
|
||||
url=self.intake_url,
|
||||
data=compressed_data, # type: ignore
|
||||
data=compressed_data,
|
||||
headers=headers,
|
||||
)
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -91,9 +91,9 @@ class DatadogMetricsLogger(CustomBatchLogger):
|
|||
metadata: Final = log.get("metadata", {}) or {}
|
||||
team_tag: Final = (
|
||||
metadata.get("user_api_key_team_alias")
|
||||
or metadata.get("team_alias") # type: ignore
|
||||
or metadata.get("team_alias")
|
||||
or metadata.get("user_api_key_team_id")
|
||||
or metadata.get("team_id") # type: ignore
|
||||
or metadata.get("team_id")
|
||||
)
|
||||
|
||||
if team_tag:
|
||||
|
|
@ -193,7 +193,7 @@ class DatadogMetricsLogger(CustomBatchLogger):
|
|||
# Extract status code from error information
|
||||
status_code = "500" # default
|
||||
error_information: Final = standard_logging_object.get("error_information", {}) or {}
|
||||
error_code: Final = error_information.get("error_code") # type: ignore
|
||||
error_code: Final = error_information.get("error_code")
|
||||
if error_code is not None:
|
||||
status_code = str(error_code)
|
||||
|
||||
|
|
@ -237,7 +237,7 @@ class DatadogMetricsLogger(CustomBatchLogger):
|
|||
response: Final = await self.async_client.post(
|
||||
self.upload_url,
|
||||
content=compressed_data,
|
||||
headers=headers, # type: ignore
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
response.raise_for_status()
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ def set_global_prompt_directory(directory: str) -> None:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
litellm.global_prompt_directory = directory # type: ignore
|
||||
litellm.global_prompt_directory = directory
|
||||
|
||||
|
||||
def _get_prompt_data_from_dotprompt_content(dotprompt_content: str) -> dict:
|
||||
|
|
|
|||
|
|
@ -311,7 +311,7 @@ class DotpromptManager(CustomPromptManagement):
|
|||
def _create_message(self, role: str, content: str) -> AllMessageValues:
|
||||
"""Create a message with the specified role and content."""
|
||||
return {
|
||||
"role": role, # type: ignore
|
||||
"role": role,
|
||||
"content": content,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -253,7 +253,7 @@ class PromptManager:
|
|||
"dict": dict,
|
||||
}
|
||||
|
||||
return type_mapping.get(schema_type.lower(), str) # type: ignore
|
||||
return type_mapping.get(schema_type.lower(), str)
|
||||
|
||||
def get_prompt(self, prompt_id: str, version: int | None = None) -> PromptTemplate | None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -42,9 +42,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils):
|
|||
batch_size=self.batch_size,
|
||||
flush_interval=self.flush_interval,
|
||||
)
|
||||
self.log_queue: asyncio.Queue[GCSLogQueueItem] = asyncio.Queue( # type: ignore[assignment]
|
||||
maxsize=LITELLM_ASYNCIO_QUEUE_MAXSIZE
|
||||
)
|
||||
self.log_queue: asyncio.Queue[GCSLogQueueItem] = asyncio.Queue(maxsize=LITELLM_ASYNCIO_QUEUE_MAXSIZE)
|
||||
asyncio.create_task(self.periodic_flush())
|
||||
AdditionalLoggingUtils.__init__(self)
|
||||
|
||||
|
|
|
|||
|
|
@ -167,12 +167,12 @@ def create_mock_gcs_client():
|
|||
|
||||
if _original_async_handler_get is None:
|
||||
_original_async_handler_get = AsyncHTTPHandler.get
|
||||
AsyncHTTPHandler.get = _mock_async_handler_get # type: ignore
|
||||
AsyncHTTPHandler.get = _mock_async_handler_get
|
||||
verbose_logger.debug("[GCS MOCK] Patched AsyncHTTPHandler.get")
|
||||
|
||||
if _original_async_handler_delete is None:
|
||||
_original_async_handler_delete = AsyncHTTPHandler.delete
|
||||
AsyncHTTPHandler.delete = _mock_async_handler_delete # type: ignore
|
||||
AsyncHTTPHandler.delete = _mock_async_handler_delete
|
||||
verbose_logger.debug("[GCS MOCK] Patched AsyncHTTPHandler.delete")
|
||||
|
||||
verbose_logger.debug(f"[GCS MOCK] Mock latency set to {_MOCK_LATENCY_SECONDS * 1000:.0f}ms")
|
||||
|
|
@ -227,9 +227,9 @@ def mock_vertex_auth_methods():
|
|||
return ("mock-gcs-token", "https://storage.googleapis.com")
|
||||
|
||||
# Patch the methods
|
||||
VertexBase._ensure_access_token_async = _mock_ensure_access_token_async # type: ignore
|
||||
VertexBase._ensure_access_token = _mock_ensure_access_token # type: ignore
|
||||
VertexBase._get_token_and_url = _mock_get_token_and_url # type: ignore
|
||||
VertexBase._ensure_access_token_async = _mock_ensure_access_token_async
|
||||
VertexBase._ensure_access_token = _mock_ensure_access_token
|
||||
VertexBase._get_token_and_url = _mock_get_token_and_url
|
||||
|
||||
verbose_logger.debug("[GCS MOCK] Patched Vertex AI auth methods")
|
||||
|
||||
|
|
|
|||
|
|
@ -382,7 +382,7 @@ class GenericAPILogger(CustomBatchLogger):
|
|||
verbose_logger.debug(
|
||||
"Generic API Logger - sent log %s, status: %s",
|
||||
idx,
|
||||
result.status_code, # type: ignore
|
||||
result.status_code,
|
||||
)
|
||||
else:
|
||||
# Format the payload based on log_format
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ def set_global_generic_prompt_config(config: dict) -> None:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
litellm.global_generic_prompt_config = config # type: ignore
|
||||
litellm.global_generic_prompt_config = config
|
||||
|
||||
|
||||
def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement":
|
||||
|
|
|
|||
|
|
@ -366,14 +366,14 @@ class GenericPromptManager(CustomPromptManagement):
|
|||
# Create a copy of the prompt template with variables applied
|
||||
updated_messages: Final[list[AllMessageValues]] = []
|
||||
for message in prompt_client["prompt_template"]:
|
||||
updated_message = dict(message) # type: ignore
|
||||
updated_message = dict(message)
|
||||
if "content" in updated_message and isinstance(updated_message["content"], str):
|
||||
content = updated_message["content"]
|
||||
for key, value in variables.items():
|
||||
content = content.replace(f"{{{key}}}", str(value))
|
||||
content = content.replace(f"{{{{{key}}}}}", str(value)) # Also support {{key}}
|
||||
updated_message["content"] = content
|
||||
updated_messages.append(updated_message) # type: ignore
|
||||
updated_messages.append(updated_message)
|
||||
|
||||
return PromptManagementClient(
|
||||
prompt_id=prompt_client["prompt_id"],
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ def set_global_gitlab_config(config: dict) -> None:
|
|||
"""
|
||||
import litellm
|
||||
|
||||
litellm.global_gitlab_config = config # type: ignore
|
||||
litellm.global_gitlab_config = config
|
||||
|
||||
|
||||
def prompt_initializer(litellm_params: "PromptLiteLLMParams", prompt_spec: "PromptSpec") -> "CustomPromptManagement":
|
||||
|
|
|
|||
|
|
@ -257,7 +257,7 @@ class GitLabTemplateManager:
|
|||
and str(f.get("path", "")).endswith(".prompt")
|
||||
and "path" in f
|
||||
):
|
||||
files.append(f["path"]) # type: ignore
|
||||
files.append(f["path"])
|
||||
|
||||
return [self._repo_path_to_id(p) for p in files]
|
||||
|
||||
|
|
@ -357,7 +357,7 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
if parsed_messages:
|
||||
final_messages: list[AllMessageValues] = parsed_messages
|
||||
else:
|
||||
final_messages = [{"role": "user", "content": rendered_prompt}] + messages # type: ignore
|
||||
final_messages = [{"role": "user", "content": rendered_prompt}] + messages
|
||||
|
||||
if litellm_params is None:
|
||||
litellm_params = {}
|
||||
|
|
@ -400,7 +400,7 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
"role": current_role,
|
||||
"content": "\n".join(current_content).strip(),
|
||||
}
|
||||
) # type: ignore
|
||||
)
|
||||
current_role = "system"
|
||||
current_content = [line[7:].strip()]
|
||||
elif low.startswith("user:"):
|
||||
|
|
@ -410,7 +410,7 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
"role": current_role,
|
||||
"content": "\n".join(current_content).strip(),
|
||||
}
|
||||
) # type: ignore
|
||||
)
|
||||
current_role = "user"
|
||||
current_content = [line[5:].strip()]
|
||||
elif low.startswith("assistant:"):
|
||||
|
|
@ -420,16 +420,16 @@ class GitLabPromptManager(CustomPromptManagement):
|
|||
"role": current_role,
|
||||
"content": "\n".join(current_content).strip(),
|
||||
}
|
||||
) # type: ignore
|
||||
)
|
||||
current_role = "assistant"
|
||||
current_content = [line[10:].strip()]
|
||||
else:
|
||||
current_content.append(line)
|
||||
|
||||
if current_role and current_content:
|
||||
messages.append({"role": current_role, "content": "\n".join(current_content).strip()}) # type: ignore
|
||||
messages.append({"role": current_role, "content": "\n".join(current_content).strip()})
|
||||
if not messages and prompt_content.strip():
|
||||
messages = [{"role": "user", "content": prompt_content.strip()}] # type: ignore
|
||||
messages = [{"role": "user", "content": prompt_content.strip()}]
|
||||
return messages
|
||||
|
||||
def post_call_hook(
|
||||
|
|
|
|||
|
|
@ -23,9 +23,9 @@ def get_utc_datetime():
|
|||
from datetime import datetime
|
||||
|
||||
if hasattr(dt, "UTC"):
|
||||
return datetime.now(dt.UTC) # type: ignore
|
||||
return datetime.now(dt.UTC)
|
||||
else:
|
||||
return datetime.utcnow() # type: ignore
|
||||
return datetime.utcnow()
|
||||
|
||||
|
||||
class LagoLogger(CustomLogger):
|
||||
|
|
@ -92,7 +92,7 @@ class LagoLogger(CustomLogger):
|
|||
"user_id",
|
||||
"team_id",
|
||||
]:
|
||||
charge_by = os.environ["LAGO_API_CHARGE_BY"] # type: ignore
|
||||
charge_by = os.environ["LAGO_API_CHARGE_BY"]
|
||||
else:
|
||||
raise Exception("invalid LAGO_API_CHARGE_BY set")
|
||||
|
||||
|
|
|
|||
|
|
@ -433,14 +433,14 @@ class LangFuseLogger:
|
|||
input,
|
||||
response_obj,
|
||||
):
|
||||
from langfuse.model import CreateGeneration, CreateTrace # type: ignore
|
||||
from langfuse.model import CreateGeneration, CreateTrace
|
||||
|
||||
verbose_logger.warning(
|
||||
"Please upgrade langfuse to v2.0.0 or higher: https://github.com/langfuse/langfuse-python/releases/tag/v2.0.1"
|
||||
)
|
||||
|
||||
trace: Final = self.Langfuse.trace( # type: ignore
|
||||
CreateTrace( # type: ignore
|
||||
trace: Final = self.Langfuse.trace(
|
||||
CreateTrace(
|
||||
name=metadata.get("generation_name", "litellm-completion"),
|
||||
input=input,
|
||||
output=output,
|
||||
|
|
@ -959,8 +959,8 @@ class LangFuseLogger:
|
|||
"guardrail_mode": guardrail_entry.get("guardrail_mode", None),
|
||||
"guardrail_masked_entity_count": guardrail_entry.get("masked_entity_count", None),
|
||||
},
|
||||
start_time=guardrail_entry.get("start_time", None), # type: ignore
|
||||
end_time=guardrail_entry.get("end_time", None), # type: ignore
|
||||
start_time=guardrail_entry.get("start_time", None),
|
||||
end_time=guardrail_entry.get("end_time", None),
|
||||
)
|
||||
|
||||
verbose_logger.debug("Logged guardrail information as span: %s", span)
|
||||
|
|
@ -1006,7 +1006,7 @@ def _add_prompt_to_generation_params(
|
|||
if "labels" in prompt_text_params and "tags" in prompt_text_params:
|
||||
_data["labels"] = user_prompt.get("labels", []) or []
|
||||
_data["tags"] = user_prompt.get("tags", []) or []
|
||||
_prompt_obj = Prompt_Text(**_data) # type: ignore
|
||||
_prompt_obj = Prompt_Text(**_data)
|
||||
generation_params["prompt"] = TextPromptClient(prompt=_prompt_obj)
|
||||
|
||||
elif isinstance(user_prompt["prompt"], list):
|
||||
|
|
@ -1021,7 +1021,7 @@ def _add_prompt_to_generation_params(
|
|||
_data["labels"] = user_prompt.get("labels", []) or []
|
||||
_data["tags"] = user_prompt.get("tags", []) or []
|
||||
|
||||
_prompt_obj = Prompt_Chat(**_data) # type: ignore
|
||||
_prompt_obj = Prompt_Chat(**_data)
|
||||
|
||||
generation_params["prompt"] = ChatPromptClient(prompt=_prompt_obj)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -74,7 +74,7 @@ class LangfuseOtelLogger(OpenTelemetry):
|
|||
LangFuseLogger as _LFLogger,
|
||||
)
|
||||
|
||||
metadata = _LFLogger.add_metadata_from_header(litellm_params, metadata) # type: ignore
|
||||
metadata = _LFLogger.add_metadata_from_header(litellm_params, metadata)
|
||||
except Exception:
|
||||
# Fallback silently if import fails; header enrichment just won't happen
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from datetime import datetime, timezone
|
|||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel # type: ignore
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -63,9 +63,9 @@ class LangsmithLogger(CustomBatchLogger):
|
|||
langsmith_tenant_id=langsmith_tenant_id,
|
||||
)
|
||||
self.sampling_rate: float = (
|
||||
langsmith_sampling_rate or float(os.getenv("LANGSMITH_SAMPLING_RATE")) # type: ignore
|
||||
langsmith_sampling_rate or float(os.getenv("LANGSMITH_SAMPLING_RATE"))
|
||||
if os.getenv("LANGSMITH_SAMPLING_RATE") is not None
|
||||
and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit() # type: ignore
|
||||
and os.getenv("LANGSMITH_SAMPLING_RATE").strip().isdigit()
|
||||
else 1.0
|
||||
)
|
||||
self.langsmith_default_run_name = os.getenv("LANGSMITH_DEFAULT_RUN_NAME", "LLMRun")
|
||||
|
|
|
|||
|
|
@ -80,9 +80,9 @@ class LunaryLogger:
|
|||
try:
|
||||
import lunary
|
||||
|
||||
version: Final = importlib.metadata.version("lunary") # type: ignore
|
||||
version: Final = importlib.metadata.version("lunary")
|
||||
# if version < 0.1.43 then raise ImportError
|
||||
if packaging.version.Version(version) < packaging.version.Version("0.1.43"): # type: ignore
|
||||
if packaging.version.Version(version) < packaging.version.Version("0.1.43"):
|
||||
print( # noqa: T201
|
||||
"Lunary version outdated. Required: >= 0.1.43. Upgrade via 'pip install lunary --upgrade'"
|
||||
)
|
||||
|
|
@ -151,7 +151,7 @@ class LunaryLogger:
|
|||
else:
|
||||
error_obj = None
|
||||
|
||||
self.lunary_client.track_event( # type: ignore
|
||||
self.lunary_client.track_event(
|
||||
type,
|
||||
"start",
|
||||
run_id,
|
||||
|
|
@ -167,7 +167,7 @@ class LunaryLogger:
|
|||
params=extra,
|
||||
)
|
||||
|
||||
self.lunary_client.track_event( # type: ignore
|
||||
self.lunary_client.track_event(
|
||||
type,
|
||||
event,
|
||||
run_id,
|
||||
|
|
|
|||
|
|
@ -275,7 +275,7 @@ class MavvrikFocusLogger(FocusLogger):
|
|||
|
||||
logger: Final = loggers[0]
|
||||
trigger_kwargs: Final = logger._build_scheduler_trigger()
|
||||
scheduler.add_job( # type: ignore[attr-defined]
|
||||
scheduler.add_job(
|
||||
logger.initialize_mavvrik_focus_export_job,
|
||||
id=MAVVRIK_FOCUS_EXPORT_JOB_NAME,
|
||||
replace_existing=True,
|
||||
|
|
|
|||
|
|
@ -54,8 +54,8 @@ class MlflowLogger(CustomLogger):
|
|||
def _extract_and_set_chat_attributes(self, span, kwargs, response_obj):
|
||||
try:
|
||||
from mlflow.tracing.utils import (
|
||||
set_span_chat_messages, # type: ignore
|
||||
set_span_chat_tools, # type: ignore
|
||||
set_span_chat_messages,
|
||||
set_span_chat_tools,
|
||||
)
|
||||
except ImportError:
|
||||
return
|
||||
|
|
@ -88,7 +88,7 @@ class MlflowLogger(CustomLogger):
|
|||
|
||||
# Record exception info as event
|
||||
if exception := kwargs.get("exception"):
|
||||
span.add_event(SpanEvent.from_exception(exception)) # type: ignore
|
||||
span.add_event(SpanEvent.from_exception(exception))
|
||||
|
||||
self._extract_and_set_chat_attributes(span, kwargs, response_obj)
|
||||
self._end_span_or_trace(
|
||||
|
|
@ -244,7 +244,7 @@ class MlflowLogger(CustomLogger):
|
|||
inputs: Final = self._construct_input(kwargs)
|
||||
attributes: Final = self._extract_attributes(kwargs)
|
||||
|
||||
if active_span := mlflow.get_current_active_span(): # type: ignore
|
||||
if active_span := mlflow.get_current_active_span():
|
||||
return self._client.start_span(
|
||||
name=span_name,
|
||||
trace_id=active_span.request_id,
|
||||
|
|
|
|||
|
|
@ -242,19 +242,19 @@ def create_mock_client_factory(config: MockClientConfig):
|
|||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
_original_async_handler_post = AsyncHTTPHandler.post
|
||||
AsyncHTTPHandler.post = _mock_async_handler_post # type: ignore
|
||||
AsyncHTTPHandler.post = _mock_async_handler_post
|
||||
verbose_logger.debug("[%s MOCK] Patched AsyncHTTPHandler.post", config.name)
|
||||
|
||||
if config.patch_sync_client and _original_sync_client_post is None:
|
||||
_original_sync_client_post = httpx.Client.post
|
||||
httpx.Client.post = _mock_sync_client_post # type: ignore
|
||||
httpx.Client.post = _mock_sync_client_post
|
||||
verbose_logger.debug("[%s MOCK] Patched httpx.Client.post", config.name)
|
||||
|
||||
if config.patch_http_handler and _original_http_handler_post is None:
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
_original_http_handler_post = HTTPHandler.post
|
||||
HTTPHandler.post = _mock_http_handler_post # type: ignore
|
||||
HTTPHandler.post = _mock_http_handler_post
|
||||
verbose_logger.debug("[%s MOCK] Patched HTTPHandler.post", config.name)
|
||||
|
||||
verbose_logger.debug(f"[{config.name} MOCK] Mock latency set to {_MOCK_LATENCY_SECONDS * 1000:.0f}ms")
|
||||
|
|
|
|||
|
|
@ -60,7 +60,7 @@ from litellm.types.utils import Message, ModelResponse, StandardLoggingPayload
|
|||
try:
|
||||
import newrelic.agent as _newrelic_agent
|
||||
except ImportError:
|
||||
_newrelic_agent = None # type: ignore
|
||||
_newrelic_agent = None
|
||||
|
||||
|
||||
class NewRelicLogger(CustomLogger):
|
||||
|
|
|
|||
|
|
@ -21,9 +21,9 @@ def get_utc_datetime():
|
|||
from datetime import datetime
|
||||
|
||||
if hasattr(dt, "UTC"):
|
||||
return datetime.now(dt.UTC) # type: ignore
|
||||
return datetime.now(dt.UTC)
|
||||
else:
|
||||
return datetime.utcnow() # type: ignore
|
||||
return datetime.utcnow()
|
||||
|
||||
|
||||
class OpenMeterLogger(CustomLogger):
|
||||
|
|
|
|||
|
|
@ -370,7 +370,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
"model_id": config.model_id or config.service_name,
|
||||
}
|
||||
|
||||
base_resource: Final = Resource.create(base_attributes) # type: ignore[arg-type]
|
||||
base_resource: Final = Resource.create(base_attributes)
|
||||
otel_resource_detector: Final = OTELResourceDetector()
|
||||
env_resource: Final = otel_resource_detector.detect()
|
||||
return base_resource.merge(env_resource)
|
||||
|
|
@ -640,9 +640,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
def create_logger_provider():
|
||||
provider: Final = OTLoggerProvider(resource=self._get_litellm_resource(self.config))
|
||||
log_exporter: Final = self._get_log_exporter()
|
||||
provider.add_log_record_processor(
|
||||
BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
|
||||
)
|
||||
provider.add_log_record_processor(BatchLogRecordProcessor(log_exporter))
|
||||
return provider
|
||||
|
||||
self._logger_provider = self._get_or_create_provider(
|
||||
|
|
@ -2455,7 +2453,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
message = choice.get("message")
|
||||
tool_calls = message.get("tool_calls")
|
||||
if tool_calls:
|
||||
kv_pairs = OpenTelemetry._tool_calls_kv_pair(tool_calls) # type: ignore
|
||||
kv_pairs = OpenTelemetry._tool_calls_kv_pair(tool_calls)
|
||||
for key, value in kv_pairs.items():
|
||||
self.safe_set_attribute(
|
||||
span=span,
|
||||
|
|
@ -2495,7 +2493,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
}
|
||||
)
|
||||
if tool_calls:
|
||||
kv_pairs = OpenTelemetry._tool_calls_kv_pair(tool_calls) # type: ignore
|
||||
kv_pairs = OpenTelemetry._tool_calls_kv_pair(tool_calls)
|
||||
for key, value in kv_pairs.items():
|
||||
self.safe_set_attribute(
|
||||
span=span,
|
||||
|
|
@ -2616,10 +2614,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
return obj
|
||||
if hasattr(obj, "get"):
|
||||
# BaseLiteLLMOpenAIResponseObject duck-type
|
||||
return obj # type: ignore[return-value]
|
||||
return obj
|
||||
if hasattr(obj, "model_dump"):
|
||||
# Raw Pydantic v2 model (e.g. openai SDK types)
|
||||
return obj.model_dump() # type: ignore[union-attr]
|
||||
return obj.model_dump()
|
||||
return None
|
||||
|
||||
def _transform_responses_api_output_to_otel(self, output: list) -> list[dict]:
|
||||
|
|
|
|||
|
|
@ -168,7 +168,7 @@ class OpikLogger(CustomBatchLogger):
|
|||
response: Final = self.sync_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=batch, # type: ignore
|
||||
json=batch,
|
||||
)
|
||||
response.raise_for_status()
|
||||
if response.status_code != 204:
|
||||
|
|
@ -252,7 +252,7 @@ class OpikLogger(CustomBatchLogger):
|
|||
response: Final = await self.async_httpx_client.post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json=batch, # type: ignore
|
||||
json=batch,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
|
|
|
|||
|
|
@ -194,7 +194,7 @@ def resolve_parent_context(threaded: Span | None = None) -> Context:
|
|||
"""
|
||||
ctx = get_current()
|
||||
if is_recordable_span(threaded) and not is_recordable_span(get_current_span(ctx)):
|
||||
ctx = context_from_span(threaded, context=ctx) # type: ignore[arg-type]
|
||||
ctx = context_from_span(threaded, context=ctx)
|
||||
return ctx
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1352,7 +1352,7 @@ class PrometheusLogger(CustomLogger):
|
|||
# why type ignore below?
|
||||
# 1. We just checked if isinstance(standard_logging_payload, dict). Pyright complains.
|
||||
# 2. Pyright does not allow us to run isinstance(standard_logging_payload, StandardLoggingPayload) <- this would be ideal
|
||||
standard_logging_payload=standard_logging_payload, # type: ignore
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
end_user_id=end_user_id,
|
||||
user_api_key=user_api_key,
|
||||
user_api_key_alias=user_api_key_alias,
|
||||
|
|
@ -1416,14 +1416,14 @@ class PrometheusLogger(CustomLogger):
|
|||
# model_group, derive remaining from configured-limit minus current usage so
|
||||
# the same metric is populated for any provider.
|
||||
await self._async_set_router_remaining_metrics(
|
||||
standard_logging_payload=standard_logging_payload, # type: ignore
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
)
|
||||
|
||||
# cache metrics
|
||||
self._increment_cache_metrics(
|
||||
standard_logging_payload=standard_logging_payload, # type: ignore
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
)
|
||||
|
|
@ -3050,7 +3050,7 @@ class PrometheusLogger(CustomLogger):
|
|||
try:
|
||||
from litellm.exceptions import BudgetExceededError
|
||||
except ImportError:
|
||||
BudgetExceededError = None # type: ignore[assignment,misc]
|
||||
BudgetExceededError = None
|
||||
|
||||
if BudgetExceededError is not None and isinstance(exception, BudgetExceededError):
|
||||
return "BudgetExceededError"
|
||||
|
|
|
|||
|
|
@ -14,8 +14,8 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
httpxSpecialProvider,
|
||||
)
|
||||
|
||||
PROMETHEUS_URL: Final[str | None] = get_secret("PROMETHEUS_URL") # type: ignore
|
||||
PROMETHEUS_SELECTED_INSTANCE: Final[str | None] = get_secret("PROMETHEUS_SELECTED_INSTANCE") # type: ignore
|
||||
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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -882,7 +882,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
|||
payload["id"] = self._correlation_id(call_details) or f"chatcmpl-{uuid.uuid4()}"
|
||||
self._prepend_system_prompt(payload, call_details)
|
||||
|
||||
return payload # type: ignore[return-value]
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _caller_metadata(user_api_key_dict: "UserAPIKeyAuth") -> StandardLoggingUserAPIKeyMetadata:
|
||||
|
|
|
|||
|
|
@ -28,9 +28,7 @@ class Supabase:
|
|||
raise ValueError(
|
||||
"LiteLLM Error, trying to use Supabase but url or key not passed. Create a table and set `litellm.supabase_url=<your-url>` and `litellm.supabase_key=<your-key>`"
|
||||
)
|
||||
self.supabase_client = supabase.create_client( # type: ignore
|
||||
self.supabase_url, self.supabase_key
|
||||
)
|
||||
self.supabase_client = supabase.create_client(self.supabase_url, self.supabase_key)
|
||||
|
||||
def input_log_event(self, model, messages, end_user, litellm_call_id, print_verbose):
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -85,7 +85,7 @@ class TraceloopLogger:
|
|||
)
|
||||
if "temperature" in optional_params:
|
||||
span.set_attribute(
|
||||
SpanAttributes.LLM_REQUEST_TEMPERATURE, # type: ignore
|
||||
SpanAttributes.LLM_REQUEST_TEMPERATURE,
|
||||
kwargs.get("temperature"),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ try:
|
|||
K = TypeVar("K", bound=str)
|
||||
V = TypeVar("V")
|
||||
|
||||
class OpenAIResponse(Protocol[K, V]): # type: ignore
|
||||
class OpenAIResponse(Protocol[K, V]):
|
||||
# contains a (known) object attribute
|
||||
object: Literal["chat.completion", "edit", "text_completion"]
|
||||
|
||||
|
|
@ -70,7 +70,7 @@ try:
|
|||
end_time_ms: Final = start_time_ms + int(round(time_elapsed * 1000))
|
||||
span: Final = trace_tree.Span(
|
||||
name=f"{response.get('model', 'openai')}_{response['object']}_{response.get('created')}",
|
||||
attributes=dict(response), # type: ignore
|
||||
attributes=dict(response),
|
||||
start_time_ms=start_time_ms,
|
||||
end_time_ms=end_time_ms,
|
||||
span_kind=trace_tree.SpanKind.LLM,
|
||||
|
|
|
|||
|
|
@ -77,7 +77,7 @@ def _make_logging_obj(
|
|||
call_type: str,
|
||||
optional_params: dict[str, Any],
|
||||
) -> LiteLLMLoggingObj:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
|
|
|
|||
|
|
@ -171,7 +171,7 @@ async def acreate(
|
|||
else:
|
||||
response = init_response
|
||||
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
@ -255,7 +255,7 @@ def create(
|
|||
local_vars: Final = locals()
|
||||
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acreate_interaction", False) is True
|
||||
|
||||
|
|
@ -378,7 +378,7 @@ async def aget(
|
|||
else:
|
||||
response = init_response
|
||||
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=None,
|
||||
|
|
@ -402,7 +402,7 @@ def get(
|
|||
custom_llm_provider = custom_llm_provider or "gemini"
|
||||
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("aget_interaction", False) is True
|
||||
|
||||
|
|
@ -480,7 +480,7 @@ async def adelete(
|
|||
else:
|
||||
response = init_response
|
||||
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=None,
|
||||
|
|
@ -504,7 +504,7 @@ def delete(
|
|||
custom_llm_provider = custom_llm_provider or "gemini"
|
||||
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("adelete_interaction", False) is True
|
||||
|
||||
|
|
@ -582,7 +582,7 @@ async def acancel(
|
|||
else:
|
||||
response = init_response
|
||||
|
||||
return response # type: ignore
|
||||
return response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=None,
|
||||
|
|
@ -606,7 +606,7 @@ def cancel(
|
|||
custom_llm_provider = custom_llm_provider or "gemini"
|
||||
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("acancel_interaction", False) is True
|
||||
|
||||
|
|
|
|||
|
|
@ -99,10 +99,10 @@ def process_audio_file(audio_file: FileTypes) -> ProcessedAudioFile:
|
|||
elif hasattr(audio_file, "read") and not isinstance(audio_file, (str, bytes, bytearray, tuple, os.PathLike)):
|
||||
# File-like object (IO) - check this after all other types
|
||||
filename = getattr(audio_file, "name", "audio.wav")
|
||||
file_content = audio_file.read() # type: ignore
|
||||
file_content = audio_file.read()
|
||||
# Reset file pointer if possible
|
||||
if hasattr(audio_file, "seek"):
|
||||
audio_file.seek(0) # type: ignore
|
||||
audio_file.seek(0)
|
||||
else:
|
||||
raise ValueError(f"Unsupported audio_file type: {type(audio_file)}")
|
||||
|
||||
|
|
@ -211,9 +211,9 @@ def get_audio_file_content_hash(file_obj: FileTypes) -> str:
|
|||
current_position: Final = file_content_obj.tell() if hasattr(file_content_obj, "tell") else None
|
||||
if hasattr(file_content_obj, "seek"):
|
||||
file_content_obj.seek(0)
|
||||
file_content = file_content_obj.read() # type: ignore
|
||||
file_content = file_content_obj.read()
|
||||
if current_position is not None and hasattr(file_content_obj, "seek"):
|
||||
file_content_obj.seek(current_position) # type: ignore
|
||||
file_content_obj.seek(current_position)
|
||||
except (OSError, AttributeError):
|
||||
file_content = None
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -65,6 +65,6 @@ class CompletionTimeout:
|
|||
float(read_timeout) if read_timeout is not None else COMPLETION_HTTP_FALLBACK_SECONDS
|
||||
) # default 10 min timeout
|
||||
elif not isinstance(resolved, httpx.Timeout):
|
||||
resolved = float(resolved) # type: ignore
|
||||
resolved = float(resolved)
|
||||
|
||||
return resolved
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ try:
|
|||
filename = str(resources.files(litellm).joinpath("litellm_core_utils/tokenizers"))
|
||||
except (ImportError, AttributeError):
|
||||
# Old way to access resources, which setuptools deprecated some time ago
|
||||
import pkg_resources # type: ignore
|
||||
import pkg_resources
|
||||
|
||||
filename = pkg_resources.resource_filename(__name__, "litellm_core_utils/tokenizers")
|
||||
|
||||
|
|
|
|||
|
|
@ -1109,7 +1109,7 @@ def _map_vertex_exception(
|
|||
response=httpx.Response(
|
||||
status_code=500,
|
||||
content=str(original_exception),
|
||||
request=httpx.Request(method="completion", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="completion", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
|
|
@ -1270,7 +1270,7 @@ def _map_vertex_exception(
|
|||
response=httpx.Response(
|
||||
status_code=500,
|
||||
content=str(original_exception),
|
||||
request=httpx.Request(method="completion", url="https://github.com/BerriAI/litellm"), # type: ignore
|
||||
request=httpx.Request(method="completion", url="https://github.com/BerriAI/litellm"),
|
||||
),
|
||||
)
|
||||
if original_exception.status_code == 502:
|
||||
|
|
@ -1872,15 +1872,13 @@ def _map_azure_exception(
|
|||
body_dict: Final = getattr(original_exception, "body", None) or {}
|
||||
if isinstance(body_dict, dict):
|
||||
if isinstance(body_dict.get("error"), dict):
|
||||
azure_error_code = body_dict["error"].get("code") # type: ignore[index]
|
||||
azure_error_code = body_dict["error"].get("code")
|
||||
# Also check inner_error for
|
||||
# ResponsibleAIPolicyViolation which indicates a
|
||||
# content policy violation even when the top-level
|
||||
# code is generic (e.g. "invalid_request_error").
|
||||
if azure_error_code != "content_policy_violation":
|
||||
_inner: Final = body_dict["error"].get("inner_error") or body_dict[ # type: ignore[index]
|
||||
"error"
|
||||
].get("innererror") # type: ignore[index]
|
||||
_inner: Final = body_dict["error"].get("inner_error") or body_dict["error"].get("innererror")
|
||||
if isinstance(_inner, dict) and _inner.get("code") == "ResponsibleAIPolicyViolation":
|
||||
azure_error_code = "content_policy_violation"
|
||||
else:
|
||||
|
|
@ -2156,7 +2154,7 @@ def _map_openrouter_exception(
|
|||
)
|
||||
|
||||
|
||||
def exception_type( # type: ignore
|
||||
def exception_type(
|
||||
model,
|
||||
original_exception,
|
||||
custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -354,7 +354,7 @@ def get_llm_provider(
|
|||
raise Exception(f"api base needs to be a string. api_base={api_base}")
|
||||
if dynamic_api_key is not None and not isinstance(dynamic_api_key, str):
|
||||
raise Exception(f"dynamic_api_key needs to be a string. dynamic_api_key={dynamic_api_key}")
|
||||
return model, custom_llm_provider, dynamic_api_key, api_base # type: ignore
|
||||
return model, custom_llm_provider, dynamic_api_key, api_base
|
||||
|
||||
# check if model in known model provider list -> for huggingface models, raise exception as they don't have a fixed provider (can be togetherai, anyscale, baseten, runpod, et.)
|
||||
## openai - chatcompletion + text completion
|
||||
|
|
@ -412,7 +412,7 @@ def get_llm_provider(
|
|||
## ai21
|
||||
elif model in litellm.ai21_chat_models or model in litellm.ai21_models:
|
||||
custom_llm_provider = "ai21_chat"
|
||||
api_base = api_base or get_secret("AI21_API_BASE") or "https://api.ai21.com/studio/v1" # type: ignore
|
||||
api_base = api_base or get_secret("AI21_API_BASE") or "https://api.ai21.com/studio/v1"
|
||||
dynamic_api_key = api_key or get_secret("AI21_API_KEY")
|
||||
## aleph_alpha
|
||||
elif model in litellm.aleph_alpha_models:
|
||||
|
|
@ -486,7 +486,7 @@ def get_llm_provider(
|
|||
print() # noqa: T201
|
||||
error_str = f"LLM Provider NOT provided. Pass in the LLM provider you are trying to call. You passed model={model}\n Pass model as E.g. For 'Huggingface' inference endpoints pass in `completion(model='huggingface/starcoder',..)` Learn more: https://docs.litellm.ai/docs/providers"
|
||||
# maps to openai.NotFoundError, this is raised when openai does not recognize the llm
|
||||
raise litellm.exceptions.BadRequestError( # type: ignore
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=error_str,
|
||||
model=model,
|
||||
response=None,
|
||||
|
|
@ -502,7 +502,7 @@ def get_llm_provider(
|
|||
raise e
|
||||
else:
|
||||
error_str = f"GetLLMProvider Exception - {e}\n\noriginal model: {model}"
|
||||
raise litellm.exceptions.BadRequestError( # type: ignore
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message=f"GetLLMProvider Exception - {e}\n\noriginal model: {model}",
|
||||
model=model,
|
||||
response=None,
|
||||
|
|
@ -551,7 +551,7 @@ def _get_openai_compatible_provider_info(
|
|||
return model, "aiohttp_openai", api_key, api_base
|
||||
elif custom_llm_provider == "anyscale":
|
||||
# anyscale is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1
|
||||
api_base = api_base or get_secret_str("ANYSCALE_API_BASE") or "https://api.endpoints.anyscale.com/v1" # type: ignore
|
||||
api_base = api_base or get_secret_str("ANYSCALE_API_BASE") or "https://api.endpoints.anyscale.com/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("ANYSCALE_API_KEY")
|
||||
elif custom_llm_provider == "deepinfra":
|
||||
(
|
||||
|
|
@ -559,7 +559,7 @@ def _get_openai_compatible_provider_info(
|
|||
dynamic_api_key,
|
||||
) = litellm.DeepInfraConfig()._get_openai_compatible_provider_info(api_base, api_key)
|
||||
elif custom_llm_provider == "empower":
|
||||
api_base = api_base or get_secret("EMPOWER_API_BASE") or "https://app.empower.dev/api/v1" # type: ignore
|
||||
api_base = api_base or get_secret("EMPOWER_API_BASE") or "https://app.empower.dev/api/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("EMPOWER_API_KEY")
|
||||
elif custom_llm_provider == "groq":
|
||||
(
|
||||
|
|
@ -575,13 +575,13 @@ def _get_openai_compatible_provider_info(
|
|||
)
|
||||
elif custom_llm_provider == "nvidia_nim":
|
||||
# nvidia_nim is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1
|
||||
api_base = api_base or get_secret("NVIDIA_NIM_API_BASE") or "https://integrate.api.nvidia.com/v1" # type: ignore
|
||||
api_base = api_base or get_secret("NVIDIA_NIM_API_BASE") or "https://integrate.api.nvidia.com/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("NVIDIA_NIM_API_KEY")
|
||||
elif custom_llm_provider == "nvidia_riva":
|
||||
# NVIDIA Riva is gRPC-based; api_base must be a host:port like
|
||||
# `grpc.nvcf.nvidia.com:443` or `localhost:50051`. There is no
|
||||
# public-default endpoint, so we do not fill one in here.
|
||||
api_base = api_base or get_secret_str("NVIDIA_RIVA_API_BASE") # type: ignore
|
||||
api_base = api_base or get_secret_str("NVIDIA_RIVA_API_BASE")
|
||||
# Fall back to NVIDIA_NIM_API_KEY because users running both NVCF
|
||||
# services typically reuse the same nvapi-* key.
|
||||
dynamic_api_key = api_key or get_secret_str("NVIDIA_RIVA_API_KEY") or get_secret_str("NVIDIA_NIM_API_KEY")
|
||||
|
|
@ -589,7 +589,7 @@ def _get_openai_compatible_provider_info(
|
|||
api_base = api_base or get_secret_str("SONIOX_API_BASE") or "https://api.soniox.com"
|
||||
dynamic_api_key = api_key or get_secret_str("SONIOX_API_KEY")
|
||||
elif custom_llm_provider == "cerebras":
|
||||
api_base = api_base or get_secret("CEREBRAS_API_BASE") or "https://api.cerebras.ai/v1" # type: ignore
|
||||
api_base = api_base or get_secret("CEREBRAS_API_BASE") or "https://api.cerebras.ai/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("CEREBRAS_API_KEY")
|
||||
elif custom_llm_provider == "baseten":
|
||||
# Use BasetenConfig to determine the appropriate API base URL
|
||||
|
|
@ -599,28 +599,28 @@ def _get_openai_compatible_provider_info(
|
|||
api_base = api_base or get_secret_str("BASETEN_API_BASE") or "https://inference.baseten.co/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("BASETEN_API_KEY")
|
||||
elif custom_llm_provider == "sambanova":
|
||||
api_base = api_base or get_secret("SAMBANOVA_API_BASE") or "https://api.sambanova.ai/v1" # type: ignore
|
||||
api_base = api_base or get_secret("SAMBANOVA_API_BASE") or "https://api.sambanova.ai/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("SAMBANOVA_API_KEY")
|
||||
elif custom_llm_provider == "meta_llama":
|
||||
api_base = api_base or get_secret("LLAMA_API_BASE") or "https://api.llama.com/compat/v1" # type: ignore
|
||||
api_base = api_base or get_secret("LLAMA_API_BASE") or "https://api.llama.com/compat/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("LLAMA_API_KEY")
|
||||
elif custom_llm_provider == "nebius":
|
||||
api_base = api_base or get_secret("NEBIUS_API_BASE") or "https://api.studio.nebius.ai/v1" # type: ignore
|
||||
api_base = api_base or get_secret("NEBIUS_API_BASE") or "https://api.studio.nebius.ai/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("NEBIUS_API_KEY")
|
||||
elif custom_llm_provider == "ollama":
|
||||
api_base = api_base or get_secret("OLLAMA_API_BASE") or "http://localhost:11434" # type: ignore
|
||||
api_base = api_base or get_secret("OLLAMA_API_BASE") or "http://localhost:11434"
|
||||
dynamic_api_key = api_key or get_secret_str("OLLAMA_API_KEY")
|
||||
elif (custom_llm_provider == "ai21_chat") or (custom_llm_provider == "ai21" and model in litellm.ai21_chat_models):
|
||||
api_base = api_base or get_secret("AI21_API_BASE") or "https://api.ai21.com/studio/v1" # type: ignore
|
||||
api_base = api_base or get_secret("AI21_API_BASE") or "https://api.ai21.com/studio/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("AI21_API_KEY")
|
||||
custom_llm_provider = "ai21_chat"
|
||||
elif custom_llm_provider == "volcengine":
|
||||
# volcengine is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1
|
||||
api_base = api_base or get_secret("VOLCENGINE_API_BASE") or "https://ark.cn-beijing.volces.com/api/v3" # type: ignore
|
||||
api_base = api_base or get_secret("VOLCENGINE_API_BASE") or "https://ark.cn-beijing.volces.com/api/v3"
|
||||
dynamic_api_key = api_key or get_secret_str("VOLCENGINE_API_KEY")
|
||||
elif custom_llm_provider == "codestral":
|
||||
# codestral is openai compatible, we just need to set this to custom_openai and have the api_base be https://codestral.mistral.ai/v1
|
||||
api_base = api_base or get_secret("CODESTRAL_API_BASE") or "https://codestral.mistral.ai/v1" # type: ignore
|
||||
api_base = api_base or get_secret("CODESTRAL_API_BASE") or "https://codestral.mistral.ai/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("CODESTRAL_API_KEY")
|
||||
elif custom_llm_provider == "hosted_vllm":
|
||||
# vllm is openai compatible, we just need to set this to custom_openai
|
||||
|
|
@ -648,7 +648,7 @@ def _get_openai_compatible_provider_info(
|
|||
) = litellm.LMStudioChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
|
||||
elif custom_llm_provider == "deepseek":
|
||||
# deepseek is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.deepseek.com/v1
|
||||
api_base = api_base or get_secret("DEEPSEEK_API_BASE") or "https://api.deepseek.com/beta" # type: ignore
|
||||
api_base = api_base or get_secret("DEEPSEEK_API_BASE") or "https://api.deepseek.com/beta"
|
||||
|
||||
dynamic_api_key = api_key or get_secret_str("DEEPSEEK_API_KEY")
|
||||
elif custom_llm_provider == "tencent":
|
||||
|
|
@ -704,7 +704,7 @@ def _get_openai_compatible_provider_info(
|
|||
dynamic_api_key,
|
||||
) = litellm.ZAIChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
|
||||
elif custom_llm_provider == "together_ai":
|
||||
api_base = api_base or get_secret_str("TOGETHER_AI_API_BASE") or "https://api.together.xyz/v1" # type: ignore
|
||||
api_base = api_base or get_secret_str("TOGETHER_AI_API_BASE") or "https://api.together.xyz/v1"
|
||||
dynamic_api_key = api_key or (
|
||||
get_secret_str("TOGETHER_API_KEY")
|
||||
or get_secret_str("TOGETHER_AI_API_KEY")
|
||||
|
|
@ -712,10 +712,10 @@ def _get_openai_compatible_provider_info(
|
|||
or get_secret_str("TOGETHER_AI_TOKEN")
|
||||
)
|
||||
elif custom_llm_provider == "friendliai":
|
||||
api_base = api_base or get_secret("FRIENDLI_API_BASE") or "https://api.friendli.ai/serverless/v1" # type: ignore
|
||||
api_base = api_base or get_secret("FRIENDLI_API_BASE") or "https://api.friendli.ai/serverless/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("FRIENDLIAI_API_KEY") or get_secret_str("FRIENDLI_TOKEN")
|
||||
elif custom_llm_provider == "galadriel":
|
||||
api_base = api_base or get_secret("GALADRIEL_API_BASE") or "https://api.galadriel.com/v1" # type: ignore
|
||||
api_base = api_base or get_secret("GALADRIEL_API_BASE") or "https://api.galadriel.com/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("GALADRIEL_API_KEY")
|
||||
elif custom_llm_provider == "github_copilot":
|
||||
(
|
||||
|
|
@ -732,7 +732,7 @@ def _get_openai_compatible_provider_info(
|
|||
custom_llm_provider,
|
||||
) = litellm.ChatGPTConfig()._get_openai_compatible_provider_info(model, api_base, api_key, custom_llm_provider)
|
||||
elif custom_llm_provider == "novita":
|
||||
api_base = api_base or get_secret("NOVITA_API_BASE") or "https://api.novita.ai/v3/openai" # type: ignore
|
||||
api_base = api_base or get_secret("NOVITA_API_BASE") or "https://api.novita.ai/v3/openai"
|
||||
dynamic_api_key = api_key or get_secret_str("NOVITA_API_KEY")
|
||||
elif custom_llm_provider == "snowflake":
|
||||
(
|
||||
|
|
@ -816,7 +816,7 @@ def _get_openai_compatible_provider_info(
|
|||
dynamic_api_key,
|
||||
) = litellm.AIMLChatConfig()._get_openai_compatible_provider_info(api_base, api_key)
|
||||
elif custom_llm_provider == "wandb":
|
||||
api_base = api_base or get_secret("WANDB_API_BASE") or "https://api.inference.wandb.ai/v1" # type: ignore
|
||||
api_base = api_base or get_secret("WANDB_API_BASE") or "https://api.inference.wandb.ai/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("WANDB_API_KEY")
|
||||
elif custom_llm_provider == "lemonade":
|
||||
(
|
||||
|
|
|
|||
|
|
@ -121,7 +121,7 @@ def initialize_standard_callback_dynamic_params(
|
|||
if param in kwargs:
|
||||
_param_value = kwargs.get(param)
|
||||
validate_no_callback_env_reference(param, _param_value, source="request body")
|
||||
standard_callback_dynamic_params[param] = _param_value # type: ignore
|
||||
standard_callback_dynamic_params[param] = _param_value
|
||||
|
||||
for slot_label, metadata in iter_client_callback_metadata_dicts(kwargs):
|
||||
for param in _supported_callback_params:
|
||||
|
|
@ -130,6 +130,6 @@ def initialize_standard_callback_dynamic_params(
|
|||
if param not in standard_callback_dynamic_params and param in metadata:
|
||||
_param_value = metadata.get(param)
|
||||
validate_no_callback_env_reference(param, _param_value, source=slot_label)
|
||||
standard_callback_dynamic_params[param] = _param_value # type: ignore
|
||||
standard_callback_dynamic_params[param] = _param_value
|
||||
|
||||
return standard_callback_dynamic_params
|
||||
|
|
|
|||
|
|
@ -196,12 +196,12 @@ try:
|
|||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("[Non-Blocking] Unable to import GenericAPILogger - LiteLLM Enterprise Feature - %s", e)
|
||||
GenericAPILogger = CustomLogger # type: ignore
|
||||
ResendEmailLogger = CustomLogger # type: ignore
|
||||
SendGridEmailLogger = CustomLogger # type: ignore
|
||||
SMTPEmailLogger = CustomLogger # type: ignore
|
||||
PagerDutyAlerting = CustomLogger # type: ignore
|
||||
EnterpriseCallbackControls = None # type: ignore
|
||||
GenericAPILogger = CustomLogger
|
||||
ResendEmailLogger = CustomLogger
|
||||
SendGridEmailLogger = CustomLogger
|
||||
SMTPEmailLogger = CustomLogger
|
||||
PagerDutyAlerting = CustomLogger
|
||||
EnterpriseCallbackControls = None
|
||||
EnterpriseStandardLoggingPayloadSetupVAR = None
|
||||
_in_memory_loggers: Final[list[Any]] = []
|
||||
|
||||
|
|
@ -462,9 +462,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
_custom_logger_init_args = {k: v for k, v in self._trusted_callback_vars if k.startswith("dd_")}
|
||||
|
||||
callback_class = _init_custom_logger_compatible_class(
|
||||
callback, # type: ignore[arg-type]
|
||||
callback,
|
||||
internal_usage_cache=None,
|
||||
llm_router=None, # type: ignore
|
||||
llm_router=None,
|
||||
custom_logger_init_args=_custom_logger_init_args,
|
||||
)
|
||||
if callback_class is not None:
|
||||
|
|
@ -1756,7 +1756,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details["litellm_params"]["metadata"] = {}
|
||||
self.model_call_details["litellm_params"]["metadata"]["hidden_params"] = getattr(
|
||||
logging_result, "_hidden_params", {}
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
if self.model_call_details.get("cache_hit") is True:
|
||||
self.model_call_details["response_cost"] = 0.0
|
||||
|
|
@ -1815,7 +1815,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
result = result.model_copy()
|
||||
transformed_usage = TranscriptionUsageObjectTransformation.transform_transcription_usage_object(
|
||||
result.usage
|
||||
) # type: ignore
|
||||
)
|
||||
setattr(result, "usage", transformed_usage)
|
||||
return result
|
||||
|
||||
|
|
@ -2137,7 +2137,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=print_verbose,
|
||||
level=LogfireLevel.INFO.value, # type: ignore
|
||||
level=LogfireLevel.INFO.value,
|
||||
)
|
||||
|
||||
if callback == "lunary" and lunaryLogger is not None:
|
||||
|
|
@ -2699,7 +2699,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
for callback_obj in all_callbacks:
|
||||
if hasattr(callback_obj, "increment_callback_logging_failure"):
|
||||
callback_obj.increment_callback_logging_failure(callback_name=callback_name) # type: ignore
|
||||
callback_obj.increment_callback_logging_failure(callback_name=callback_name)
|
||||
break # Only increment once
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -2779,7 +2779,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
exception=exception,
|
||||
original_model_group=model_group,
|
||||
kwargs=self.model_call_details,
|
||||
) # type: ignore
|
||||
)
|
||||
|
||||
def failure_handler(self, exception, traceback_exception, start_time=None, end_time=None):
|
||||
verbose_logger.debug("Logging Details LiteLLM-Failure Call: %s", litellm.failure_callback)
|
||||
|
|
@ -2934,7 +2934,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
response_obj=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
level=LogfireLevel.ERROR.value, # type: ignore
|
||||
level=LogfireLevel.ERROR.value,
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
|
||||
|
|
@ -2988,7 +2988,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
response_obj=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
) # type: ignore
|
||||
)
|
||||
if callable(callback): # custom logger functions
|
||||
global customLogger
|
||||
if customLogger is None:
|
||||
|
|
@ -3478,7 +3478,7 @@ def set_callbacks(callback_list, function_id=None):
|
|||
)
|
||||
sentry_sdk_instance.init(
|
||||
dsn=os.environ.get("SENTRY_DSN"),
|
||||
traces_sample_rate=float(sentry_trace_rate), # type: ignore
|
||||
traces_sample_rate=float(sentry_trace_rate),
|
||||
sample_rate=float(sentry_sample_rate if sentry_sample_rate else 1.0),
|
||||
send_default_pii=False, # Prevent sending Personal Identifiable Information
|
||||
event_scrubber=EventScrubber(denylist=SENTRY_DENYLIST, pii_denylist=SENTRY_PII_DENYLIST),
|
||||
|
|
@ -3552,90 +3552,90 @@ def _init_custom_logger_compatible_class(
|
|||
if logging_integration == "agentops": # Add AgentOps initialization
|
||||
_v2 = _maybe_construct_otel_v2("agentops", _in_memory_loggers)
|
||||
if _v2 is not None:
|
||||
return _v2 # type: ignore
|
||||
return _v2
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, AgentOps):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
agentops_logger: Final = AgentOps()
|
||||
_in_memory_loggers.append(agentops_logger)
|
||||
return agentops_logger # type: ignore
|
||||
return agentops_logger
|
||||
elif logging_integration == "lago":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LagoLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
lago_logger: Final = LagoLogger()
|
||||
_in_memory_loggers.append(lago_logger)
|
||||
return lago_logger # type: ignore
|
||||
return lago_logger
|
||||
elif logging_integration == "openmeter":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, OpenMeterLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_openmeter_logger: Final = OpenMeterLogger()
|
||||
_in_memory_loggers.append(_openmeter_logger)
|
||||
return _openmeter_logger # type: ignore
|
||||
return _openmeter_logger
|
||||
elif logging_integration == "posthog":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PostHogLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_posthog_logger: Final = PostHogLogger()
|
||||
_in_memory_loggers.append(_posthog_logger)
|
||||
return _posthog_logger # type: ignore
|
||||
return _posthog_logger
|
||||
elif logging_integration == "braintrust":
|
||||
from litellm.integrations.braintrust_logging import BraintrustLogger
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, BraintrustLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
braintrust_logger: Final = BraintrustLogger()
|
||||
_in_memory_loggers.append(braintrust_logger)
|
||||
return braintrust_logger # type: ignore
|
||||
return braintrust_logger
|
||||
elif logging_integration == "langsmith":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LangsmithLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_langsmith_logger: Final = LangsmithLogger()
|
||||
_in_memory_loggers.append(_langsmith_logger)
|
||||
return _langsmith_logger # type: ignore
|
||||
return _langsmith_logger
|
||||
elif logging_integration == "argilla":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ArgillaLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_argilla_logger: Final = ArgillaLogger()
|
||||
_in_memory_loggers.append(_argilla_logger)
|
||||
return _argilla_logger # type: ignore
|
||||
return _argilla_logger
|
||||
elif logging_integration == "literalai":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LiteralAILogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_literalai_logger: Final = LiteralAILogger()
|
||||
_in_memory_loggers.append(_literalai_logger)
|
||||
return _literalai_logger # type: ignore
|
||||
return _literalai_logger
|
||||
elif logging_integration == "litellm_agent":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LiteLLMAgentModelResolver):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_litellm_agent_resolver: Final = LiteLLMAgentModelResolver()
|
||||
_in_memory_loggers.append(_litellm_agent_resolver)
|
||||
return _litellm_agent_resolver # type: ignore
|
||||
return _litellm_agent_resolver
|
||||
elif logging_integration == "prometheus":
|
||||
PrometheusLogger: Final = _get_cached_prometheus_logger()
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PrometheusLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_prometheus_logger: Final = PrometheusLogger()
|
||||
_in_memory_loggers.append(_prometheus_logger)
|
||||
return _prometheus_logger # type: ignore
|
||||
return _prometheus_logger
|
||||
elif logging_integration == "datadog":
|
||||
# Check if team-scoped credentials are provided
|
||||
_dd_api_key: Final = custom_logger_init_args.get("dd_api_key")
|
||||
|
|
@ -3650,82 +3650,82 @@ def _init_custom_logger_compatible_class(
|
|||
)
|
||||
|
||||
return DataDogHandler.get_datadog_logger_for_request(
|
||||
standard_callback_dynamic_params=custom_logger_init_args, # type: ignore
|
||||
standard_callback_dynamic_params=custom_logger_init_args,
|
||||
in_memory_dynamic_logger_cache=in_memory_dynamic_logger_cache,
|
||||
)
|
||||
|
||||
# Global (env-var based): reuse cached instance
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, DataDogLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_datadog_logger: Final = DataDogLogger()
|
||||
_in_memory_loggers.append(_datadog_logger)
|
||||
return _datadog_logger # type: ignore
|
||||
return _datadog_logger
|
||||
elif logging_integration == "datadog_metrics":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, DatadogMetricsLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_datadog_metrics_logger: Final = DatadogMetricsLogger()
|
||||
_in_memory_loggers.append(_datadog_metrics_logger)
|
||||
return _datadog_metrics_logger # type: ignore
|
||||
return _datadog_metrics_logger
|
||||
elif logging_integration == "datadog_llm_observability":
|
||||
_datadog_llm_obs_logger: Final = DataDogLLMObsLogger()
|
||||
_in_memory_loggers.append(_datadog_llm_obs_logger)
|
||||
return _datadog_llm_obs_logger # type: ignore
|
||||
return _datadog_llm_obs_logger
|
||||
elif logging_integration == "azure_sentinel":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, AzureSentinelLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_azure_sentinel_logger: Final = AzureSentinelLogger()
|
||||
_in_memory_loggers.append(_azure_sentinel_logger)
|
||||
return _azure_sentinel_logger # type: ignore
|
||||
return _azure_sentinel_logger
|
||||
elif logging_integration == "gcs_bucket":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, GCSBucketLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_gcs_bucket_logger: Final = GCSBucketLogger()
|
||||
_in_memory_loggers.append(_gcs_bucket_logger)
|
||||
return _gcs_bucket_logger # type: ignore
|
||||
return _gcs_bucket_logger
|
||||
elif logging_integration == "s3_v2":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, S3V2Logger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_s3_v2_logger: Final = S3V2Logger()
|
||||
_in_memory_loggers.append(_s3_v2_logger)
|
||||
return _s3_v2_logger # type: ignore
|
||||
return _s3_v2_logger
|
||||
elif logging_integration == "aws_sqs":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, SQSLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_aws_sqs_logger: Final = SQSLogger()
|
||||
_in_memory_loggers.append(_aws_sqs_logger)
|
||||
return _aws_sqs_logger # type: ignore
|
||||
return _aws_sqs_logger
|
||||
elif logging_integration == "azure_storage":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, AzureBlobStorageLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_azure_storage_logger: Final = AzureBlobStorageLogger()
|
||||
_in_memory_loggers.append(_azure_storage_logger)
|
||||
return _azure_storage_logger # type: ignore
|
||||
return _azure_storage_logger
|
||||
elif logging_integration == "opik":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, OpikLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_opik_logger: Final = OpikLogger()
|
||||
_in_memory_loggers.append(_opik_logger)
|
||||
return _opik_logger # type: ignore
|
||||
return _opik_logger
|
||||
elif logging_integration == "arize":
|
||||
_v2 = _maybe_construct_otel_v2("arize", _in_memory_loggers)
|
||||
if _v2 is not None:
|
||||
return _v2 # type: ignore
|
||||
return _v2
|
||||
from litellm.integrations.opentelemetry import (
|
||||
OpenTelemetry,
|
||||
OpenTelemetryConfig,
|
||||
|
|
@ -3747,14 +3747,14 @@ def _init_custom_logger_compatible_class(
|
|||
)
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ArizeLogger) and callback.callback_name == "arize":
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
_arize_otel_logger: Final = ArizeLogger(config=otel_config, callback_name="arize")
|
||||
_in_memory_loggers.append(_arize_otel_logger)
|
||||
return _arize_otel_logger # type: ignore
|
||||
return _arize_otel_logger
|
||||
elif logging_integration == "arize_phoenix":
|
||||
_v2 = _maybe_construct_otel_v2("arize_phoenix", _in_memory_loggers)
|
||||
if _v2 is not None:
|
||||
return _v2 # type: ignore
|
||||
return _v2
|
||||
from litellm.integrations.opentelemetry import (
|
||||
OpenTelemetry,
|
||||
OpenTelemetryConfig,
|
||||
|
|
@ -3773,14 +3773,14 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ArizePhoenixLogger) and callback.callback_name == "arize_phoenix":
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
_arize_phoenix_otel_logger: Final = ArizePhoenixLogger(config=otel_config, callback_name="arize_phoenix")
|
||||
_in_memory_loggers.append(_arize_phoenix_otel_logger)
|
||||
return _arize_phoenix_otel_logger # type: ignore
|
||||
return _arize_phoenix_otel_logger
|
||||
elif logging_integration == "levo":
|
||||
_v2 = _maybe_construct_otel_v2("levo", _in_memory_loggers)
|
||||
if _v2 is not None:
|
||||
return _v2 # type: ignore
|
||||
return _v2
|
||||
from litellm.integrations.levo.levo import LevoLogger
|
||||
from litellm.integrations.opentelemetry import (
|
||||
OpenTelemetry,
|
||||
|
|
@ -3797,11 +3797,11 @@ def _init_custom_logger_compatible_class(
|
|||
# Check if LevoLogger instance already exists
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LevoLogger) and callback.callback_name == "levo":
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_levo_otel_logger: Final = LevoLogger(config=otel_config, callback_name="levo")
|
||||
_in_memory_loggers.append(_levo_otel_logger)
|
||||
return _levo_otel_logger # type: ignore
|
||||
return _levo_otel_logger
|
||||
elif logging_integration == "otel":
|
||||
# Gate the new typed V2 adapter behind LITELLM_OTEL_V2. When off,
|
||||
# the legacy 3,227-line god-class is used unchanged. The two are
|
||||
|
|
@ -3815,19 +3815,19 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if type(callback) is OpenTelemetryV2:
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
otel_logger_v2: Final = OpenTelemetryV2(
|
||||
**_get_custom_logger_settings_from_proxy_server(callback_name=logging_integration)
|
||||
)
|
||||
_in_memory_loggers.append(otel_logger_v2)
|
||||
_maybe_auto_initialize_arize_phoenix(_in_memory_loggers)
|
||||
return otel_logger_v2 # type: ignore
|
||||
return otel_logger_v2
|
||||
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if type(callback) is OpenTelemetry:
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
otel_logger: Final = OpenTelemetry(
|
||||
**_get_custom_logger_settings_from_proxy_server(callback_name=logging_integration)
|
||||
)
|
||||
|
|
@ -3838,34 +3838,34 @@ def _init_custom_logger_compatible_class(
|
|||
# by only specifying "otel" in callbacks
|
||||
_maybe_auto_initialize_arize_phoenix(_in_memory_loggers)
|
||||
|
||||
return otel_logger # type: ignore
|
||||
return otel_logger
|
||||
|
||||
elif logging_integration == "galileo":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, GalileoObserve):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
galileo_logger: Final = GalileoObserve()
|
||||
_in_memory_loggers.append(galileo_logger)
|
||||
return galileo_logger # type: ignore
|
||||
return galileo_logger
|
||||
elif logging_integration == "cloudzero":
|
||||
from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, CloudZeroLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
cloudzero_logger: Final = CloudZeroLogger()
|
||||
_in_memory_loggers.append(cloudzero_logger)
|
||||
return cloudzero_logger # type: ignore
|
||||
return cloudzero_logger
|
||||
elif logging_integration == "focus":
|
||||
from litellm.integrations.focus.focus_logger import FocusLogger
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if type(callback) is FocusLogger: # exact match; exclude subclasses like VantageLogger
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
focus_logger: Final = FocusLogger()
|
||||
_in_memory_loggers.append(focus_logger)
|
||||
return focus_logger # type: ignore
|
||||
return focus_logger
|
||||
elif logging_integration == "mavvrik":
|
||||
from litellm.integrations.mavvrik_focus.mavvrik_focus_logger import (
|
||||
MavvrikFocusLogger,
|
||||
|
|
@ -3873,26 +3873,26 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if type(callback) is MavvrikFocusLogger:
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
mavvrik_focus_logger: Final = MavvrikFocusLogger()
|
||||
_in_memory_loggers.append(mavvrik_focus_logger)
|
||||
return mavvrik_focus_logger # type: ignore
|
||||
return mavvrik_focus_logger
|
||||
elif logging_integration == "vantage":
|
||||
from litellm.integrations.vantage.vantage_logger import VantageLogger
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, VantageLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
vantage_logger: Final = VantageLogger()
|
||||
_in_memory_loggers.append(vantage_logger)
|
||||
return vantage_logger # type: ignore
|
||||
return vantage_logger
|
||||
elif logging_integration == "deepeval":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, DeepEvalLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
deepeval_logger: Final = DeepEvalLogger()
|
||||
_in_memory_loggers.append(deepeval_logger)
|
||||
return deepeval_logger # type: ignore
|
||||
return deepeval_logger
|
||||
|
||||
elif logging_integration == "logfire":
|
||||
if "LOGFIRE_TOKEN" not in os.environ:
|
||||
|
|
@ -3911,10 +3911,10 @@ def _init_custom_logger_compatible_class(
|
|||
for callback in _in_memory_loggers:
|
||||
# Use exact type check to avoid matching ArizePhoenixLogger (subclass)
|
||||
if type(callback) is OpenTelemetry:
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
_otel_logger = OpenTelemetry(config=otel_config)
|
||||
_in_memory_loggers.append(_otel_logger)
|
||||
return _otel_logger # type: ignore
|
||||
return _otel_logger
|
||||
elif logging_integration == "dynamic_rate_limiter":
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter import (
|
||||
_PROXY_DynamicRateLimitHandler,
|
||||
|
|
@ -3922,7 +3922,7 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, _PROXY_DynamicRateLimitHandler):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
if internal_usage_cache is None:
|
||||
raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}")
|
||||
|
|
@ -3932,7 +3932,7 @@ def _init_custom_logger_compatible_class(
|
|||
if llm_router is not None and isinstance(llm_router, litellm.Router):
|
||||
dynamic_rate_limiter_obj.update_variables(llm_router=llm_router)
|
||||
_in_memory_loggers.append(dynamic_rate_limiter_obj)
|
||||
return dynamic_rate_limiter_obj # type: ignore
|
||||
return dynamic_rate_limiter_obj
|
||||
elif logging_integration == "dynamic_rate_limiter_v3":
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
|
||||
_PROXY_DynamicRateLimitHandlerV3,
|
||||
|
|
@ -3940,7 +3940,7 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
if internal_usage_cache is None:
|
||||
raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}")
|
||||
|
|
@ -3950,13 +3950,13 @@ def _init_custom_logger_compatible_class(
|
|||
if llm_router is not None and isinstance(llm_router, litellm.Router):
|
||||
dynamic_rate_limiter_obj_v3.update_variables(llm_router=llm_router)
|
||||
_in_memory_loggers.append(dynamic_rate_limiter_obj_v3)
|
||||
return dynamic_rate_limiter_obj_v3 # type: ignore
|
||||
return dynamic_rate_limiter_obj_v3
|
||||
elif logging_integration == "langtrace":
|
||||
if "LANGTRACE_API_KEY" not in os.environ:
|
||||
raise ValueError("LANGTRACE_API_KEY not found in environment variables")
|
||||
_v2 = _maybe_construct_otel_v2("langtrace", _in_memory_loggers)
|
||||
if _v2 is not None:
|
||||
return _v2 # type: ignore
|
||||
return _v2
|
||||
|
||||
from litellm.integrations.opentelemetry import (
|
||||
OpenTelemetry,
|
||||
|
|
@ -3970,19 +3970,19 @@ def _init_custom_logger_compatible_class(
|
|||
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = f"api_key={os.getenv('LANGTRACE_API_KEY')}"
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, OpenTelemetry) and callback.callback_name == "langtrace":
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
_otel_logger = OpenTelemetry(config=otel_config, callback_name="langtrace")
|
||||
_in_memory_loggers.append(_otel_logger)
|
||||
return _otel_logger # type: ignore
|
||||
return _otel_logger
|
||||
|
||||
elif logging_integration == "mlflow":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, MlflowLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
_mlflow_logger: Final = MlflowLogger()
|
||||
_in_memory_loggers.append(_mlflow_logger)
|
||||
return _mlflow_logger # type: ignore
|
||||
return _mlflow_logger
|
||||
elif logging_integration == "langfuse":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LangfusePromptManagement):
|
||||
|
|
@ -3990,25 +3990,25 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
langfuse_logger: Final = LangfusePromptManagement()
|
||||
_in_memory_loggers.append(langfuse_logger)
|
||||
return langfuse_logger # type: ignore
|
||||
return langfuse_logger
|
||||
elif logging_integration == "langfuse_otel":
|
||||
_v2 = _maybe_construct_otel_v2("langfuse_otel", _in_memory_loggers)
|
||||
if _v2 is not None:
|
||||
return _v2 # type: ignore
|
||||
return _v2
|
||||
from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, LangfuseOtelLogger) and callback.callback_name == "langfuse_otel":
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
# Allow LangfuseOtelLogger to initialize its own config safely
|
||||
# This prevents startup crashes if LANGFUSE keys are not in env (e.g. for dynamic usage)
|
||||
_otel_logger = LangfuseOtelLogger(config=None, callback_name="langfuse_otel")
|
||||
_in_memory_loggers.append(_otel_logger)
|
||||
return _otel_logger # type: ignore
|
||||
return _otel_logger
|
||||
elif logging_integration == "weave_otel":
|
||||
_v2 = _maybe_construct_otel_v2("weave_otel", _in_memory_loggers)
|
||||
if _v2 is not None:
|
||||
return _v2 # type: ignore
|
||||
return _v2
|
||||
from litellm.integrations.opentelemetry import OpenTelemetryConfig
|
||||
from litellm.integrations.weave.weave_otel import (
|
||||
WeaveOtelLogger,
|
||||
|
|
@ -4025,24 +4025,24 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, WeaveOtelLogger) and callback.callback_name == "weave_otel":
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
_otel_logger = WeaveOtelLogger(config=otel_config, callback_name="weave_otel")
|
||||
_in_memory_loggers.append(_otel_logger)
|
||||
return _otel_logger # type: ignore
|
||||
return _otel_logger
|
||||
elif logging_integration == "pagerduty":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PagerDutyAlerting):
|
||||
return callback
|
||||
pagerduty_logger: Final = PagerDutyAlerting(**custom_logger_init_args)
|
||||
_in_memory_loggers.append(pagerduty_logger)
|
||||
return pagerduty_logger # type: ignore
|
||||
return pagerduty_logger
|
||||
elif logging_integration == "anthropic_cache_control_hook":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, AnthropicCacheControlHook):
|
||||
return callback
|
||||
anthropic_cache_control_hook: Final = AnthropicCacheControlHook()
|
||||
_in_memory_loggers.append(anthropic_cache_control_hook)
|
||||
return anthropic_cache_control_hook # type: ignore
|
||||
return anthropic_cache_control_hook
|
||||
elif logging_integration == "vector_store_pre_call_hook":
|
||||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
VectorStorePreCallHook,
|
||||
|
|
@ -4053,42 +4053,42 @@ def _init_custom_logger_compatible_class(
|
|||
return callback
|
||||
vector_store_pre_call_hook: Final = VectorStorePreCallHook()
|
||||
_in_memory_loggers.append(vector_store_pre_call_hook)
|
||||
return vector_store_pre_call_hook # type: ignore
|
||||
return vector_store_pre_call_hook
|
||||
elif logging_integration == "gcs_pubsub":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, GcsPubSubLogger):
|
||||
return callback
|
||||
_gcs_pubsub_logger: Final = GcsPubSubLogger()
|
||||
_in_memory_loggers.append(_gcs_pubsub_logger)
|
||||
return _gcs_pubsub_logger # type: ignore
|
||||
return _gcs_pubsub_logger
|
||||
elif logging_integration == "generic_api":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, GenericAPILogger):
|
||||
return callback
|
||||
generic_api_logger: Final = GenericAPILogger()
|
||||
_in_memory_loggers.append(generic_api_logger)
|
||||
return generic_api_logger # type: ignore
|
||||
return generic_api_logger
|
||||
elif logging_integration == "resend_email":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ResendEmailLogger):
|
||||
return callback
|
||||
resend_email_logger: Final = ResendEmailLogger()
|
||||
_in_memory_loggers.append(resend_email_logger)
|
||||
return resend_email_logger # type: ignore
|
||||
return resend_email_logger
|
||||
elif logging_integration == "sendgrid_email":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, SendGridEmailLogger):
|
||||
return callback
|
||||
sendgrid_email_logger: Final = SendGridEmailLogger()
|
||||
_in_memory_loggers.append(sendgrid_email_logger)
|
||||
return sendgrid_email_logger # type: ignore
|
||||
return sendgrid_email_logger
|
||||
elif logging_integration == "smtp_email":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, SMTPEmailLogger):
|
||||
return callback
|
||||
smtp_email_logger: Final = SMTPEmailLogger()
|
||||
_in_memory_loggers.append(smtp_email_logger)
|
||||
return smtp_email_logger # type: ignore
|
||||
return smtp_email_logger
|
||||
elif logging_integration == "humanloop":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, HumanloopLogger):
|
||||
|
|
@ -4096,7 +4096,7 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
humanloop_logger: Final = HumanloopLogger()
|
||||
_in_memory_loggers.append(humanloop_logger)
|
||||
return humanloop_logger # type: ignore
|
||||
return humanloop_logger
|
||||
elif logging_integration == "dotprompt":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, DotpromptManager):
|
||||
|
|
@ -4104,7 +4104,7 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
dotprompt_logger: Final = DotpromptManager()
|
||||
_in_memory_loggers.append(dotprompt_logger)
|
||||
return dotprompt_logger # type: ignore
|
||||
return dotprompt_logger
|
||||
elif logging_integration == "bitbucket":
|
||||
from litellm.integrations.bitbucket.bitbucket_prompt_manager import (
|
||||
BitBucketPromptManager,
|
||||
|
|
@ -4121,7 +4121,7 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
bitbucket_logger: Final = BitBucketPromptManager(bitbucket_config=bitbucket_config)
|
||||
_in_memory_loggers.append(bitbucket_logger)
|
||||
return bitbucket_logger # type: ignore
|
||||
return bitbucket_logger
|
||||
elif logging_integration == "gitlab":
|
||||
from litellm.integrations.gitlab.gitlab_prompt_manager import (
|
||||
GitLabPromptManager,
|
||||
|
|
@ -4138,14 +4138,14 @@ def _init_custom_logger_compatible_class(
|
|||
|
||||
gitlab_logger: Final = GitLabPromptManager(gitlab_config=gitlab_config)
|
||||
_in_memory_loggers.append(gitlab_logger)
|
||||
return gitlab_logger # type: ignore
|
||||
return gitlab_logger
|
||||
elif logging_integration == "newrelic":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, NewRelicLogger):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
newrelic_logger: Final = NewRelicLogger()
|
||||
_in_memory_loggers.append(newrelic_logger)
|
||||
return newrelic_logger # type: ignore
|
||||
return newrelic_logger
|
||||
return None
|
||||
except Exception as e:
|
||||
verbose_logger.exception("[Non-Blocking Error] Error initializing custom logger: %s", e)
|
||||
|
|
@ -4322,7 +4322,7 @@ def get_custom_logger_compatible_class(
|
|||
return callback
|
||||
_aws_sqs_logger: Final = SQSLogger()
|
||||
_in_memory_loggers.append(_aws_sqs_logger)
|
||||
return _aws_sqs_logger # type: ignore
|
||||
return _aws_sqs_logger
|
||||
elif logging_integration == "azure_storage":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, AzureBlobStorageLogger):
|
||||
|
|
@ -4356,7 +4356,7 @@ def get_custom_logger_compatible_class(
|
|||
for callback in _in_memory_loggers:
|
||||
# Use exact type check to avoid matching ArizePhoenixLogger (subclass)
|
||||
if type(callback) is OpenTelemetry:
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
elif logging_integration == "dynamic_rate_limiter":
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter import (
|
||||
|
|
@ -4365,7 +4365,7 @@ def get_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, _PROXY_DynamicRateLimitHandler):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
elif logging_integration == "dynamic_rate_limiter_v3":
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
|
||||
_PROXY_DynamicRateLimitHandlerV3,
|
||||
|
|
@ -4373,7 +4373,7 @@ def get_custom_logger_compatible_class(
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3):
|
||||
return callback # type: ignore
|
||||
return callback
|
||||
|
||||
elif logging_integration == "langtrace":
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
|
|
@ -4663,7 +4663,7 @@ class StandardLoggingPayloadSetup:
|
|||
)
|
||||
if isinstance(metadata, dict):
|
||||
for key in metadata.keys() & _STANDARD_LOGGING_METADATA_KEYS:
|
||||
clean_metadata[key] = metadata[key] # type: ignore
|
||||
clean_metadata[key] = metadata[key]
|
||||
|
||||
user_api_key: Final = metadata.get("user_api_key")
|
||||
if user_api_key and isinstance(user_api_key, str) and is_valid_sha256_hash(user_api_key):
|
||||
|
|
@ -4763,7 +4763,7 @@ class StandardLoggingPayloadSetup:
|
|||
) -> StandardLoggingModelInformation:
|
||||
model_cost_name: Final = _select_model_name_for_cost_calc(
|
||||
model=base_model if custom_pricing else None,
|
||||
completion_response=init_response_obj, # type: ignore
|
||||
completion_response=init_response_obj,
|
||||
base_model=base_model,
|
||||
custom_pricing=custom_pricing,
|
||||
)
|
||||
|
|
@ -4832,14 +4832,14 @@ class StandardLoggingPayloadSetup:
|
|||
typed_keys[_key] = key
|
||||
if _key in additiona_headers:
|
||||
try:
|
||||
additional_logging_headers[key] = int(additiona_headers[_key]) # type: ignore
|
||||
additional_logging_headers[key] = int(additiona_headers[_key])
|
||||
except (ValueError, TypeError):
|
||||
additional_logging_headers[key] = additiona_headers[_key] # type: ignore
|
||||
additional_logging_headers[key] = additiona_headers[_key]
|
||||
|
||||
# Preserve all remaining headers verbatim (e.g. llm_provider-x-request-id)
|
||||
for k, v in additiona_headers.items():
|
||||
if k.lower() not in typed_keys:
|
||||
additional_logging_headers[k] = v # type: ignore
|
||||
additional_logging_headers[k] = v
|
||||
|
||||
return additional_logging_headers
|
||||
|
||||
|
|
@ -4866,7 +4866,7 @@ class StandardLoggingPayloadSetup:
|
|||
hidden_params[key]
|
||||
)
|
||||
else:
|
||||
clean_hidden_params[key] = hidden_params[key] # type: ignore
|
||||
clean_hidden_params[key] = hidden_params[key]
|
||||
return clean_hidden_params
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -5310,7 +5310,7 @@ def get_standard_logging_object_payload(
|
|||
saved_cache_cost = (
|
||||
logging_obj._response_cost_calculator(
|
||||
result=init_response_obj,
|
||||
cache_hit=False, # type: ignore
|
||||
cache_hit=False,
|
||||
)
|
||||
or 0.0
|
||||
)
|
||||
|
|
@ -5503,7 +5503,7 @@ def get_standard_logging_metadata(
|
|||
# Update the clean_metadata with values from input metadata that match StandardLoggingMetadata fields
|
||||
for key in StandardLoggingMetadata.__annotations__.keys():
|
||||
if key in metadata:
|
||||
clean_metadata[key] = metadata[key] # type: ignore
|
||||
clean_metadata[key] = metadata[key]
|
||||
|
||||
if metadata.get("user_api_key") is not None:
|
||||
if is_valid_sha256_hash(str(metadata.get("user_api_key"))):
|
||||
|
|
@ -5555,7 +5555,7 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
|
|||
# First create the nested objects with proper typing
|
||||
model_info: Final = StandardLoggingModelInformation(model_map_key="gpt-3.5-turbo", model_map_value=None)
|
||||
|
||||
metadata: Final = StandardLoggingMetadata( # type: ignore
|
||||
metadata: Final = StandardLoggingMetadata(
|
||||
user_api_key_hash="test_hash",
|
||||
user_api_key_alias="test_alias",
|
||||
user_api_key_team_id="test_team",
|
||||
|
|
@ -5596,7 +5596,7 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
|
|||
response: Final[dict[str, list[dict[str, dict[str, str]]]]] = {"choices": [{"message": {"content": "Hi there!"}}]}
|
||||
|
||||
# Main payload initialization
|
||||
return StandardLoggingPayload( # type: ignore
|
||||
return StandardLoggingPayload(
|
||||
id="test_id",
|
||||
call_type="completion",
|
||||
stream=False,
|
||||
|
|
|
|||
|
|
@ -188,13 +188,13 @@ class StandardBuiltInToolCostTracking:
|
|||
|
||||
if storage_gb_val is not None:
|
||||
try:
|
||||
storage_gb = float(storage_gb_val) # type: ignore
|
||||
storage_gb = float(storage_gb_val)
|
||||
except (TypeError, ValueError):
|
||||
storage_gb = None
|
||||
|
||||
if days_val is not None:
|
||||
try:
|
||||
days = float(days_val) # type: ignore
|
||||
days = float(days_val)
|
||||
except (TypeError, ValueError):
|
||||
days = None
|
||||
|
||||
|
|
@ -286,7 +286,7 @@ class StandardBuiltInToolCostTracking:
|
|||
"""Safely convert a value to int."""
|
||||
if value is not None:
|
||||
try:
|
||||
return int(value) # type: ignore
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -146,7 +146,7 @@ def _clear_later_replay_slice_metadata(choice: StreamingChoices) -> None:
|
|||
# streamed deltas collect it once per slice, and a field added to Delta
|
||||
# later can't silently re-introduce the duplication.
|
||||
choice.delta = Delta(content=choice.delta.content)
|
||||
choice.logprobs = None # type: ignore[assignment]
|
||||
choice.logprobs = None
|
||||
if hasattr(choice, "enhancements"):
|
||||
del choice.enhancements
|
||||
|
||||
|
|
@ -270,9 +270,7 @@ async def convert_to_streaming_response_async(
|
|||
slice_chunk.choices[0].delta.content = piece
|
||||
if i > 0:
|
||||
_clear_later_replay_slice_metadata(slice_chunk.choices[0])
|
||||
slice_chunk.choices[0].finish_reason = (
|
||||
original_finish_reason if i == last_idx else None # type: ignore[assignment]
|
||||
)
|
||||
slice_chunk.choices[0].finish_reason = original_finish_reason if i == last_idx else None
|
||||
if i == last_idx and original_usage is not None:
|
||||
setattr(slice_chunk, "usage", original_usage)
|
||||
yield slice_chunk
|
||||
|
|
@ -322,9 +320,9 @@ def convert_to_streaming_response(
|
|||
|
||||
if "usage" in response_object and response_object["usage"] is not None:
|
||||
setattr(model_response_object, "usage", Usage())
|
||||
model_response_object.usage.completion_tokens = response_object["usage"].get("completion_tokens", 0) # type: ignore
|
||||
model_response_object.usage.prompt_tokens = response_object["usage"].get("prompt_tokens", 0) # type: ignore
|
||||
model_response_object.usage.total_tokens = response_object["usage"].get("total_tokens", 0) # type: ignore
|
||||
model_response_object.usage.completion_tokens = response_object["usage"].get("completion_tokens", 0)
|
||||
model_response_object.usage.prompt_tokens = response_object["usage"].get("prompt_tokens", 0)
|
||||
model_response_object.usage.total_tokens = response_object["usage"].get("total_tokens", 0)
|
||||
|
||||
if "id" in response_object:
|
||||
model_response_object.id = response_object["id"]
|
||||
|
|
@ -358,9 +356,7 @@ def convert_to_streaming_response(
|
|||
slice_chunk.choices[0].delta.content = piece
|
||||
if i > 0:
|
||||
_clear_later_replay_slice_metadata(slice_chunk.choices[0])
|
||||
slice_chunk.choices[0].finish_reason = (
|
||||
original_finish_reason if i == last_idx else None # type: ignore[assignment]
|
||||
)
|
||||
slice_chunk.choices[0].finish_reason = original_finish_reason if i == last_idx else None
|
||||
if i == last_idx and original_usage is not None:
|
||||
setattr(slice_chunk, "usage", original_usage)
|
||||
yield slice_chunk
|
||||
|
|
@ -715,7 +711,7 @@ def convert_to_model_response_object(
|
|||
provider_specific_fields=provider_specific_fields,
|
||||
)
|
||||
choice_list.append(choice)
|
||||
model_response_object.choices = choice_list # type: ignore
|
||||
model_response_object.choices = choice_list
|
||||
|
||||
if "usage" in response_object and response_object["usage"] is not None:
|
||||
usage_object: Final = litellm.Usage(**response_object["usage"])
|
||||
|
|
@ -740,9 +736,7 @@ def convert_to_model_response_object(
|
|||
|
||||
if start_time is not None and end_time is not None:
|
||||
if isinstance(start_time, type(end_time)):
|
||||
model_response_object._response_ms = ( # type: ignore
|
||||
end_time - start_time
|
||||
).total_seconds() * 1000
|
||||
model_response_object._response_ms = (end_time - start_time).total_seconds() * 1000
|
||||
|
||||
if hidden_params is not None:
|
||||
if model_response_object._hidden_params is None:
|
||||
|
|
@ -775,12 +769,12 @@ def convert_to_model_response_object(
|
|||
model_response_object.data = response_object["data"]
|
||||
|
||||
if "usage" in response_object and response_object["usage"] is not None:
|
||||
model_response_object.usage.completion_tokens = response_object["usage"].get("completion_tokens", 0) # type: ignore
|
||||
model_response_object.usage.prompt_tokens = response_object["usage"].get("prompt_tokens", 0) # type: ignore
|
||||
model_response_object.usage.total_tokens = response_object["usage"].get("total_tokens", 0) # type: ignore
|
||||
model_response_object.usage.completion_tokens = response_object["usage"].get("completion_tokens", 0)
|
||||
model_response_object.usage.prompt_tokens = response_object["usage"].get("prompt_tokens", 0)
|
||||
model_response_object.usage.total_tokens = response_object["usage"].get("total_tokens", 0)
|
||||
|
||||
if start_time is not None and end_time is not None:
|
||||
model_response_object._response_ms = ( # type: ignore
|
||||
model_response_object._response_ms = (
|
||||
end_time - start_time
|
||||
).total_seconds() * 1000 # return response latency in ms like openai
|
||||
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ class LoggingCallbackManager:
|
|||
"""
|
||||
self._safe_add_callback_to_list(
|
||||
callback=callback,
|
||||
parent_list=litellm.callbacks, # type: ignore
|
||||
parent_list=litellm.callbacks,
|
||||
)
|
||||
|
||||
def add_litellm_success_callback(self, callback: CustomLogger | str | Callable):
|
||||
|
|
|
|||
|
|
@ -99,7 +99,7 @@ def strip_name_from_message(message: AllMessageValues, allowed_name_roles: list[
|
|||
"""
|
||||
msg_copy: Final = message.copy()
|
||||
if msg_copy.get("role") not in allowed_name_roles:
|
||||
msg_copy.pop("name", None) # type: ignore
|
||||
msg_copy.pop("name", None)
|
||||
return msg_copy
|
||||
|
||||
|
||||
|
|
@ -114,7 +114,7 @@ def strip_name_from_messages(
|
|||
msg_role = message.get("role")
|
||||
msg_copy = message.copy()
|
||||
if msg_role not in allowed_name_roles:
|
||||
msg_copy.pop("name", None) # type: ignore
|
||||
msg_copy.pop("name", None)
|
||||
new_messages.append(msg_copy)
|
||||
return new_messages
|
||||
|
||||
|
|
@ -1511,9 +1511,7 @@ def convert_prefix_message_to_non_prefix_messages(
|
|||
"content": "You are a helpful assistant. You are given a message and you need to respond to it. You are also given a generated content. You need to respond to the message in continuation of the generated content. Do not repeat the same content. Your response should be in continuation of this text: ",
|
||||
}
|
||||
)
|
||||
new_messages.append(
|
||||
{**{k: v for k, v in message.items() if k != "prefix"}} # type: ignore
|
||||
)
|
||||
new_messages.append({**{k: v for k, v in message.items() if k != "prefix"}})
|
||||
else:
|
||||
new_messages.append(message)
|
||||
return new_messages
|
||||
|
|
|
|||
|
|
@ -380,7 +380,7 @@ def _render_chat_template(env, chat_template: str, bos_token: str, eos_token: st
|
|||
Rendered template string
|
||||
"""
|
||||
try:
|
||||
template: Final = env.from_string(chat_template) # type: ignore
|
||||
template: Final = env.from_string(chat_template)
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
|
@ -471,7 +471,7 @@ async def _afetch_and_extract_template(
|
|||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
and "chat_template" in tokenizer_config["tokenizer"]
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"]
|
||||
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
|
||||
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
|
||||
chat_template = tokenizer_data["chat_template"]
|
||||
|
|
@ -486,13 +486,13 @@ async def _afetch_and_extract_template(
|
|||
and "tokenizer" in tokenizer_config
|
||||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"]
|
||||
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
|
||||
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
|
||||
else:
|
||||
raise Exception("No chat template found")
|
||||
|
||||
return chat_template, bos_token, eos_token # type: ignore
|
||||
return chat_template, bos_token, eos_token
|
||||
|
||||
|
||||
def _fetch_and_extract_template(
|
||||
|
|
@ -525,7 +525,7 @@ def _fetch_and_extract_template(
|
|||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
and "chat_template" in tokenizer_config["tokenizer"]
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"]
|
||||
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
|
||||
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
|
||||
chat_template = tokenizer_data["chat_template"]
|
||||
|
|
@ -540,13 +540,13 @@ def _fetch_and_extract_template(
|
|||
and "tokenizer" in tokenizer_config
|
||||
and isinstance(tokenizer_config["tokenizer"], dict)
|
||||
):
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"] # type: ignore
|
||||
tokenizer_data: dict = tokenizer_config["tokenizer"]
|
||||
bos_token = _extract_token_value(token_value=tokenizer_data.get("bos_token"))
|
||||
eos_token = _extract_token_value(token_value=tokenizer_data.get("eos_token"))
|
||||
else:
|
||||
raise Exception("No chat template found")
|
||||
|
||||
return chat_template, bos_token, eos_token # type: ignore
|
||||
return chat_template, bos_token, eos_token
|
||||
|
||||
|
||||
async def ahf_chat_template(model: str, messages: list, chat_template: Any | None = None):
|
||||
|
|
@ -1067,9 +1067,7 @@ def anthropic_messages_pt_xml(messages: list):
|
|||
while msg_i < len(messages) and messages[msg_i]["role"] == "assistant":
|
||||
assistant_text = messages[msg_i].get("content") or "" # either string or none
|
||||
if messages[msg_i].get("tool_calls", []): # support assistant tool invoke conversion
|
||||
assistant_text += convert_to_anthropic_tool_invoke_xml( # type: ignore
|
||||
messages[msg_i]["tool_calls"]
|
||||
)
|
||||
assistant_text += convert_to_anthropic_tool_invoke_xml(messages[msg_i]["tool_calls"])
|
||||
|
||||
assistant_content.append({"type": "text", "text": assistant_text})
|
||||
msg_i += 1
|
||||
|
|
@ -1124,7 +1122,7 @@ def convert_to_azure_openai_messages(
|
|||
if m["role"] == "user" and isinstance(m.get("content"), list):
|
||||
for content in m.get("content", []):
|
||||
if isinstance(content, dict) and content.get("type") == "image_url":
|
||||
_azure_image_url_helper(content) # type: ignore
|
||||
_azure_image_url_helper(content)
|
||||
return messages
|
||||
|
||||
|
||||
|
|
@ -1475,7 +1473,7 @@ def convert_to_gemini_tool_call_result(
|
|||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to process file in tool response: %s", e)
|
||||
name: str | None = message.get("name", "") # type: ignore
|
||||
name: str | None = message.get("name", "")
|
||||
|
||||
# Recover name from last message with tool calls
|
||||
if last_message_with_tool_calls:
|
||||
|
|
@ -1521,7 +1519,7 @@ def convert_to_gemini_tool_call_result(
|
|||
# error call result so default to the successful result template
|
||||
_function_response: Final = VertexFunctionResponse(
|
||||
name=name,
|
||||
response=response_data, # type: ignore
|
||||
response=response_data,
|
||||
)
|
||||
if gemini_call_id:
|
||||
_function_response["id"] = gemini_call_id
|
||||
|
|
@ -1693,7 +1691,7 @@ def convert_to_anthropic_tool_result(
|
|||
if anthropic_tool_result is None:
|
||||
raise Exception(f"Unable to parse anthropic tool result for message: {message}")
|
||||
if cache_control is not None:
|
||||
anthropic_tool_result["cache_control"] = cache_control # type: ignore
|
||||
anthropic_tool_result["cache_control"] = cache_control
|
||||
return anthropic_tool_result
|
||||
|
||||
|
||||
|
|
@ -1841,7 +1839,7 @@ def add_cache_control_to_content(
|
|||
):
|
||||
cache_control_param: Final = original_content_element.get("cache_control")
|
||||
if cache_control_param is not None and isinstance(cache_control_param, dict):
|
||||
transformed_param: Final = ChatCompletionCachedContent(**cache_control_param) # type: ignore
|
||||
transformed_param: Final = ChatCompletionCachedContent(**cache_control_param)
|
||||
|
||||
anthropic_content_element["cache_control"] = transformed_param
|
||||
|
||||
|
|
@ -2020,7 +2018,7 @@ def _sanitize_empty_text_content(
|
|||
|
||||
if rewrote_any:
|
||||
message = cast(AllMessageValues, dict(message)) # Make a copy
|
||||
message["content"] = new_blocks # type: ignore
|
||||
message["content"] = new_blocks
|
||||
verbose_logger.debug(
|
||||
"_sanitize_empty_text_content: Replaced empty text block(s) in %s message", message.get("role")
|
||||
)
|
||||
|
|
@ -2396,12 +2394,12 @@ def anthropic_messages_pt(
|
|||
user_content: list[AnthropicMessagesUserMessageValues] = []
|
||||
init_msg_i = msg_i
|
||||
if isinstance(messages[msg_i], BaseModel):
|
||||
messages[msg_i] = dict(messages[msg_i]) # type: ignore
|
||||
messages[msg_i] = dict(messages[msg_i])
|
||||
## MERGE CONSECUTIVE USER CONTENT ##
|
||||
while msg_i < len(messages) and messages[msg_i]["role"] in user_message_types:
|
||||
user_message_types_block: (
|
||||
ChatCompletionToolMessage | ChatCompletionUserMessage | ChatCompletionFunctionMessage
|
||||
) = messages[msg_i] # type: ignore
|
||||
) = messages[msg_i]
|
||||
if user_message_types_block["role"] == "user":
|
||||
if isinstance(user_message_types_block["content"], list):
|
||||
for m in user_message_types_block["content"]:
|
||||
|
|
@ -2507,7 +2505,7 @@ def anthropic_messages_pt(
|
|||
assistant_content: list[AnthropicMessagesAssistantMessageValues] = []
|
||||
## MERGE CONSECUTIVE ASSISTANT CONTENT ##
|
||||
while msg_i < len(messages) and messages[msg_i]["role"] == "assistant":
|
||||
assistant_content_block: ChatCompletionAssistantMessage = messages[msg_i] # type: ignore
|
||||
assistant_content_block: ChatCompletionAssistantMessage = messages[msg_i]
|
||||
|
||||
# Extract compaction_blocks from provider_specific_fields and add them first
|
||||
_provider_specific_fields_raw = assistant_content_block.get("provider_specific_fields")
|
||||
|
|
@ -2515,7 +2513,7 @@ def anthropic_messages_pt(
|
|||
_compaction_blocks = _provider_specific_fields_raw.get("compaction_blocks")
|
||||
if _compaction_blocks and isinstance(_compaction_blocks, list):
|
||||
# Add compaction blocks at the beginning of assistant content : https://platform.claude.com/docs/en/build-with-claude/compaction
|
||||
assistant_content.extend(_compaction_blocks) # type: ignore
|
||||
assistant_content.extend(_compaction_blocks)
|
||||
|
||||
_raw_thinking_blocks = assistant_content_block.get("thinking_blocks", None)
|
||||
thinking_blocks = (
|
||||
|
|
@ -2555,7 +2553,7 @@ def anthropic_messages_pt(
|
|||
_web_search_results_tc = _provider_specific_fields_tc.get("web_search_results")
|
||||
_tool_results_tc = _provider_specific_fields_tc.get("tool_results")
|
||||
tool_invoke_results = convert_to_anthropic_tool_invoke(
|
||||
assistant_tool_calls, # type: ignore
|
||||
assistant_tool_calls,
|
||||
web_search_results=_web_search_results_tc,
|
||||
tool_results=_tool_results_tc,
|
||||
)
|
||||
|
|
@ -2706,7 +2704,7 @@ def anthropic_messages_pt(
|
|||
# handle server_tool_use blocks (tool search, web search, etc.)
|
||||
# Pass through as-is since these are Anthropic-native content types
|
||||
elif m.get("type", "") == "server_tool_use" or m.get("type", "").endswith("_tool_result"):
|
||||
assistant_content.append(m) # type: ignore
|
||||
assistant_content.append(m)
|
||||
elif (
|
||||
"content" in assistant_content_block
|
||||
and isinstance(assistant_content_block["content"], str)
|
||||
|
|
@ -2834,10 +2832,10 @@ def parse_xml_params(xml_content, json_schema: dict | None = None):
|
|||
if child is not None and child.text is not None:
|
||||
try:
|
||||
# Attempt to decode the element's text as JSON
|
||||
params[child.tag] = json.loads(child.text) # type: ignore
|
||||
params[child.tag] = json.loads(child.text)
|
||||
except json.JSONDecodeError:
|
||||
# If JSON decoding fails, use the original text
|
||||
params[child.tag] = child.text # type: ignore
|
||||
params[child.tag] = child.text
|
||||
|
||||
return params
|
||||
|
||||
|
|
@ -3282,7 +3280,7 @@ def gemini_text_image_pt(messages: list):
|
|||
}
|
||||
"""
|
||||
try:
|
||||
pass # type: ignore
|
||||
pass
|
||||
except Exception:
|
||||
raise Exception("Importing google.generativeai failed, please run 'pip install -q google-generativeai")
|
||||
|
||||
|
|
@ -4063,9 +4061,7 @@ def get_user_message_block_or_continue_message(
|
|||
if content_block.strip():
|
||||
return message
|
||||
else:
|
||||
return ChatCompletionUserMessage(
|
||||
**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE) # type: ignore
|
||||
)
|
||||
return ChatCompletionUserMessage(**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE))
|
||||
|
||||
# Handle list case
|
||||
if isinstance(content_block, list):
|
||||
|
|
@ -4079,9 +4075,7 @@ def get_user_message_block_or_continue_message(
|
|||
],
|
||||
"""
|
||||
if not content_block:
|
||||
return ChatCompletionUserMessage(
|
||||
**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE) # type: ignore
|
||||
)
|
||||
return ChatCompletionUserMessage(**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE))
|
||||
# Create a copy of the message to avoid modifying the original
|
||||
modified_content_block: Final = content_block.copy()
|
||||
|
||||
|
|
@ -4091,7 +4085,7 @@ def get_user_message_block_or_continue_message(
|
|||
if not item["text"].strip():
|
||||
# Replace empty text with continue message
|
||||
_user_continue_message = ChatCompletionUserMessage(
|
||||
**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE) # type: ignore
|
||||
**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE)
|
||||
)
|
||||
text = convert_content_list_to_str(_user_continue_message)
|
||||
item["text"] = text
|
||||
|
|
@ -4178,14 +4172,12 @@ def skip_empty_text_blocks(
|
|||
|
||||
# Type-specific casting based on message role
|
||||
if message["role"] == "assistant":
|
||||
modified_message_alt["content"] = cast( # type: ignore
|
||||
modified_message_alt["content"] = cast(
|
||||
list[OpenAIMessageContentListBlock] | None,
|
||||
modified_content_block or None,
|
||||
)
|
||||
elif message["role"] == "user" and modified_content_block is not None:
|
||||
modified_message_alt["content"] = cast( # type: ignore
|
||||
list[ChatCompletionTextObject] | None, modified_content_block
|
||||
)
|
||||
modified_message_alt["content"] = cast(list[ChatCompletionTextObject] | None, modified_content_block)
|
||||
|
||||
return modified_message_alt
|
||||
|
||||
|
|
@ -4356,10 +4348,10 @@ class BedrockConverseMessagesProcessor:
|
|||
format = element["image_url"].get("format")
|
||||
else:
|
||||
image_url = element["image_url"]
|
||||
_part = await BedrockImageProcessor.process_image_async( # type: ignore
|
||||
_part = await BedrockImageProcessor.process_image_async(
|
||||
image_url=image_url, format=format
|
||||
)
|
||||
_parts.append(_part) # type: ignore
|
||||
_parts.append(_part)
|
||||
elif element["type"] == "file":
|
||||
_part = await BedrockConverseMessagesProcessor._async_process_file_message(
|
||||
message=cast(ChatCompletionFileObject, element)
|
||||
|
|
@ -4501,9 +4493,7 @@ class BedrockConverseMessagesProcessor:
|
|||
image_url = element["image_url"]["url"]
|
||||
else:
|
||||
image_url = element["image_url"]
|
||||
assistants_part = await BedrockImageProcessor.process_image_async( # type: ignore
|
||||
image_url=image_url
|
||||
)
|
||||
assistants_part = await BedrockImageProcessor.process_image_async(image_url=image_url)
|
||||
assistants_parts.append(assistants_part)
|
||||
# Add cache point block for assistant content elements
|
||||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
|
|
@ -4730,11 +4720,11 @@ def _bedrock_converse_messages_pt(
|
|||
format = element["image_url"].get("format")
|
||||
else:
|
||||
image_url = element["image_url"]
|
||||
_part = BedrockImageProcessor.process_image_sync( # type: ignore
|
||||
_part = BedrockImageProcessor.process_image_sync(
|
||||
image_url=image_url,
|
||||
format=format,
|
||||
)
|
||||
_parts.append(_part) # type: ignore
|
||||
_parts.append(_part)
|
||||
elif element["type"] == "file":
|
||||
_part = BedrockConverseMessagesProcessor._process_file_message(
|
||||
message=cast(ChatCompletionFileObject, element)
|
||||
|
|
@ -4881,9 +4871,7 @@ def _bedrock_converse_messages_pt(
|
|||
image_url = element["image_url"]["url"]
|
||||
else:
|
||||
image_url = element["image_url"]
|
||||
assistants_part = BedrockImageProcessor.process_image_sync( # type: ignore
|
||||
image_url=image_url
|
||||
)
|
||||
assistants_part = BedrockImageProcessor.process_image_sync(image_url=image_url)
|
||||
assistants_parts.append(assistants_part)
|
||||
# Add cache point block for assistant content elements
|
||||
_cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block(
|
||||
|
|
@ -5060,7 +5048,7 @@ def _bedrock_tools_pt(tools: list, model: str | None = None) -> list[BedrockTool
|
|||
# Check if tool is already a BedrockToolBlock (e.g., systemTool for Nova grounding)
|
||||
if _is_bedrock_tool_block(tool):
|
||||
# Already a BedrockToolBlock, pass it through
|
||||
tool_block_list.append(tool) # type: ignore
|
||||
tool_block_list.append(tool)
|
||||
continue
|
||||
|
||||
# Responses built-in tools (web_search, image_generation, namespace, tool_search,
|
||||
|
|
@ -5539,8 +5527,5 @@ def resolve_structured_messages(
|
|||
for handler in handlers_to_try:
|
||||
structured = handler.get_structured_messages(request_kwargs)
|
||||
if structured:
|
||||
return [
|
||||
msg if isinstance(msg, dict) else msg.model_dump() # type: ignore
|
||||
for msg in structured
|
||||
]
|
||||
return [msg if isinstance(msg, dict) else msg.model_dump() for msg in structured]
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -173,13 +173,13 @@ class RealTimeStreaming:
|
|||
try:
|
||||
event_type: Final = message_obj.get("type", "")
|
||||
if event_type in self._SESSION_EVENT_TYPES:
|
||||
typed_obj: OpenAIRealtimeEvents = OpenAIRealtimeStreamSessionEvents(**message_obj) # type: ignore
|
||||
typed_obj: OpenAIRealtimeEvents = OpenAIRealtimeStreamSessionEvents(**message_obj)
|
||||
else:
|
||||
# Catch-all base object so unknown/new event names never raise.
|
||||
typed_obj = OpenAIRealtimeStreamResponseBaseObject(**message_obj) # type: ignore
|
||||
typed_obj = OpenAIRealtimeStreamResponseBaseObject(**message_obj)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error parsing message for logging: %s", e)
|
||||
self.messages.append(message_obj) # type: ignore[arg-type]
|
||||
self.messages.append(message_obj)
|
||||
return
|
||||
self.messages.append(typed_obj)
|
||||
|
||||
|
|
@ -346,7 +346,7 @@ class RealTimeStreaming:
|
|||
verbose_logger.debug("Dropping follow-up setup after content was already sent to backend")
|
||||
continue
|
||||
msg = self._maybe_inject_guardrail_auto_response_disable(msg)
|
||||
await self.backend_ws.send(msg) # type: ignore[union-attr, attr-defined]
|
||||
await self.backend_ws.send(msg)
|
||||
self._cache_session_configuration_request(msg)
|
||||
sent = True
|
||||
else:
|
||||
|
|
@ -357,13 +357,13 @@ class RealTimeStreaming:
|
|||
# content before send would leave the session believing the
|
||||
# backend received a setup/content frame it never got, causing
|
||||
# subsequent client session.update messages to be dropped.
|
||||
await self.backend_ws.send(msg) # type: ignore[union-attr, attr-defined]
|
||||
await self.backend_ws.send(msg)
|
||||
self._cache_session_configuration_request(msg)
|
||||
if is_content_message:
|
||||
self._content_sent_after_setup = True
|
||||
sent = True
|
||||
return sent
|
||||
await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined]
|
||||
await self.backend_ws.send(message)
|
||||
return True
|
||||
|
||||
def _enforce_transcription_session_model(self, message: str) -> str:
|
||||
|
|
@ -816,7 +816,7 @@ class RealTimeStreaming:
|
|||
"[realtime guardrail] ending session after violation %d",
|
||||
self._violation_count,
|
||||
)
|
||||
await self.backend_ws.close() # type: ignore[union-attr, attr-defined]
|
||||
await self.backend_ws.close()
|
||||
|
||||
verbose_logger.warning(
|
||||
"[realtime guardrail] BLOCKED transcript (violation %d): %r",
|
||||
|
|
@ -828,7 +828,7 @@ class RealTimeStreaming:
|
|||
|
||||
async def _handle_provider_config_message(self, raw_response) -> None:
|
||||
"""Process a backend message when a provider_config is set (transformed path)."""
|
||||
returned_object: Final = self.provider_config.transform_realtime_response( # type: ignore[union-attr]
|
||||
returned_object: Final = self.provider_config.transform_realtime_response(
|
||||
raw_response,
|
||||
self.model,
|
||||
self.logging_obj,
|
||||
|
|
@ -964,11 +964,9 @@ class RealTimeStreaming:
|
|||
try:
|
||||
while True:
|
||||
try:
|
||||
raw_response = await self.backend_ws.recv( # type: ignore[union-attr]
|
||||
decode=False
|
||||
)
|
||||
raw_response = await self.backend_ws.recv(decode=False)
|
||||
except TypeError:
|
||||
raw_response = await self.backend_ws.recv() # type: ignore[union-attr, assignment]
|
||||
raw_response = await self.backend_ws.recv()
|
||||
|
||||
if isinstance(raw_response, bytes):
|
||||
try:
|
||||
|
|
@ -1007,7 +1005,7 @@ class RealTimeStreaming:
|
|||
continue
|
||||
await self.websocket.send_text(json.dumps(translated))
|
||||
|
||||
except websockets.exceptions.ConnectionClosed as e: # type: ignore
|
||||
except websockets.exceptions.ConnectionClosed as e:
|
||||
verbose_logger.exception("Connection closed in backend to client send messages - %s", e)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error in backend to client send messages: %s", e)
|
||||
|
|
@ -1410,7 +1408,7 @@ class RealTimeStreaming:
|
|||
forward_task: Final = asyncio.create_task(self.backend_to_client_send_messages())
|
||||
try:
|
||||
await self.client_ack_messages()
|
||||
except self.websocket.exceptions.ConnectionClosed: # type: ignore
|
||||
except self.websocket.exceptions.ConnectionClosed:
|
||||
verbose_logger.debug("Connection closed")
|
||||
forward_task.cancel()
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -35,7 +35,7 @@ class Rules:
|
|||
message="LLM Response failed post-call-rule check",
|
||||
llm_provider="",
|
||||
model=model,
|
||||
) # type: ignore
|
||||
)
|
||||
return True
|
||||
|
||||
def post_call_rules(self, input: str | None, model: str) -> bool:
|
||||
|
|
@ -50,10 +50,10 @@ class Rules:
|
|||
message="LLM Response failed post-call-rule check",
|
||||
llm_provider="",
|
||||
model=model,
|
||||
) # type: ignore
|
||||
)
|
||||
elif isinstance(decision, dict):
|
||||
decision_val = decision.get("decision", True)
|
||||
decision_message = decision.get("message", "LLM Response failed post-call-rule check")
|
||||
if decision_val is False:
|
||||
raise litellm.APIResponseValidationError(message=decision_message, llm_provider="", model=model) # type: ignore
|
||||
raise litellm.APIResponseValidationError(message=decision_message, llm_provider="", model=model)
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -893,7 +893,7 @@ class CustomStreamWrapper:
|
|||
for choice in original_chunk.choices:
|
||||
try:
|
||||
if isinstance(choice, BaseModel):
|
||||
choice_json = choice.model_dump() # type: ignore
|
||||
choice_json = choice.model_dump()
|
||||
choice_json.pop(
|
||||
"finish_reason", None
|
||||
) # for mistral etc. which return a value in their last chunk (not-openai compatible).
|
||||
|
|
@ -1050,7 +1050,7 @@ class CustomStreamWrapper:
|
|||
# Strip finish_reason from the content chunk so it appears
|
||||
# only on the trailing empty-delta chunk (OpenAI spec).
|
||||
# finish_reason_handler() will emit the proper terminal chunk.
|
||||
chunk.choices[0].finish_reason = None # type: ignore[assignment]
|
||||
chunk.choices[0].finish_reason = None
|
||||
return _ProviderChunkEarlyReturn(chunk)
|
||||
|
||||
if (
|
||||
|
|
@ -1139,19 +1139,17 @@ class CustomStreamWrapper:
|
|||
self.received_finish_reason = "stop"
|
||||
elif self.custom_llm_provider == "vertex_ai" and not isinstance(chunk, ModelResponseStream):
|
||||
chunk = cast(Any, chunk)
|
||||
import proto # type: ignore
|
||||
import proto
|
||||
|
||||
if hasattr(chunk, "candidates") is True:
|
||||
try:
|
||||
try:
|
||||
completion_obj["content"] = chunk.text # type: ignore
|
||||
completion_obj["content"] = chunk.text
|
||||
except Exception as e:
|
||||
original_exception: Final = e
|
||||
if "Part has no text." in str(e):
|
||||
## check for function calling
|
||||
function_call: Final = (
|
||||
chunk.candidates[0].content.parts[0].function_call # type: ignore
|
||||
)
|
||||
function_call: Final = chunk.candidates[0].content.parts[0].function_call
|
||||
|
||||
args_dict: Final = {}
|
||||
|
||||
|
|
@ -1159,7 +1157,7 @@ class CustomStreamWrapper:
|
|||
for key, val in function_call.args.items():
|
||||
if isinstance(
|
||||
val,
|
||||
proto.marshal.collections.repeated.RepeatedComposite, # type: ignore
|
||||
proto.marshal.collections.repeated.RepeatedComposite,
|
||||
):
|
||||
# If so, convert to list
|
||||
args_dict[key] = [v for v in val]
|
||||
|
|
@ -1190,15 +1188,12 @@ class CustomStreamWrapper:
|
|||
else:
|
||||
raise original_exception
|
||||
if (
|
||||
hasattr(chunk.candidates[0], "finish_reason") # type: ignore
|
||||
and chunk.candidates[0].finish_reason.name # type: ignore
|
||||
!= "FINISH_REASON_UNSPECIFIED"
|
||||
hasattr(chunk.candidates[0], "finish_reason")
|
||||
and chunk.candidates[0].finish_reason.name != "FINISH_REASON_UNSPECIFIED"
|
||||
): # every non-final chunk in vertex ai has this
|
||||
self.received_finish_reason = map_finish_reason( # type: ignore
|
||||
chunk.candidates[0].finish_reason.name
|
||||
)
|
||||
self.received_finish_reason = map_finish_reason(chunk.candidates[0].finish_reason.name)
|
||||
except Exception:
|
||||
if chunk.candidates[0].finish_reason.name == "SAFETY": # type: ignore
|
||||
if chunk.candidates[0].finish_reason.name == "SAFETY":
|
||||
raise Exception(f"The response was blocked by VertexAI. {chunk}")
|
||||
else:
|
||||
completion_obj["content"] = str(chunk)
|
||||
|
|
@ -1352,7 +1347,7 @@ class CustomStreamWrapper:
|
|||
)
|
||||
return _ProviderChunkParsed(response_obj)
|
||||
|
||||
def chunk_creator(self, chunk: Any): # type: ignore
|
||||
def chunk_creator(self, chunk: Any):
|
||||
if hasattr(chunk, "id"):
|
||||
self.response_id = chunk.id
|
||||
model_response = self.model_response_creator()
|
||||
|
|
@ -1460,7 +1455,7 @@ class CustomStreamWrapper:
|
|||
## RETURN ARG
|
||||
result: Final = self.return_processed_chunk_logic(
|
||||
completion_obj=completion_obj,
|
||||
model_response=model_response, # type: ignore
|
||||
model_response=model_response,
|
||||
response_obj=response_obj,
|
||||
)
|
||||
return result
|
||||
|
|
@ -1702,7 +1697,7 @@ class CustomStreamWrapper:
|
|||
):
|
||||
chunk = self.completion_stream
|
||||
else:
|
||||
chunk = next(self.completion_stream) # type: ignore[arg-type]
|
||||
chunk = next(self.completion_stream)
|
||||
if chunk is not None and chunk != b"":
|
||||
print_verbose(
|
||||
f"PROCESSED CHUNK PRE CHUNK CREATOR: {chunk.decode('utf-8', errors='replace') if isinstance(chunk, bytes) else chunk}; custom_llm_provider: {self.custom_llm_provider}"
|
||||
|
|
@ -1951,7 +1946,7 @@ class CustomStreamWrapper:
|
|||
if self.sent_last_chunk is True:
|
||||
processed_chunk = await self._call_post_streaming_deployment_hook(processed_chunk)
|
||||
# Add MCP metadata to final chunk if present (after hooks)
|
||||
processed_chunk = self._add_mcp_metadata_to_final_chunk(processed_chunk) # type: ignore[reportArgumentType]
|
||||
processed_chunk = self._add_mcp_metadata_to_final_chunk(processed_chunk)
|
||||
|
||||
return processed_chunk
|
||||
raise StopAsyncIteration
|
||||
|
|
@ -1961,7 +1956,7 @@ class CustomStreamWrapper:
|
|||
if isinstance(self.completion_stream, str) or isinstance(self.completion_stream, bytes):
|
||||
chunk = self.completion_stream
|
||||
else:
|
||||
chunk = await asyncio.to_thread(_next_sync_or_exhausted, self.completion_stream) # type: ignore[arg-type]
|
||||
chunk = await asyncio.to_thread(_next_sync_or_exhausted, self.completion_stream)
|
||||
if chunk is _SYNC_ITER_EXHAUSTED:
|
||||
raise StopAsyncIteration
|
||||
if chunk is not None and chunk != b"":
|
||||
|
|
@ -2069,7 +2064,7 @@ class CustomStreamWrapper:
|
|||
# end-of-stream blocks complete. Scheduling here via
|
||||
# create_task would race with unified_guardrail's
|
||||
# end-of-stream block for short-stream providers.
|
||||
self.logging_obj._deferred_stream_complete_args = ( # type: ignore[attr-defined]
|
||||
self.logging_obj._deferred_stream_complete_args = (
|
||||
complete_streaming_response,
|
||||
cache_hit,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -524,7 +524,7 @@ def _get_count_function(
|
|||
from litellm.utils import _select_tokenizer, print_verbose
|
||||
|
||||
if model is not None or custom_tokenizer is not None:
|
||||
tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model) # type: ignore
|
||||
tokenizer_json: Final = custom_tokenizer or _select_tokenizer(model)
|
||||
if tokenizer_json["type"] == "huggingface_tokenizer":
|
||||
|
||||
def count_tokens(text: str) -> int:
|
||||
|
|
@ -532,7 +532,7 @@ def _get_count_function(
|
|||
return len(enc.ids)
|
||||
|
||||
elif tokenizer_json["type"] == "openai_tokenizer":
|
||||
model_to_use: Final = _fix_model_name(model) # type: ignore
|
||||
model_to_use: Final = _fix_model_name(model)
|
||||
try:
|
||||
if "gpt-4o" in model_to_use:
|
||||
encoding = tiktoken.get_encoding("o200k_base")
|
||||
|
|
@ -561,7 +561,7 @@ def _fix_model_name(model: str) -> str:
|
|||
# azure llms use gpt-35-turbo instead of gpt-3.5-turbo 🙃
|
||||
return model.replace("-35", "-3.5")
|
||||
elif model in litellm.open_ai_chat_completion_models:
|
||||
return model # type: ignore
|
||||
return model
|
||||
else:
|
||||
return "gpt-3.5-turbo"
|
||||
|
||||
|
|
@ -592,7 +592,7 @@ def _count_image_tokens(
|
|||
raise ValueError("Missing required key 'url' in image_url dict.")
|
||||
return calculate_img_tokens(
|
||||
data=url,
|
||||
mode=detail, # type: ignore
|
||||
mode=detail,
|
||||
use_default_image_token_count=use_default_image_token_count,
|
||||
)
|
||||
elif isinstance(image_url, str):
|
||||
|
|
@ -669,7 +669,7 @@ def _count_anthropic_content(
|
|||
elif isinstance(field_value, list):
|
||||
tokens += _count_content_list(
|
||||
count_function,
|
||||
field_value, # type: ignore
|
||||
field_value,
|
||||
use_default_image_token_count,
|
||||
default_token_count,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -53,7 +53,7 @@ def convert_messages_to_prompt(messages: list[AllMessageValues]) -> str:
|
|||
elif isinstance(msg, dict):
|
||||
role = msg.get("role", "user")
|
||||
else:
|
||||
role = dict(msg).get("role", "user") # type: ignore
|
||||
role = dict(msg).get("role", "user")
|
||||
|
||||
if content_text:
|
||||
conversation_parts.append(f"{role}: {content_text}")
|
||||
|
|
|
|||
|
|
@ -15,6 +15,6 @@ class AIMLChatConfig(OpenAIGPTConfig):
|
|||
# AIML is openai compatible, we just need to set the api_base
|
||||
api_base = (
|
||||
api_base or get_secret_str("AIML_API_BASE") or "https://api.aimlapi.com/v1" # Default AIML API base URL
|
||||
) # type: ignore
|
||||
)
|
||||
dynamic_api_key: Final = api_key or get_secret_str("AIML_API_KEY")
|
||||
return api_base, dynamic_api_key
|
||||
|
|
|
|||
|
|
@ -56,7 +56,7 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig):
|
|||
) -> dict:
|
||||
return {"Authorization": f"Bearer {api_key}"}
|
||||
|
||||
async def transform_response( # type: ignore
|
||||
async def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: ClientResponse,
|
||||
|
|
|
|||
|
|
@ -52,7 +52,7 @@ class AmazonNovaChatConfig(OpenAILikeChatConfig):
|
|||
self, api_base: str | None, api_key: str | None
|
||||
) -> tuple[str | None, str | None]:
|
||||
# Amazon Nova is openai compatible, we just need to set this to custom_openai and have the api_base be Nova's endpoint
|
||||
api_base = api_base or get_secret_str("AMAZON_NOVA_API_BASE") or "https://api.nova.amazon.com/v1" # type: ignore
|
||||
api_base = api_base or get_secret_str("AMAZON_NOVA_API_BASE") or "https://api.nova.amazon.com/v1"
|
||||
|
||||
# Get API key from multiple sources
|
||||
key: Final = api_key or litellm.amazon_nova_api_key or get_secret_str("AMAZON_NOVA_API_KEY") or litellm.api_key
|
||||
|
|
|
|||
|
|
@ -469,7 +469,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
openai_tools: Final = self.adapter.translate_anthropic_tools_to_openai(
|
||||
tools=cast(list[AllAnthropicToolsValues], tools)
|
||||
)
|
||||
tools_to_check.extend(openai_tools) # type: ignore
|
||||
tools_to_check.extend(openai_tools)
|
||||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
self,
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue