mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
546a508c8c
commit
58f958f30a
13 changed files with 221 additions and 39 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -391,6 +391,7 @@ def test_select_azure_base_url_called(setup_mocks):
|
|||
"add_message",
|
||||
"arun_thread_stream",
|
||||
"aresponses",
|
||||
"acreate_fine_tuning_job",
|
||||
]
|
||||
],
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue