返回提交历史
Added
g4f/Provider/PollinationsAI.py
+69
-0
Modified
g4f/Provider/__init__.py
+1
-0
Modified
g4f/Provider/needs_auth/OpenaiChat.py
+13
-52
Modified
g4f/providers/base_provider.py
+1
-0
XFEstudio/gpt4free
Fix image generation in OpenaiChat (#2390)
* Fix image generation in OpenaiChat * Add PollinationsAI provider with image and text generation
dba41cda
代码差异
4 个文件
+84
-52
@@ -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
@@ -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
@@ -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
"""
@@ -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]: