Fix: Encoding cancel batch response

This commit is contained in:
Sameer Kankute 2026-01-29 12:18:43 +05:30
parent 654edbd15c
commit 8966852c86
3 changed files with 16 additions and 26 deletions

View file

@ -5,7 +5,6 @@ Azure Batches API Handler
from typing import Any, Coroutine, Optional, Union, cast
import httpx
from openai import AsyncOpenAI, OpenAI
from litellm.llms.azure.azure import AsyncAzureOpenAI, AzureOpenAI
@ -130,9 +129,9 @@ class AzureBatchesAPI(BaseAzureLLM):
self,
cancel_batch_data: CancelBatchRequest,
client: Union[AsyncAzureOpenAI, AsyncOpenAI],
) -> Batch:
) -> LiteLLMBatch:
response = await client.batches.cancel(**cancel_batch_data)
return response
return LiteLLMBatch(**response.model_dump())
def cancel_batch(
self,
@ -161,7 +160,7 @@ class AzureBatchesAPI(BaseAzureLLM):
"OpenAI client is not initialized. Make sure api_key is passed or OPENAI_API_KEY is set in the environment."
)
response = azure_client.batches.cancel(**cancel_batch_data)
return response
return LiteLLMBatch(**response.model_dump())
async def alist_batches(
self,

View file

@ -1923,10 +1923,10 @@ class OpenAIBatchesAPI(BaseLLM):
self,
cancel_batch_data: CancelBatchRequest,
openai_client: AsyncOpenAI,
) -> Batch:
) -> LiteLLMBatch:
verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data)
response = await openai_client.batches.cancel(**cancel_batch_data)
return response
return LiteLLMBatch(**response.model_dump())
def cancel_batch(
self,
@ -1963,7 +1963,7 @@ class OpenAIBatchesAPI(BaseLLM):
)
response = openai_client.batches.cancel(**cancel_batch_data)
return response
return LiteLLMBatch(**response.model_dump())
async def alist_batches(
self,

View file

@ -382,25 +382,6 @@ async def retrieve_batch(
**data # type: ignore
)
# Re-encode all IDs in the response
if response:
if hasattr(response, "id") and response.id:
response.id = batch_id # Keep the encoded batch ID
if hasattr(response, "input_file_id") and response.input_file_id:
response.input_file_id = encode_file_id_with_model(
file_id=response.input_file_id, model=model_from_id
)
if hasattr(response, "output_file_id") and response.output_file_id:
response.output_file_id = encode_file_id_with_model(
file_id=response.output_file_id, model=model_from_id
)
if hasattr(response, "error_file_id") and response.error_file_id:
response.error_file_id = encode_file_id_with_model(
file_id=response.error_file_id, model=model_from_id
)
verbose_proxy_logger.debug(
f"Retrieved batch using model: {model_from_id}, original_id: {original_batch_id}"
@ -769,6 +750,11 @@ async def cancel_batch(
# Hook has already extracted model and unwrapped batch_id into data dict
response = await llm_router.acancel_batch(**data) # type: ignore
response._hidden_params["unified_batch_id"] = unified_batch_id
# Ensure model_id is set for the post_call_success_hook to re-encode IDs
if not response._hidden_params.get("model_id") and data.get("model"):
response._hidden_params["model_id"] = data["model"]
# SCENARIO 3: Fallback to custom_llm_provider (uses env variables)
else:
@ -782,6 +768,11 @@ async def cancel_batch(
**_cancel_batch_data,
)
### CALL HOOKS ### - modify outgoing data
response = await proxy_logging_obj.post_call_success_hook(
data=data, user_api_key_dict=user_api_key_dict, response=response
)
### ALERTING ###
asyncio.create_task(
proxy_logging_obj.update_request_status(