Litellm dev 05 21 2025 p2 (#11039)

* 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

* test: update test

* fix: fix linting error

* fix: fix ruff linting error

* test: fix check
This commit is contained in:
Krish Dholakia 2025-05-21 21:40:53 -07:00 • committed by GitHub
parent 546a508c8c
commit 58f958f30a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 221 additions and 39 deletions

View file

@ -23,7 +23,12 @@ from litellm.types.llms.openai import (
OpenAIFileObject,
OpenAIFilesPurpose,
)
from litellm.types.utils import LiteLLMBatch, LLMResponseTypes, SpecialEnums
from litellm.types.utils import (
LiteLLMBatch,
LiteLLMFineTuningJob,
LLMResponseTypes,
SpecialEnums,
)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -138,12 +143,16 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
"acreate_batch",
"aretrieve_batch",
"afile_content",
"acreate_fine_tuning_job",
],
) -> Union[Exception, str, Dict, None]:
"""
- Detect litellm_proxy/ file_id
- add dictionary of mappings of litellm_proxy/ file_id -> provider_file_id => {litellm_proxy/file_id: {"model_id": id, "file_id": provider_file_id}}
"""
print(
"CALLS ASYNC PRE CALL HOOK - DATA={}, CALL_TYPE={}".format(data, call_type)
)
if call_type == CallTypes.completion.value:
messages = data.get("messages")
if messages:
@ -196,7 +205,15 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
data["batch_id"] = self.get_batch_id_from_unified_batch_id(
potential_batch_id
)
elif call_type == CallTypes.acreate_fine_tuning_job.value:
input_file_id = cast(Optional[str], data.get("training_file"))
if input_file_id:
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
async def async_pre_call_deployment_hook(
@ -205,8 +222,21 @@ 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:
input_file_id = cast(Optional[str], kwargs.get("input_file_id"))
accessor_key = "input_file_id"
elif call_type and call_type == CallTypes.acreate_fine_tuning_job:
accessor_key = "training_file"
else:
return kwargs
if accessor_key:
input_file_id = cast(Optional[str], kwargs.get(accessor_key))
model_file_id_mapping = cast(
Optional[Dict[str, Dict[str, str]]], kwargs.get("model_file_id_mapping")
)
@ -217,7 +247,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_id, None
)
if mapped_file_id:
kwargs["input_file_id"] = mapped_file_id
kwargs[accessor_key] = mapped_file_id
return kwargs
def get_file_ids_from_messages(self, messages: List[AllMessageValues]) -> List[str]:
@ -383,6 +414,20 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return response
def get_unified_generic_response_id(
self, model_id: str, generic_response_id: str
) -> str:
unified_generic_response_id = (
SpecialEnums.LITELLM_MANAGED_GENERIC_RESPONSE_COMPLETE_STR.value.format(
model_id, generic_response_id
)
)
return (
base64.urlsafe_b64encode(unified_generic_response_id.encode())
.decode()
.rstrip("=")
)
def get_unified_batch_id(self, batch_id: str, model_id: str) -> str:
unified_batch_id = SpecialEnums.LITELLM_MANAGED_BATCH_COMPLETE_STR.value.format(
model_id, batch_id
@ -455,7 +500,21 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
model_id=model_id,
model_name=model_name,
)
return response
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
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:
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
)

View file

@ -28,6 +28,7 @@ from litellm.types.llms.openai import (
Hyperparameters,
)
from litellm.types.router import *
from litellm.types.utils import LiteLLMFineTuningJob
from litellm.utils import client, supports_httpx_timeout
####### ENVIRONMENT VARIABLES ###################
@ -50,7 +51,7 @@ async def acreate_fine_tuning_job(
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,
) -> FineTuningJob:
) -> LiteLLMFineTuningJob:
"""
Async: Creates and executes a batch from an uploaded file of request
@ -104,7 +105,7 @@ def create_fine_tuning_job(
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]]:
"""
Creates a fine-tuning job which begins the process of creating a new model from a given dataset.

View file

@ -1,4 +1,4 @@
from typing import Optional, Union
from typing import List, Optional, Union
import httpx
@ -67,7 +67,7 @@ class OllamaModelInfo(BaseLLMModelInfo):
# env var OLLAMA_API_BASE or default
return api_base or get_secret_str("OLLAMA_API_BASE") or "http://localhost:11434"
def get_models(self, api_key=None, api_base: Optional[str] = None) -> list[str]:
def get_models(self, api_key=None, api_base: Optional[str] = None) -> List[str]:
"""
List all models available on the Ollama server via /api/tags endpoint.
"""

View file

@ -1,10 +1,11 @@
from typing import Any, Coroutine, Optional, Union
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
class OpenAIFineTuningAPI:
@ -55,11 +56,12 @@ class OpenAIFineTuningAPI:
self,
create_fine_tuning_job_data: dict,
openai_client: Union[AsyncOpenAI, AsyncAzureOpenAI],
) -> FineTuningJob:
) -> LiteLLMFineTuningJob:
response = await openai_client.fine_tuning.jobs.create(
**create_fine_tuning_job_data
)
return response
return LiteLLMFineTuningJob(**response.model_dump())
def create_fine_tuning_job(
self,
@ -74,7 +76,7 @@ class OpenAIFineTuningAPI:
client: Optional[
Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
] = None,
) -> Union[FineTuningJob, Coroutine[Any, Any, FineTuningJob]]:
) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
openai_client: Optional[
Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]
] = self.get_openai_client(
@ -104,8 +106,10 @@ class OpenAIFineTuningAPI:
verbose_logger.debug(
"creating fine tuning job, args= %s", create_fine_tuning_job_data
)
response = openai_client.fine_tuning.jobs.create(**create_fine_tuning_job_data)
return response
response = cast(OpenAI, openai_client).fine_tuning.jobs.create(
**create_fine_tuning_job_data
)
return LiteLLMFineTuningJob(**response.model_dump())
async def acancel_fine_tuning_job(
self,

View file

@ -1,10 +1,9 @@
import json
import traceback
from datetime import datetime
from typing import Literal, Optional, Union
from typing import Any, Coroutine, Literal, Optional, Union
import httpx
from openai.types.fine_tuning.fine_tuning_job import FineTuningJob
import litellm
from litellm._logging import verbose_logger
@ -20,6 +19,7 @@ from litellm.types.llms.vertex_ai import (
ResponseSupervisedTuningSpec,
ResponseTuningJob,
)
from litellm.types.utils import LiteLLMFineTuningJob
class VertexFineTuningAPI(VertexLLM):
@ -113,7 +113,7 @@ class VertexFineTuningAPI(VertexLLM):
def convert_vertex_response_to_open_ai_response(
self, response: ResponseTuningJob
) -> FineTuningJob:
) -> LiteLLMFineTuningJob:
status: Literal[
"validating_files", "queued", "running", "succeeded", "failed", "cancelled"
] = "queued"
@ -134,7 +134,7 @@ class VertexFineTuningAPI(VertexLLM):
response.get("supervisedTuningSpec", None) or {}
)
training_uri: str = _supervisedTuningSpec.get("trainingDatasetUri", "") or ""
return FineTuningJob(
return LiteLLMFineTuningJob(
id=response.get("name", "") or "",
created_at=created_at,
fine_tuned_model=response.get("tunedModelDisplayName", ""),
@ -226,7 +226,7 @@ class VertexFineTuningAPI(VertexLLM):
timeout: Union[float, httpx.Timeout],
kwargs: Optional[dict] = None,
original_hyperparameters: Optional[dict] = {},
):
) -> Union[LiteLLMFineTuningJob, Coroutine[Any, Any, LiteLLMFineTuningJob]]:
verbose_logger.debug(
"creating fine tuning job, args= %s", create_fine_tuning_job_data
)

View file

@ -2,9 +2,9 @@ model_list:
- model_name: "gemini-2.0-flash"
litellm_params:
model: gemini/gemini-2.0-flash-live-001
- model_name: "gpt-4o-mini-openai"
- model_name: "gpt-4.1-openai"
litellm_params:
model: gpt-4o-mini
model: gpt-4.1-mini-2025-04-14
api_key: os.environ/OPENAI_API_KEY
model_info:
access_groups: ["default-openai-models"]

View file

@ -117,6 +117,7 @@ class ProxyBaseLLMRequestProcessing:
"acreate_batch",
"aretrieve_batch",
"afile_content",
"acreate_fine_tuning_job",
],
version: Optional[str] = None,
user_model: Optional[str] = None,

View file

@ -7,16 +7,20 @@
import asyncio
import traceback
from typing import Optional
from typing import Optional, cast
from fastapi import APIRouter, Depends, Request, Response
from fastapi import APIRouter, Depends, HTTPException, Request, Response
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
from litellm.proxy.utils import handle_exception_on_proxy
from litellm.types.utils import LiteLLMFineTuningJob
router = APIRouter()
@ -96,8 +100,8 @@ async def create_fine_tuning_job(
```
"""
from litellm.proxy.proxy_server import (
add_litellm_data_to_request,
general_settings,
llm_router,
premium_user,
proxy_config,
proxy_logging_obj,
@ -117,25 +121,68 @@ async def create_fine_tuning_job(
)
# 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="acreate_fine_tuning_job",
)
# get configs for custom_llm_provider
llm_provider_config = get_fine_tuning_provider_config(
custom_llm_provider=fine_tuning_request.custom_llm_provider,
## CHECK IF MANAGED FILE ID
unified_file_id: Union[str, Literal[False]] = False
training_file = fine_tuning_request.training_file
response: Optional[LiteLLMFineTuningJob] = None
if training_file:
unified_file_id = _is_base64_encoded_unified_file_id(training_file)
## IF SO, Route based on that
if unified_file_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.acreate_fine_tuning_job(**data)
)
response.training_file = unified_file_id
response._hidden_params["unified_file_id"] = unified_file_id
## ELSE, Route based on custom_llm_provider
elif fine_tuning_request.custom_llm_provider:
# get configs for custom_llm_provider
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)
response = await litellm.acreate_fine_tuning_job(**data)
if response is None:
raise ValueError(
"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,
)
# add llm_provider_config to data
if llm_provider_config is not None:
data.update(llm_provider_config)
response = await litellm.acreate_fine_tuning_job(**data)
if _response is not None and isinstance(_response, LiteLLMFineTuningJob):
response = _response
### ALERTING ###
asyncio.create_task(
@ -166,12 +213,11 @@ async def create_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(
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.create_fine_tuning_job(): Exception occurred - {}".format(
str(e)
)
)
verbose_proxy_logger.debug(traceback.format_exc())
raise handle_exception_on_proxy(e)

View file

@ -752,6 +752,9 @@ class Router:
self._arealtime = self.factory_function(
litellm._arealtime, call_type="_arealtime"
)
self.acreate_fine_tuning_job = self.factory_function(
litellm.acreate_fine_tuning_job, call_type="acreate_fine_tuning_job"
)
def validate_fallbacks(self, fallback_param: Optional[List]):
"""
@ -3159,6 +3162,7 @@ class Router:
"afile_delete",
"afile_content",
"_arealtime",
"acreate_fine_tuning_job",
] = "assistants",
):
"""
@ -3207,6 +3211,7 @@ class Router:
"anthropic_messages",
"aresponses",
"_arealtime",
"acreate_fine_tuning_job",
):
return await self._ageneric_api_call_with_fallbacks(
original_function=original_function,

View file

@ -878,7 +878,7 @@ class FineTuningJobCreate(BaseModel):
class LiteLLMFineTuningJobCreate(FineTuningJobCreate):
custom_llm_provider: Literal["openai", "azure", "vertex_ai"]
custom_llm_provider: Optional[Literal["openai", "azure", "vertex_ai"]] = None
model_config = {
"extra": "allow"

View file

@ -44,6 +44,7 @@ from .llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionUsageBlock,
FileSearchTool,
FineTuningJob,
OpenAIChatCompletionChunk,
OpenAIFileObject,
OpenAIRealtimeStreamList,
@ -2256,6 +2257,18 @@ class SelectTokenizerResponse(TypedDict):
tokenizer: Any
class LiteLLMFineTuningJob(FineTuningJob):
_hidden_params: dict = {}
def __init__(self, **kwargs):
if "error" in kwargs and kwargs["error"] is not None:
# check if error is all None - if so, set error to None
if all(value is None for value in kwargs["error"].values()):
kwargs["error"] = None
super().__init__(**kwargs)
self._hidden_params = kwargs.get("_hidden_params", {})
class LiteLLMBatch(Batch):
_hidden_params: dict = {}
usage: Optional[Usage] = None
@ -2360,9 +2373,16 @@ class SpecialEnums(Enum):
LITELLM_MANAGED_BATCH_COMPLETE_STR = "litellm_proxy;model_id:{};llm_batch_id:{}"
LITELLM_MANAGED_GENERIC_RESPONSE_COMPLETE_STR = "litellm_proxy;model_id:{};generic_response_id:{}" # generic implementation of 'managed batches' - used for finetuning and any future work.
LLMResponseTypes = Union[
ModelResponse, EmbeddingResponse, ImageResponse, OpenAIFileObject, LiteLLMBatch
ModelResponse,
EmbeddingResponse,
ImageResponse,
OpenAIFileObject,
LiteLLMBatch,
LiteLLMFineTuningJob,
]

View file

@ -13,6 +13,9 @@ from unittest.mock import MagicMock
from enterprise.enterprise_hooks.managed_files import _PROXY_LiteLLMManagedFiles
from litellm.caching import DualCache
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
from litellm.types.utils import SpecialEnums
@ -161,3 +164,45 @@ async def test_async_pre_call_hook_batch_retrieve():
# assert len(batch_files) == 1
# assert assistant_files[0].id == file1.id
# assert batch_files[0].id == file2.id
@pytest.mark.asyncio
async def test_async_post_call_success_hook_for_unified_finetuning_job():
from litellm.types.utils import LiteLLMFineTuningJob
unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxiZTQ0ZDVlYi1mNDU3LTRiNzktOWM4My01N2QxMTMxYWM0YzY7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00LjEtb3BlbmFpO2xsbV9vdXRwdXRfZmlsZV9pZCxmaWxlLURKMnQ0OWZlQ2NTQk5vNG9oekZ6NGc7bGxtX291dHB1dF9maWxlX21vZGVsX2lkLGRiNjY5ODcwNzdkZTdmYzZjNzAzY2Y1MDczMGU2MmNkOWQ3YTU1N2NlNjVmMDUzNTFkYTM4YTA3ZjBlZDEyNzQ"
provider_ft_job = LiteLLMFineTuningJob(
object="fine_tuning.job",
id="ftjob-0kEBV5b4sPrFcMnuzmYSzU1G",
model="gpt-3.5-turbo-0613",
created_at=1692779769,
finished_at=None,
fine_tuned_model=None,
organization_id="org-dUVLhaAQ37YCGwVC2QVY8sdB",
result_files=[],
status="validating_files",
validation_file=None,
training_file="file-azQuKMLAmiFdEjxpCcbI11zF",
hyperparameters={"n_epochs": 8},
trained_tokens=None,
seed=0,
)
provider_ft_job._hidden_params = {
"unified_file_id": unified_file_id,
"model_id": "gpt-3.5-turbo-0613",
}
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=MagicMock()
)
data = {
"user_api_key_dict": {"parent_otel_span": MagicMock()},
}
response = await proxy_managed_files.async_post_call_success_hook(
data=data,
user_api_key_dict=MagicMock(),
response=provider_ft_job,
)
assert isinstance(response, LiteLLMFineTuningJob)
assert _is_base64_encoded_unified_file_id(response.id)

View file

@ -391,6 +391,7 @@ def test_select_azure_base_url_called(setup_mocks):
"add_message",
"arun_thread_stream",
"aresponses",
"acreate_fine_tuning_job",
]
],
)