返回提交历史
Modified
g4f/Provider/needs_auth/OpenaiChat.py
+133
-95
Modified
g4f/providers/base_provider.py
+1
-1
Modified
g4f/providers/types.py
+1
-0
XFEstudio/gpt4free
Add websocket support in OpenaiChat
ac86e576
代码差异
3 个文件
+135
-96
@@ -4,6 +4,8 @@ import asyncio
4
4
import uuid
5
5
import json
6
6
import os
7
import base64
8
from aiohttp import ClientWebSocketResponse
7
9
8
10
try:
9
11
from py_arkose_generator.arkose import get_values_for_request
@@ -22,7 +24,7 @@ except ImportError:
22
24
from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
23
25
from ..helper import get_cookies
24
26
from ...webdriver import get_browser
25
from ...typing import AsyncResult, Messages, Cookies, ImageType, Union
27
from ...typing import AsyncResult, Messages, Cookies, ImageType, Union, AsyncIterator
26
28
from ...requests import get_args_from_browser
27
29
from ...requests.aiohttp import StreamSession
28
30
from ...image import to_image, to_bytes, ImageResponse, ImageRequest
@@ -38,10 +40,14 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
38
40
supports_gpt_35_turbo = True
39
41
supports_gpt_4 = True
40
42
supports_message_history = True
43
supports_system_message = True
41
44
default_model = None
42
45
models = ["gpt-3.5-turbo", "gpt-4", "gpt-4-gizmo"]
43
model_aliases = {"text-davinci-002-render-sha": "gpt-3.5-turbo"}
44
_args: dict = None
46
model_aliases = {"text-davinci-002-render-sha": "gpt-3.5-turbo", "": "gpt-3.5-turbo"}
47
_api_key: str = None
48
_headers: dict = None
49
_cookies: Cookies = None
50
_last_message: int = 0
45
51
46
52
@classmethod
47
53
async def create(
@@ -299,6 +305,7 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
299
305
conversation_id: str = None,
300
306
parent_id: str = None,
301
307
image: ImageType = None,
308
image_name: str = None,
302
309
response_fields: bool = False,
303
310
**kwargs
304
311
) -> AsyncResult:
@@ -332,67 +339,64 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
332
339
if not parent_id:
333
340
parent_id = str(uuid.uuid4())
334
341
335
# Read api_key from args
342
# Read api_key from arguments
336
343
api_key = kwargs["access_token"] if "access_token" in kwargs else api_key
337
# If no cached args
338
if cls._args is None:
339
if api_key is None:
340
# Read api_key from cookies
341
cookies = get_cookies("chat.openai.com", False) if cookies is None else cookies
342
api_key = cookies["access_token"] if "access_token" in cookies else api_key
343
cls._args = cls._create_request_args(cookies)
344
else:
345
# Read api_key from cache
346
api_key = cls._args["headers"]["Authorization"] if "Authorization" in cls._args["headers"] else None
347
344
348
345
async with StreamSession(
349
346
proxies={"https": proxy},
350
347
impersonate="chrome",
351
348
timeout=timeout
352
349
) as session:
353
# Read api_key from session cookies
350
# Read api_key and cookies from cache / browser config
351
if cls._headers is None:
352
if api_key is None:
353
# Read api_key from cookies
354
cookies = get_cookies("chat.openai.com", False) if cookies is None else cookies
355
api_key = cookies["access_token"] if "access_token" in cookies else api_key
356
cls._create_request_args(cookies)
357
else:
358
api_key = cls._api_key if api_key is None else api_key
359
# Read api_key with session cookies
354
360
if api_key is None and cookies:
355
api_key = await cls.fetch_access_token(session, cls._args["headers"])
361
api_key = await cls.fetch_access_token(session, cls._headers)
356
362
# Load default model
357
if cls.default_model is None:
363
if cls.default_model is None and api_key is not None:
358
364
try:
359
if cookies and not model and api_key is not None:
360
cls._args["headers"]["Authorization"] = api_key
361
cls.default_model = cls.get_model(await cls.get_default_model(session, cls._args["headers"]))
362
elif api_key:
363
cls.default_model = cls.get_model(model or "gpt-3.5-turbo")
365
if not model:
366
cls._set_api_key(api_key)
367
cls.default_model = cls.get_model(await cls.get_default_model(session, cls._headers))
368
else:
369
cls.default_model = cls.get_model(model)
364
370
except Exception as e:
365
371
if debug.logging:
366
372
print("OpenaiChat: Load default_model failed")
367
373
print(f"{e.__class__.__name__}: {e}")
368
# Browse api_key and update default model
374
# Browse api_key and default model
369
375
if api_key is None or cls.default_model is None:
370
376
login_url = os.environ.get("G4F_LOGIN_URL")
371
377
if login_url:
372
378
yield f"Please login: [ChatGPT]({login_url})\n\n"
373
379
try:
374
cls._args = cls.browse_access_token(proxy)
380
cls.browse_access_token(proxy)
375
381
except MissingRequirementsError:
376
382
raise MissingAuthError(f'Missing "access_token". Add a "api_key" please')
377
cls.default_model = cls.get_model(await cls.get_default_model(session, cls._args["headers"]))
383
cls.default_model = cls.get_model(await cls.get_default_model(session, cls._headers))
378
384
else:
379
cls._args["headers"]["Authorization"] = api_key
385
cls._set_api_key(api_key)
380
386
381
387
try:
382
image_response = await cls.upload_image(
383
session,
384
cls._args["headers"],
385
image,
386
kwargs.get("image_name")
387
) if image else None
388
image_request = await cls.upload_image(session, cls._headers, image, image_name) if image else None
388
389
except Exception as e:
389
yield e
390
if debug.logging:
391
print("OpenaiChat: Upload image failed")
392
print(f"{e.__class__.__name__}: {e}")
390
393
391
end_turn = EndTurn()
392
model = cls.get_model(model)
393
model = "text-davinci-002-render-sha" if model == "gpt-3.5-turbo" else model
394
while not end_turn.is_end:
394
model = cls.get_model(model).replace("gpt-3.5-turbo", "text-davinci-002-render-sha")
395
fields = ResponseFields()
396
while fields.finish_reason is None:
395
397
arkose_token = await cls.get_arkose_token(session)
398
conversation_id = conversation_id if fields.conversation_id is None else fields.conversation_id
399
parent_id = parent_id if fields.message_id is None else fields.message_id
396
400
data = {
397
401
"action": action,
398
402
"arkose_token": arkose_token,
@@ -405,8 +409,8 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
405
409
"history_and_training_disabled": history_disabled and not auto_continue,
406
410
}
407
411
if action != "continue":
408
messages = messages if not conversation_id else [messages[-1]]
409
data["messages"] = cls.create_messages(messages, image_response)
412
messages = messages if conversation_id is None else [messages[-1]]
413
data["messages"] = cls.create_messages(messages, image_request)
410
414
411
415
async with session.post(
412
416
f"{cls.url}/backend-api/conversation",
@@ -414,63 +418,88 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
414
418
headers={
415
419
"Accept": "text/event-stream",
416
420
"OpenAI-Sentinel-Arkose-Token": arkose_token,
417
**cls._args["headers"]
421
**cls._headers
418
422
}
419
423
) as response:
420
424
cls._update_request_args(session)
421
425
if not response.ok:
422
message = f"{await response.text()} headers:\n{json.dumps(cls._args['headers'], indent=4)}"
423
raise RuntimeError(f"Response {response.status}: {message}")
424
last_message: int = 0
425
async for line in response.iter_lines():
426
if not line.startswith(b"data: "):
427
continue
428
elif line.startswith(b"data: [DONE]"):
429
break
430
try:
431
line = json.loads(line[6:])
432
except:
433
continue
434
if "message" not in line:
435
continue
436
if "error" in line and line["error"]:
437
raise RuntimeError(line["error"])
438
if "message_type" not in line["message"]["metadata"]:
439
continue
440
try:
441
image_response = await cls.get_generated_image(session, cls._args["headers"], line)
442
if image_response is not None:
443
yield image_response
444
except Exception as e:
445
yield e
446
if line["message"]["author"]["role"] != "assistant":
447
continue
448
if line["message"]["content"]["content_type"] != "text":
449
continue
450
if line["message"]["metadata"]["message_type"] not in ("next", "continue", "variant"):
451
continue
452
conversation_id = line["conversation_id"]
453
parent_id = line["message"]["id"]
426
raise RuntimeError(f"Response {response.status}: {await response.text()}")
427
async for chunk in cls.iter_messages_chunk(response.iter_lines(), session, fields):
454
428
if response_fields:
455
429
response_fields = False
456
yield ResponseFields(conversation_id, parent_id, end_turn)
457
if "parts" in line["message"]["content"]:
458
new_message = line["message"]["content"]["parts"][0]
459
if len(new_message) > last_message:
460
yield new_message[last_message:]
461
last_message = len(new_message)
462
if "finish_details" in line["message"]["metadata"]:
463
if line["message"]["metadata"]["finish_details"]["type"] == "stop":
464
end_turn.end()
430
yield fields
431
yield chunk
465
432
if not auto_continue:
466
433
break
467
434
action = "continue"
468
435
await asyncio.sleep(5)
469
436
if history_disabled and auto_continue:
470
await cls.delete_conversation(session, cls._args["headers"], conversation_id)
437
await cls.delete_conversation(session, cls._headers, conversation_id)
438
439
@staticmethod
440
async def iter_messages_ws(ws: ClientWebSocketResponse) -> AsyncIterator:
441
while True:
442
yield base64.b64decode((await ws.receive_json())["body"])
443
444
@classmethod
445
async def iter_messages_chunk(cls, messages: AsyncIterator, session: StreamSession, fields: ResponseFields) -> AsyncIterator:
446
last_message: int = 0
447
async for message in messages:
448
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):
451
yield chunk
452
break
453
async for chunk in cls.iter_messages_line(session, message, fields):
454
if fields.finish_reason is not None:
455
break
456
elif isinstance(chunk, str):
457
if len(chunk) > last_message:
458
yield chunk[last_message:]
459
last_message = len(chunk)
460
else:
461
yield chunk
462
if fields.finish_reason is not None:
463
break
464
465
@classmethod
466
async def iter_messages_line(cls, session: StreamSession, line: bytes, fields: ResponseFields) -> AsyncIterator:
467
if not line.startswith(b"data: "):
468
return
469
elif line.startswith(b"data: [DONE]"):
470
return
471
try:
472
line = json.loads(line[6:])
473
except:
474
return
475
if "message" not in line:
476
return
477
if "error" in line and line["error"]:
478
raise RuntimeError(line["error"])
479
if "message_type" not in line["message"]["metadata"]:
480
return
481
try:
482
image_response = await cls.get_generated_image(session, cls._headers, line)
483
if image_response is not None:
484
yield image_response
485
except Exception as e:
486
yield e
487
if line["message"]["author"]["role"] != "assistant":
488
return
489
if line["message"]["content"]["content_type"] != "text":
490
return
491
if line["message"]["metadata"]["message_type"] not in ("next", "continue", "variant"):
492
return
493
if fields.conversation_id is None:
494
fields.conversation_id = line["conversation_id"]
495
fields.message_id = line["message"]["id"]
496
if "parts" in line["message"]["content"]:
497
yield line["message"]["content"]["parts"][0]
498
if "finish_details" in line["message"]["metadata"]:
499
fields.finish_reason = line["message"]["metadata"]["finish_details"]["type"]
471
500
472
501
@classmethod
473
def browse_access_token(cls, proxy: str = None, timeout: int = 1200) -> tuple[str, dict]:
502
def browse_access_token(cls, proxy: str = None, timeout: int = 1200) -> None:
474
503
"""
475
504
Browse to obtain an access token.
476
505
@@ -493,9 +522,10 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
493
522
"return accessToken;"
494
523
)
495
524
args = get_args_from_browser(f"{cls.url}/", driver, do_bypass_cloudflare=False)
496
args["headers"]["Authorization"] = f"Bearer {access_token}"
497
args["headers"]["Cookie"] = cls._format_cookies(args["cookies"])
498
return args
525
cls._headers = args["headers"]
526
cls._cookies = args["cookies"]
527
cls._update_cookie_header()
528
cls._set_api_key(access_token)
499
529
finally:
500
530
driver.close()
501
531
@@ -546,16 +576,24 @@ class OpenaiChat(AsyncGeneratorProvider, ProviderModelMixin):
546
576
547
577
@classmethod
548
578
def _create_request_args(cls, cookies: Union[Cookies, None]):
549
return {
550
"headers": {} if cookies is None else {"Cookie": cls._format_cookies(cookies)},
551
"cookies": {} if cookies is None else cookies
552
}
579
cls._headers = {}
580
cls._cookies = {} if cookies is None else cookies
581
cls._update_cookie_header()
553
582
554
583
@classmethod
555
584
def _update_request_args(cls, session: StreamSession):
556
585
for c in session.cookie_jar if hasattr(session, "cookie_jar") else session.cookies.jar:
557
cls._args["cookies"][c.name if hasattr(c, "name") else c.key] = c.value
558
cls._args["headers"]["Cookie"] = cls._format_cookies(cls._args["cookies"])
586
cls._cookies[c.name if hasattr(c, "name") else c.key] = c.value
587
cls._update_cookie_header()
588
589
@classmethod
590
def _set_api_key(cls, api_key: str):
591
cls._api_key = api_key
592
cls._headers["Authorization"] = f"Bearer {api_key}"
593
594
@classmethod
595
def _update_cookie_header(cls):
596
cls._headers["Cookie"] = cls._format_cookies(cls._cookies)
559
597
560
598
class EndTurn:
561
599
"""
@@ -571,10 +609,10 @@ class ResponseFields:
571
609
"""
572
610
Class to encapsulate response fields.
573
611
"""
574
def __init__(self, conversation_id: str, message_id: str, end_turn: EndTurn):
612
def __init__(self, conversation_id: str = None, message_id: str = None, finish_reason: str = None):
575
613
self.conversation_id = conversation_id
576
614
self.message_id = message_id
577
self._end_turn = end_turn
615
self.finish_reason = finish_reason
578
616
579
617
class Response():
580
618
"""
@@ -608,7 +646,7 @@ class Response():
608
646
self._message = "".join(chunks)
609
647
if not self._fields:
610
648
raise RuntimeError("Missing response fields")
611
self.is_end = self._fields._end_turn.is_end
649
self.is_end = self._fields.end_turn
612
650
613
651
def __aiter__(self):
614
652
return self.generator()
@@ -270,7 +270,7 @@ class ProviderModelMixin:
270
270
271
271
@classmethod
272
272
def get_model(cls, model: str) -> str:
273
if not model:
273
if not model and cls.default_model is not None:
274
274
model = cls.default_model
275
275
elif model in cls.model_aliases:
276
276
model = cls.model_aliases[model]
@@ -26,6 +26,7 @@ class BaseProvider(ABC):
26
26
supports_gpt_35_turbo: bool = False
27
27
supports_gpt_4: bool = False
28
28
supports_message_history: bool = False
29
supports_system_message: bool = False
29
30
params: str
30
31
31
32
@classmethod