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

XFEstudio/gpt4free

Expire cache, Fix multiple websocket conversations in OpenaiChat Map system messages to user messages in GeminiPro

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

代码差异

2 个文件 +19 -16
Modified g4f/Provider/GeminiPro.py +8 -10
@@ -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
Modified g4f/Provider/needs_auth/OpenaiChat.py +11 -6
@@ -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