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

XFEstudio/gpt4free

refactor: Refactor image generation parameters handling

- Updated the `PollinationsAI` class in `g4f/Provider/PollinationsAI.py`: - Changed `aspect_ratio` parameter handling to conditionally use default "1:1" if not specified. - Enhanced media handling by introducing `media` parameter in `_generate_image` method. - Updated parameter processing in `_generate_image_async` method for `model == "gptimage"`. - Updated `Api` class in `g4f/api/__init__.py`: - Simplified handling of `credentials` for `config.api_key`. - Updated `Images` class in `g4f/client/__init__.py`: - Added `download_media` parameter to `_process_image_response` method. - Enhanced `_process_image_response` method to conditionally download media based on `download_media` flag. - Updated `_process_image_response` method in `Images` class in `g4f/client/__init__.py`: - Enhanced handling of media response based on `download_media` flag. - Updated `is_valid_media` function in `g4f/image/__init__.py`: - Added typing annotations for clarity. - Updated `AnyProvider` class in `g4f/providers/any_provider.py`: - Improved handling of `api_key` dictionary to set `extra_body["api_key"]`. - Updated `IterListProvider` class in `g4f/providers/retry_provider.py`: - Enhanced handling of `model` and `api_key` parameters. - Updated `BaseProvider` class in `g4f/providers/types.py`: - Added `create_function` and `async_create_function` methods. - Updated `BaseRetryProvider` class in `g4f/providers/types.py`: - Enhanced handling of `model` and `api_key` parameters in provider iteration.

963c3a58
hlohaus <983577+hlohaus@users.noreply.github.com>
提交于

代码差异

7 个文件 +93 -26
Modified g4f/Provider/PollinationsAI.py +20 -8
@@ -252,7 +252,7 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
252 252 extra_body: dict = None,
253 253 # Image generation parameters
254 254 prompt: str = None,
255 aspect_ratio: str = "1:1",
255 aspect_ratio: str = None,
256 256 width: int = None,
257 257 height: int = None,
258 258 seed: Optional[int] = None,
@@ -294,6 +294,7 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
294 294 async for chunk in cls._generate_image(
295 295 model=model,
296 296 prompt=format_media_prompt(messages, prompt),
297 media=media,
297 298 proxy=proxy,
298 299 aspect_ratio=aspect_ratio,
299 300 width=width,
@@ -347,6 +348,7 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
347 348 cls,
348 349 model: str,
349 350 prompt: str,
351 media: MediaListType,
350 352 proxy: str,
351 353 aspect_ratio: str,
352 354 width: int,
@@ -362,20 +364,30 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
362 364 api_key: str,
363 365 timeout: int = 120
364 366 ) -> AsyncResult:
365 if model == "gptimage":
366 n = 1
367 params = use_aspect_ratio({
368 "width": width,
369 "height": height,
367 params = {
370 368 "model": model,
371 369 "nologo": str(nologo).lower(),
372 370 "private": str(private).lower(),
373 371 "enhance": str(enhance).lower(),
374 372 "safe": str(safe).lower(),
375 }, aspect_ratio)
373 }
374 if model == "gptimage":
375 n = 1
376 # Only remote images are supported
377 image = [item[0] for item in media if isinstance(item[0], str) and item[0].startswith("http")]
378 params = {
379 **params,
380 "image": ",".join(image) if image else "",
381 }
382 else:
383 params = use_aspect_ratio({
384 "width": width,
385 "height": height,
386 **params
387 }, "1:1" if aspect_ratio is None else aspect_ratio)
376 388 query = "&".join(f"{k}={quote_plus(str(v))}" for k, v in params.items() if v is not None)
377 389 encoded_prompt = prompt
378 if model == "gptimage" and aspect_ratio != "1:1":
390 if model == "gptimage" and aspect_ratio is not None:
379 391 encoded_prompt = f"{encoded_prompt} aspect-ratio: {aspect_ratio}"
380 392 encoded_prompt = quote_plus(encoded_prompt)[:2048-len(cls.image_api_endpoint)-len(query)-8]
381 393 url = f"{cls.image_api_endpoint}prompt/{encoded_prompt}?{query}"
Modified g4f/api/__init__.py +2 -4
@@ -425,7 +425,7 @@ class Api:
425 425 try:
426 426 if config.provider is None:
427 427 config.provider = AppConfig.provider if provider is None else provider
428 if credentials is not None and credentials.credentials != "secret":
428 if config.api_key is None and credentials is not None and credentials.credentials != "secret":
429 429 config.api_key = credentials.credentials
430 430
431 431 conversation = None
@@ -618,9 +618,7 @@ class Api:
618 618 provider=config.provider if provider is None else provider,
619 619 prompt=config.input,
620 620 audio=filter_none(voice=config.voice, format=config.response_format, language=config.language),
621 **filter_none(
622 api_key=api_key,
623 )
621 api_key=api_key,
624 622 )
625 623 if isinstance(response.choices[0].message.content, AudioResponse):
626 624 response = response.choices[0].message.content.data
Modified g4f/client/__init__.py +14 -4
@@ -479,6 +479,7 @@ class Images:
479 479 response,
480 480 model,
481 481 provider_name,
482 kwargs.get("download_media", True),
482 483 response_format,
483 484 proxy
484 485 )
@@ -531,7 +532,7 @@ class Images:
531 532 urls.extend(item.urls)
532 533 if not urls:
533 534 return None
534 alt = getattr(items[0], "alt", items[0].options.get("text"))
535 alt = getattr(items[0], "alt", "")
535 536 return MediaResponse(urls, alt, items[0].options)
536 537
537 538 def create_variation(
@@ -580,13 +581,21 @@ class Images:
580 581 if error is not None:
581 582 raise error
582 583 raise NoMediaResponseError(f"No media response from {provider_name}")
583 return await self._process_image_response(response, model, provider_name, response_format, proxy)
584 return await self._process_image_response(
585 response,
586 model,
587 provider_name,
588 kwargs.get("download_media", True),
589 response_format,
590 proxy
591 )
584 592
585 593 async def _process_image_response(
586 594 self,
587 595 response: MediaResponse,
588 596 model: str,
589 597 provider: str,
598 download_media: bool,
590 599 response_format: Optional[str] = None,
591 600 proxy: str = None
592 601 ) -> ImagesResponse:
@@ -609,9 +618,10 @@ class Images:
609 618 images = await asyncio.gather(*[get_b64_from_url(image) for image in response.get_list()])
610 619 else:
611 620 # Save locally for None (default) case
612 images = await copy_media(response.get_list(), response.get("cookies"), response.get("headers"), proxy, response.alt)
621 if download_media or response.get("cookies"):
622 images = await copy_media(response.get_list(), response.get("cookies"), response.get("headers"), proxy, response.alt)
613 623 images = [Image.model_construct(url=image, revised_prompt=response.alt) for image in images]
614
624
615 625 return ImagesResponse.model_construct(
616 626 created=int(time.time()),
617 627 data=images,
Modified g4f/image/__init__.py +2 -2
@@ -14,7 +14,7 @@ try:
14 14 except ImportError:
15 15 has_requirements = False
16 16
17 from ..typing import ImageType, Union, Image
17 from ..typing import ImageType, Image
18 18 from ..errors import MissingRequirementsError
19 19
20 20 EXTENSIONS_MAP: dict[str, str] = {
@@ -107,7 +107,7 @@ def is_data_an_media(data, filename: str = None) -> str:
107 107 return is_accepted_format(data)
108 108 return is_data_uri_an_image(data)
109 109
110 def is_valid_media(data, filename: str = None) -> str:
110 def is_valid_media(data: ImageType = None, filename: str = None) -> str:
111 111 if is_valid_audio(data, filename):
112 112 return True
113 113 if filename:
Modified g4f/providers/any_provider.py +3 -2
@@ -319,8 +319,9 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
319 319 extra_providers = []
320 320 if isinstance(api_key, dict):
321 321 for provider in api_key:
322 if provider in __map__ and __map__[provider] not in MAIN_PROVIERS:
323 extra_providers.append(__map__[provider])
322 if api_key.get(provider):
323 if provider in __map__ and __map__[provider] not in MAIN_PROVIERS:
324 extra_providers.append(__map__[provider])
324 325 for provider in MAIN_PROVIERS + extra_providers:
325 326 if provider.working:
326 327 if not model or model in provider.get_models() or model in provider.model_aliases:
Modified g4f/providers/retry_provider.py +16 -6
@@ -53,11 +53,16 @@ class IterListProvider(BaseRetryProvider):
53 53
54 54 for provider in self.get_providers(stream and not ignore_stream, ignored):
55 55 self.last_provider = provider
56 debug.log(f"Using {provider.__name__} provider")
57 yield ProviderInfo(**provider.get_dict(), model=model if model else getattr(provider, "default_model"))
56 if not model:
57 model = getattr(provider, "default_model", None)
58 model = provider.model_aliases.get(model, model) if hasattr(provider, "model_aliases") else model
59 debug.log(f"Using {provider.__name__} provider with model {model}")
60 yield ProviderInfo(**provider.get_dict(), model=model)
58 61 extra_body = kwargs.copy()
59 62 if isinstance(api_key, dict):
60 extra_body["api_key"] = api_key.get(provider.get_parent())
63 api_key = api_key.get(provider.get_parent())
64 if api_key:
65 extra_body["api_key"] = api_key
61 66 try:
62 67 response = provider.create_function(model, messages, stream=stream, **extra_body)
63 68 for chunk in response:
@@ -92,11 +97,16 @@ class IterListProvider(BaseRetryProvider):
92 97
93 98 for provider in self.get_providers(stream and not ignore_stream, ignored):
94 99 self.last_provider = provider
95 debug.log(f"Using {provider.__name__} provider" + (f" and {model} model" if model else ""))
96 yield ProviderInfo(**provider.get_dict(), model=model if model else getattr(provider, "default_model"))
100 if not model:
101 model = getattr(provider, "default_model", None)
102 model = provider.model_aliases.get(model, model) if hasattr(provider, "model_aliases") else model
103 debug.log(f"Using {provider.__name__} provider with model {model}")
104 yield ProviderInfo(**provider.get_dict(), model=model)
97 105 extra_body = kwargs.copy()
98 106 if isinstance(api_key, dict):
99 extra_body["api_key"] = api_key.get(provider.get_parent())
107 api_key = api_key.get(provider.get_parent())
108 if api_key:
109 extra_body["api_key"] = api_key
100 110 if conversation is not None and hasattr(conversation, provider.__name__):
101 111 extra_body["conversation"] = JsonConversation(**getattr(conversation, provider.__name__))
102 112 try:
Modified g4f/providers/types.py +36 -0
@@ -42,6 +42,42 @@ class BaseProvider(ABC):
42 42 def get_parent(cls) -> str:
43 43 return getattr(cls, "parent", cls.__name__)
44 44
45 @abstractmethod
46 def create_function(
47 *args,
48 **kwargs
49 ) -> CreateResult:
50 """
51 Create a function to generate a response based on the model and messages.
52
53 Args:
54 model (str): The model to use.
55 messages (Messages): The messages to process.
56 stream (bool): Whether to stream the response.
57
58 Returns:
59 CreateResult: The result of the creation.
60 """
61 raise NotImplementedError()
62
63 @staticmethod
64 def async_create_function(
65 *args,
66 **kwargs
67 ) -> CreateResult:
68 """
69 Asynchronously create a function to generate a response based on the model and messages.
70
71 Args:
72 model (str): The model to use.
73 messages (Messages): The messages to process.
74 stream (bool): Whether to stream the response.
75
76 Returns:
77 CreateResult: The result of the creation.
78 """
79 raise NotImplementedError()
80
45 81 class BaseRetryProvider(BaseProvider):
46 82 """
47 83 Base class for a provider that implements retry logic.