XFE Git
XFE Studio Git
Git 首页 全局搜索
XFE 主站 文档 NuGet
公开
关注 0 Fork 0 Star 0
返回提交历史

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
hlohaus <983577+hlohaus@users.noreply.github.com>
提交于

代码差异

4 个文件 +64 -14
Modified g4f/Provider/needs_auth/Azure.py +53 -3
@@ -3,9 +3,13 @@ from __future__ import annotations
3 3 import os
4 4 import json
5 5
6 from ...typing import Messages, AsyncResult
6 from ...typing import Messages, AsyncResult, MediaListType
7 7 from ...errors import MissingAuthError, ModelNotFoundError
8 from ...requests import StreamSession, FormData, raise_for_status
9 from ...image import get_width_height, to_bytes
10 from ...image.copy_images import save_response_media
8 11 from ..template import OpenaiTemplate
12 from ..helper import format_media_prompt
9 13
10 14 class Azure(OpenaiTemplate):
11 15 url = "https://ai.azure.com"
@@ -26,9 +30,16 @@ class Azure(OpenaiTemplate):
26 30 "modalities": ["text", "audio"],
27 31 }
28 32 }
33 api_keys: dict[str, str] = {}
29 34
30 35 @classmethod
31 36 def get_models(cls, api_key: str = None, **kwargs) -> list[str]:
37 api_keys = os.environ.get("AZURE_API_KEYS")
38 if api_keys:
39 try:
40 cls.api_keys = json.loads(api_keys)
41 except json.JSONDecodeError:
42 raise ValueError(f"Invalid AZURE_API_KEYS environment variable")
32 43 routes = os.environ.get("AZURE_ROUTES")
33 44 if routes:
34 45 try:
@@ -46,6 +57,7 @@ class Azure(OpenaiTemplate):
46 57 model: str,
47 58 messages: Messages,
48 59 stream: bool = True,
60 media: MediaListType = None,
49 61 extra_body: dict = None,
50 62 api_key: str = None,
51 63 api_endpoint: str = None,
@@ -53,8 +65,6 @@ class Azure(OpenaiTemplate):
53 65 ) -> AsyncResult:
54 66 if not model:
55 67 model = os.environ.get("AZURE_DEFAULT_MODEL", cls.default_model)
56 if not api_key:
57 raise ValueError(f"API key is required for Azure provider. Ask for API key in the {cls.login_url} Discord server.")
58 68 if not api_endpoint:
59 69 if not cls.routes:
60 70 cls.get_models()
@@ -63,6 +73,45 @@ class Azure(OpenaiTemplate):
63 73 raise ModelNotFoundError(f"No API endpoint found for model: {model}")
64 74 if not api_endpoint:
65 75 api_endpoint = os.environ.get("AZURE_API_ENDPOINT")
76 if not api_key:
77 api_key = cls.api_keys.get(model, cls.api_keys.get("default"))
78 if not api_key:
79 raise ValueError(f"API key is required for Azure provider. Ask for API key in the {cls.login_url} Discord server.")
80 if "/images/" in api_endpoint:
81 prompt = format_media_prompt(messages, kwargs.get("prompt"))
82 width, height = get_width_height(kwargs.get("aspect_ratio", "1:1"), kwargs.get("width"), kwargs.get("height"))
83 output_format = kwargs.get("output_format", "webp")
84 form = None
85 data = None
86 if media:
87 form = FormData()
88 form.add_field("prompt", prompt)
89 form.add_field("size", f"{width}x{height}")
90 output_format = "png"
91 for i in range(len(media)):
92 if media[i][1] is None and isinstance(media[i][0], str):
93 media[i] = media[i][0], os.path.basename(media[i][0])
94 media[i] = (to_bytes(media[i][0]), media[i][1])
95 for image, image_name in media:
96 form.add_field(f"image", image, filename=image_name)
97 else:
98 data = {
99 "prompt": prompt,
100 "n": 1,
101 "size": f"{width}x{height}",
102 "output_format": output_format,
103 }
104 async with StreamSession(proxy=kwargs.get("proxy"), headers={
105 "Authorization": f"Bearer {api_key}",
106 "x-ms-model-mesh-model-name": model,
107 }) as session:
108 async with session.post(api_endpoint, data=form, json=data) as response:
109 data = await response.json()
110 cls.raise_error(data, response.status)
111 await raise_for_status(response, data)
112 async for chunk in save_response_media(data["data"][0]["b64_json"], prompt, content_type=f"image/{output_format}"):
113 yield chunk
114 return
66 115 if extra_body is None:
67 116 if model in cls.model_extra_body:
68 117 extra_body = cls.model_extra_body[model]
@@ -76,6 +125,7 @@ class Azure(OpenaiTemplate):
76 125 model=model,
77 126 messages=messages,
78 127 stream=stream,
128 media=media,
79 129 api_key=api_key,
80 130 api_endpoint=api_endpoint,
81 131 extra_body=extra_body,
Modified g4f/image/__init__.py +2 -2
@@ -324,9 +324,9 @@ def get_width_height(
324 324 if aspect_ratio == "1:1":
325 325 return width or 1024, height or 1024
326 326 elif aspect_ratio == "16:9":
327 return width or 832, height or 480
327 return width or 1792, height or 1024
328 328 elif aspect_ratio == "9:16":
329 return width or 480, height or 832,
329 return width or 1024, height or 1792
330 330 return width, height
331 331
332 332 class ImageRequest:
Modified g4f/image/copy_images.py +4 -4
@@ -60,15 +60,15 @@ def update_filename(response, filename: str) -> str:
60 60 timestamp = datetime.strptime(date, '%a, %d %b %Y %H:%M:%S %Z').timestamp()
61 61 return str(int(timestamp)) + "_" + filename.split("_", maxsplit=1)[-1]
62 62
63 async def save_response_media(response, prompt: str, tags: list[str] = [], transcript: str = None) -> AsyncIterator:
63 async def save_response_media(response, prompt: str, tags: list[str] = [], transcript: str = None, content_type: str = None) -> AsyncIterator:
64 64 """Save media from response to local file and return URL"""
65 65 if isinstance(response, dict):
66 content_type = response.get("mimeType", "audio/mpeg")
66 content_type = response.get("mimeType", content_type or "audio/mpeg")
67 67 transcript = response.get("transcript")
68 68 response = response.get("data")
69 69 elif hasattr(response, "headers"):
70 content_type = response.headers["content-type"]
71 else:
70 content_type = response.headers.get("content-type", content_type)
71 elif not content_type:
72 72 raise ValueError("Response must be a dict or have headers")
73 73
74 74 if isinstance(response, str):
Modified g4f/requests/curl_cffi.py +5 -5
@@ -107,14 +107,14 @@ class StreamSession(AsyncSession):
107 107 delete = partialmethod(request, "DELETE")
108 108 options = partialmethod(request, "OPTIONS")
109 109
110 if has_curl_mime:
111 class FormData(CurlMime):
112 def add_field(self, name, data=None, content_type: str = None, filename: str = None) -> None:
113 self.addpart(name, content_type=content_type, filename=filename, data=data)
114 else:
110 if not has_curl_mime:
115 111 class FormData():
116 112 def __init__(self) -> None:
117 113 raise RuntimeError("CurlMimi in curl_cffi is missing | pip install -U curl_cffi")
114 else:
115 class FormData(CurlMime):
116 def add_field(self, name, data=None, content_type: str = None, filename: str = None) -> None:
117 self.addpart(name, content_type=content_type, filename=filename, data=data)
118 118
119 119 class WebSocket():
120 120 def __init__(self, session, url, **kwargs) -> None: