mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(openai): forward project to Batches and Files API clients
OpenAI project scoping (OpenAI-Project header / client project=) was dropped by the Batches and Files API paths: the OpenAI SDK client was constructed without the project argument, so multi-project keys failed with "Cannot find file file-XXX, or project proj_YYY does not have access to it". - GenericLiteLLMParams accepts project alongside organization - OpenAIFilesAPI / OpenAIBatchesAPI get_openai_client forwards project to OpenAI(**kwargs); every sync method takes a project parameter - batches/files main resolve project from optional_params or OPENAI_PROJECT and pass it through get_openai_credentials - proxy /batches and /files accept project in the body and fall back to the OpenAI-Project header - unit tests cover param/env resolution and client construction Fixes #41803
This commit is contained in:
parent
d6115679bc
commit
ea94d40e55
8 changed files with 222 additions and 0 deletions
|
|
@ -293,6 +293,7 @@ def create_batch(
|
|||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
|
||||
)
|
||||
project: Final = optional_params.project or os.getenv("OPENAI_PROJECT", None) or None
|
||||
# set API KEY
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -305,6 +306,7 @@ def create_batch(
|
|||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
organization=organization,
|
||||
project=project,
|
||||
create_batch_data=_create_batch_request,
|
||||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
|
|
@ -455,6 +457,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
|
||||
)
|
||||
project: Final = optional_params.project or os.getenv("OPENAI_PROJECT", None) or None
|
||||
# set API KEY
|
||||
api_key = (
|
||||
optional_params.api_key
|
||||
|
|
@ -469,6 +472,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
organization=organization,
|
||||
project=project,
|
||||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
|
|
@ -807,6 +811,7 @@ def list_batches(
|
|||
or os.getenv("OPENAI_ORGANIZATION", None)
|
||||
or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
|
||||
)
|
||||
project: Final = optional_params.project or os.getenv("OPENAI_PROJECT", None) or None
|
||||
|
||||
response = openai_batches_instance.list_batches(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -815,6 +820,7 @@ def list_batches(
|
|||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
organization=organization,
|
||||
project=project,
|
||||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
|
|
@ -1004,6 +1010,7 @@ def cancel_batch(
|
|||
organization: Final = (
|
||||
optional_params.organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None
|
||||
)
|
||||
project: Final = optional_params.project or os.getenv("OPENAI_PROJECT", None) or None
|
||||
api_key = optional_params.api_key or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY")
|
||||
|
||||
response = openai_batches_instance.cancel_batch(
|
||||
|
|
@ -1012,6 +1019,7 @@ def cancel_batch(
|
|||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
organization=organization,
|
||||
project=project,
|
||||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -253,6 +253,7 @@ def create_file(
|
|||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
project=optional_params.project,
|
||||
)
|
||||
response = openai_files_instance.create_file(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -261,6 +262,7 @@ def create_file(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
organization=openai_creds.organization,
|
||||
project=openai_creds.project,
|
||||
create_file_data=_create_file_request,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
|
|
@ -374,6 +376,7 @@ def file_retrieve(
|
|||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
project=optional_params.project,
|
||||
)
|
||||
response = openai_files_instance.retrieve_file(
|
||||
file_id=file_id,
|
||||
|
|
@ -383,6 +386,7 @@ def file_retrieve(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
organization=openai_creds.organization,
|
||||
project=openai_creds.project,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
|
|
@ -555,6 +559,7 @@ def file_delete(
|
|||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
project=optional_params.project,
|
||||
)
|
||||
response = openai_files_instance.delete_file(
|
||||
file_id=file_id,
|
||||
|
|
@ -564,6 +569,7 @@ def file_delete(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
organization=openai_creds.organization,
|
||||
project=openai_creds.project,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
|
|
@ -761,6 +767,7 @@ def file_list(
|
|||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
project=optional_params.project,
|
||||
)
|
||||
response = openai_files_instance.list_files(
|
||||
purpose=purpose,
|
||||
|
|
@ -770,6 +777,7 @@ def file_list(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
organization=openai_creds.organization,
|
||||
project=openai_creds.project,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
|
|
@ -958,6 +966,7 @@ def file_content(
|
|||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
project=optional_params.project,
|
||||
)
|
||||
response = openai_files_instance.file_content(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -967,6 +976,7 @@ def file_content(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
organization=openai_creds.organization,
|
||||
project=openai_creds.project,
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
azure_creds: Final = get_azure_credentials(
|
||||
|
|
@ -1077,6 +1087,7 @@ def file_content_streaming(
|
|||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
organization=optional_params.organization,
|
||||
project=optional_params.project,
|
||||
)
|
||||
response = openai_files_instance.file_content_streaming(
|
||||
_is_async=_is_async,
|
||||
|
|
@ -1086,6 +1097,7 @@ def file_content_streaming(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
organization=openai_creds.organization,
|
||||
project=openai_creds.project,
|
||||
chunk_size=chunk_size,
|
||||
client=client if isinstance(client, (OpenAI, AsyncOpenAI)) else None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -353,12 +353,14 @@ class OpenAICredentials(NamedTuple):
|
|||
api_base: str
|
||||
api_key: str | None
|
||||
organization: str | None
|
||||
project: str | None = None
|
||||
|
||||
|
||||
def get_openai_credentials(
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
organization: str | None = None,
|
||||
project: str | None = None,
|
||||
) -> OpenAICredentials:
|
||||
"""Resolve OpenAI credentials from params, litellm globals, and env vars."""
|
||||
resolved_api_base: Final = (
|
||||
|
|
@ -369,9 +371,11 @@ def get_openai_credentials(
|
|||
or "https://api.openai.com/v1"
|
||||
)
|
||||
resolved_organization = organization or litellm.organization or os.getenv("OPENAI_ORGANIZATION", None) or None
|
||||
resolved_project = project or os.getenv("OPENAI_PROJECT", None) or None
|
||||
resolved_api_key: Final = api_key or litellm.api_key or litellm.openai_key or os.getenv("OPENAI_API_KEY")
|
||||
return OpenAICredentials(
|
||||
api_base=resolved_api_base,
|
||||
api_key=resolved_api_key,
|
||||
organization=resolved_organization,
|
||||
project=resolved_project,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1689,6 +1689,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
_is_async: bool = False,
|
||||
) -> OpenAI | AsyncOpenAI | None:
|
||||
|
|
@ -1729,6 +1730,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
) -> OpenAIFileObject | Coroutine[None, None, OpenAIFileObject]:
|
||||
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
|
||||
|
|
@ -1737,6 +1739,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -1771,6 +1774,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
) -> HttpxBinaryResponseContent | Coroutine[None, None, HttpxBinaryResponseContent]:
|
||||
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
|
||||
|
|
@ -1779,6 +1783,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -1835,6 +1840,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
chunk_size: int = 1024 * 1024,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
) -> FileContentStreamingResult:
|
||||
|
|
@ -1844,6 +1850,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -1899,6 +1906,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
):
|
||||
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
|
||||
|
|
@ -1907,6 +1915,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -1945,6 +1954,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
):
|
||||
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
|
||||
|
|
@ -1953,6 +1963,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -1993,6 +2004,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
purpose: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
):
|
||||
|
|
@ -2002,6 +2014,7 @@ class OpenAIFilesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -2047,6 +2060,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
_is_async: bool = False,
|
||||
) -> OpenAI | AsyncOpenAI | None:
|
||||
|
|
@ -2087,6 +2101,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | AsyncOpenAI | None = None,
|
||||
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
|
||||
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
|
||||
|
|
@ -2095,6 +2110,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -2131,6 +2147,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | None = None,
|
||||
):
|
||||
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
|
||||
|
|
@ -2139,6 +2156,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -2174,6 +2192,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
client: OpenAI | None = None,
|
||||
):
|
||||
openai_client: Final[OpenAI | AsyncOpenAI | None] = self.get_openai_client(
|
||||
|
|
@ -2182,6 +2201,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
@ -2221,6 +2241,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout: float | httpx.Timeout,
|
||||
max_retries: int | None,
|
||||
organization: str | None,
|
||||
project: str | None = None,
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
client: OpenAI | None = None,
|
||||
|
|
@ -2231,6 +2252,7 @@ class OpenAIBatchesAPI(BaseLLM):
|
|||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
project=project,
|
||||
client=client,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -268,6 +268,8 @@ async def create_batch(
|
|||
"Request received by LiteLLM:\n%s",
|
||||
json.dumps(data, indent=4),
|
||||
)
|
||||
if data.get("project") is None and request.headers.get("OpenAI-Project"):
|
||||
data["project"] = request.headers.get("OpenAI-Project")
|
||||
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
(
|
||||
data,
|
||||
|
|
|
|||
|
|
@ -758,6 +758,9 @@ async def create_file(
|
|||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
if data.get("project") is None and request.headers.get("OpenAI-Project"):
|
||||
data["project"] = request.headers.get("OpenAI-Project")
|
||||
|
||||
uploaded_file_info: Final[UploadedFileInfo] = {
|
||||
"filename": file.filename,
|
||||
"content_type": file.content_type,
|
||||
|
|
|
|||
|
|
@ -490,6 +490,7 @@ class CreateFileRequest(TypedDict, total=False):
|
|||
file: Required[FileTypes]
|
||||
purpose: Required[CREATE_FILE_REQUESTS_PURPOSE]
|
||||
expires_after: FileExpiresAfter | None
|
||||
project: str | None
|
||||
extra_headers: dict[str, str] | None
|
||||
extra_body: dict[str, str] | None
|
||||
timeout: float | None
|
||||
|
|
@ -546,6 +547,7 @@ class CreateBatchRequest(TypedDict, total=False):
|
|||
class LiteLLMBatchCreateRequest(CreateBatchRequest, total=False):
|
||||
model: str
|
||||
disable_fallbacks: ReadOnly[bool]
|
||||
project: str | None
|
||||
|
||||
|
||||
class RetrieveBatchRequest(TypedDict, total=False):
|
||||
|
|
|
|||
169
tests/batches_tests/test_openai_project_passthrough.py
Normal file
169
tests/batches_tests/test_openai_project_passthrough.py
Normal file
|
|
@ -0,0 +1,169 @@
|
|||
# What is this?
|
||||
## Unit tests: OpenAI `project` passthrough for the Batches and Files APIs.
|
||||
## https://github.com/BerriAI/litellm/issues/41803
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.openai.common_utils import get_openai_credentials
|
||||
|
||||
|
||||
def test_get_openai_credentials_project_param():
|
||||
creds = get_openai_credentials(api_key="sk-test", organization="org-abc", project="proj-xyz")
|
||||
assert creds.organization == "org-abc"
|
||||
assert creds.project == "proj-xyz"
|
||||
|
||||
|
||||
def test_get_openai_credentials_project_env_fallback(monkeypatch):
|
||||
monkeypatch.setenv("OPENAI_PROJECT", "proj-from-env")
|
||||
creds = get_openai_credentials(api_key="sk-test")
|
||||
assert creds.project == "proj-from-env"
|
||||
|
||||
|
||||
def test_openai_batches_api_create_batch_forwards_project():
|
||||
from litellm.llms.openai.openai import OpenAIBatchesAPI
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
batch_response = LiteLLMBatch.model_validate(
|
||||
{
|
||||
"id": "batch_abc123",
|
||||
"object": "batch",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"input_file_id": "file-abc123",
|
||||
"status": "validating",
|
||||
"completion_window": "24h",
|
||||
"created_at": 1700000000,
|
||||
}
|
||||
)
|
||||
with patch("litellm.llms.openai.openai.OpenAI") as mock_openai:
|
||||
mock_openai.return_value.batches.create.return_value = batch_response
|
||||
OpenAIBatchesAPI().create_batch(
|
||||
_is_async=False,
|
||||
create_batch_data={
|
||||
"completion_window": "24h",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"input_file_id": "file-abc123",
|
||||
},
|
||||
api_key="sk-test",
|
||||
api_base=None,
|
||||
timeout=600.0,
|
||||
max_retries=2,
|
||||
organization="org-abc",
|
||||
project="proj-xyz",
|
||||
)
|
||||
|
||||
kwargs = mock_openai.call_args.kwargs
|
||||
assert kwargs["project"] == "proj-xyz"
|
||||
assert kwargs["organization"] == "org-abc"
|
||||
|
||||
|
||||
def test_openai_batches_api_create_batch_omits_project_when_none():
|
||||
from litellm.llms.openai.openai import OpenAIBatchesAPI
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
batch_response = LiteLLMBatch.model_validate(
|
||||
{
|
||||
"id": "batch_abc123",
|
||||
"object": "batch",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"input_file_id": "file-abc123",
|
||||
"status": "validating",
|
||||
"completion_window": "24h",
|
||||
"created_at": 1700000000,
|
||||
}
|
||||
)
|
||||
with patch("litellm.llms.openai.openai.OpenAI") as mock_openai:
|
||||
mock_openai.return_value.batches.create.return_value = batch_response
|
||||
OpenAIBatchesAPI().create_batch(
|
||||
_is_async=False,
|
||||
create_batch_data={
|
||||
"completion_window": "24h",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"input_file_id": "file-abc123",
|
||||
},
|
||||
api_key="sk-test",
|
||||
api_base=None,
|
||||
timeout=600.0,
|
||||
max_retries=2,
|
||||
organization=None,
|
||||
project=None,
|
||||
)
|
||||
|
||||
kwargs = mock_openai.call_args.kwargs
|
||||
assert "project" not in kwargs
|
||||
assert "organization" not in kwargs
|
||||
|
||||
|
||||
def test_openai_files_api_retrieve_file_forwards_project():
|
||||
from litellm.llms.openai.openai import OpenAIFilesAPI
|
||||
|
||||
with patch("litellm.llms.openai.openai.OpenAI") as mock_openai:
|
||||
OpenAIFilesAPI().retrieve_file(
|
||||
_is_async=False,
|
||||
file_id="file-abc123",
|
||||
api_base=None,
|
||||
api_key="sk-test",
|
||||
timeout=600.0,
|
||||
max_retries=2,
|
||||
organization=None,
|
||||
project="proj-xyz",
|
||||
)
|
||||
|
||||
kwargs = mock_openai.call_args.kwargs
|
||||
assert kwargs["project"] == "proj-xyz"
|
||||
|
||||
|
||||
def test_litellm_create_batch_project_param(monkeypatch):
|
||||
from litellm.batches import main as batches_main
|
||||
|
||||
monkeypatch.delenv("OPENAI_PROJECT", raising=False)
|
||||
mock_instance = MagicMock()
|
||||
with patch.object(batches_main, "openai_batches_instance", mock_instance):
|
||||
litellm.create_batch(
|
||||
completion_window="24h",
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-abc123",
|
||||
custom_llm_provider="openai",
|
||||
api_key="sk-test",
|
||||
project="proj-xyz",
|
||||
)
|
||||
|
||||
kwargs = mock_instance.create_batch.call_args.kwargs
|
||||
assert kwargs["project"] == "proj-xyz"
|
||||
|
||||
|
||||
def test_litellm_create_batch_project_env_fallback(monkeypatch):
|
||||
from litellm.batches import main as batches_main
|
||||
|
||||
monkeypatch.setenv("OPENAI_PROJECT", "proj-from-env")
|
||||
mock_instance = MagicMock()
|
||||
with patch.object(batches_main, "openai_batches_instance", mock_instance):
|
||||
litellm.create_batch(
|
||||
completion_window="24h",
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-abc123",
|
||||
custom_llm_provider="openai",
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
kwargs = mock_instance.create_batch.call_args.kwargs
|
||||
assert kwargs["project"] == "proj-from-env"
|
||||
|
||||
|
||||
def test_litellm_file_retrieve_project_param(monkeypatch):
|
||||
from litellm.files import main as files_main
|
||||
|
||||
monkeypatch.delenv("OPENAI_PROJECT", raising=False)
|
||||
mock_instance = MagicMock()
|
||||
with patch.object(files_main, "openai_files_instance", mock_instance):
|
||||
files_main.file_retrieve(
|
||||
file_id="file-abc123",
|
||||
custom_llm_provider="openai",
|
||||
api_key="sk-test",
|
||||
project="proj-xyz",
|
||||
)
|
||||
|
||||
kwargs = mock_instance.retrieve_file.call_args.kwargs
|
||||
assert kwargs["project"] == "proj-xyz"
|
||||
Loading…
Add table
Reference in a new issue