Litellm managed file updates combined (#11040)

* Add LiteLLM Managed file support for `retrieve`, `list` and `cancel` finetuning jobs (#11033)

* feat: initial commit adding managed file support to fine tuning endpoints

* feat(fine_tuning/endpoints.py): working call to openai finetuning route

Uses litellm managed files for finetuning api support

* feat(fine-tuning/main.py): refactor to use LiteLLMFineTuningJob pydantic object

includes 'hidden_params'

* fix: initial commit adding unified finetuning id support

return a unified finetuning id we can use to understand which deployment to route the ft request to

* test: fix test

* feat(managed_files.py): return unified finetuning job id on create finetuning job

enables retrieve, delete to work with litellm managed files

* feat(managed_files.py): support managed files for cancel ft job endpoint

* feat(managed_files.py): support managed files for cancel ft job endpoint

* feat(fine_tuning_endpoints/endpoints.py): add managed files support to list finetuning jobs

* feat(finetuning_endpoints/main): add managed files support for retrieving ft job

Makes it easier to control permissions for ft endpoint

* LiteLLM Managed Files - Enforce validation check if user can access finetuning job (#11034)

* feat: initial commit adding managed file support to fine tuning endpoints

* feat(fine_tuning/endpoints.py): working call to openai finetuning route

Uses litellm managed files for finetuning api support

* feat(fine-tuning/main.py): refactor to use LiteLLMFineTuningJob pydantic object

includes 'hidden_params'

* fix: initial commit adding unified finetuning id support

return a unified finetuning id we can use to understand which deployment to route the ft request to

* test: fix test

* feat(managed_files.py): return unified finetuning job id on create finetuning job

enables retrieve, delete to work with litellm managed files

* feat(managed_files.py): support managed files for cancel ft job endpoint

* feat(managed_files.py): support managed files for cancel ft job endpoint

* feat(fine_tuning_endpoints/endpoints.py): add managed files support to list finetuning jobs

* feat(finetuning_endpoints/main): add managed files support for retrieving ft job

Makes it easier to control permissions for ft endpoint

* feat(managed_files.py): store create fine-tune / batch response object in db

storing this allows us to filter files returned on list based on what user created

* feat(managed_files.py): Ensures users can't retrieve / modify each others jobs

* fix: fix check

* fix: fix ruff check errors

* test: update to handle testing

* fix: suppress linting warning - openai 'seed' is none on azure

* test: update tests

* test: update test
This commit is contained in:
Krish Dholakia 2025-05-22 17:20:41 -07:00 • committed by GitHub
parent 469d395177
commit 70f32154c5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 444 additions and 113 deletions

View file

@ -1,4 +1,3 @@
import os
from typing import Dict, Literal, Type, Union
from litellm.integrations.custom_logger import CustomLogger

View file

@ -1,17 +1,25 @@
# What is this?
## This hook is used to check for LiteLLM managed files in the request body, and replace them with model-specific file id
import asyncio
import base64
import json
import uuid
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
from fastapi import HTTPException
from litellm import Router, verbose_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.proxy._types import CallTypes, LiteLLM_ManagedFileTable, UserAPIKeyAuth
from litellm.proxy._types import (
CallTypes,
LiteLLM_ManagedFileTable,
LiteLLM_ManagedObjectTable,
UserAPIKeyAuth,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
convert_b64_uid_to_unified_uid,
@ -82,6 +90,41 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
}
)
async def store_unified_object_id(
self,
unified_object_id: str,
file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob],
litellm_parent_otel_span: Optional[Span],
model_object_id: str,
file_purpose: Literal["batch", "fine-tune"],
user_api_key_dict: UserAPIKeyAuth,
) -> None:
verbose_logger.info(
f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache"
)
litellm_managed_object = LiteLLM_ManagedObjectTable(
unified_object_id=unified_object_id,
model_object_id=model_object_id,
file_purpose=file_purpose,
file_object=file_object,
)
await self.internal_usage_cache.async_set_cache(
key=unified_object_id,
value=litellm_managed_object.model_dump(),
litellm_parent_otel_span=litellm_parent_otel_span,
)
await self.prisma_client.db.litellm_managedobjecttable.create(
data={
"unified_object_id": unified_object_id,
"file_object": file_object.model_dump_json(),
"model_object_id": model_object_id,
"file_purpose": file_purpose,
"created_by": user_api_key_dict.user_id,
"updated_by": user_api_key_dict.user_id,
}
)
async def get_unified_file_id(
self, file_id: str, litellm_parent_otel_span: Optional[Span] = None
) -> Optional[LiteLLM_ManagedFileTable]:
@ -126,6 +169,21 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
return initial_value.file_object
async def can_user_call_unified_object_id(
self, unified_object_id: str, user_api_key_dict: UserAPIKeyAuth
) -> bool:
## check if the user has access to the unified object id
## check if the user has access to the unified object id
user_id = user_api_key_dict.user_id
managed_object = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"unified_object_id": unified_object_id}
)
)
if managed_object:
return managed_object.created_by == user_id
return False
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -144,6 +202,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"aretrieve_batch",
"afile_content",
"acreate_fine_tuning_job",
"aretrieve_fine_tuning_job",
"alist_fine_tuning_jobs",
"acancel_fine_tuning_job",
],
) -> Union[Exception, str, Dict, None]:
"""
@ -185,25 +246,50 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
data["model_file_id_mapping"] = model_file_id_mapping
elif call_type == CallTypes.aretrieve_batch.value:
retrieve_batch_id = cast(Optional[str], data.get("batch_id"))
potential_batch_id = (
_is_base64_encoded_unified_file_id(retrieve_batch_id)
if retrieve_batch_id
elif (
call_type == CallTypes.aretrieve_batch.value
or call_type == CallTypes.acancel_fine_tuning_job.value
or call_type == CallTypes.aretrieve_fine_tuning_job.value
):
accessor_key: Optional[str] = None
retrieve_object_id: Optional[str] = None
if call_type == CallTypes.aretrieve_batch.value:
accessor_key = "batch_id"
elif (
call_type == CallTypes.acancel_fine_tuning_job.value
or call_type == CallTypes.aretrieve_fine_tuning_job.value
):
accessor_key = "fine_tuning_job_id"
if accessor_key:
retrieve_object_id = cast(Optional[str], data.get(accessor_key))
potential_llm_object_id = (
_is_base64_encoded_unified_file_id(retrieve_object_id)
if retrieve_object_id
else False
)
if potential_batch_id:
if potential_llm_object_id and retrieve_object_id:
## VALIDATE USER HAS ACCESS TO THE OBJECT ##
if not await self.can_user_call_unified_object_id(
retrieve_object_id, user_api_key_dict
):
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to the object {retrieve_object_id}",
)
## for managed batch id - get the model id
potential_model_id = self.get_model_id_from_unified_batch_id(
potential_batch_id
potential_llm_object_id
)
if potential_model_id is None:
raise Exception(
f"LiteLLM Managed Batch ID with id={retrieve_batch_id} is invalid - does not contain encoded model_id."
f"LiteLLM Managed {accessor_key} with id={retrieve_object_id} is invalid - does not contain encoded model_id."
)
data["model"] = potential_model_id
data["batch_id"] = self.get_batch_id_from_unified_batch_id(
potential_batch_id
data[accessor_key] = self.get_batch_id_from_unified_batch_id(
potential_llm_object_id
)
elif call_type == CallTypes.acreate_fine_tuning_job.value:
input_file_id = cast(Optional[str], data.get("training_file"))
@ -211,7 +297,6 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_file_id_mapping = await self.get_model_file_id_mapping(
[input_file_id], user_api_key_dict.parent_otel_span
)
data["model_file_id_mapping"] = model_file_id_mapping
print("DATA={}".format(data))
return data
@ -222,11 +307,6 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"""
Allow modifying the request just before it's sent to the deployment.
"""
print(
"CALLS ASYNC PRE CALL DEPLOYMENT HOOK - KWARGS={}, CALL_TYPE={}".format(
kwargs, call_type
)
)
accessor_key: Optional[str] = None
if call_type and call_type == CallTypes.acreate_batch:
accessor_key = "input_file_id"
@ -472,7 +552,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
def get_batch_id_from_unified_batch_id(self, file_id: str) -> str:
## use regex to get the batch_id from the file_id
return file_id.split("llm_batch_id:")[1].split(",")[0]
if "llm_batch_id" in file_id:
return file_id.split("llm_batch_id:")[1].split(",")[0]
else:
return file_id.split("generic_response_id:")[1].split(",")[0]
async def async_post_call_success_hook(
self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
@ -487,6 +570,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
) # managed batch id
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
original_response_id = response.id
if (unified_batch_id or unified_file_id) and model_id:
response.id = self.get_unified_batch_id(
batch_id=response.id, model_id=model_id
@ -500,24 +584,42 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_id=model_id,
model_name=model_name,
)
return response
asyncio.create_task(
self.store_unified_object_id(
unified_object_id=response.id,
file_object=response,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
model_object_id=original_response_id,
file_purpose="batch",
user_api_key_dict=user_api_key_dict,
)
)
elif isinstance(response, LiteLLMFineTuningJob):
## Check if unified_file_id is in the response
print(f"hidden params={response._hidden_params}")
unified_file_id = response._hidden_params.get(
"unified_file_id"
) # managed file id
unified_finetuning_job_id = response._hidden_params.get(
"unified_finetuning_job_id"
) # managed finetuning job id
model_id = cast(Optional[str], response._hidden_params.get("model_id"))
print("MODEL_ID={}".format(model_id))
model_name = cast(Optional[str], response._hidden_params.get("model_name"))
if unified_file_id and model_id:
original_response_id = response.id
if (unified_file_id or unified_finetuning_job_id) and model_id:
response.id = self.get_unified_generic_response_id(
model_id=model_id, generic_response_id=response.id
)
return response
return await super().async_post_call_success_hook(
data, user_api_key_dict, response
)
asyncio.create_task(
self.store_unified_object_id(
unified_object_id=response.id,
file_object=response,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
model_object_id=original_response_id,
file_purpose="fine-tune",
user_api_key_dict=user_api_key_dict,
)
)
return response
async def afile_retrieve(
self, file_id: str, litellm_parent_otel_span: Optional[Span]

View file

@ -22,11 +22,7 @@ from litellm.llms.azure.fine_tuning.handler import AzureOpenAIFineTuningAPI
from litellm.llms.openai.fine_tuning.handler import OpenAIFineTuningAPI
from litellm.llms.vertex_ai.fine_tuning.handler import VertexFineTuningAPI
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
FineTuningJob,
FineTuningJobCreate,
Hyperparameters,
)
from litellm.types.llms.openai import FineTuningJobCreate, Hyperparameters
from litellm.types.router import *
from litellm.types.utils import LiteLLMFineTuningJob
from litellm.utils import client, supports_httpx_timeout
@ -289,13 +285,14 @@ def create_fine_tuning_job(
raise e
@client
async def acancel_fine_tuning_job(
fine_tuning_job_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,
) -> FineTuningJob:
) -> LiteLLMFineTuningJob:
"""
Async: Immediately cancel a fine-tune job.
"""
@ -326,13 +323,14 @@ async def acancel_fine_tuning_job(
raise e
@client
def cancel_fine_tuning_job(
fine_tuning_job_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,
) -> Union[FineTuningJob, Coroutine[Any, Any, FineTuningJob]]:
) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
"""
Immediately cancel a fine-tune job.
@ -610,13 +608,14 @@ def list_fine_tuning_jobs(
raise e
@client
async def aretrieve_fine_tuning_job(
fine_tuning_job_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,
) -> FineTuningJob:
) -> LiteLLMFineTuningJob:
"""
Async: Get info about a fine-tuning job.
"""
@ -647,13 +646,14 @@ async def aretrieve_fine_tuning_job(
raise e
@client
def retrieve_fine_tuning_job(
fine_tuning_job_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,
) -> Union[FineTuningJob, Coroutine[Any, Any, FineTuningJob]]:
) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
"""
Get info about a fine-tuning job.
"""

View file

@ -2,7 +2,6 @@ from typing import Any, Coroutine, Optional, Union, cast
import httpx
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from openai.types.fine_tuning import FineTuningJob
from litellm._logging import verbose_logger
from litellm.types.utils import LiteLLMFineTuningJob
@ -115,11 +114,11 @@ class OpenAIFineTuningAPI:
self,
fine_tuning_job_id: str,
openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI],
) -> FineTuningJob:
) -> LiteLLMFineTuningJob:
response = await openai_client.fine_tuning.jobs.cancel(
fine_tuning_job_id=fine_tuning_job_id
)
return response
return LiteLLMFineTuningJob(**response.model_dump())
def cancel_fine_tuning_job(
self,
@ -134,7 +133,7 @@ class OpenAIFineTuningAPI:
client: Optional[
Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
] = None,
):
) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
openai_client: Optional[
Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
] = self.get_openai_client(
@ -162,10 +161,10 @@ class OpenAIFineTuningAPI:
openai_client=openai_client,
)
verbose_logger.debug("canceling fine tuning job, args= %s", fine_tuning_job_id)
response = openai_client.fine_tuning.jobs.cancel(
response = cast(OpenAI, openai_client).fine_tuning.jobs.cancel(
fine_tuning_job_id=fine_tuning_job_id
)
return response
return LiteLLMFineTuningJob(**response.model_dump())
async def alist_fine_tuning_jobs(
self,
@ -226,11 +225,11 @@ class OpenAIFineTuningAPI:
self,
fine_tuning_job_id: str,
openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI],
) -> FineTuningJob:
) -> LiteLLMFineTuningJob:
response = await openai_client.fine_tuning.jobs.retrieve(
fine_tuning_job_id=fine_tuning_job_id
)
return response
return LiteLLMFineTuningJob(**response.model_dump())
def retrieve_fine_tuning_job(
self,
@ -245,7 +244,7 @@ class OpenAIFineTuningAPI:
client: Optional[
Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
] = None,
):
) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
openai_client: Optional[
Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
] = self.get_openai_client(
@ -273,7 +272,7 @@ class OpenAIFineTuningAPI:
openai_client=openai_client,
)
verbose_logger.debug("retrieving fine tuning job, id= %s", fine_tuning_job_id)
response = openai_client.fine_tuning.jobs.retrieve(
response = cast(OpenAI, openai_client).fine_tuning.jobs.retrieve(
fine_tuning_job_id=fine_tuning_job_id
)
return response
return LiteLLMFineTuningJob(**response.model_dump())

View file

@ -5,13 +5,13 @@ model_list:
- model_name: "gpt-4o-mini-openai"
litellm_params:
model: gpt-4.1-mini-2025-04-14
api_key: os.environ/OPENAI_API_KEY
api_key: os.environ/OPENAI_API_KEY_2
model_info:
access_groups: ["default-openai-models"]
- model_name: "gpt-4o-realtime-preview"
litellm_params:
model: gpt-4o-realtime-preview-2024-10-01
api_key: os.environ/OPENAI_API_KEY
api_key: os.environ/OPENAI_API_KEY_2
- model_name: "bedrock-nova"
litellm_params:
model: us.amazon.nova-pro-v1:0

View file

@ -23,6 +23,8 @@ from litellm.types.utils import (
EmbeddingResponse,
GenericBudgetConfigType,
ImageResponse,
LiteLLMBatch,
LiteLLMFineTuningJob,
LiteLLMPydanticObjectBase,
ModelResponse,
ProviderField,
@ -2879,3 +2881,10 @@ class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase):
unified_file_id: str
file_object: OpenAIFileObject
model_mappings: Dict[str, str]
class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase):
unified_object_id: str
model_object_id: str
file_purpose: Literal["batch", "fine-tune"]
file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob]

View file

@ -118,6 +118,9 @@ class ProxyBaseLLMRequestProcessing:
"aretrieve_batch",
"afile_content",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",
"aretrieve_fine_tuning_job",
],
version: Optional[str] = None,
user_model: Optional[str] = None,

View file

@ -6,10 +6,9 @@
##########################################################################
import asyncio
import traceback
from typing import Optional, cast
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response
import litellm
from litellm._logging import verbose_proxy_logger
@ -163,7 +162,6 @@ async def create_fine_tuning_job(
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=fine_tuning_request.custom_llm_provider,
)
# add llm_provider_config to data
if llm_provider_config is not None:
data.update(llm_provider_config)
@ -237,7 +235,7 @@ async def retrieve_fine_tuning_job(
request: Request,
fastapi_response: Response,
fine_tuning_job_id: str,
custom_llm_provider: Literal["openai", "azure"],
custom_llm_provider: Optional[Literal["openai", "azure"]] = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
@ -249,41 +247,99 @@ async def retrieve_fine_tuning_job(
- `fine_tuning_job_id`: The ID of the fine-tuning job to retrieve.
"""
from litellm.proxy.proxy_server import (
add_litellm_data_to_request,
general_settings,
llm_router,
premium_user,
proxy_config,
proxy_logging_obj,
version,
)
data: dict = {}
data: dict = {"fine_tuning_job_id": fine_tuning_job_id}
try:
if premium_user is not True:
raise ValueError(
f"Only premium users can use this endpoint + {CommonProxyErrors.not_premium_user.value}"
)
# Include original request and headers in the data
data = await add_litellm_data_to_request(
data=data,
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
(
data,
litellm_logging_obj,
) = await base_llm_response_processor.common_processing_pre_call_logic(
request=request,
general_settings=general_settings,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_logging_obj=proxy_logging_obj,
proxy_config=proxy_config,
route_type=CallTypes.aretrieve_fine_tuning_job.value,
)
# get configs for custom_llm_provider
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=custom_llm_provider
try:
request_body = await request.json()
except Exception:
request_body = {}
custom_llm_provider = request_body.get("custom_llm_provider", None)
## CHECK IF MANAGED FILE ID
unified_finetuning_job_id: Union[str, Literal[False]] = False
response: Optional[LiteLLMFineTuningJob] = None
if fine_tuning_job_id:
unified_finetuning_job_id = _is_base64_encoded_unified_file_id(
fine_tuning_job_id
)
if unified_finetuning_job_id:
if llm_router is None:
raise HTTPException(
status_code=500,
detail={
"error": "LLM Router not initialized. Ensure models added to proxy."
},
)
response = cast(
LiteLLMFineTuningJob,
await llm_router.aretrieve_fine_tuning_job(
**data,
),
)
response._hidden_params[
"unified_finetuning_job_id"
] = unified_finetuning_job_id
elif custom_llm_provider:
# get configs for custom_llm_provider
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=custom_llm_provider
)
if llm_provider_config is not None:
data.update(llm_provider_config)
response = await litellm.aretrieve_fine_tuning_job(
**data,
)
if response is None:
raise HTTPException(
status_code=400,
detail="Invalid request, No litellm managed file id or custom_llm_provider provided.",
)
### 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,
)
if _response is not None and isinstance(_response, LiteLLMFineTuningJob):
response = _response
if llm_provider_config is not None:
data.update(llm_provider_config)
response = await litellm.aretrieve_fine_tuning_job(
**data,
fine_tuning_job_id=fine_tuning_job_id,
### ALERTING ###
asyncio.create_task(
proxy_logging_obj.update_request_status(
litellm_call_id=data.get("litellm_call_id", ""), status="success"
)
)
### RESPONSE HEADERS ###
@ -309,12 +365,11 @@ async def retrieve_fine_tuning_job(
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
)
verbose_proxy_logger.error(
"litellm.proxy.proxy_server.list_fine_tuning_jobs(): Exception occurred - {}".format(
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.retrieve_fine_tuning_job(): Exception occurred - {}".format(
str(e)
)
)
verbose_proxy_logger.debug(traceback.format_exc())
raise handle_exception_on_proxy(e)
@ -333,7 +388,11 @@ async def retrieve_fine_tuning_job(
async def list_fine_tuning_jobs(
request: Request,
fastapi_response: Response,
custom_llm_provider: Literal["openai", "azure"],
custom_llm_provider: Optional[Literal["openai", "azure"]] = None,
target_model_names: Optional[str] = Query(
default=None,
description="Comma separated list of model names to filter by. Example: 'gpt-4o,gpt-4o-mini'",
),
after: Optional[str] = None,
limit: Optional[int] = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -348,8 +407,8 @@ async def list_fine_tuning_jobs(
- `limit`: Number of fine-tuning jobs to retrieve (default is 20).
"""
from litellm.proxy.proxy_server import (
add_litellm_data_to_request,
general_settings,
llm_router,
premium_user,
proxy_config,
proxy_logging_obj,
@ -363,28 +422,60 @@ async def list_fine_tuning_jobs(
f"Only premium users can use this endpoint + {CommonProxyErrors.not_premium_user.value}"
)
# Include original request and headers in the data
data = await add_litellm_data_to_request(
data=data,
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
(
data,
litellm_logging_obj,
) = await base_llm_response_processor.common_processing_pre_call_logic(
request=request,
general_settings=general_settings,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_logging_obj=proxy_logging_obj,
proxy_config=proxy_config,
route_type=CallTypes.alist_fine_tuning_jobs.value,
)
# get configs for custom_llm_provider
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=custom_llm_provider
)
response: Optional[Any] = None
if target_model_names and isinstance(target_model_names, str):
target_model_names_list = target_model_names.split(",")
if len(target_model_names_list) != 1:
raise HTTPException(
status_code=400,
detail="target_model_names on list fine-tuning jobs must be a list of one model name. Example: ['gpt-4o']",
)
## Use router to list fine-tuning jobs for that model
if llm_router is None:
raise HTTPException(
status_code=500,
detail="LLM Router not initialized. Ensure models added to proxy.",
)
data["model"] = target_model_names_list[0]
response = await llm_router.alist_fine_tuning_jobs(
**data,
after=after,
limit=limit,
)
return response
elif custom_llm_provider:
# get configs for custom_llm_provider
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=custom_llm_provider
)
if llm_provider_config is not None:
data.update(llm_provider_config)
if llm_provider_config is not None:
data.update(llm_provider_config)
response = await litellm.alist_fine_tuning_jobs(
**data,
after=after,
limit=limit,
)
response = await litellm.alist_fine_tuning_jobs(
**data,
after=after,
limit=limit,
)
if response is None:
raise HTTPException(
status_code=400,
detail="Invalid request, No litellm managed file id or custom_llm_provider provided.",
)
### RESPONSE HEADERS ###
hidden_params = getattr(response, "_hidden_params", {}) or {}
@ -409,12 +500,11 @@ async def list_fine_tuning_jobs(
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
)
verbose_proxy_logger.error(
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.list_fine_tuning_jobs(): Exception occurred - {}".format(
str(e)
)
)
verbose_proxy_logger.debug(traceback.format_exc())
raise handle_exception_on_proxy(e)
@ -446,45 +536,99 @@ async def cancel_fine_tuning_job(
- `fine_tuning_job_id`: The ID of the fine-tuning job to cancel.
"""
from litellm.proxy.proxy_server import (
add_litellm_data_to_request,
general_settings,
llm_router,
premium_user,
proxy_config,
proxy_logging_obj,
version,
)
data: dict = {}
data: dict = {"fine_tuning_job_id": fine_tuning_job_id}
try:
if premium_user is not True:
raise ValueError(
f"Only premium users can use this endpoint + {CommonProxyErrors.not_premium_user.value}"
)
# Include original request and headers in the data
data = await add_litellm_data_to_request(
data=data,
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
(
data,
litellm_logging_obj,
) = await base_llm_response_processor.common_processing_pre_call_logic(
request=request,
general_settings=general_settings,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_logging_obj=proxy_logging_obj,
proxy_config=proxy_config,
route_type=CallTypes.acancel_fine_tuning_job.value,
)
request_body = await request.json()
try:
request_body = await request.json()
except Exception:
request_body = {}
custom_llm_provider = request_body.get("custom_llm_provider", None)
# get configs for custom_llm_provider
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=custom_llm_provider
## CHECK IF MANAGED FILE ID
unified_finetuning_job_id: Union[str, Literal[False]] = False
response: Optional[LiteLLMFineTuningJob] = None
if fine_tuning_job_id:
unified_finetuning_job_id = _is_base64_encoded_unified_file_id(
fine_tuning_job_id
)
if unified_finetuning_job_id:
if llm_router is None:
raise HTTPException(
status_code=500,
detail={
"error": "LLM Router not initialized. Ensure models added to proxy."
},
)
response = cast(
LiteLLMFineTuningJob,
await llm_router.acancel_fine_tuning_job(
**data,
),
)
response._hidden_params[
"unified_finetuning_job_id"
] = unified_finetuning_job_id
else:
# get configs for custom_llm_provider
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=custom_llm_provider
)
if llm_provider_config is not None:
data.update(llm_provider_config)
response = await litellm.acancel_fine_tuning_job(
**data,
)
if response is None:
raise HTTPException(
status_code=400,
detail="Invalid request, No litellm managed file id or custom_llm_provider provided.",
)
### 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,
)
if _response is not None and isinstance(_response, LiteLLMFineTuningJob):
response = _response
if llm_provider_config is not None:
data.update(llm_provider_config)
response = await litellm.acancel_fine_tuning_job(
**data,
fine_tuning_job_id=fine_tuning_job_id,
### ALERTING ###
asyncio.create_task(
proxy_logging_obj.update_request_status(
litellm_call_id=data.get("litellm_call_id", ""), status="success"
)
)
### RESPONSE HEADERS ###
@ -510,10 +654,9 @@ async def cancel_fine_tuning_job(
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
)
verbose_proxy_logger.error(
"litellm.proxy.proxy_server.list_fine_tuning_jobs(): Exception occurred - {}".format(
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.cancel_fine_tuning_job(): Exception occurred - {}".format(
str(e)
)
)
verbose_proxy_logger.debug(traceback.format_exc())
raise handle_exception_on_proxy(e)

View file

@ -453,13 +453,27 @@ model LiteLLM_ManagedFileTable {
id String @id @default(uuid())
unified_file_id String @unique // The base64 encoded unified file ID
file_object Json // Stores the OpenAIFileObject
model_mappings Json // Stores the mapping of model_id -> provider_file_id
model_mappings Json
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@index([unified_file_id])
}
model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use the
id String @id @default(uuid())
unified_object_id String @unique // The base64 encoded unified file ID
model_object_id String @unique // the id returned by the backend API provider
file_object Json // Stores the OpenAIFileObject
file_purpose String // either 'batch' or 'fine-tune'
created_at DateTime @default(now())
created_by String?
updated_at DateTime @updatedAt
updated_by String?
@@index([unified_object_id])
@@index([model_object_id])
}
model LiteLLM_ManagedVectorStoresTable {
vector_store_id String @id

View file

@ -755,6 +755,15 @@ class Router:
self.acreate_fine_tuning_job = self.factory_function(
litellm.acreate_fine_tuning_job, call_type="acreate_fine_tuning_job"
)
self.acancel_fine_tuning_job = self.factory_function(
litellm.acancel_fine_tuning_job, call_type="acancel_fine_tuning_job"
)
self.alist_fine_tuning_jobs = self.factory_function(
litellm.alist_fine_tuning_jobs, call_type="alist_fine_tuning_jobs"
)
self.aretrieve_fine_tuning_job = self.factory_function(
litellm.aretrieve_fine_tuning_job, call_type="aretrieve_fine_tuning_job"
)
def validate_fallbacks(self, fallback_param: Optional[List]):
"""
@ -2439,6 +2448,7 @@ class Router:
messages=kwargs.get("messages", None),
specific_deployment=kwargs.pop("specific_deployment", None),
)
self._update_kwargs_with_deployment(
deployment=deployment, kwargs=kwargs, function_name="generic_api_call"
)
@ -3172,6 +3182,9 @@ class Router:
"afile_content",
"_arealtime",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",
"aretrieve_fine_tuning_job",
] = "assistants",
):
"""
@ -3221,6 +3234,9 @@ class Router:
"aresponses",
"_arealtime",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",
"aretrieve_fine_tuning_job",
):
return await self._ageneric_api_call_with_fallbacks(
original_function=original_function,

View file

@ -96,6 +96,8 @@ class LowestTPMLoggingHandler(CustomLogger):
if kwargs["litellm_params"].get("metadata") is None:
pass
else:
if "litellm_params" not in kwargs:
return
model_group = kwargs["litellm_params"]["metadata"].get(
"model_group", None
)

View file

@ -2293,6 +2293,7 @@ class SelectTokenizerResponse(TypedDict):
class LiteLLMFineTuningJob(FineTuningJob):
_hidden_params: dict = {}
seed: Optional[int] = None # type: ignore
def __init__(self, **kwargs):
if "error" in kwargs and kwargs["error"] is not None:

View file

@ -525,9 +525,12 @@ async def test_mock_openai_cancel_fine_tune_job():
client = AsyncOpenAI(api_key="fake-api-key")
with patch.object(client.fine_tuning.jobs, "cancel") as mock_cancel:
await litellm.acancel_fine_tuning_job(
fine_tuning_job_id="ft-123", client=client
)
try:
await litellm.acancel_fine_tuning_job(
fine_tuning_job_id="ft-123", client=client
)
except Exception as e:
print("error=", e)
# Only verify that the client was called with correct parameters
mock_cancel.assert_called_once_with(fine_tuning_job_id="ft-123")
@ -541,10 +544,13 @@ async def test_mock_openai_retrieve_fine_tune_job():
client = AsyncOpenAI(api_key="fake-api-key")
with patch.object(client.fine_tuning.jobs, "retrieve") as mock_retrieve:
try:
response = await litellm.aretrieve_fine_tuning_job(
fine_tuning_job_id="ft-123", client=client
)
except Exception as e:
print("error=", e)
response = await litellm.aretrieve_fine_tuning_job(
fine_tuning_job_id="ft-123", client=client
)
# Verify the request
mock_retrieve.assert_called_once_with(fine_tuning_job_id="ft-123")

View file

@ -9,7 +9,7 @@ sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
from enterprise.enterprise_hooks.managed_files import _PROXY_LiteLLMManagedFiles
from litellm.caching import DualCache
@ -45,11 +45,19 @@ def test_get_file_ids_from_messages():
@pytest.mark.asyncio
async def test_async_pre_call_hook_batch_retrieve():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
return_value = MagicMock()
return_value.created_by = "123"
prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=MagicMock()
DualCache(), prisma_client=prisma_client
)
data = {
"user_api_key_dict": {"parent_otel_span": MagicMock()},
"user_api_key_dict": UserAPIKeyAuth(
user_id="123", parent_otel_span=MagicMock()
),
"data": {
"batch_id": "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDpteS1nZW5lcmFsLWF6dXJlLWRlcGxveW1lbnQ7bGxtX2JhdGNoX2lkOmJhdGNoX2EzMjJiNmJhLWFjN2UtNDg4OC05MjljLTFhZDM0NDJmMDZlZA",
},
@ -206,3 +214,29 @@ async def test_async_post_call_success_hook_for_unified_finetuning_job():
assert isinstance(response, LiteLLMFineTuningJob)
assert _is_base64_encoded_unified_file_id(response.id)
@pytest.mark.asyncio
async def test_async_pre_call_hook_for_unified_finetuning_job():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
return_value = MagicMock()
return_value.created_by = "123"
prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
data = {
"user_api_key_dict": UserAPIKeyAuth(
user_id="123", parent_otel_span=MagicMock()
),
"data": {
"fine_tuning_job_id": "bGl0ZWxsbV9wcm94eTttb2RlbF9pZDo0OTIxODU4MWY3OGViZTllZjE4NDE0ZmE0ZjdmYjlmYTc0YzA5NWVkMTEyY2E4NDBkZDU2ZGZmZTliZDMwZGQxO2dlbmVyaWNfcmVzcG9uc2VfaWQ6ZnRqb2ItalRCeXM3YlZzYnlaRE93TDlHbHBZcVhS",
},
"call_type": "acancel_fine_tuning_job",
"cache": MagicMock(),
}
response = await proxy_managed_files.async_pre_call_hook(**data)
assert response["fine_tuning_job_id"] == "ftjob-jTBys7bVsbyZDOwL9GlpYqXR"

View file

@ -392,6 +392,9 @@ def test_select_azure_base_url_called(setup_mocks):
"arun_thread_stream",
"aresponses",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
"alist_fine_tuning_jobs",
"aretrieve_fine_tuning_job",
]
],
)