返回提交历史
Modified
g4f/Provider/needs_auth/OpenaiChat.py
+305
-18
XFEstudio/gpt4free
fix: use ref_index and ref_type to match sources. add image, video, forecast references support
8d7a31a3
代码差异
1 个文件
+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"})"
508
509
if is_video_embedding:
510
if reference.get("url", "") and reference.get("thumbnail_url", ""):
511
return f"[]({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