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

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

代码差异

16 个文件 +392 -136
Modified etc/unittest/__main__.py +1 -0
@@ -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
Added etc/unittest/image_client.py +44 -0
@@ -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()
Modified etc/unittest/mocks.py +21 -0
@@ -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
Modified g4f/Provider/AmigoChat.py +5 -5
@@ -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
Modified g4f/Provider/Copilot.py +37 -42
@@ -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
Added g4f/Provider/needs_auth/MicrosoftDesigner.py +167 -0
@@ -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
Modified g4f/Provider/needs_auth/OpenaiChat.py +2 -2
@@ -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
Modified g4f/Provider/needs_auth/__init__.py +26 -24
@@ -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
Modified g4f/Provider/openai/har_file.py +1 -3
@@ -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
Modified g4f/Provider/you/har_file.py +1 -4
@@ -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
Modified g4f/client/__init__.py +71 -49
@@ -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)}")
Modified g4f/errors.py +3 -0
@@ -44,4 +44,7 @@ class ResponseError(Exception):
44 44 ...
45 45
46 46 class ResponseStatusError(Exception):
47 ...
48
49 class NoValidHarFileError(Exception):
47 50 ...
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