返回提交历史
Modified
g4f/Provider/DeepInfra.py
+1
-1
Modified
g4f/tools/run_tools.py
+4
-52
XFEstudio/gpt4free
Show provider label
1dcb85ba
代码差异
2 个文件
+5
-53
@@ -132,7 +132,7 @@ def _get_turnstile_token_sync(model: str) -> str:
132
132
133
133
async def get_turnstile_token_async(model: str = None) -> str:
134
134
"""Run the synchronous Turnstile solver in a thread pool executor."""
135
if model is None:
135
if not model:
136
136
model = DeepInfra.default_model
137
137
loop = asyncio.get_running_loop()
138
138
return await loop.run_in_executor(None, _get_turnstile_token_sync, model)
@@ -116,8 +116,6 @@ Instruction: Make sure to add the sources of cites using [[domain]](Url) notatio
116
116
117
117
TOOL_NAMES = {
118
118
"SEARCH": "search_tool",
119
"CONTINUE": "continue_tool",
120
"BUCKET": "bucket_tool",
121
119
}
122
120
123
121
def is_provider_api_key(api_key: str) -> bool:
@@ -162,47 +160,6 @@ class ToolHandler:
162
160
)
163
161
return messages, sources
164
162
165
@staticmethod
166
def process_continue_tool(
167
messages: Messages, tool: dict, provider: Any
168
) -> Tuple[Messages, Dict[str, Any]]:
169
"""Process continue tool requests"""
170
kwargs = {}
171
if provider not in ("OpenaiAccount", "HuggingFaceAPI"):
172
messages = messages.copy()
173
last_line = messages[-1]["content"].strip().splitlines()[-1]
174
content = f"Carry on from this point:\n{last_line}"
175
messages.append({"role": "user", "content": content})
176
else:
177
# Enable provider native continue
178
kwargs["action"] = "continue"
179
return messages, kwargs
180
181
@staticmethod
182
def process_bucket_tool(messages: Messages, tool: dict) -> Messages:
183
"""Process bucket tool requests"""
184
messages = messages.copy()
185
186
def on_bucket(match):
187
return "".join(read_bucket(get_bucket_dir(match.group(1))))
188
189
has_bucket = False
190
for message in messages:
191
if "content" in message and isinstance(message["content"], str):
192
new_message_content = re.sub(
193
r'{"bucket_id":\s*"([^"]*)"}', on_bucket, message["content"]
194
)
195
if new_message_content != message["content"]:
196
has_bucket = True
197
message["content"] = new_message_content
198
199
last_message_content = messages[-1]["content"]
200
if has_bucket and isinstance(last_message_content, str):
201
if "\nSource: " in last_message_content:
202
messages[-1]["content"] = last_message_content + BUCKET_INSTRUCTIONS
203
204
return messages
205
206
163
@staticmethod
207
164
async def process_tools(
208
165
messages: Messages, tool_calls: List[dict], provider: Any
@@ -227,15 +184,6 @@ class ToolHandler:
227
184
messages, tool
228
185
)
229
186
230
elif function_name == TOOL_NAMES["CONTINUE"]:
231
messages, kwargs = ToolHandler.process_continue_tool(
232
messages, tool, provider
233
)
234
extra_kwargs.update(kwargs)
235
236
elif function_name == TOOL_NAMES["BUCKET"]:
237
messages = ToolHandler.process_bucket_tool(messages, tool)
238
239
187
return messages, sources, extra_kwargs
240
188
241
189
@@ -424,6 +372,7 @@ async def async_iter_run_tools(
424
372
try:
425
373
usage_model = model or getattr(provider, "default_model", model)
426
374
usage_provider = provider.__name__
375
usage_label = getattr(provider, "label", usage_provider)
427
376
completion_tokens = 0
428
377
usage = None
429
378
async for chunk in response:
@@ -457,6 +406,7 @@ async def async_iter_run_tools(
457
406
"user": kwargs.get("user"),
458
407
"model": usage_model,
459
408
"provider": usage_provider,
409
"label": usage_label,
460
410
**usage.get_dict(),
461
411
}
462
412
if saved_tokens:
@@ -630,6 +580,7 @@ def iter_run_tools(
630
580
processor = ThinkingProcessor()
631
581
usage_model = model or getattr(provider, "default_model", model)
632
582
usage_provider = provider.__name__
583
usage_label = getattr(provider, "label", usage_provider)
633
584
completion_tokens = 0
634
585
usage = None
635
586
method = get_provider_method(provider)
@@ -674,6 +625,7 @@ def iter_run_tools(
674
625
"user": kwargs.get("user"),
675
626
"model": usage_model,
676
627
"provider": usage_provider,
628
"label": usage_label,
677
629
**usage.get_dict(),
678
630
}
679
631
if saved_tokens: