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

XFEstudio/gpt4free

Fix provider selection in images generate Improve image generation in Airforce provider

442185ea
Heiner Lohaus <hlohaus@users.noreply.github.com>
提交于

代码差异

2 个文件 +39 -34
Modified g4f/Provider/Airforce.py +9 -9
@@ -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
Modified g4f/client/__init__.py +30 -25
@@ -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(