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

XFEstudio/gpt4free

Add Cerebras and HuggingFace2 provider, Fix RubiksAI provider Add support for image generation in Copilot provider

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

代码差异

14 个文件 +239 -137
Modified g4f/Provider/Blackbox.py +10 -2
@@ -28,6 +28,9 @@ class Blackbox(AsyncGeneratorProvider, ProviderModelMixin):
28 28 image_models = [default_image_model, 'repomap']
29 29 text_models = [default_model, 'gpt-4o', 'gemini-pro', 'claude-sonnet-3.5', 'blackboxai-pro']
30 30 vision_models = [default_model, 'gpt-4o', 'gemini-pro', 'blackboxai-pro']
31 model_aliases = {
32 "claude-3.5-sonnet": "claude-sonnet-3.5",
33 }
31 34 agentMode = {
32 35 default_image_model: {'mode': True, 'id': "ImageGenerationLV45LJp", 'name': "Image Generation"},
33 36 }
@@ -198,6 +201,7 @@ class Blackbox(AsyncGeneratorProvider, ProviderModelMixin):
198 201 async with ClientSession(headers=headers) as session:
199 202 async with session.post(cls.api_endpoint, json=data, proxy=proxy) as response:
200 203 response.raise_for_status()
204 is_first = False
201 205 async for chunk in response.content.iter_any():
202 206 text_chunk = chunk.decode(errors="ignore")
203 207 if model in cls.image_models:
@@ -217,5 +221,9 @@ class Blackbox(AsyncGeneratorProvider, ProviderModelMixin):
217 221 for i, result in enumerate(search_results, 1):
218 222 formatted_response += f"\n{i}. {result['title']}: {result['link']}"
219 223 yield formatted_response
220 else:
221 yield text_chunk.strip()
224 elif text_chunk:
225 if is_first:
226 is_first = False
227 yield text_chunk.lstrip()
228 else:
229 yield text_chunk
Modified g4f/Provider/Copilot.py +16 -7
@@ -21,8 +21,9 @@ from .helper import format_prompt
21 21 from ..typing import CreateResult, Messages, ImageType
22 22 from ..errors import MissingRequirementsError
23 23 from ..requests.raise_for_status import raise_for_status
24 from ..providers.helper import format_cookies
24 25 from ..requests import get_nodriver
25 from ..image import to_bytes, is_accepted_format
26 from ..image import ImageResponse, to_bytes, is_accepted_format
26 27 from .. import debug
27 28
28 29 class Conversation(BaseConversation):
@@ -70,18 +71,21 @@ class Copilot(AbstractProvider):
70 71 access_token, cookies = asyncio.run(cls.get_access_token_and_cookies(proxy))
71 72 else:
72 73 access_token = conversation.access_token
73 websocket_url = f"{websocket_url}&acessToken={quote(access_token)}"
74 headers = {"Authorization": f"Bearer {access_token}"}
74 debug.log(f"Copilot: Access token: {access_token[:7]}...{access_token[-5:]}")
75 debug.log(f"Copilot: Cookies: {';'.join([*cookies])}")
76 websocket_url = f"{websocket_url}&accessToken={quote(access_token)}"
77 headers = {"authorization": f"Bearer {access_token}", "cookie": format_cookies(cookies)}
75 78
76 79 with Session(
77 80 timeout=timeout,
78 81 proxy=proxy,
79 82 impersonate="chrome",
80 83 headers=headers,
81 cookies=cookies
84 cookies=cookies,
82 85 ) as session:
83 response = session.get(f"{cls.url}/")
86 response = session.get("https://copilot.microsoft.com/c/api/user")
84 87 raise_for_status(response)
88 debug.log(f"Copilot: User: {response.json().get('firstName', 'null')}")
85 89 if conversation is None:
86 90 response = session.post(cls.conversation_url)
87 91 raise_for_status(response)
@@ -119,6 +123,7 @@ class Copilot(AbstractProvider):
119 123
120 124 is_started = False
121 125 msg = None
126 image_prompt: str = None
122 127 while True:
123 128 try:
124 129 msg = wss.recv()[0]
@@ -128,7 +133,11 @@ class Copilot(AbstractProvider):
128 133 if msg.get("event") == "appendText":
129 134 is_started = True
130 135 yield msg.get("text")
131 elif msg.get("event") in ["done", "partCompleted"]:
136 elif msg.get("event") == "generatingImage":
137 image_prompt = msg.get("prompt")
138 elif msg.get("event") == "imageGenerated":
139 yield ImageResponse(msg.get("url"), image_prompt, {"preview": msg.get("thumbnailUrl")})
140 elif msg.get("event") == "done":
132 141 break
133 142 if not is_started:
134 143 raise RuntimeError(f"Last message: {msg}")
@@ -152,7 +161,7 @@ class Copilot(AbstractProvider):
152 161 })()
153 162 """)
154 163 if access_token is None:
155 asyncio.sleep(1)
164 await asyncio.sleep(1)
156 165 cookies = {}
157 166 for c in await page.send(nodriver.cdp.network.get_cookies([cls.url])):
158 167 cookies[c.name] = c.value
Modified g4f/Provider/RubiksAI.py +47 -77
@@ -1,7 +1,6 @@
1
1 2 from __future__ import annotations
2 3
3 import asyncio
4 import aiohttp
5 4 import random
6 5 import string
7 6 import json
@@ -11,34 +10,24 @@ from aiohttp import ClientSession
11 10
12 11 from ..typing import AsyncResult, Messages
13 12 from .base_provider import AsyncGeneratorProvider, ProviderModelMixin
14 from .helper import format_prompt
15
13 from ..requests.raise_for_status import raise_for_status
16 14
17 15 class RubiksAI(AsyncGeneratorProvider, ProviderModelMixin):
18 16 label = "Rubiks AI"
19 17 url = "https://rubiks.ai"
20 api_endpoint = "https://rubiks.ai/search/api.php"
18 api_endpoint = "https://rubiks.ai/search/api/"
21 19 working = True
22 20 supports_stream = True
23 21 supports_system_message = True
24 22 supports_message_history = True
25 23
26 default_model = 'llama-3.1-70b-versatile'
27 models = [default_model, 'gpt-4o-mini']
24 default_model = 'gpt-4o-mini'
25 models = [default_model, 'gpt-4o', 'o1-mini', 'claude-3.5-sonnet', 'grok-beta', 'gemini-1.5-pro', 'nova-pro']
28 26
29 27 model_aliases = {
30 28 "llama-3.1-70b": "llama-3.1-70b-versatile",
31 29 }
32 30
33 @classmethod
34 def get_model(cls, model: str) -> str:
35 if model in cls.models:
36 return model
37 elif model in cls.model_aliases:
38 return cls.model_aliases[model]
39 else:
40 return cls.default_model
41
42 31 @staticmethod
43 32 def generate_mid() -> str:
44 33 """
@@ -70,7 +59,8 @@ class RubiksAI(AsyncGeneratorProvider, ProviderModelMixin):
70 59 model: str,
71 60 messages: Messages,
72 61 proxy: str = None,
73 websearch: bool = False,
62 web_search: bool = False,
63 temperature: float = 0.6,
74 64 **kwargs
75 65 ) -> AsyncResult:
76 66 """
@@ -80,20 +70,18 @@ class RubiksAI(AsyncGeneratorProvider, ProviderModelMixin):
80 70 - model (str): The model to use in the request.
81 71 - messages (Messages): The messages to send as a prompt.
82 72 - proxy (str, optional): Proxy URL, if needed.
83 - websearch (bool, optional): Indicates whether to include search sources in the response. Defaults to False.
73 - web_search (bool, optional): Indicates whether to include search sources in the response. Defaults to False.
84 74 """
85 75 model = cls.get_model(model)
86 prompt = format_prompt(messages)
87 q_value = prompt
88 76 mid_value = cls.generate_mid()
89 referer = cls.create_referer(q=q_value, mid=mid_value, model=model)
90
91 url = cls.api_endpoint
92 params = {
93 'q': q_value,
94 'model': model,
95 'id': '',
96 'mid': mid_value
77 referer = cls.create_referer(q=messages[-1]["content"], mid=mid_value, model=model)
78
79 data = {
80 "messages": messages,
81 "model": model,
82 "search": web_search,
83 "stream": True,
84 "temperature": temperature
97 85 }
98 86
99 87 headers = {
@@ -111,52 +99,34 @@ class RubiksAI(AsyncGeneratorProvider, ProviderModelMixin):
111 99 'sec-ch-ua-mobile': '?0',
112 100 'sec-ch-ua-platform': '"Linux"'
113 101 }
114
115 try:
116 timeout = aiohttp.ClientTimeout(total=None)
117 async with ClientSession(timeout=timeout) as session:
118 async with session.get(url, headers=headers, params=params, proxy=proxy) as response:
119 if response.status != 200:
120 yield f"Request ended with status code {response.status}"
121 return
122
123 assistant_text = ''
124 sources = []
125
126 async for line in response.content:
127 decoded_line = line.decode('utf-8').strip()
128 if not decoded_line.startswith('data: '):
129 continue
130 data = decoded_line[6:]
131 if data in ('[DONE]', '{"done": ""}'):
132 break
133 try:
134 json_data = json.loads(data)
135 except json.JSONDecodeError:
136 continue
137
138 if 'url' in json_data and 'title' in json_data:
139 if websearch:
140 sources.append({'title': json_data['title'], 'url': json_data['url']})
141
142 elif 'choices' in json_data:
143 for choice in json_data['choices']:
144 delta = choice.get('delta', {})
145 content = delta.get('content', '')
146 role = delta.get('role', '')
147 if role == 'assistant':
148 continue
149 assistant_text += content
150
151 if websearch and sources:
152 sources_text = '\n'.join([f"{i+1}. [{s['title']}]: {s['url']}" for i, s in enumerate(sources)])
153 assistant_text += f"\n\n**Source:**\n{sources_text}"
154
155 yield assistant_text
156
157 except asyncio.CancelledError:
158 yield "The request was cancelled."
159 except aiohttp.ClientError as e:
160 yield f"An error occurred during the request: {e}"
161 except Exception as e:
162 yield f"An unexpected error occurred: {e}"
102 async with ClientSession() as session:
103 async with session.post(cls.api_endpoint, headers=headers, json=data, proxy=proxy) as response:
104 await raise_for_status(response)
105
106 sources = []
107 async for line in response.content:
108 decoded_line = line.decode('utf-8').strip()
109 if not decoded_line.startswith('data: '):
110 continue
111 data = decoded_line[6:]
112 if data in ('[DONE]', '{"done": ""}'):
113 break
114 try:
115 json_data = json.loads(data)
116 except json.JSONDecodeError:
117 continue
118
119 if 'url' in json_data and 'title' in json_data:
120 if web_search:
121 sources.append({'title': json_data['title'], 'url': json_data['url']})
122
123 elif 'choices' in json_data:
124 for choice in json_data['choices']:
125 delta = choice.get('delta', {})
126 content = delta.get('content', '')
127 if content:
128 yield content
129
130 if web_search and sources:
131 sources_text = '\n'.join([f"{i+1}. [{s['title']}]: {s['url']}" for i, s in enumerate(sources)])
132 yield f"\n\n**Source:**\n{sources_text}"
Added g4f/Provider/needs_auth/Cerebras.py +65 -0
@@ -0,0 +1,65 @@
1 from __future__ import annotations
2
3 import requests
4 from aiohttp import ClientSession
5
6 from .OpenaiAPI import OpenaiAPI
7 from ...typing import AsyncResult, Messages, Cookies
8 from ...requests.raise_for_status import raise_for_status
9 from ...cookies import get_cookies
10
11 class Cerebras(OpenaiAPI):
12 label = "Cerebras Inference"
13 url = "https://inference.cerebras.ai/"
14 working = True
15 default_model = "llama3.1-70b"
16 fallback_models = [
17 "llama3.1-70b",
18 "llama3.1-8b",
19 ]
20 model_aliases = {"llama-3.1-70b": "llama3.1-70b", "llama-3.1-8b": "llama3.1-8b"}
21
22 @classmethod
23 def get_models(cls, api_key: str = None):
24 if not cls.models:
25 try:
26 headers = {}
27 if api_key:
28 headers["authorization"] = f"Bearer ${api_key}"
29 response = requests.get(f"https://api.cerebras.ai/v1/models", headers=headers)
30 raise_for_status(response)
31 data = response.json()
32 cls.models = [model.get("model") for model in data.get("models")]
33 except Exception:
34 cls.models = cls.fallback_models
35 return cls.models
36
37 @classmethod
38 async def create_async_generator(
39 cls,
40 model: str,
41 messages: Messages,
42 api_base: str = "https://api.cerebras.ai/v1",
43 api_key: str = None,
44 cookies: Cookies = None,
45 **kwargs
46 ) -> AsyncResult:
47 if api_key is None and cookies is None:
48 cookies = get_cookies(".cerebras.ai")
49 async with ClientSession(cookies=cookies) as session:
50 async with session.get("https://inference.cerebras.ai/api/auth/session") as response:
51 raise_for_status(response)
52 data = await response.json()
53 if data:
54 api_key = data.get("user", {}).get("demoApiKey")
55 async for chunk in super().create_async_generator(
56 model, messages,
57 api_base=api_base,
58 impersonate="chrome",
59 api_key=api_key,
60 headers={
61 "User-Agent": "ex/JS 1.5.0",
62 },
63 **kwargs
64 ):
65 yield chunk
Modified g4f/Provider/needs_auth/CopilotAccount.py +5 -2
@@ -1,9 +1,12 @@
1 1 from __future__ import annotations
2 2
3 from ..base_provider import ProviderModelMixin
3 4 from ..Copilot import Copilot
4 5
5 class CopilotAccount(Copilot):
6 class CopilotAccount(Copilot, ProviderModelMixin):
6 7 needs_auth = True
7 8 parent = "Copilot"
8 9 default_model = "Copilot"
9 default_vision_model = default_model
10 default_vision_model = default_model
11 models = [default_model]
12 image_models = models
Added g4f/Provider/needs_auth/HuggingFace2.py +28 -0
@@ -0,0 +1,28 @@
1 from __future__ import annotations
2
3 from .OpenaiAPI import OpenaiAPI
4 from ..HuggingChat import HuggingChat
5 from ...typing import AsyncResult, Messages
6
7 class HuggingFace2(OpenaiAPI):
8 label = "HuggingFace (Inference API)"
9 url = "https://huggingface.co"
10 working = True
11 default_model = "meta-llama/Llama-3.2-11B-Vision-Instruct"
12 default_vision_model = default_model
13 models = [
14 *HuggingChat.models
15 ]
16
17 @classmethod
18 def create_async_generator(
19 cls,
20 model: str,
21 messages: Messages,
22 api_base: str = "https://api-inference.huggingface.co/v1",
23 max_tokens: int = 500,
24 **kwargs
25 ) -> AsyncResult:
26 return super().create_async_generator(
27 model, messages, api_base=api_base, max_tokens=max_tokens, **kwargs
28 )
Modified g4f/Provider/needs_auth/OpenaiAPI.py +3 -1
@@ -34,6 +34,7 @@ class OpenaiAPI(AsyncGeneratorProvider, ProviderModelMixin):
34 34 stop: Union[str, list[str]] = None,
35 35 stream: bool = False,
36 36 headers: dict = None,
37 impersonate: str = None,
37 38 extra_data: dict = {},
38 39 **kwargs
39 40 ) -> AsyncResult:
@@ -55,7 +56,8 @@ class OpenaiAPI(AsyncGeneratorProvider, ProviderModelMixin):
55 56 async with StreamSession(
56 57 proxies={"all": proxy},
57 58 headers=cls.get_headers(stream, api_key, headers),
58 timeout=timeout
59 timeout=timeout,
60 impersonate=impersonate,
59 61 ) as session:
60 62 data = filter_none(
61 63 messages=messages,
Modified g4f/Provider/needs_auth/__init__.py +2 -0
@@ -1,6 +1,7 @@
1 1 from .gigachat import *
2 2
3 3 from .BingCreateImages import BingCreateImages
4 from .Cerebras import Cerebras
4 5 from .CopilotAccount import CopilotAccount
5 6 from .DeepInfra import DeepInfra
6 7 from .DeepInfraImage import DeepInfraImage
@@ -8,6 +9,7 @@ from .Gemini import Gemini
8 9 from .GeminiPro import GeminiPro
9 10 from .Groq import Groq
10 11 from .HuggingFace import HuggingFace
12 from .HuggingFace2 import HuggingFace2
11 13 from .MetaAI import MetaAI
12 14 from .MetaAIAccount import MetaAIAccount
13 15 from .OpenaiAPI import OpenaiAPI
Modified g4f/gui/client/index.html +6 -2
@@ -128,6 +128,10 @@
128 128 <label for="BingCreateImages-api_key" class="label" title="">Microsoft Designer in Bing:</label>
129 129 <textarea id="BingCreateImages-api_key" name="BingCreateImages[api_key]" placeholder="&quot;_U&quot; cookie"></textarea>
130 130 </div>
131 <div class="field box">
132 <label for="Cerebras-api_key" class="label" title="">Cerebras Inference:</label>
133 <textarea id="Cerebras-api_key" name="Cerebras[api_key]" placeholder="api_key"></textarea>
134 </div>
131 135 <div class="field box">
132 136 <label for="DeepInfra-api_key" class="label" title="">DeepInfra:</label>
133 137 <textarea id="DeepInfra-api_key" name="DeepInfra[api_key]" class="DeepInfraImage-api_key" placeholder="api_key"></textarea>
@@ -142,7 +146,7 @@
142 146 </div>
143 147 <div class="field box">
144 148 <label for="HuggingFace-api_key" class="label" title="">HuggingFace:</label>
145 <textarea id="HuggingFace-api_key" name="HuggingFace[api_key]" placeholder="api_key"></textarea>
149 <textarea id="HuggingFace-api_key" name="HuggingFace[api_key]" class="HuggingFace2-api_key" placeholder="api_key"></textarea>
146 150 </div>
147 151 <div class="field box">
148 152 <label for="Openai-api_key" class="label" title="">OpenAI API:</label>
@@ -192,7 +196,7 @@
192 196 <div class="stop_generating stop_generating-hidden">
193 197 <button id="cancelButton">
194 198 <span>Stop Generating</span>
195 <i class="fa-regular fa-stop"></i>
199 <i class="fa-solid fa-stop"></i>
196 200 </button>
197 201 </div>
198 202 <div class="regenerate">
Modified g4f/gui/client/static/css/style.css +1 -3
@@ -512,9 +512,7 @@ body {
512 512
513 513 @media only screen and (min-width: 40em) {
514 514 .stop_generating {
515 left: 50%;
516 transform: translateX(-50%);
517 right: auto;
515 right: 4px;
518 516 }
519 517 .toolbar .regenerate span {
520 518 display: block;
Modified g4f/gui/client/static/js/chat.v1.js +48 -39
@@ -215,7 +215,6 @@ const register_message_buttons = async () => {
215 215 const message_el = el.parentElement.parentElement.parentElement;
216 216 el.classList.add("clicked");
217 217 setTimeout(() => el.classList.remove("clicked"), 1000);
218 await hide_message(window.conversation_id, message_el.dataset.index);
219 218 await ask_gpt(message_el.dataset.index, get_message_id());
220 219 })
221 220 }
@@ -317,6 +316,7 @@ async function remove_cancel_button() {
317 316
318 317 regenerate.addEventListener("click", async () => {
319 318 regenerate.classList.add("regenerate-hidden");
319 setTimeout(()=>regenerate.classList.remove("regenerate-hidden"), 3000);
320 320 stop_generating.classList.remove("stop_generating-hidden");
321 321 await hide_message(window.conversation_id);
322 322 await ask_gpt(-1, get_message_id());
@@ -383,12 +383,12 @@ const prepare_messages = (messages, message_index = -1) => {
383 383 return new_messages;
384 384 }
385 385
386 async function add_message_chunk(message, message_index) {
387 content_map = content_storage[message_index];
386 async function add_message_chunk(message, message_id) {
387 content_map = content_storage[message_id];
388 388 if (message.type == "conversation") {
389 389 console.info("Conversation used:", message.conversation)
390 390 } else if (message.type == "provider") {
391 provider_storage[message_index] = message.provider;
391 provider_storage[message_id] = message.provider;
392 392 content_map.content.querySelector('.provider').innerHTML = `
393 393 <a href="${message.provider.url}" target="_blank">
394 394 ${message.provider.label ? message.provider.label : message.provider.name}
@@ -398,7 +398,7 @@ async function add_message_chunk(message, message_index) {
398 398 } else if (message.type == "message") {
399 399 console.error(message.message)
400 400 } else if (message.type == "error") {
401 error_storage[message_index] = message.error
401 error_storage[message_id] = message.error
402 402 console.error(message.error);
403 403 content_map.inner.innerHTML += `<p><strong>An error occured:</strong> ${message.error}</p>`;
404 404 let p = document.createElement("p");
@@ -407,8 +407,8 @@ async function add_message_chunk(message, message_index) {
407 407 } else if (message.type == "preview") {
408 408 content_map.inner.innerHTML = markdown_render(message.preview);
409 409 } else if (message.type == "content") {
410 message_storage[message_index] += message.content;
411 html = markdown_render(message_storage[message_index]);
410 message_storage[message_id] += message.content;
411 html = markdown_render(message_storage[message_id]);
412 412 let lastElement, lastIndex = null;
413 413 for (element of ['</p>', '</code></pre>', '</p>\n</li>\n</ol>', '</li>\n</ol>', '</li>\n</ul>']) {
414 414 const index = html.lastIndexOf(element)
@@ -421,7 +421,7 @@ async function add_message_chunk(message, message_index) {
421 421 html = html.substring(0, lastIndex) + '<span class="cursor"></span>' + lastElement;
422 422 }
423 423 content_map.inner.innerHTML = html;
424 content_map.count.innerText = count_words_and_tokens(message_storage[message_index], provider_storage[message_index]?.model);
424 content_map.count.innerText = count_words_and_tokens(message_storage[message_id], provider_storage[message_id]?.model);
425 425 highlight(content_map.inner);
426 426 } else if (message.type == "log") {
427 427 let p = document.createElement("p");
@@ -453,7 +453,7 @@ const ask_gpt = async (message_index = -1, message_id) => {
453 453 let total_messages = messages.length;
454 454 messages = prepare_messages(messages, message_index);
455 455 message_index = total_messages
456 message_storage[message_index] = "";
456 message_storage[message_id] = "";
457 457 stop_generating.classList.remove(".stop_generating-hidden");
458 458
459 459 message_box.scrollTop = message_box.scrollHeight;
@@ -477,10 +477,10 @@ const ask_gpt = async (message_index = -1, message_id) => {
477 477 </div>
478 478 `;
479 479
480 controller_storage[message_index] = new AbortController();
480 controller_storage[message_id] = new AbortController();
481 481
482 482 let content_el = document.getElementById(`gpt_${message_id}`)
483 let content_map = content_storage[message_index] = {
483 let content_map = content_storage[message_id] = {
484 484 content: content_el,
485 485 inner: content_el.querySelector('.content_inner'),
486 486 count: content_el.querySelector('.count'),
@@ -492,12 +492,7 @@ const ask_gpt = async (message_index = -1, message_id) => {
492 492 const file = input && input.files.length > 0 ? input.files[0] : null;
493 493 const provider = providerSelect.options[providerSelect.selectedIndex].value;
494 494 const auto_continue = document.getElementById("auto_continue")?.checked;
495 let api_key = null;
496 if (provider) {
497 api_key = document.getElementById(`${provider}-api_key`)?.value || null;
498 if (api_key == null)
499 api_key = document.querySelector(`.${provider}-api_key`)?.value || null;
500 }
495 let api_key = get_api_key_by_provider(provider);
501 496 await api("conversation", {
502 497 id: message_id,
503 498 conversation_id: window.conversation_id,
@@ -506,10 +501,10 @@ const ask_gpt = async (message_index = -1, message_id) => {
506 501 provider: provider,
507 502 messages: messages,
508 503 auto_continue: auto_continue,
509 api_key: api_key
510 }, file, message_index);
511 if (!error_storage[message_index]) {
512 html = markdown_render(message_storage[message_index]);
504 api_key: api_key,
505 }, file, message_id);
506 if (!error_storage[message_id]) {
507 html = markdown_render(message_storage[message_id]);
513 508 content_map.inner.innerHTML = html;
514 509 highlight(content_map.inner);
515 510
@@ -520,14 +515,14 @@ const ask_gpt = async (message_index = -1, message_id) => {
520 515 } catch (e) {
521 516 console.error(e);
522 517 if (e.name != "AbortError") {
523 error_storage[message_index] = true;
518 error_storage[message_id] = true;
524 519 content_map.inner.innerHTML += `<p><strong>An error occured:</strong> ${e}</p>`;
525 520 }
526 521 }
527 delete controller_storage[message_index];
528 if (!error_storage[message_index] && message_storage[message_index]) {
529 const message_provider = message_index in provider_storage ? provider_storage[message_index] : null;
530 await add_message(window.conversation_id, "assistant", message_storage[message_index], message_provider);
522 delete controller_storage[message_id];
523 if (!error_storage[message_id] && message_storage[message_id]) {
524 const message_provider = message_id in provider_storage ? provider_storage[message_id] : null;
525 await add_message(window.conversation_id, "assistant", message_storage[message_id], message_provider);
531 526 await safe_load_conversation(window.conversation_id);
532 527 } else {
533 528 let cursorDiv = message_box.querySelector(".cursor");
@@ -1156,7 +1151,7 @@ async function on_api() {
1156 1151 evt.preventDefault();
1157 1152 console.log("pressed enter");
1158 1153 prompt_lock = true;
1159 setTimeout(()=>prompt_lock=false, 3);
1154 setTimeout(()=>prompt_lock=false, 3000);
1160 1155 await handle_ask();
1161 1156 } else {
1162 1157 messageInput.style.removeProperty("height");
@@ -1167,7 +1162,7 @@ async function on_api() {
1167 1162 console.log("clicked send");
1168 1163 if (prompt_lock) return;
1169 1164 prompt_lock = true;
1170 setTimeout(()=>prompt_lock=false, 3);
1165 setTimeout(()=>prompt_lock=false, 3000);
1171 1166 await handle_ask();
1172 1167 });
1173 1168 messageInput.focus();
@@ -1189,8 +1184,8 @@ async function on_api() {
1189 1184 providerSelect.appendChild(option);
1190 1185 })
1191 1186
1192 await load_provider_models(appStorage.getItem("provider"));
1193 1187 await load_settings_storage()
1188 await load_provider_models(appStorage.getItem("provider"));
1194 1189
1195 1190 const hide_systemPrompt = document.getElementById("hide-systemPrompt")
1196 1191 const slide_systemPrompt_icon = document.querySelector(".slide-systemPrompt i");
@@ -1316,7 +1311,7 @@ function get_selected_model() {
1316 1311 }
1317 1312 }
1318 1313
1319 async function api(ressource, args=null, file=null, message_index=null) {
1314 async function api(ressource, args=null, file=null, message_id=null) {
1320 1315 if (window?.pywebview) {
1321 1316 if (args !== null) {
1322 1317 if (ressource == "models") {
@@ -1326,15 +1321,19 @@ async function api(ressource, args=null, file=null, message_index=null) {
1326 1321 }
1327 1322 return pywebview.api[`get_${ressource}`]();
1328 1323 }
1324 let api_key;
1329 1325 if (ressource == "models" && args) {
1326 api_key = get_api_key_by_provider(args);
1330 1327 ressource = `${ressource}/${args}`;
1331 1328 }
1332 1329 const url = `/backend-api/v2/${ressource}`;
1330 const headers = {};
1331 if (api_key) {
1332 headers.authorization = `Bearer ${api_key}`;
1333 }
1333 1334 if (ressource == "conversation") {
1334 1335 let body = JSON.stringify(args);
1335 const headers = {
1336 accept: 'text/event-stream'
1337 }
1336 headers.accept = 'text/event-stream';
1338 1337 if (file !== null) {
1339 1338 const formData = new FormData();
1340 1339 formData.append('file', file);
@@ -1345,17 +1344,17 @@ async function api(ressource, args=null, file=null, message_index=null) {
1345 1344 }
1346 1345 response = await fetch(url, {
1347 1346 method: 'POST',
1348 signal: controller_storage[message_index].signal,
1347 signal: controller_storage[message_id].signal,
1349 1348 headers: headers,
1350 body: body
1349 body: body,
1351 1350 });
1352 return read_response(response, message_index);
1351 return read_response(response, message_id);
1353 1352 }
1354 response = await fetch(url);
1353 response = await fetch(url, {headers: headers});
1355 1354 return await response.json();
1356 1355 }
1357 1356
1358 async function read_response(response, message_index) {
1357 async function read_response(response, message_id) {
1359 1358 const reader = response.body.pipeThrough(new TextDecoderStream()).getReader();
1360 1359 let buffer = ""
1361 1360 while (true) {
@@ -1368,7 +1367,7 @@ async function read_response(response, message_index) {
1368 1367 continue;
1369 1368 }
1370 1369 try {
1371 add_message_chunk(JSON.parse(buffer + line), message_index);
1370 add_message_chunk(JSON.parse(buffer + line), message_id);
1372 1371 buffer = "";
1373 1372 } catch {
1374 1373 buffer += line
@@ -1377,6 +1376,16 @@ async function read_response(response, message_index) {
1377 1376 }
1378 1377 }
1379 1378
1379 function get_api_key_by_provider(provider) {
1380 let api_key = null;
1381 if (provider) {
1382 api_key = document.getElementById(`${provider}-api_key`)?.value || null;
1383 if (api_key == null)
1384 api_key = document.querySelector(`.${provider}-api_key`)?.value || null;
1385 }
1386 return api_key;
1387 }
1388
1380 1389 async function load_provider_models(providerIndex=null) {
1381 1390 if (!providerIndex) {
1382 1391 providerIndex = providerSelect.selectedIndex;
Modified g4f/gui/server/api.py +3 -2
@@ -38,10 +38,11 @@ class Api:
38 38 return models._all_models
39 39
40 40 @staticmethod
41 def get_provider_models(provider: str) -> list[dict]:
41 def get_provider_models(provider: str, api_key: str = None) -> list[dict]:
42 42 if provider in __map__:
43 43 provider: ProviderType = __map__[provider]
44 44 if issubclass(provider, ProviderModelMixin):
45 models = provider.get_models() if api_key is None else provider.get_models(api_key=api_key)
45 46 return [
46 47 {
47 48 "model": model,
@@ -49,7 +50,7 @@ class Api:
49 50 "vision": getattr(provider, "default_vision_model", None) == model or model in getattr(provider, "vision_models", []),
50 51 "image": model in getattr(provider, "image_models", []),
51 52 }
52 for model in provider.get_models()
53 for model in models
53 54 ]
54 55 return []
55 56
Modified g4f/gui/server/backend.py +2 -1
Modified g4f/requests/raise_for_status.py +3 -1