返回提交历史
Modified
g4f/Provider/GeminiPro.py
+8
-10
Modified
g4f/Provider/needs_auth/OpenaiChat.py
+11
-6
XFEstudio/gpt4free
Expire cache, Fix multiple websocket conversations in OpenaiChat Map system messages to user messages in GeminiPro
cfa45e70
代码差异
2 个文件
+19
-16
@@ -26,38 +26,35 @@ class GeminiPro(AsyncGeneratorProvider, ProviderModelMixin):
26
26
stream: bool = False,
27
27
proxy: str = None,
28
28
api_key: str = None,
29
api_base: str = None,
30
use_auth_header: bool = True,
29
api_base: str = "https://generativelanguage.googleapis.com/v1beta",
30
use_auth_header: bool = False,
31
31
image: ImageType = None,
32
32
connector: BaseConnector = None,
33
33
**kwargs
34
34
) -> AsyncResult:
35
model = "gemini-pro-vision" if not model and image else model
35
model = "gemini-pro-vision" if model is None and image is not None else model
36
36
model = cls.get_model(model)
37
37
38
38
if not api_key:
39
39
raise MissingAuthError('Missing "api_key"')
40
40
41
41
headers = params = None
42
if api_base and use_auth_header:
42
if use_auth_header:
43
43
headers = {"Authorization": f"Bearer {api_key}"}
44
44
else:
45
45
params = {"key": api_key}
46
46
47
if not api_base:
48
api_base = f"https://generativelanguage.googleapis.com/v1beta"
49
50
47
method = "streamGenerateContent" if stream else "generateContent"
51
48
url = f"{api_base.rstrip('/')}/models/{model}:{method}"
52
49
async with ClientSession(headers=headers, connector=get_connector(connector, proxy)) as session:
53
50
contents = [
54
51
{
55
"role": "model" if message["role"] == "assistant" else message["role"],
52
"role": "model" if message["role"] == "assistant" else "user",
56
53
"parts": [{"text": message["content"]}]
57
54
}
58
55
for message in messages
59
56
]
60
if image:
57
if image is not None:
61
58
image = to_bytes(image)
62
59
contents[-1]["parts"].append({
63
60
"inline_data": {
@@ -87,7 +84,8 @@ class GeminiPro(AsyncGeneratorProvider, ProviderModelMixin):
87
84
lines = [b"{\n"]
88
85
elif chunk == b",\r\n" or chunk == b"]":
89
86
try:
90
data = json.loads(b"".join(lines))
87
data = b"".join(lines)
88
data = json.loads(data)
91
89
yield data["candidates"][0]["content"]["parts"][0]["text"]
92
90
except:
93
91
data = data.decode() if isinstance(data, bytes) else data
@@ -5,6 +5,7 @@ import uuid
5
5
import json
6
6
import os
7
7
import base64
8
import time
8
9
from aiohttp import ClientWebSocketResponse
9
10
10
11
try:
@@ -47,7 +48,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
47
48
_api_key: str = None
48
49
_headers: dict = None
49
50
_cookies: Cookies = None
50
_last_message: int = 0
51
_expires: int = None
51
52
52
53
@classmethod
53
54
async def create(
@@ -348,7 +349,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
348
349
timeout=timeout
349
350
) as session:
350
351
# Read api_key and cookies from cache / browser config
351
if cls._headers is None:
352
if cls._headers is None or time.time() > cls._expires:
352
353
if api_key is None:
353
354
# Read api_key from cookies
354
355
cookies = get_cookies("chat.openai.com", False) if cookies is None else cookies
@@ -437,17 +438,20 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
437
438
await cls.delete_conversation(session, cls._headers, fields.conversation_id)
438
439
439
440
@staticmethod
440
async def iter_messages_ws(ws: ClientWebSocketResponse) -> AsyncIterator:
441
async def iter_messages_ws(ws: ClientWebSocketResponse, conversation_id: str) -> AsyncIterator:
441
442
while True:
442
yield base64.b64decode((await ws.receive_json())["body"])
443
message = await ws.receive_json()
444
if message["conversation_id"] == conversation_id:
445
yield base64.b64decode(message["body"])
443
446
444
447
@classmethod
445
448
async def iter_messages_chunk(cls, messages: AsyncIterator, session: StreamSession, fields: ResponseFields) -> AsyncIterator:
446
449
last_message: int = 0
447
450
async for message in messages:
448
451
if message.startswith(b'{"wss_url":'):
449
async with session.ws_connect(json.loads(message)["wss_url"]) as ws:
450
async for chunk in cls.iter_messages_chunk(cls.iter_messages_ws(ws), session, fields):
452
message = json.loads(message)
453
async with session.ws_connect(message["wss_url"]) as ws:
454
async for chunk in cls.iter_messages_chunk(cls.iter_messages_ws(ws, message["conversation_id"]), session, fields):
451
455
yield chunk
452
456
break
453
457
async for chunk in cls.iter_messages_line(session, message, fields):
@@ -589,6 +593,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
589
593
@classmethod
590
594
def _set_api_key(cls, api_key: str):
591
595
cls._api_key = api_key
596
cls._expires = int(time.time()) + 60 * 60 * 4
592
597
cls._headers["Authorization"] = f"Bearer {api_key}"
593
598
594
599
@classmethod