XFEstudio/gpt4free
feat: add image generation support to Azure provider
- Updated `Azure.create_completion` to support media uploads and image generation via `/images/` endpoint - Added `media` parameter to `Azure.create_completion` and handled image-related request formatting - Imported `StreamSession`, `FormData`, `raise_for_status`, `get_width_height`, `to_bytes`, `save_response_media`, and `format_media_prompt` in `Azure.py` - Modified `get_models` to load `AZURE_API_KEYS` from environment and parse it into `cls.api_keys` - Adjusted `get_width_height` in `image/__init__.py` to return higher default resolutions for "16:9" and "9:16" aspect ratios - Modified `save_response_media` in `image/copy_images.py` to accept optional `content_type` parameter and use it when provided - Updated `FormData` class logic in `requests/curl_cffi.py` to define it only when `has_curl_mime` is True and raise an error otherwise
00bd517f
代码差异
@@ -3,9 +3,13 @@ from __future__ import annotations
import os
import json
from ...typing import Messages, AsyncResult
from ...typing import Messages, AsyncResult, MediaListType
from ...errors import MissingAuthError, ModelNotFoundError
from ...requests import StreamSession, FormData, raise_for_status
from ...image import get_width_height, to_bytes
from ...image.copy_images import save_response_media
from ..template import OpenaiTemplate
from ..helper import format_media_prompt
class Azure(OpenaiTemplate):
url = "https://ai.azure.com"
@@ -26,9 +30,16 @@ class Azure(OpenaiTemplate):
"modalities": ["text", "audio"],
}
}
api_keys: dict[str, str] = {}
@classmethod
def get_models(cls, api_key: str = None, **kwargs) -> list[str]:
api_keys = os.environ.get("AZURE_API_KEYS")
if api_keys:
try:
cls.api_keys = json.loads(api_keys)
except json.JSONDecodeError:
raise ValueError(f"Invalid AZURE_API_KEYS environment variable")
routes = os.environ.get("AZURE_ROUTES")
if routes:
try:
@@ -46,6 +57,7 @@ class Azure(OpenaiTemplate):
model: str,
messages: Messages,
stream: bool = True,
media: MediaListType = None,
extra_body: dict = None,
api_key: str = None,
api_endpoint: str = None,
@@ -53,8 +65,6 @@ class Azure(OpenaiTemplate):
) -> AsyncResult:
if not model:
model = os.environ.get("AZURE_DEFAULT_MODEL", cls.default_model)
if not api_key:
raise ValueError(f"API key is required for Azure provider. Ask for API key in the {cls.login_url} Discord server.")
if not api_endpoint:
if not cls.routes:
cls.get_models()
@@ -63,6 +73,45 @@ class Azure(OpenaiTemplate):
raise ModelNotFoundError(f"No API endpoint found for model: {model}")
if not api_endpoint:
api_endpoint = os.environ.get("AZURE_API_ENDPOINT")
if not api_key:
api_key = cls.api_keys.get(model, cls.api_keys.get("default"))
if not api_key:
raise ValueError(f"API key is required for Azure provider. Ask for API key in the {cls.login_url} Discord server.")
if "/images/" in api_endpoint:
prompt = format_media_prompt(messages, kwargs.get("prompt"))
width, height = get_width_height(kwargs.get("aspect_ratio", "1:1"), kwargs.get("width"), kwargs.get("height"))
output_format = kwargs.get("output_format", "webp")
form = None
data = None
if media:
form = FormData()
form.add_field("prompt", prompt)
form.add_field("size", f"{width}x{height}")
output_format = "png"
for i in range(len(media)):
if media[i][1] is None and isinstance(media[i][0], str):
media[i] = media[i][0], os.path.basename(media[i][0])
media[i] = (to_bytes(media[i][0]), media[i][1])
for image, image_name in media:
form.add_field(f"image", image, filename=image_name)
else:
data = {
"prompt": prompt,
"n": 1,
"size": f"{width}x{height}",
"output_format": output_format,
}
async with StreamSession(proxy=kwargs.get("proxy"), headers={
"Authorization": f"Bearer {api_key}",
"x-ms-model-mesh-model-name": model,
}) as session:
async with session.post(api_endpoint, data=form, json=data) as response:
data = await response.json()
cls.raise_error(data, response.status)
await raise_for_status(response, data)
async for chunk in save_response_media(data["data"][0]["b64_json"], prompt, content_type=f"image/{output_format}"):
yield chunk
return
if extra_body is None:
if model in cls.model_extra_body:
extra_body = cls.model_extra_body[model]
@@ -76,6 +125,7 @@ class Azure(OpenaiTemplate):
model=model,
messages=messages,
stream=stream,
media=media,
api_key=api_key,
api_endpoint=api_endpoint,
extra_body=extra_body,
@@ -324,9 +324,9 @@ def get_width_height(
if aspect_ratio == "1:1":
return width or 1024, height or 1024
elif aspect_ratio == "16:9":
return width or 832, height or 480
return width or 1792, height or 1024
elif aspect_ratio == "9:16":
return width or 480, height or 832,
return width or 1024, height or 1792
return width, height
class ImageRequest:
@@ -60,15 +60,15 @@ def update_filename(response, filename: str) -> str:
timestamp = datetime.strptime(date, '%a, %d %b %Y %H:%M:%S %Z').timestamp()
return str(int(timestamp)) + "_" + filename.split("_", maxsplit=1)[-1]
async def save_response_media(response, prompt: str, tags: list[str] = [], transcript: str = None) -> AsyncIterator:
async def save_response_media(response, prompt: str, tags: list[str] = [], transcript: str = None, content_type: str = None) -> AsyncIterator:
"""Save media from response to local file and return URL"""
if isinstance(response, dict):
content_type = response.get("mimeType", "audio/mpeg")
content_type = response.get("mimeType", content_type or "audio/mpeg")
transcript = response.get("transcript")
response = response.get("data")
elif hasattr(response, "headers"):
content_type = response.headers["content-type"]
else:
content_type = response.headers.get("content-type", content_type)
elif not content_type:
raise ValueError("Response must be a dict or have headers")
if isinstance(response, str):
@@ -107,14 +107,14 @@ class StreamSession(AsyncSession):
delete = partialmethod(request, "DELETE")
options = partialmethod(request, "OPTIONS")
if has_curl_mime:
class FormData(CurlMime):
def add_field(self, name, data=None, content_type: str = None, filename: str = None) -> None:
self.addpart(name, content_type=content_type, filename=filename, data=data)
else:
if not has_curl_mime:
class FormData():
def __init__(self) -> None:
raise RuntimeError("CurlMimi in curl_cffi is missing | pip install -U curl_cffi")
else:
class FormData(CurlMime):
def add_field(self, name, data=None, content_type: str = None, filename: str = None) -> None:
self.addpart(name, content_type=content_type, filename=filename, data=data)
class WebSocket():
def __init__(self, session, url, **kwargs) -> None: