mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
litellm_fix(azure): Fix acancel_batch not using Azure SDK client initialization (#20168)
- Fixed model parameter being overwritten to None in acancel_batch function - Added dedicated acancel_batch/\_acancel_batch methods in Router - Properly extracts custom_llm_provider from deployment like acreate_batch This fixes test_ensure_initialize_azure_sdk_client_always_used[acancel_batch] which expected azure_batches_instance.initialize_azure_sdk_client to be called. Co-authored-by: shin-bot-litellm <shin-bot-litellm@users.noreply.github.com>
This commit is contained in:
parent
14f31a0df9
commit
bcc05a67b2
2 changed files with 119 additions and 6 deletions
|
|
@ -876,7 +876,9 @@ async def acancel_batch(
|
|||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["acancel_batch"] = True
|
||||
model = kwargs.pop("model", None)
|
||||
# Preserve model parameter - only pop from kwargs if it exists there
|
||||
# (to avoid passing it twice), otherwise keep the function parameter value
|
||||
model = kwargs.pop("model", None) or model
|
||||
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(
|
||||
|
|
|
|||
|
|
@ -880,9 +880,8 @@ class Router:
|
|||
self.allm_passthrough_route = self.factory_function(
|
||||
litellm.allm_passthrough_route, call_type="allm_passthrough_route"
|
||||
)
|
||||
self.acancel_batch = self.factory_function(
|
||||
litellm.acancel_batch, call_type="acancel_batch"
|
||||
)
|
||||
# Note: acancel_batch is defined as a method on the Router class (not using factory_function)
|
||||
# to properly handle model-to-provider mapping like acreate_batch and aretrieve_batch
|
||||
|
||||
def _initialize_vector_store_endpoints(self):
|
||||
"""Initialize vector store endpoints."""
|
||||
|
|
@ -4021,6 +4020,120 @@ class Router:
|
|||
)
|
||||
raise e
|
||||
|
||||
async def acancel_batch(
|
||||
self,
|
||||
model: str,
|
||||
**kwargs,
|
||||
) -> LiteLLMBatch:
|
||||
"""
|
||||
Cancel a batch through the router with proper model-to-provider mapping.
|
||||
"""
|
||||
try:
|
||||
kwargs["model"] = model
|
||||
kwargs["original_function"] = self._acancel_batch
|
||||
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
|
||||
metadata_variable_name = _get_router_metadata_variable_name(
|
||||
function_name="_acancel_batch"
|
||||
)
|
||||
self._update_kwargs_before_fallbacks(
|
||||
model=model,
|
||||
kwargs=kwargs,
|
||||
metadata_variable_name=metadata_variable_name,
|
||||
)
|
||||
response = await self.async_function_with_fallbacks(**kwargs)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
asyncio.create_task(
|
||||
send_llm_exception_alert(
|
||||
litellm_router_instance=self,
|
||||
request_kwargs=kwargs,
|
||||
error_traceback_str=traceback.format_exc(),
|
||||
original_exception=e,
|
||||
)
|
||||
)
|
||||
raise e
|
||||
|
||||
async def _acancel_batch(
|
||||
self,
|
||||
model: str,
|
||||
**kwargs,
|
||||
) -> LiteLLMBatch:
|
||||
try:
|
||||
verbose_router_logger.debug(
|
||||
f"Inside _acancel_batch()- model: {model}; kwargs: {kwargs}"
|
||||
)
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
deployment = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "batch-api-fake-text"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
|
||||
data = deployment["litellm_params"].copy()
|
||||
model_name = data["model"]
|
||||
self._update_kwargs_with_deployment(
|
||||
deployment=deployment, kwargs=kwargs, function_name="_acancel_batch"
|
||||
)
|
||||
|
||||
model_client = self._get_async_openai_model_client(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
self.total_calls[model_name] += 1
|
||||
|
||||
## SET CUSTOM PROVIDER TO SELECTED DEPLOYMENT ##
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(model=data["model"])
|
||||
|
||||
response = litellm.acancel_batch(
|
||||
**{
|
||||
**data,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"caching": self.cache_responses,
|
||||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
)
|
||||
|
||||
rpm_semaphore = self._get_client(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
client_type="max_parallel_requests",
|
||||
)
|
||||
|
||||
if rpm_semaphore is not None and isinstance(
|
||||
rpm_semaphore, asyncio.Semaphore
|
||||
):
|
||||
async with rpm_semaphore:
|
||||
"""
|
||||
- Check rpm limits before making the call
|
||||
- If allowed, increment the rpm limit (allows global value to be updated, concurrency-safe)
|
||||
"""
|
||||
await self.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment, parent_otel_span=parent_otel_span
|
||||
)
|
||||
response = await response # type: ignore
|
||||
else:
|
||||
await self.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment, parent_otel_span=parent_otel_span
|
||||
)
|
||||
response = await response # type: ignore
|
||||
|
||||
self.success_calls[model_name] += 1
|
||||
verbose_router_logger.info(
|
||||
f"litellm.acancel_batch(model={model_name})\033[32m 200 OK\033[0m"
|
||||
)
|
||||
|
||||
return response # type: ignore
|
||||
except Exception as e:
|
||||
verbose_router_logger.exception(
|
||||
f"litellm._acancel_batch(model={model}, {kwargs})\033[31m Exception {str(e)}\033[0m"
|
||||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
raise e
|
||||
|
||||
async def alist_batches(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -4114,7 +4227,6 @@ class Router:
|
|||
"afile_delete",
|
||||
"afile_content",
|
||||
"_arealtime",
|
||||
"acancel_batch",
|
||||
"acreate_fine_tuning_job",
|
||||
"acancel_fine_tuning_job",
|
||||
"alist_fine_tuning_jobs",
|
||||
|
|
@ -4302,7 +4414,6 @@ class Router:
|
|||
"avideo_status",
|
||||
"avideo_content",
|
||||
"avideo_remix",
|
||||
"acancel_batch",
|
||||
"acreate_skill",
|
||||
"alist_skills",
|
||||
"aget_skill",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue