返回提交历史
Modified
etc/unittest/__main__.py
+1
-0
Added
etc/unittest/image_client.py
+44
-0
Modified
etc/unittest/mocks.py
+21
-0
Modified
g4f/Provider/AmigoChat.py
+5
-5
Modified
g4f/Provider/Copilot.py
+37
-42
Added
g4f/Provider/needs_auth/MicrosoftDesigner.py
+167
-0
Modified
g4f/Provider/needs_auth/OpenaiChat.py
+2
-2
Modified
g4f/Provider/needs_auth/__init__.py
+26
-24
Modified
g4f/Provider/openai/har_file.py
+1
-3
Modified
g4f/Provider/you/har_file.py
+1
-4
Modified
g4f/client/__init__.py
+71
-49
Modified
g4f/errors.py
+3
-0
Modified
g4f/models.py
+8
-3
Modified
g4f/providers/asyncio.py
+2
-2
Modified
g4f/providers/retry_provider.py
+1
-1
Modified
g4f/typing.py
+2
-1
XFEstudio/gpt4free
IterListProvider support for generating images (#2441)
* IterListProvider support for generating images * Add missing get_har_files import in Copilot * Fix typo in dall-e-3 model name * Add image client unittests * Add MicrosoftDesigner provider * Import MicrosoftDesigner and add it to the model list
79c407b9
代码差异
16 个文件
+392
-136
@@ -5,6 +5,7 @@ from .backend import *
5
5
from .main import *
6
6
from .model import *
7
7
from .client import *
8
from .image_client import *
8
9
from .include import *
9
10
from .retry_provider import *
10
11
@@ -0,0 +1,44 @@
1
from __future__ import annotations
2
3
import asyncio
4
import unittest
5
6
from g4f.client import AsyncClient, ImagesResponse
7
from g4f.providers.retry_provider import IterListProvider
8
from .mocks import (
9
YieldImageResponseProviderMock,
10
MissingAuthProviderMock,
11
AsyncRaiseExceptionProviderMock,
12
YieldNoneProviderMock
13
)
14
15
DEFAULT_MESSAGES = [{'role': 'user', 'content': 'Hello'}]
16
17
class TestIterListProvider(unittest.IsolatedAsyncioTestCase):
18
19
async def test_skip_provider(self):
20
client = AsyncClient(image_provider=IterListProvider([MissingAuthProviderMock, YieldImageResponseProviderMock], False))
21
response = await client.images.generate("Hello", "", response_format="orginal")
22
self.assertIsInstance(response, ImagesResponse)
23
self.assertEqual("Hello", response.data[0].url)
24
25
async def test_only_one_result(self):
26
client = AsyncClient(image_provider=IterListProvider([YieldImageResponseProviderMock, YieldImageResponseProviderMock], False))
27
response = await client.images.generate("Hello", "", response_format="orginal")
28
self.assertIsInstance(response, ImagesResponse)
29
self.assertEqual("Hello", response.data[0].url)
30
31
async def test_skip_none(self):
32
client = AsyncClient(image_provider=IterListProvider([YieldNoneProviderMock, YieldImageResponseProviderMock], False))
33
response = await client.images.generate("Hello", "", response_format="orginal")
34
self.assertIsInstance(response, ImagesResponse)
35
self.assertEqual("Hello", response.data[0].url)
36
37
def test_raise_exception(self):
38
async def run_exception():
39
client = AsyncClient(image_provider=IterListProvider([YieldNoneProviderMock, AsyncRaiseExceptionProviderMock], False))
40
await client.images.generate("Hello", "")
41
self.assertRaises(RuntimeError, asyncio.run, run_exception())
42
43
if __name__ == '__main__':
44
unittest.main()
@@ -1,4 +1,6 @@
1
1
from g4f.providers.base_provider import AbstractProvider, AsyncProvider, AsyncGeneratorProvider
2
from g4f.image import ImageResponse
3
from g4f.errors import MissingAuthError
2
4
3
5
class ProviderMock(AbstractProvider):
4
6
working = True
@@ -41,6 +43,25 @@ class YieldProviderMock(AsyncGeneratorProvider):
41
43
for message in messages:
42
44
yield message["content"]
43
45
46
class YieldImageResponseProviderMock(AsyncGeneratorProvider):
47
working = True
48
49
@classmethod
50
async def create_async_generator(
51
cls, model, messages, stream, prompt: str, **kwargs
52
):
53
yield ImageResponse(prompt, "")
54
55
class MissingAuthProviderMock(AbstractProvider):
56
working = True
57
58
@classmethod
59
def create_completion(
60
cls, model, messages, stream, **kwargs
61
):
62
raise MissingAuthError(cls.__name__)
63
yield cls.__name__
64
44
65
class RaiseExceptionProviderMock(AbstractProvider):
45
66
working = True
46
67
@@ -65,9 +65,9 @@ MODELS = {
65
65
'flux-pro/v1.1-ultra': {'persona_id': "flux-pro-v1.1-ultra"}, # Amigo, your balance is not enough to make the request, wait until 12 UTC or upgrade your plan
66
66
'flux-pro/v1.1-ultra-raw': {'persona_id': "flux-pro-v1.1-ultra-raw"}, # Amigo, your balance is not enough to make the request, wait until 12 UTC or upgrade your plan
67
67
'flux/dev': {'persona_id': "flux-dev"},
68
69
'dalle-e-3': {'persona_id': "dalle-three"},
70
68
69
'dall-e-3': {'persona_id': "dalle-three"},
70
71
71
'recraft-v3': {'persona_id': "recraft"}
72
72
}
73
73
}
@@ -129,8 +129,8 @@ class AmigoChat(AsyncGeneratorProvider, ProviderModelMixin):
129
129
### image ###
130
130
"flux-realism": "flux-realism",
131
131
"flux-dev": "flux/dev",
132
133
"dalle-3": "dalle-e-3",
132
133
"dalle-3": "dall-e-3",
134
134
}
135
135
136
136
@classmethod
@@ -1,6 +1,5 @@
1
1
from __future__ import annotations
2
2
3
import os
4
3
import json
5
4
import asyncio
6
5
from http.cookiejar import CookieJar
@@ -20,10 +19,10 @@ except ImportError:
20
19
from .base_provider import AbstractProvider, ProviderModelMixin, BaseConversation
21
20
from .helper import format_prompt
22
21
from ..typing import CreateResult, Messages, ImageType
23
from ..errors import MissingRequirementsError
22
from ..errors import MissingRequirementsError, NoValidHarFileError
24
23
from ..requests.raise_for_status import raise_for_status
25
24
from ..providers.asyncio import get_running_loop
26
from .openai.har_file import NoValidHarFileError, get_headers
25
from .openai.har_file import get_headers, get_har_files
27
26
from ..requests import get_nodriver
28
27
from ..image import ImageResponse, to_bytes, is_accepted_format
29
28
from .. import debug
@@ -76,12 +75,12 @@ class Copilot(AbstractProvider, ProviderModelMixin):
76
75
if cls.needs_auth or image is not None:
77
76
if conversation is None or conversation.access_token is None:
78
77
try:
79
access_token, cookies = readHAR()
78
access_token, cookies = readHAR(cls.url)
80
79
except NoValidHarFileError as h:
81
80
debug.log(f"Copilot: {h}")
82
81
try:
83
82
get_running_loop(check_nested=True)
84
access_token, cookies = asyncio.run(cls.get_access_token_and_cookies(proxy))
83
access_token, cookies = asyncio.run(get_access_token_and_cookies(cls.url, proxy))
85
84
except MissingRequirementsError:
86
85
raise h
87
86
else:
@@ -162,35 +161,34 @@ class Copilot(AbstractProvider, ProviderModelMixin):
162
161
if not is_started:
163
162
raise RuntimeError(f"Invalid response: {last_msg}")
164
163
165
@classmethod
166
async def get_access_token_and_cookies(cls, proxy: str = None):
167
browser = await get_nodriver(proxy=proxy)
168
page = await browser.get(cls.url)
169
access_token = None
170
while access_token is None:
171
access_token = await page.evaluate("""
172
(() => {
173
for (var i = 0; i < localStorage.length; i++) {
174
try {
175
item = JSON.parse(localStorage.getItem(localStorage.key(i)));
176
if (item.credentialType == "AccessToken"
177
&& item.expiresOn > Math.floor(Date.now() / 1000)
178
&& item.target.includes("ChatAI")) {
179
return item.secret;
180
}
181
} catch(e) {}
182
}
183
})()
184
""")
185
if access_token is None:
186
await asyncio.sleep(1)
187
cookies = {}
188
for c in await page.send(nodriver.cdp.network.get_cookies([cls.url])):
189
cookies[c.name] = c.value
190
await page.close()
191
return access_token, cookies
192
193
def readHAR():
164
async def get_access_token_and_cookies(url: str, proxy: str = None, target: str = "ChatAI",):
165
browser = await get_nodriver(proxy=proxy)
166
page = await browser.get(url)
167
access_token = None
168
while access_token is None:
169
access_token = await page.evaluate("""
170
(() => {
171
for (var i = 0; i < localStorage.length; i++) {
172
try {
173
item = JSON.parse(localStorage.getItem(localStorage.key(i)));
174
if (item.credentialType == "AccessToken"
175
&& item.expiresOn > Math.floor(Date.now() / 1000)
176
&& item.target.includes("target")) {
177
return item.secret;
178
}
179
} catch(e) {}
180
}
181
})()
182
""".replace('"target"', json.dumps(target)))
183
if access_token is None:
184
await asyncio.sleep(1)
185
cookies = {}
186
for c in await page.send(nodriver.cdp.network.get_cookies([url])):
187
cookies[c.name] = c.value
188
await page.close()
189
return access_token, cookies
190
191
def readHAR(url: str):
194
192
api_key = None
195
193
cookies = None
196
194
for path in get_har_files():
@@ -201,16 +199,13 @@ def readHAR():
201
199
# Error: not a HAR file!
202
200
continue
203
201
for v in harFile['log']['entries']:
204
v_headers = get_headers(v)
205
if v['request']['url'].startswith(Copilot.url):
206
try:
207
if "authorization" in v_headers:
208
api_key = v_headers["authorization"].split(maxsplit=1).pop()
209
except Exception as e:
210
debug.log(f"Error on read headers: {e}")
202
if v['request']['url'].startswith(url):
203
v_headers = get_headers(v)
204
if "authorization" in v_headers:
205
api_key = v_headers["authorization"].split(maxsplit=1).pop()
211
206
if v['request']['cookies']:
212
207
cookies = {c['name']: c['value'] for c in v['request']['cookies']}
213
208
if api_key is None:
214
209
raise NoValidHarFileError("No access token found in .har files")
215
210
216
return api_key, cookies
211
return api_key, cookies
@@ -0,0 +1,167 @@
1
from __future__ import annotations
2
3
import uuid
4
import aiohttp
5
import random
6
import asyncio
7
import json
8
9
from ...image import ImageResponse
10
from ...errors import MissingRequirementsError, NoValidHarFileError
11
from ...typing import AsyncResult, Messages
12
from ...requests.raise_for_status import raise_for_status
13
from ...requests.aiohttp import get_connector
14
from ...requests import get_nodriver
15
from ..Copilot import get_headers, get_har_files
16
from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
17
from ..helper import get_random_hex
18
from ... import debug
19
20
class MicrosoftDesigner(AsyncGeneratorProvider, ProviderModelMixin):
21
label = "Microsoft Designer"
22
url = "https://designer.microsoft.com"
23
working = True
24
needs_auth = True
25
default_image_model = "dall-e-3"
26
image_models = [default_image_model, "1024x1024", "1024x1792", "1792x1024"]
27
models = image_models
28
29
@classmethod
30
async def create_async_generator(
31
cls,
32
model: str,
33
messages: Messages,
34
prompt: str = None,
35
proxy: str = None,
36
**kwargs
37
) -> AsyncResult:
38
image_size = "1024x1024"
39
if model != cls.default_image_model and model in cls.image_models:
40
image_size = model
41
yield await cls.generate(messages[-1]["content"] if prompt is None else prompt, image_size, proxy)
42
43
@classmethod
44
async def generate(cls, prompt: str, image_size: str, proxy: str = None) -> ImageResponse:
45
try:
46
access_token, user_agent = readHAR("https://designerapp.officeapps.live.com")
47
except NoValidHarFileError as h:
48
debug.log(f"{cls.__name__}: {h}")
49
try:
50
access_token, user_agent = await get_access_token_and_user_agent(cls.url, proxy)
51
except MissingRequirementsError:
52
raise h
53
images = await create_images(prompt, access_token, user_agent, image_size, proxy)
54
return ImageResponse(images, prompt)
55
56
async def create_images(prompt: str, access_token: str, user_agent: str, image_size: str, proxy: str = None, seed: int = None):
57
url = 'https://designerapp.officeapps.live.com/designerapp/DallE.ashx?action=GetDallEImagesCogSci'
58
if seed is None:
59
seed = random.randint(0, 10000)
60
61
headers = {
62
"User-Agent": user_agent,
63
"Accept": "application/json, text/plain, */*",
64
"Accept-Language": "en-US",
65
'Authorization': f'Bearer {access_token}',
66
"AudienceGroup": "Production",
67
"Caller": "DesignerApp",
68
"ClientId": "b5c2664a-7e9b-4a7a-8c9a-cd2c52dcf621",
69
"SessionId": str(uuid.uuid4()),
70
"UserId": get_random_hex(16),
71
"ContainerId": "1e2843a7-2a98-4a6c-93f2-42002de5c478",
72
"FileToken": "9f1a4cb7-37e7-4c90-b44d-cb61cfda4bb8",
73
"x-upload-to-storage-das": "1",
74
"traceparent": "",
75
"X-DC-Hint": "FranceCentral",
76
"Platform": "Web",
77
"HostApp": "DesignerApp",
78
"ReleaseChannel": "",
79
"IsSignedInUser": "true",
80
"Locale": "de-DE",
81
"UserType": "MSA",
82
"x-req-start": "2615401",
83
"ClientBuild": "1.0.20241120.9",
84
"ClientName": "DesignerApp",
85
"Sec-Fetch-Dest": "empty",
86
"Sec-Fetch-Mode": "cors",
87
"Sec-Fetch-Site": "cross-site",
88
"Pragma": "no-cache",
89
"Cache-Control": "no-cache",
90
"Referer": "https://designer.microsoft.com/"
91
}
92
93
form_data = aiohttp.FormData()
94
form_data.add_field('dalle-caption', prompt)
95
form_data.add_field('dalle-scenario-name', 'TextToImage')
96
form_data.add_field('dalle-batch-size', '4')
97
form_data.add_field('dalle-image-response-format', 'UrlWithBase64Thumbnail')
98
form_data.add_field('dalle-seed', seed)
99
form_data.add_field('ClientFlights', 'EnableBICForDALLEFlight')
100
form_data.add_field('dalle-hear-back-in-ms', 1000)
101
form_data.add_field('dalle-include-b64-thumbnails', 'true')
102
form_data.add_field('dalle-aspect-ratio-scaling-factor-b64-thumbnails', 0.3)
103
form_data.add_field('dalle-image-size', image_size)
104
105
async with aiohttp.ClientSession(connector=get_connector(proxy=proxy)) as session:
106
async with session.post(url, headers=headers, data=form_data) as response:
107
await raise_for_status(response)
108
response_data = await response.json()
109
form_data.add_field('dalle-boost-count', response_data.get('dalle-boost-count', 0))
110
polling_meta_data = response_data.get('polling_response', {}).get('polling_meta_data', {})
111
form_data.add_field('dalle-poll-url', polling_meta_data.get('poll_url', ''))
112
113
while True:
114
await asyncio.sleep(polling_meta_data.get('poll_interval', 1000) / 1000)
115
async with session.post(url, headers=headers, data=form_data) as response:
116
await raise_for_status(response)
117
response_data = await response.json()
118
images = [image["ImageUrl"] for image in response_data.get('image_urls_thumbnail', [])]
119
if images:
120
return images
121
122
def readHAR(url: str) -> tuple[str, str]:
123
api_key = None
124
user_agent = None
125
for path in get_har_files():
126
with open(path, 'rb') as file:
127
try:
128
harFile = json.loads(file.read())
129
except json.JSONDecodeError:
130
# Error: not a HAR file!
131
continue
132
for v in harFile['log']['entries']:
133
if v['request']['url'].startswith(url):
134
v_headers = get_headers(v)
135
if "authorization" in v_headers:
136
api_key = v_headers["authorization"].split(maxsplit=1).pop()
137
if "user-agent" in v_headers:
138
user_agent = v_headers["user-agent"]
139
if api_key is None:
140
raise NoValidHarFileError("No access token found in .har files")
141
142
return api_key, user_agent
143
144
async def get_access_token_and_user_agent(url: str, proxy: str = None):
145
browser = await get_nodriver(proxy=proxy)
146
page = await browser.get(url)
147
user_agent = await page.evaluate("navigator.userAgent")
148
access_token = None
149
while access_token is None:
150
access_token = await page.evaluate("""
151
(() => {
152
for (var i = 0; i < localStorage.length; i++) {
153
try {
154
item = JSON.parse(localStorage.getItem(localStorage.key(i)));
155
if (item.credentialType == "AccessToken"
156
&& item.expiresOn > Math.floor(Date.now() / 1000)
157
&& item.target.includes("designerappservice")) {
158
return item.secret;
159
}
160
} catch(e) {}
161
}
162
})()
163
""")
164
if access_token is None:
165
await asyncio.sleep(1)
166
await page.close()
167
return access_token, user_agent
@@ -22,10 +22,10 @@ from ...requests.raise_for_status import raise_for_status
22
22
from ...requests import StreamSession
23
23
from ...requests import get_nodriver
24
24
from ...image import ImageResponse, ImageRequest, to_image, to_bytes, is_accepted_format
25
from ...errors import MissingAuthError
25
from ...errors import MissingAuthError, NoValidHarFileError
26
26
from ...providers.response import BaseConversation, FinishReason, SynthesizeData
27
27
from ..helper import format_cookies
28
from ..openai.har_file import get_request_config, NoValidHarFileError
28
from ..openai.har_file import get_request_config
29
29
from ..openai.har_file import RequestConfig, arkReq, arkose_url, start_url, conversation_url, backend_url, backend_anon_url
30
30
from ..openai.proofofwork import generate_proof_token
31
31
from ..openai.new import get_requirements_token
@@ -1,25 +1,27 @@
1
from .gigachat import *
1
from .gigachat import *
2
2
3
from .BingCreateImages import BingCreateImages
4
from .Cerebras import Cerebras
5
from .CopilotAccount import CopilotAccount
6
from .DeepInfra import DeepInfra
7
from .DeepInfraImage import DeepInfraImage
8
from .Gemini import Gemini
9
from .GeminiPro import GeminiPro
10
from .GithubCopilot import GithubCopilot
11
from .Groq import Groq
12
from .HuggingFace import HuggingFace
13
from .HuggingFace2 import HuggingFace2
14
from .MetaAI import MetaAI
15
from .MetaAIAccount import MetaAIAccount
16
from .OpenaiAPI import OpenaiAPI
17
from .OpenaiChat import OpenaiChat
18
from .PerplexityApi import PerplexityApi
19
from .Poe import Poe
20
from .PollinationsAI import PollinationsAI
21
from .Raycast import Raycast
22
from .Replicate import Replicate
23
from .Theb import Theb
24
from .ThebApi import ThebApi
25
from .WhiteRabbitNeo import WhiteRabbitNeo
3
from .BingCreateImages import BingCreateImages
4
from .Cerebras import Cerebras
5
from .CopilotAccount import CopilotAccount
6
from .DeepInfra import DeepInfra
7
from .DeepInfraImage import DeepInfraImage
8
from .Gemini import Gemini
9
from .GeminiPro import GeminiPro
10
from .GithubCopilot import GithubCopilot
11
from .Groq import Groq
12
from .HuggingFace import HuggingFace
13
from .HuggingFace2 import HuggingFace2
14
from .MetaAI import MetaAI
15
from .MetaAIAccount import MetaAIAccount
16
from .MicrosoftDesigner import MicrosoftDesigner
17
from .OpenaiAccount import OpenaiAccount
18
from .OpenaiAPI import OpenaiAPI
19
from .OpenaiChat import OpenaiChat
20
from .PerplexityApi import PerplexityApi
21
from .Poe import Poe
22
from .PollinationsAI import PollinationsAI
23
from .Raycast import Raycast
24
from .Replicate import Replicate
25
from .Theb import Theb
26
from .ThebApi import ThebApi
27
from .WhiteRabbitNeo import WhiteRabbitNeo
@@ -13,6 +13,7 @@ from copy import deepcopy
13
13
from .crypt import decrypt, encrypt
14
14
from ...requests import StreamSession
15
15
from ...cookies import get_cookies_dir
16
from ...errors import NoValidHarFileError
16
17
from ... import debug
17
18
18
19
arkose_url = "https://tcr9i.chat.openai.com/fc/gt2/public_key/35536E1E-65B4-4D96-9D97-6ADB7EFF8147"
@@ -21,9 +22,6 @@ backend_anon_url = "https://chatgpt.com/backend-anon/conversation"
21
22
start_url = "https://chatgpt.com/"
22
23
conversation_url = "https://chatgpt.com/c/"
23
24
24
class NoValidHarFileError(Exception):
25
pass
26
27
25
class RequestConfig:
28
26
cookies: dict = None
29
27
headers: dict = None
@@ -8,14 +8,11 @@ import logging
8
8
9
9
from ...requests import StreamSession, raise_for_status
10
10
from ...cookies import get_cookies_dir
11
from ...errors import MissingRequirementsError
11
from ...errors import MissingRequirementsError, NoValidHarFileError
12
12
from ... import debug
13
13
14
14
logger = logging.getLogger(__name__)
15
15
16
class NoValidHarFileError(Exception):
17
...
18
19
16
class arkReq:
20
17
def __init__(self, arkURL, arkHeaders, arkBody, arkCookies, userAgent):
21
18
self.arkURL = arkURL
@@ -12,15 +12,17 @@ from ..image import ImageResponse, copy_images, images_dir
12
12
from ..typing import Messages, ImageType
13
13
from ..providers.types import ProviderType
14
14
from ..providers.response import ResponseType, FinishReason, BaseConversation, SynthesizeData
15
from ..errors import NoImageResponseError, ModelNotFoundError
15
from ..errors import NoImageResponseError, MissingAuthError, NoValidHarFileError
16
16
from ..providers.retry_provider import IterListProvider
17
from ..providers.asyncio import get_running_loop, to_sync_generator, async_generator_to_list
17
from ..providers.asyncio import to_sync_generator, async_generator_to_list
18
18
from ..Provider.needs_auth import BingCreateImages, OpenaiAccount
19
from ..image import to_bytes
19
20
from .stubs import ChatCompletion, ChatCompletionChunk, Image, ImagesResponse
20
21
from .image_models import ImageModels
21
22
from .types import IterResponse, ImageProvider, Client as BaseClient
22
23
from .service import get_model_and_provider, get_last_provider, convert_to_provider
23
24
from .helper import find_stop, filter_json, filter_none, safe_aclose, to_async_iterator
25
from .. import debug
24
26
25
27
ChatCompletionResponseType = Iterator[Union[ChatCompletion, ChatCompletionChunk, BaseConversation]]
26
28
AsyncChatCompletionResponseType = AsyncIterator[Union[ChatCompletion, ChatCompletionChunk, BaseConversation]]
@@ -274,11 +276,6 @@ class Images:
274
276
provider_handler = provider
275
277
if provider_handler is None:
276
278
return default
277
if isinstance(provider_handler, IterListProvider):
278
if provider_handler.providers:
279
provider_handler = provider_handler.providers[0]
280
else:
281
raise ModelNotFoundError(f"IterListProvider for model {model} has no providers")
282
279
return provider_handler
283
280
284
281
async def async_generate(
@@ -291,33 +288,23 @@ class Images:
291
288
**kwargs
292
289
) -> ImagesResponse:
293
290
provider_handler = await self.get_provider_handler(model, provider, BingCreateImages)
294
provider_name = provider.__name__ if hasattr(provider, "__name__") else type(provider).__name__
291
provider_name = provider_handler.__name__ if hasattr(provider_handler, "__name__") else type(provider_handler).__name__
295
292
if proxy is None:
296
293
proxy = self.client.proxy
297
294
298
295
response = None
299
if hasattr(provider_handler, "create_async_generator"):
300
messages = [{"role": "user", "content": f"Generate a image: {prompt}"}]
301
async for item in provider_handler.create_async_generator(model, messages, prompt=prompt, **kwargs):
302
if isinstance(item, ImageResponse):
303
response = item
304
break
305
elif hasattr(provider_handler, 'create'):
306
if asyncio.iscoroutinefunction(provider_handler.create):
307
response = await provider_handler.create(prompt)
308
else:
309
response = provider_handler.create(prompt)
310
if isinstance(response, str):
311
response = ImageResponse([response], prompt)
312
elif hasattr(provider_handler, "create_completion"):
313
get_running_loop(check_nested=True)
314
messages = [{"role": "user", "content": f"Generate a image: {prompt}"}]
315
for item in provider_handler.create_completion(model, messages, prompt=prompt, **kwargs):
316
if isinstance(item, ImageResponse):
317
response = item
318
break
296
if isinstance(provider_handler, IterListProvider):
297
for provider in provider_handler.providers:
298
try:
299
response = await self._generate_image_response(provider, provider.__name__, model, prompt, **kwargs)
300
if response is not None:
301
provider_name = provider.__name__
302
break
303
except (MissingAuthError, NoValidHarFileError) as e:
304
debug.log(f"Image provider {provider.__name__}: {e}")
319
305
else:
320
raise ValueError(f"Provider {provider_name} does not support image generation")
306
response = await self._generate_image_response(provider_handler, provider_name, model, prompt, **kwargs)
307
321
308
if isinstance(response, ImageResponse):
322
309
return await self._process_image_response(
323
310
response,
@@ -330,6 +317,46 @@ class Images:
330
317
raise NoImageResponseError(f"No image response from {provider_name}")
331
318
raise NoImageResponseError(f"Unexpected response type: {type(response)}")
332
319
320
async def _generate_image_response(
321
self,
322
provider_handler,
323
provider_name,
324
model: str,
325
prompt: str,
326
prompt_prefix: str = "Generate a image: ",
327
image: ImageType = None,
328
**kwargs
329
) -> ImageResponse:
330
messages = [{"role": "user", "content": f"{prompt_prefix}{prompt}"}]
331
response = None
332
if hasattr(provider_handler, "create_async_generator"):
333
async for item in provider_handler.create_async_generator(
334
model,
335
messages,
336
stream=True,
337
prompt=prompt,
338
image=image,
339
**kwargs
340
):
341
if isinstance(item, ImageResponse):
342
response = item
343
break
344
elif hasattr(provider_handler, "create_completion"):
345
for item in provider_handler.create_completion(
346
model,
347
messages,
348
True,
349
prompt=prompt,
350
image=image,
351
**kwargs
352
):
353
if isinstance(item, ImageResponse):
354
response = item
355
break
356
else:
357
raise ValueError(f"Provider {provider_name} does not support image generation")
358
return response
359
333
360
def create_variation(
334
361
self,
335
362
image: ImageType,
@@ -352,33 +379,28 @@ class Images:
352
379
**kwargs
353
380
) -> ImagesResponse:
354
381
provider_handler = await self.get_provider_handler(model, provider, OpenaiAccount)
355
provider_name = provider.__name__ if hasattr(provider, "__name__") else type(provider).__name__
382
provider_name = provider_handler.__name__ if hasattr(provider_handler, "__name__") else type(provider_handler).__name__
356
383
if proxy is None:
357
384
proxy = self.client.proxy
385
prompt = "create a variation of this image"
358
386
359
if hasattr(provider_handler, "create_async_generator"):
360
messages = [{"role": "user", "content": "create a variation of this image"}]
361
generator = None
362
try:
363
generator = provider_handler.create_async_generator(model, messages, image=image, response_format=response_format, proxy=proxy, **kwargs)
364
async for chunk in generator:
365
if isinstance(chunk, ImageResponse):
366
response = chunk
387
response = None
388
if isinstance(provider_handler, IterListProvider):
389
# File pointer can be read only once, so we need to convert it to bytes
390
image = to_bytes(image)
391
for provider in provider_handler.providers:
392
try:
393
response = await self._generate_image_response(provider, provider.__name__, model, prompt, image=image, **kwargs)
394
if response is not None:
395
provider_name = provider.__name__
367
396
break
368
finally:
369
await safe_aclose(generator)
370
elif hasattr(provider_handler, 'create_variation'):
371
if asyncio.iscoroutinefunction(provider.provider_handler):
372
response = await provider_handler.create_variation(image, model=model, response_format=response_format, proxy=proxy, **kwargs)
373
else:
374
response = provider_handler.create_variation(image, model=model, response_format=response_format, proxy=proxy, **kwargs)
397
except (MissingAuthError, NoValidHarFileError) as e:
398
debug.log(f"Image provider {provider.__name__}: {e}")
375
399
else:
376
raise NoImageResponseError(f"Provider {provider_name} does not support image variation")
400
response = await self._generate_image_response(provider_handler, provider_name, model, prompt, image=image, **kwargs)
377
401
378
if isinstance(response, str):
379
response = ImageResponse([response])
380
402
if isinstance(response, ImageResponse):
381
return self._process_image_response(response, response_format, proxy, model, provider_name)
403
return await self._process_image_response(response, response_format, proxy, model, provider_name)
382
404
if response is None:
383
405
raise NoImageResponseError(f"No image response from {provider_name}")
384
406
raise NoImageResponseError(f"Unexpected response type: {type(response)}")
@@ -44,4 +44,7 @@ class ResponseError(Exception):
44
44
...
45
45
46
46
class ResponseStatusError(Exception):
47
...
48
49
class NoValidHarFileError(Exception):
47
50
...