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:
shin-bot-litellm 2026-01-31 11:45:25 -08:00 • committed by GitHub
parent 14f31a0df9
commit bcc05a67b2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 119 additions and 6 deletions

View file

@ -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(

View file

@ -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",