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

XFEstudio/gpt4free

Fix image generation in OpenaiChat (#2390)

* Fix image generation in OpenaiChat * Add PollinationsAI provider with image and text generation

dba41cda
H Lohaus <hlohaus@users.noreply.github.com>
提交于

代码差异

4 个文件 +84 -52
Added g4f/Provider/PollinationsAI.py +69 -0
@@ -0,0 +1,69 @@
1 from __future__ import annotations
2
3 from urllib.parse import quote
4 import random
5 import requests
6 from sys import maxsize
7 from aiohttp import ClientSession
8
9 from ..typing import AsyncResult, Messages
10 from ..image import ImageResponse
11 from ..requests.raise_for_status import raise_for_status
12 from ..requests.aiohttp import get_connector
13 from .needs_auth.OpenaiAPI import OpenaiAPI
14 from .helper import format_prompt
15
16 class PollinationsAI(OpenaiAPI):
17 label = "Pollinations.AI"
18 url = "https://pollinations.ai"
19 working = True
20 supports_stream = True
21 default_model = "openai"
22
23 @classmethod
24 def get_models(cls):
25 if not cls.image_models:
26 url = "https://image.pollinations.ai/models"
27 response = requests.get(url)
28 raise_for_status(response)
29 cls.image_models = response.json()
30 if not cls.models:
31 url = "https://text.pollinations.ai/models"
32 response = requests.get(url)
33 raise_for_status(response)
34 cls.models = [model.get("name") for model in response.json()]
35 cls.models.extend(cls.image_models)
36 return cls.models
37
38 @classmethod
39 async def create_async_generator(
40 cls,
41 model: str,
42 messages: Messages,
43 api_base: str = "https://text.pollinations.ai/openai",
44 api_key: str = None,
45 proxy: str = None,
46 seed: str = None,
47 **kwargs
48 ) -> AsyncResult:
49 if model:
50 model = cls.get_model(model)
51 if model in cls.image_models:
52 prompt = messages[-1]["content"]
53 if seed is None:
54 seed = random.randint(0, maxsize)
55 image = f"https://image.pollinations.ai/prompt/{quote(prompt)}?width=1024&height=1024&seed={int(seed)}&nofeed=true&nologo=true&model={quote(model)}"
56 yield ImageResponse(image, prompt)
57 return
58 if api_key is None:
59 async with ClientSession(connector=get_connector(proxy=proxy)) as session:
60 prompt = format_prompt(messages)
61 async with session.get(f"https://text.pollinations.ai/{quote(prompt)}?model={quote(model)}") as response:
62 await raise_for_status(response)
63 async for line in response.content.iter_any():
64 yield line.decode(errors="ignore")
65 else:
66 async for chunk in super().create_async_generator(
67 model, messages, api_base=api_base, proxy=proxy, **kwargs
68 ):
69 yield chunk
Modified g4f/Provider/__init__.py +1 -0
@@ -32,6 +32,7 @@ from .MagickPen import MagickPen
32 32 from .PerplexityLabs import PerplexityLabs
33 33 from .Pi import Pi
34 34 from .Pizzagpt import Pizzagpt
35 from .PollinationsAI import PollinationsAI
35 36 from .Prodia import Prodia
36 37 from .Reka import Reka
37 38 from .ReplicateHome import ReplicateHome
Modified g4f/Provider/needs_auth/OpenaiChat.py +13 -52
@@ -65,6 +65,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
65 65 default_vision_model = "gpt-4o"
66 66 fallback_models = ["auto", "gpt-4", "gpt-4o", "gpt-4o-mini", "gpt-4o-canmore", "o1-preview", "o1-mini"]
67 67 vision_models = fallback_models
68 image_models = fallback_models
68 69
69 70 _api_key: str = None
70 71 _headers: dict = None
@@ -330,7 +331,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
330 331 api_key: str = None,
331 332 cookies: Cookies = None,
332 333 auto_continue: bool = False,
333 history_disabled: bool = True,
334 history_disabled: bool = False,
334 335 action: str = "next",
335 336 conversation_id: str = None,
336 337 conversation: Conversation = None,
@@ -425,12 +426,6 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
425 426 f"Arkose: {'False' if not need_arkose else RequestConfig.arkose_token[:12]+'...'}",
426 427 f"Proofofwork: {'False' if proofofwork is None else proofofwork[:12]+'...'}",
427 428 )]
428 ws = None
429 if need_arkose:
430 async with session.post(f"{cls.url}/backend-api/register-websocket", headers=cls._headers) as response:
431 wss_url = (await response.json()).get("wss_url")
432 if wss_url:
433 ws = await session.ws_connect(wss_url)
434 429 data = {
435 430 "action": action,
436 431 "messages": None,
@@ -474,7 +469,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
474 469 await asyncio.sleep(5)
475 470 continue
476 471 await raise_for_status(response)
477 async for chunk in cls.iter_messages_chunk(response.iter_lines(), session, conversation, ws):
472 async for chunk in cls.iter_messages_chunk(response.iter_lines(), session, conversation):
478 473 if return_conversation:
479 474 history_disabled = False
480 475 return_conversation = False
@@ -489,44 +484,16 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
489 484 if history_disabled and auto_continue:
490 485 await cls.delete_conversation(session, cls._headers, conversation.conversation_id)
491 486
492 @staticmethod
493 async def iter_messages_ws(ws: ClientWebSocketResponse, conversation_id: str, is_curl: bool) -> AsyncIterator:
494 while True:
495 if is_curl:
496 message = json.loads(ws.recv()[0])
497 else:
498 message = await ws.receive_json()
499 if message["conversation_id"] == conversation_id:
500 yield base64.b64decode(message["body"])
501
502 487 @classmethod
503 488 async def iter_messages_chunk(
504 489 cls,
505 490 messages: AsyncIterator,
506 491 session: StreamSession,
507 492 fields: Conversation,
508 ws = None
509 493 ) -> AsyncIterator:
510 494 async for message in messages:
511 if message.startswith(b'{"wss_url":'):
512 message = json.loads(message)
513 ws = await session.ws_connect(message["wss_url"]) if ws is None else ws
514 try:
515 async for chunk in cls.iter_messages_chunk(
516 cls.iter_messages_ws(ws, message["conversation_id"], hasattr(ws, "recv")),
517 session, fields
518 ):
519 yield chunk
520 finally:
521 await ws.aclose() if hasattr(ws, "aclose") else await ws.close()
522 break
523 495 async for chunk in cls.iter_messages_line(session, message, fields):
524 if fields.finish_reason is not None:
525 break
526 else:
527 yield chunk
528 if fields.finish_reason is not None:
529 break
496 yield chunk
530 497
531 498 @classmethod
532 499 async def iter_messages_line(cls, session: StreamSession, line: bytes, fields: Conversation) -> AsyncIterator:
@@ -542,9 +509,9 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
542 509 return
543 510 if isinstance(line, dict) and "v" in line:
544 511 v = line.get("v")
545 if isinstance(v, str):
512 if isinstance(v, str) and fields.is_recipient:
546 513 yield v
547 elif isinstance(v, list):
514 elif isinstance(v, list) and fields.is_recipient:
548 515 for m in v:
549 516 if m.get("p") == "/message/content/parts/0":
550 517 yield m.get("v")
@@ -556,25 +523,20 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
556 523 fields.conversation_id = v.get("conversation_id")
557 524 debug.log(f"OpenaiChat: New conversation: {fields.conversation_id}")
558 525 m = v.get("message", {})
559 if m.get("author", {}).get("role") == "assistant":
560 fields.message_id = v.get("message", {}).get("id")
526 fields.is_recipient = m.get("recipient") == "all"
527 if fields.is_recipient:
561 528 c = m.get("content", {})
562 529 if c.get("content_type") == "multimodal_text":
563 530 generated_images = []
564 531 for element in c.get("parts"):
565 if isinstance(element, str):
566 debug.log(f"No image or text: {line}")
567 elif element.get("content_type") == "image_asset_pointer":
532 if isinstance(element, dict) and element.get("content_type") == "image_asset_pointer":
568 533 generated_images.append(
569 534 cls.get_generated_image(session, cls._headers, element)
570 535 )
571 elif element.get("content_type") == "text":
572 for part in element.get("parts", []):
573 yield part
574 536 for image_response in await asyncio.gather(*generated_images):
575 537 yield image_response
576 else:
577 debug.log(f"OpenaiChat: {line}")
538 if m.get("author", {}).get("role") == "assistant":
539 fields.message_id = v.get("message", {}).get("id")
578 540 return
579 541 if "error" in line and line.get("error"):
580 542 raise RuntimeError(line.get("error"))
@@ -652,7 +614,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
652 614 cls._headers = cls.get_default_headers() if headers is None else headers
653 615 if user_agent is not None:
654 616 cls._headers["user-agent"] = user_agent
655 cls._cookies = {} if cookies is None else {k: v for k, v in cookies.items() if k != "access_token"}
617 cls._cookies = {} if cookies is None else cookies
656 618 cls._update_cookie_header()
657 619
658 620 @classmethod
@@ -671,8 +633,6 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
671 633 @classmethod
672 634 def _update_cookie_header(cls):
673 635 cls._headers["cookie"] = format_cookies(cls._cookies)
674 if "oai-did" in cls._cookies:
675 cls._headers["oai-device-id"] = cls._cookies["oai-did"]
676 636
677 637 class Conversation(BaseConversation):
678 638 """
@@ -682,6 +642,7 @@ class Conversation(BaseConversation):
682 642 self.conversation_id = conversation_id
683 643 self.message_id = message_id
684 644 self.finish_reason = finish_reason
645 self.is_recipient = False
685 646
686 647 class Response():
687 648 """
Modified g4f/providers/base_provider.py +1 -0
@@ -290,6 +290,7 @@ class ProviderModelMixin:
290 290 default_model: str = None
291 291 models: list[str] = []
292 292 model_aliases: dict[str, str] = {}
293 image_models: list = None
293 294
294 295 @classmethod
295 296 def get_models(cls) -> list[str]: