返回提交历史
Modified
g4f/Provider/Airforce.py
+9
-9
Modified
g4f/client/__init__.py
+30
-25
XFEstudio/gpt4free
Fix provider selection in images generate Improve image generation in Airforce provider
442185ea
代码差异
2 个文件
+39
-34
@@ -7,6 +7,7 @@ import re
7
7
import requests
8
8
from requests.packages.urllib3.exceptions import InsecureRequestWarning
9
9
requests.packages.urllib3.disable_warnings(InsecureRequestWarning)
10
from urllib.parse import quote
10
11
11
12
from ..typing import AsyncResult, Messages
12
13
from .base_provider import AsyncGeneratorProvider, ProviderModelMixin
@@ -95,14 +96,18 @@ class Airforce(AsyncGeneratorProvider, ProviderModelMixin):
95
96
model: str,
96
97
messages: Messages,
97
98
proxy: str = None,
99
prompt: str = None,
98
100
seed: int = None,
99
101
size: str = "1:1", # "1:1", "16:9", "9:16", "21:9", "9:21", "1:2", "2:1"
100
102
stream: bool = False,
101
103
**kwargs
102
104
) -> AsyncResult:
103
105
model = cls.get_model(model)
106
104
107
if model in cls.image_models:
105
return cls._generate_image(model, messages, proxy, seed, size)
108
if prompt is None:
109
prompt = messages[-1]['content']
110
return cls._generate_image(model, prompt, proxy, seed, size)
106
111
else:
107
112
return cls._generate_text(model, messages, proxy, stream, **kwargs)
108
113
@@ -110,7 +115,7 @@ class Airforce(AsyncGeneratorProvider, ProviderModelMixin):
110
115
async def _generate_image(
111
116
cls,
112
117
model: str,
113
messages: Messages,
118
prompt: str,
114
119
proxy: str = None,
115
120
seed: int = None,
116
121
size: str = "1:1",
@@ -125,7 +130,6 @@ class Airforce(AsyncGeneratorProvider, ProviderModelMixin):
125
130
}
126
131
if seed is None:
127
132
seed = random.randint(0, 100000)
128
prompt = messages[-1]['content']
129
133
130
134
async with StreamSession(headers=headers, proxy=proxy) as session:
131
135
params = {
@@ -140,12 +144,8 @@ class Airforce(AsyncGeneratorProvider, ProviderModelMixin):
140
144
141
145
if 'application/json' in content_type:
142
146
raise RuntimeError(await response.json().get("error", {}).get("message"))
143
elif 'image' in content_type:
144
image_data = b""
145
async for chunk in response.iter_content():
146
if chunk:
147
image_data += chunk
148
image_url = f"{cls.api_endpoint_imagine}?model={model}&prompt={prompt}&size={size}&seed={seed}"
147
elif content_type.startswith("image/"):
148
image_url = f"{cls.api_endpoint_imagine}?model={model}&prompt={quote(prompt)}&size={size}&seed={seed}"
149
149
yield ImageResponse(images=image_url, alt=prompt)
150
150
151
151
@classmethod
@@ -16,7 +16,7 @@ from ..providers.response import ResponseType, FinishReason, BaseConversation, S
16
16
from ..errors import NoImageResponseError, ModelNotFoundError
17
17
from ..providers.retry_provider import IterListProvider
18
18
from ..providers.asyncio import get_running_loop, to_sync_generator, async_generator_to_list
19
from ..Provider.needs_auth.BingCreateImages import BingCreateImages
19
from ..Provider.needs_auth import BingCreateImages, OpenaiAccount
20
20
from .stubs import ChatCompletion, ChatCompletionChunk, Image, ImagesResponse
21
21
from .image_models import ImageModels
22
22
from .types import IterResponse, ImageProvider, Client as BaseClient
@@ -264,28 +264,34 @@ class Images:
264
264
"""
265
265
return asyncio.run(self.async_generate(prompt, model, provider, response_format, proxy, **kwargs))
266
266
267
async def async_generate(
268
self,
269
prompt: str,
270
model: Optional[str] = None,
271
provider: Optional[ProviderType] = None,
272
response_format: Optional[str] = "url",
273
proxy: Optional[str] = None,
274
**kwargs
275
) -> ImagesResponse:
267
async def get_provider_handler(self, model: Optional[str], provider: Optional[ImageProvider], default: ImageProvider) -> ImageProvider:
276
268
if provider is None:
277
provider_handler = self.models.get(model, provider or self.provider or BingCreateImages)
269
provider_handler = self.provider
270
if provider_handler is None:
271
provider_handler = self.models.get(model, default)
278
272
elif isinstance(provider, str):
279
273
provider_handler = convert_to_provider(provider)
280
274
else:
281
275
provider_handler = provider
282
276
if provider_handler is None:
283
raise ModelNotFoundError(f"Unknown model: {model}")
277
return default
284
278
if isinstance(provider_handler, IterListProvider):
285
279
if provider_handler.providers:
286
280
provider_handler = provider_handler.providers[0]
287
281
else:
288
282
raise ModelNotFoundError(f"IterListProvider for model {model} has no providers")
283
return provider_handler
284
285
async def async_generate(
286
self,
287
prompt: str,
288
model: Optional[str] = None,
289
provider: Optional[ProviderType] = None,
290
response_format: Optional[str] = "url",
291
proxy: Optional[str] = None,
292
**kwargs
293
) -> ImagesResponse:
294
provider_handler = await self.get_provider_handler(model, provider, BingCreateImages)
289
295
if proxy is None:
290
296
proxy = self.client.proxy
291
297
@@ -311,7 +317,7 @@ class Images:
311
317
response = item
312
318
break
313
319
else:
314
raise ValueError(f"Provider {provider} does not support image generation")
320
raise ValueError(f"Provider {getattr(provider_handler, '__name__')} does not support image generation")
315
321
if isinstance(response, ImageResponse):
316
322
return await self._process_image_response(
317
323
response,
@@ -320,6 +326,8 @@ class Images:
320
326
model,
321
327
getattr(provider_handler, "__name__", None)
322
328
)
329
if response is None:
330
raise NoImageResponseError(f"No image response from {getattr(provider_handler, '__name__')}")
323
331
raise NoImageResponseError(f"Unexpected response type: {type(response)}")
324
332
325
333
def create_variation(
@@ -343,31 +351,26 @@ class Images:
343
351
proxy: Optional[str] = None,
344
352
**kwargs
345
353
) -> ImagesResponse:
346
if provider is None:
347
provider = self.models.get(model, provider or self.provider or BingCreateImages)
348
if provider is None:
349
raise ModelNotFoundError(f"Unknown model: {model}")
350
if isinstance(provider, str):
351
provider = convert_to_provider(provider)
354
provider_handler = await self.get_provider_handler(model, provider, OpenaiAccount)
352
355
if proxy is None:
353
356
proxy = self.client.proxy
354
357
355
if hasattr(provider, "create_async_generator"):
358
if hasattr(provider_handler, "create_async_generator"):
356
359
messages = [{"role": "user", "content": "create a variation of this image"}]
357
360
generator = None
358
361
try:
359
generator = provider.create_async_generator(model, messages, image=image, response_format=response_format, proxy=proxy, **kwargs)
362
generator = provider_handler.create_async_generator(model, messages, image=image, response_format=response_format, proxy=proxy, **kwargs)
360
363
async for chunk in generator:
361
364
if isinstance(chunk, ImageResponse):
362
365
response = chunk
363
366
break
364
367
finally:
365
368
await safe_aclose(generator)
366
elif hasattr(provider, 'create_variation'):
367
if asyncio.iscoroutinefunction(provider.create_variation):
368
response = await provider.create_variation(image, model=model, response_format=response_format, proxy=proxy, **kwargs)
369
elif hasattr(provider_handler, 'create_variation'):
370
if asyncio.iscoroutinefunction(provider.provider_handler):
371
response = await provider_handler.create_variation(image, model=model, response_format=response_format, proxy=proxy, **kwargs)
369
372
else:
370
response = provider.create_variation(image, model=model, response_format=response_format, proxy=proxy, **kwargs)
373
response = provider_handler.create_variation(image, model=model, response_format=response_format, proxy=proxy, **kwargs)
371
374
else:
372
375
raise NoImageResponseError(f"Provider {provider} does not support image variation")
373
376
@@ -375,6 +378,8 @@ class Images:
375
378
response = ImageResponse([response])
376
379
if isinstance(response, ImageResponse):
377
380
return self._process_image_response(response, response_format, proxy, model, getattr(provider, "__name__", None))
381
if response is None:
382
raise NoImageResponseError(f"No image response from {getattr(provider, '__name__')}")
378
383
raise NoImageResponseError(f"Unexpected response type: {type(response)}")
379
384
380
385
async def _process_image_response(