mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): keep image edit defaults within type-discipline budget and give request mocks a scope
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ed226626fd
commit
57e95cc66f
3 changed files with 8 additions and 3 deletions
|
|
@ -1,6 +1,8 @@
|
|||
import asyncio
|
||||
import io
|
||||
from collections.abc import Sequence
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import Final, get_type_hints
|
||||
|
||||
import orjson
|
||||
|
|
@ -36,6 +38,7 @@ from litellm.types.llms.openai import ChatCompletionUserMessage
|
|||
router: Final = APIRouter()
|
||||
|
||||
IMAGE_EDIT_NUMERIC_FORM_FIELDS: Final = numeric_form_fields(get_type_hints(ImageEditRequestParams))
|
||||
IMAGE_EDIT_OPTIONAL_FIELD_DEFAULTS: Final = MappingProxyType({"prompt": None, "image": None})
|
||||
|
||||
IMAGE_ARRAY_FIELD: Final = "image[]"
|
||||
MASK_ARRAY_FIELD: Final = "mask[]"
|
||||
|
|
@ -299,9 +302,9 @@ async def image_edit_api(
|
|||
numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
|
||||
)
|
||||
data: Final = {
|
||||
"prompt": None,
|
||||
"image": None,
|
||||
**{key: value for key, value in form_fields.items() if key not in BRACKETED_FILE_FIELDS},
|
||||
key: value
|
||||
for key, value in chain(IMAGE_EDIT_OPTIONAL_FIELD_DEFAULTS.items(), form_fields.items())
|
||||
if key not in BRACKETED_FILE_FIELDS
|
||||
}
|
||||
image_files: Final = await batch_to_bytesio(image)
|
||||
mask_files: Final = await batch_to_bytesio(mask)
|
||||
|
|
|
|||
|
|
@ -1034,6 +1034,7 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock:
|
|||
request.url.__str__.return_value = "http://localhost/v1/batches"
|
||||
request.url.path = "/v1/batches"
|
||||
request.method = "POST"
|
||||
request.scope = {"type": "http", "path": "/v1/batches", "method": "POST"}
|
||||
request.query_params = {}
|
||||
request.headers = {"Content-Type": "application/json"}
|
||||
request.client = MagicMock()
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ def mock_request(request):
|
|||
mock_req.headers = Headers({"content-type": "application/json"})
|
||||
mock_req.method = "POST"
|
||||
mock_req.url.path = request.param.get("path")
|
||||
mock_req.scope = {"type": "http", "path": request.param.get("path"), "method": "POST"}
|
||||
|
||||
async def mock_body():
|
||||
return json.dumps(request.param.get("payload", {})).encode("utf-8")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue