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

XFEstudio/gpt4free

fix: use ref_index and ref_type to match sources. add image, video, forecast references support

8d7a31a3
GravityTwoG <crytekov@gmail.com>
提交于

代码差异

1 个文件 +305 -18
Modified g4f/Provider/needs_auth/OpenaiChat.py +305 -18
@@ -8,7 +8,7 @@ import json
8 8 import base64
9 9 import time
10 10 import random
11 from typing import AsyncIterator, Iterator, Optional, Generator, Dict, Union
11 from typing import AsyncIterator, Iterator, Optional, Generator, Dict, Union, List, Any
12 12 from copy import copy
13 13
14 14 try:
@@ -24,8 +24,8 @@ from ...requests import StreamSession
24 24 from ...requests import get_nodriver
25 25 from ...image import ImageRequest, to_image, to_bytes, is_accepted_format
26 26 from ...errors import MissingAuthError, NoValidHarFileError, ModelNotFoundError
27 from ...providers.response import JsonConversation, FinishReason, SynthesizeData, AuthResult, ImageResponse, ImagePreview
28 from ...providers.response import Sources, TitleGeneration, RequestLogin, Reasoning
27 from ...providers.response import JsonConversation, FinishReason, SynthesizeData, AuthResult, ImageResponse, ImagePreview, ResponseType, format_link
28 from ...providers.response import TitleGeneration, RequestLogin, Reasoning
29 29 from ...tools.media import merge_media
30 30 from ..helper import format_cookies, format_media_prompt, to_string
31 31 from ..openai.models import default_model, default_image_model, models, image_models, text_models, model_aliases
@@ -370,7 +370,8 @@ class OpenaiChat(AsyncAuthedProvider, ProviderModelMixin):
370 370 if cls._api_key is None:
371 371 auto_continue = False
372 372 conversation.finish_reason = None
373 sources = Sources([])
373 sources = OpenAISources([])
374 references = ContentReferences()
374 375 while conversation.finish_reason is None:
375 376 async with session.post(
376 377 f"{cls.url}/backend-anon/sentinel/chat-requirements"
@@ -475,29 +476,99 @@ class OpenaiChat(AsyncAuthedProvider, ProviderModelMixin):
475 476 generated_image = await cls.get_generated_image(session, auth_result, match.group(0), prompt)
476 477 if generated_image is not None:
477 478 yield generated_image
478 async for chunk in cls.iter_messages_line(session, auth_result, line, conversation, sources):
479 async for chunk in cls.iter_messages_line(session, auth_result, line, conversation, sources, references):
479 480 if isinstance(chunk, str):
480 481 chunk = chunk.replace("\ue203", "").replace("\ue204", "").replace("\ue206", "")
481 482 buffer += chunk
482 483 if buffer.find(u"\ue200") != -1:
483 484 if buffer.find(u"\ue201") != -1:
484 buffer = buffer.replace("\ue200", "").replace("\ue202", "\n").replace("\ue201", "")
485 buffer = buffer.replace("navlist\n", "#### ")
486 def replacer(match):
487 link = None
488 if len(sources.list) > int(match.group(1)):
489 link = sources.list[int(match.group(1))]["url"]
490 return f"[[{int(match.group(1))+1}]]({link})"
491 return f" [{int(match.group(1))+1}]"
492 buffer = re.sub(r'(?:cite\nturn[0-9]+|turn[0-9]+)(?:search|news|view)(\d+)', replacer, buffer)
485 def sequence_replacer(match):
486 def citation_replacer(match: re.Match[str]):
487 ref_type = match.group(1)
488 ref_index = int(match.group(2))
489 if ((ref_type == "image" and is_image_embedding) or
490 is_video_embedding or
491 ref_type == "forecast"):
492
493 reference = references.get_reference({
494 "ref_index": ref_index,
495 "ref_type": ref_type
496 })
497 if not reference:
498 return ""
499
500 if ref_type == "forecast":
501 if reference.get("alt"):
502 return reference.get("alt")
503 if reference.get("prompt_text"):
504 return reference.get("prompt_text")
505
506 if is_image_embedding and reference.get("content_url", ""):
507 return f"![{reference.get('title', '')}]({reference.get('content_url')})"
508
509 if is_video_embedding:
510 if reference.get("url", "") and reference.get("thumbnail_url", ""):
511 return f"[![{reference.get('title', '')}]({reference['thumbnail_url']})]({reference['url']})"
512 video_match = re.match(r"video\n(.*?)\nturn[0-9]+", match.group(0))
513 if video_match:
514 return video_match.group(1)
515 return ""
516
517 source_index = sources.get_index({
518 "ref_index": ref_index,
519 "ref_type": ref_type
520 })
521 if source_index is not None and len(sources.list) > source_index:
522 link = sources.list[source_index]["url"]
523 return f"[[{source_index+1}]]({link})"
524 return f""
525
526 def products_replacer(match: re.Match[str]):
527 try:
528 products_data = json.loads(match.group(1))
529 products_str = ""
530 for idx, _ in enumerate(products_data.get("selections", []) or []):
531 name = products_data.get('selections', [])[idx][1]
532 tags = products_data.get('tags', [])[idx]
533 products_str += f"{name} - {tags}\n\n"
534
535 return products_str
536 except:
537 return ""
538
539 sequence_content = match.group(1)
540 sequence_content = sequence_content.replace("\ue200", "").replace("\ue202", "\n").replace("\ue201", "")
541 sequence_content = sequence_content.replace("navlist\n", "#### ")
542
543 # Handle search, news, view and image citations
544 is_image_embedding = sequence_content.startswith("i\nturn")
545 is_video_embedding = sequence_content.startswith("video\n")
546 sequence_content = re.sub(
547 r'(?:cite\nturn[0-9]+|forecast\nturn[0-9]+|video\n.*?\nturn[0-9]+|i?\n?turn[0-9]+)(search|news|view|image|forecast)(\d+)',
548 citation_replacer,
549 sequence_content
550 )
551 sequence_content = re.sub(r'products\n(.*)', products_replacer, sequence_content)
552 sequence_content = re.sub(r'product_entity\n\[".*","(.*)"\]', lambda x: x.group(1), sequence_content)
553 return sequence_content
554
555 # process only completed sequences and do not touch start of next not completed sequence
556 buffer = re.sub(r'\ue200(.*?)\ue201', sequence_replacer, buffer, flags=re.DOTALL)
557
558 if buffer.find(u"\ue200") != -1: # still have uncompleted sequence
559 continue
493 560 else:
561 # do not yield to consume rest part of special sequence
494 562 continue
563
495 564 yield buffer
496 565 buffer = ""
497 566 else:
498 567 yield chunk
499 568 if conversation.finish_reason is not None:
500 569 break
570 if buffer:
571 yield buffer
501 572 if sources.list:
502 573 yield sources
503 574 if conversation.generated_images:
@@ -521,7 +592,7 @@ class OpenaiChat(AsyncAuthedProvider, ProviderModelMixin):
521 592 yield FinishReason(conversation.finish_reason)
522 593
523 594 @classmethod
524 async def iter_messages_line(cls, session: StreamSession, auth_result: AuthResult, line: bytes, fields: Conversation, sources: Sources) -> AsyncIterator:
595 async def iter_messages_line(cls, session: StreamSession, auth_result: AuthResult, line: bytes, fields: Conversation, sources: OpenAISources, references: ContentReferences) -> AsyncIterator:
525 596 if not line.startswith(b"data: "):
526 597 return
527 598 elif line.startswith(b"data: [DONE]"):
@@ -553,31 +624,84 @@ class OpenaiChat(AsyncAuthedProvider, ProviderModelMixin):
553 624 elif "p" not in line or line.get("p") == "/message/content/parts/0":
554 625 yield Reasoning(token=v) if fields.is_thinking else v
555 626 elif isinstance(v, list):
627 buffer = ""
556 628 for m in v:
557 629 if m.get("p") == "/message/content/parts/0" and fields.recipient == "all":
558 yield m.get("v")
630 buffer += m.get("v")
559 631 elif m.get("p") == "/message/metadata/image_gen_title":
560 632 fields.prompt = m.get("v")
561 633 elif m.get("p") == "/message/content/parts/0/asset_pointer":
562 634 generated_images = fields.generated_images = await cls.get_generated_image(session, auth_result, m.get("v"), fields.prompt, fields.conversation_id)
563 635 if generated_images is not None:
636 if buffer:
637 yield buffer
564 638 yield generated_images
565 639 elif m.get("p") == "/message/metadata/search_result_groups":
566 640 for entry in [p.get("entries") for p in m.get("v")]:
567 641 for link in entry:
568 642 sources.add_source(link)
569 elif m.get("p") == "/message/metadata/content_references":
643 elif m.get("p") == "/message/metadata/content_references" and not isinstance(m.get("v"), int):
570 644 for entry in m.get("v"):
571 645 for link in entry.get("sources", []):
572 646 sources.add_source(link)
647 for link in entry.get("items", []):
648 sources.add_source(link)
649 for link in entry.get("fallback_items", []) or []:
650 sources.add_source(link)
651 if m.get("o", None) == "append":
652 references.add_reference(entry)
573 653 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+$", m.get("p")):
574 sources.add_source(m.get("v"))
654 if "url" in m.get("v") or "link" in m.get("v"):
655 sources.add_source(m.get("v"))
656 for link in m.get("v").get("fallback_items", []) or []:
657 sources.add_source(link)
658
659 match = re.match(r"^/message/metadata/content_references/(\d+)$", m.get("p"))
660 if match and m.get("o") == "append" and isinstance(m.get("v"), dict):
661 idx = int(match.group(1))
662 references.merge_reference(idx, m.get("v"))
663 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+/fallback_items$", m.get("p")) and isinstance(m.get("v"), list):
664 for link in m.get("v", []) or []:
665 sources.add_source(link)
666 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+/items$", m.get("p")) and isinstance(m.get("v"), list):
667 for link in m.get("v", []) or []:
668 sources.add_source(link)
669 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+/refs$", m.get("p")) and isinstance(m.get("v"), list):
670 match = re.match(r"^/message/metadata/content_references/(\d+)/refs$", m.get("p"))
671 if match:
672 idx = int(match.group(1))
673 references.update_reference(idx, m.get("o"), "refs", m.get("v"))
674 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+/alt$", m.get("p")) and isinstance(m.get("v"), list):
675 match = re.match(r"^/message/metadata/content_references/(\d+)/alt$", m.get("p"))
676 if match:
677 idx = int(match.group(1))
678 references.update_reference(idx, m.get("o"), "alt", m.get("v"))
679 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+/prompt_text$", m.get("p")) and isinstance(m.get("v"), list):
680 match = re.match(r"^/message/metadata/content_references/(\d+)/prompt_text$", m.get("p"))
681 if match:
682 idx = int(match.group(1))
683 references.update_reference(idx, m.get("o"), "prompt_text", m.get("v"))
684 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+/refs/\d+$", m.get("p")) and isinstance(m.get("v"), dict):
685 match = re.match(r"^/message/metadata/content_references/(\d+)/refs/(\d+)$", m.get("p"))
686 if match:
687 reference_idx = int(match.group(1))
688 ref_idx = int(match.group(2))
689 references.update_reference(reference_idx, m.get("o"), "refs", m.get("v"), ref_idx)
690 elif m.get("p") and re.match(r"^/message/metadata/content_references/\d+/images$", m.get("p")) and isinstance(m.get("v"), list):
691 match = re.match(r"^/message/metadata/content_references/(\d+)/images$", m.get("p"))
692 if match:
693 idx = int(match.group(1))
694 references.update_reference(idx, m.get("o"), "images", m.get("v"))
575 695 elif m.get("p") == "/message/metadata/finished_text":
576 696 fields.is_thinking = False
697 if buffer:
698 yield buffer
577 699 yield Reasoning(status=m.get("v"))
578 700 elif m.get("p") == "/message/metadata" and fields.recipient == "all":
579 701 fields.finish_reason = m.get("v", {}).get("finish_details", {}).get("type")
580 702 break
703
704 yield buffer
581 705 elif isinstance(v, dict):
582 706 if fields.conversation_id is None:
583 707 fields.conversation_id = v.get("conversation_id")
@@ -794,3 +918,166 @@ def get_cookies(
794 918 }
795 919 json = yield cmd_dict
796 920 return {c["name"]: c["value"] for c in json['cookies']} if 'cookies' in json else {}
921
922 class OpenAISources(ResponseType):
923 list: List[Dict[str, str]]
924
925 def __init__(self, sources: List[Dict[str, str]]) -> None:
926 """Initialize with a list of source dictionaries."""
927 self.list = []
928 for source in sources:
929 self.add_source(source)
930
931 def add_source(self, source: Union[Dict[str, str], str]) -> None:
932 """Add a source to the list, cleaning the URL if necessary."""
933 source = source if isinstance(source, dict) else {"url": source}
934 url = source.get("url", source.get("link", None))
935 if not url:
936 return
937
938 url = re.sub(r"[&?]utm_source=.+", "", url)
939 source["url"] = url
940
941 ref_info = self.get_ref_info(source)
942 if ref_info:
943 existing_source, idx = self.find_by_ref_info(ref_info)
944 if existing_source and idx is not None:
945 self.list[idx] = source
946 return
947
948 existing_source, idx = self.find_by_url(source["url"])
949 if existing_source and idx is not None:
950 self.list[idx] = source
951 return
952
953 self.list.append(source)
954
955 def __str__(self) -> str:
956 """Return formatted sources as a string."""
957 if not self.list:
958 return ""
959 return "\n\n\n\n" + ("\n>\n".join([
960 f"> [{idx+1}] {format_link(link['url'], link.get('title', ''))}"
961 for idx, link in enumerate(self.list)
962 ]))
963
964 def get_ref_info(self, source: Dict[str, str]) -> dict[str, str|int] | None:
965 ref_index = source.get("ref_id", {}).get("ref_index", None)
966 ref_type = source.get("ref_id", {}).get("ref_type", None)
967 if isinstance(ref_index, int):
968 return {
969 "ref_index": ref_index,
970 "ref_type": ref_type,
971 }
972
973 for ref_info in source.get('refs') or []:
974 ref_index = ref_info.get("ref_index", None)
975 ref_type = ref_info.get("ref_type", None)
976 if isinstance(ref_index, int):
977 return {
978 "ref_index": ref_index,
979 "ref_type": ref_type,
980 }
981
982 return None
983
984 def find_by_ref_info(self, ref_info: dict[str, str|int]):
985 for idx, source in enumerate(self.list):
986 source_ref_info = self.get_ref_info(source)
987 if (source_ref_info and
988 source_ref_info["ref_index"] == ref_info["ref_index"] and
989 source_ref_info["ref_type"] == ref_info["ref_type"]):
990 return source, idx
991
992 return None, None
993
994 def find_by_url(self, url: str):
995 for idx, source in enumerate(self.list):
996 if source["url"] == url:
997 return source, idx
998 return None, None
999
1000 def get_index(self, ref_info: dict[str, str|int]) -> int | None:
1001 _, index = self.find_by_ref_info(ref_info)
1002 if index is not None:
1003 return index
1004
1005 return None
1006
1007 class ContentReferences:
1008 def __init__(self) -> None:
1009 self.list: List[Dict[str, Any]] = []
1010
1011 def add_reference(self, reference_part: dict) -> None:
1012 self.list.append(reference_part)
1013
1014 def merge_reference(self, idx: int, reference_part: dict):
1015 while len(self.list) <= idx:
1016 self.list.append({})
1017
1018 self.list[idx] = {**self.list[idx], **reference_part}
1019
1020 def update_reference(self, idx: int, operation: str, field: str, value: Any, ref_idx = None) -> None:
1021 while len(self.list) <= idx:
1022 self.list.append({})
1023
1024 if operation == "append" or operation == "add":
1025 if not isinstance(self.list[idx].get(field, None), list):
1026 self.list[idx][field] = []
1027 if isinstance(value, list):
1028 self.list[idx][field].extend(value)
1029 else:
1030 self.list[idx][field].append(value)
1031
1032 if operation == "replace" and ref_idx is not None:
1033 if field == "refs" and not isinstance(self.list[idx].get(field, None), list):
1034 self.list[idx][field] = []
1035
1036 if isinstance(self.list[idx][field], list):
1037 if len(self.list[idx][field]) <= ref_idx:
1038 self.list[idx][field].append(value)
1039 else:
1040 self.list[idx][field][ref_idx] = value
1041 else:
1042 self.list[idx][field] = value
1043
1044 def get_ref_info(
1045 self,
1046 source: Dict[str, str],
1047 target_ref_info: Dict[str, Union[str, int]]
1048 ) -> dict[str, str|int] | None:
1049 for idx, ref_info in enumerate(source.get("refs", [])) or []:
1050 if not isinstance(ref_info, dict):
1051 continue
1052
1053 ref_index = ref_info.get("ref_index", None)
1054 ref_type = ref_info.get("ref_type", None)
1055 if isinstance(ref_index, int) and isinstance(ref_type, str):
1056 if (not target_ref_info or
1057 (target_ref_info["ref_index"] == ref_index and
1058 target_ref_info["ref_type"] == ref_type)):
1059 return {
1060 "ref_index": ref_index,
1061 "ref_type": ref_type,
1062 "idx": idx
1063 }
1064
1065 return None
1066
1067 def get_reference(self, ref_info: Dict[str, Union[str, int]]) -> Any:
1068 for reference in self.list:
1069 reference_ref_info = self.get_ref_info(reference, ref_info)
1070
1071 if (not reference_ref_info or
1072 reference_ref_info["ref_index"] != ref_info["ref_index"] or
1073 reference_ref_info["ref_type"] != ref_info["ref_type"]):
1074 continue
1075
1076 if ref_info["ref_type"] != "image":
1077 return reference
1078
1079 images = reference.get("images", [])
1080 if isinstance(images, list) and len(images) > reference_ref_info["idx"]:
1081 return images[reference_ref_info["idx"]]
1082
1083 return None