mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
469d395177
commit
70f32154c5
15 changed files with 444 additions and 113 deletions
|
|
@ -1,4 +1,3 @@
|
|||
import os
|
||||
from typing import Dict, Literal, Type, Union
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
],
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue