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

XFEstudio/gpt4free

Add websocket support in OpenaiChat

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

代码差异

3 个文件 +135 -96
Modified g4f/Provider/needs_auth/OpenaiChat.py +133 -95
@@ -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()
Modified g4f/providers/base_provider.py +1 -1
@@ -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]
Modified g4f/providers/types.py +1 -0
@@ -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