返回提交历史
Modified
g4f/Provider/Blackbox.py
+10
-2
Modified
g4f/Provider/Copilot.py
+16
-7
Modified
g4f/Provider/RubiksAI.py
+47
-77
Added
g4f/Provider/needs_auth/Cerebras.py
+65
-0
Modified
g4f/Provider/needs_auth/CopilotAccount.py
+5
-2
Added
g4f/Provider/needs_auth/HuggingFace2.py
+28
-0
Modified
g4f/Provider/needs_auth/OpenaiAPI.py
+3
-1
Modified
g4f/Provider/needs_auth/__init__.py
+2
-0
Modified
g4f/gui/client/index.html
+6
-2
Modified
g4f/gui/client/static/css/style.css
+1
-3
Modified
g4f/gui/client/static/js/chat.v1.js
+48
-39
Modified
g4f/gui/server/api.py
+3
-2
Modified
g4f/gui/server/backend.py
+2
-1
Modified
g4f/requests/raise_for_status.py
+3
-1
XFEstudio/gpt4free
Add Cerebras and HuggingFace2 provider, Fix RubiksAI provider Add support for image generation in Copilot provider
58fa409e
代码差异
14 个文件
+239
-137
@@ -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
@@ -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
@@ -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}"
@@ -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
@@ -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
@@ -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
)
@@ -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,
@@ -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
@@ -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=""_U" 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">
@@ -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;
@@ -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;
@@ -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