返回提交历史
Modified
etc/unittest/__main__.py
+2
-1
Modified
etc/unittest/mocks.py
+30
-2
Added
etc/unittest/retry_provider.py
+60
-0
Modified
g4f/client/__init__.py
+1
-1
Modified
g4f/client/service.py
+2
-2
Modified
g4f/providers/base_provider.py
+4
-2
Modified
g4f/providers/retry_provider.py
+62
-83
XFEstudio/gpt4free
Fix api with default providers, add unittests for RetryProvider
c31f5435
代码差异
7 个文件
+161
-91
@@ -6,5 +6,6 @@ from .main import *
6
6
from .model import *
7
7
from .client import *
8
8
from .include import *
9
from .retry_provider import *
9
10
10
unittest.main()
11
unittest.main()
@@ -34,9 +34,37 @@ class ModelProviderMock(AbstractProvider):
34
34
35
35
class YieldProviderMock(AsyncGeneratorProvider):
36
36
working = True
37
37
38
38
async def create_async_generator(
39
39
model, messages, stream, **kwargs
40
40
):
41
41
for message in messages:
42
yield message["content"]
42
yield message["content"]
43
44
class RaiseExceptionProviderMock(AbstractProvider):
45
working = True
46
47
@classmethod
48
def create_completion(
49
cls, model, messages, stream, **kwargs
50
):
51
raise RuntimeError(cls.__name__)
52
yield cls.__name__
53
54
class AsyncRaiseExceptionProviderMock(AsyncGeneratorProvider):
55
working = True
56
57
@classmethod
58
async def create_async_generator(
59
cls, model, messages, stream, **kwargs
60
):
61
raise RuntimeError(cls.__name__)
62
yield cls.__name__
63
64
class YieldNoneProviderMock(AsyncGeneratorProvider):
65
working = True
66
67
async def create_async_generator(
68
model, messages, stream, **kwargs
69
):
70
yield None
@@ -0,0 +1,60 @@
1
from __future__ import annotations
2
3
import unittest
4
5
from g4f.client import AsyncClient, ChatCompletion, ChatCompletionChunk
6
from g4f.providers.retry_provider import IterListProvider
7
from .mocks import YieldProviderMock, RaiseExceptionProviderMock, AsyncRaiseExceptionProviderMock, YieldNoneProviderMock
8
9
DEFAULT_MESSAGES = [{'role': 'user', 'content': 'Hello'}]
10
11
class TestIterListProvider(unittest.IsolatedAsyncioTestCase):
12
13
async def test_skip_provider(self):
14
client = AsyncClient(provider=IterListProvider([RaiseExceptionProviderMock, YieldProviderMock], False))
15
response = await client.chat.completions.create(DEFAULT_MESSAGES, "")
16
self.assertIsInstance(response, ChatCompletion)
17
self.assertEqual("Hello", response.choices[0].message.content)
18
19
async def test_only_one_result(self):
20
client = AsyncClient(provider=IterListProvider([YieldProviderMock, YieldProviderMock]))
21
response = await client.chat.completions.create(DEFAULT_MESSAGES, "")
22
self.assertIsInstance(response, ChatCompletion)
23
self.assertEqual("Hello", response.choices[0].message.content)
24
25
async def test_stream_skip_provider(self):
26
client = AsyncClient(provider=IterListProvider([AsyncRaiseExceptionProviderMock, YieldProviderMock], False))
27
messages = [{'role': 'user', 'content': chunk} for chunk in ["How ", "are ", "you", "?"]]
28
response = client.chat.completions.create(messages, "Hello", stream=True)
29
async for chunk in response:
30
chunk: ChatCompletionChunk = chunk
31
self.assertIsInstance(chunk, ChatCompletionChunk)
32
if chunk.choices[0].delta.content is not None:
33
self.assertIsInstance(chunk.choices[0].delta.content, str)
34
35
async def test_stream_only_one_result(self):
36
client = AsyncClient(provider=IterListProvider([YieldProviderMock, YieldProviderMock], False))
37
messages = [{'role': 'user', 'content': chunk} for chunk in ["You ", "You "]]
38
response = client.chat.completions.create(messages, "Hello", stream=True, max_tokens=2)
39
response_list = []
40
async for chunk in response:
41
response_list.append(chunk)
42
self.assertEqual(len(response_list), 3)
43
for chunk in response_list:
44
if chunk.choices[0].delta.content is not None:
45
self.assertEqual(chunk.choices[0].delta.content, "You ")
46
47
async def test_skip_none(self):
48
client = AsyncClient(provider=IterListProvider([YieldNoneProviderMock, YieldProviderMock], False))
49
response = await client.chat.completions.create(DEFAULT_MESSAGES, "")
50
self.assertIsInstance(response, ChatCompletion)
51
self.assertEqual("Hello", response.choices[0].message.content)
52
53
async def test_stream_skip_none(self):
54
client = AsyncClient(provider=IterListProvider([YieldNoneProviderMock, YieldProviderMock], False))
55
response = client.chat.completions.create(DEFAULT_MESSAGES, "", stream=True)
56
response_list = [chunk async for chunk in response]
57
self.assertEqual(len(response_list), 2)
58
for chunk in response_list:
59
if chunk.choices[0].delta.content is not None:
60
self.assertEqual(chunk.choices[0].delta.content, "Hello")
@@ -474,7 +474,7 @@ class AsyncCompletions:
474
474
**kwargs
475
475
)
476
476
477
if not isinstance(response, AsyncIterator):
477
if not hasattr(response, "__aiter__"):
478
478
response = to_async_iterator(response)
479
479
response = async_iter_response(response, stream, response_format, max_tokens, stop)
480
480
response = async_iter_append_model_and_provider(response)
@@ -7,14 +7,14 @@ from ..errors import ProviderNotFoundError, ModelNotFoundError, ProviderNotWorki
7
7
from ..models import Model, ModelUtils, default
8
8
from ..Provider import ProviderUtils
9
9
from ..providers.types import BaseRetryProvider, ProviderType
10
from ..providers.retry_provider import IterProvider
10
from ..providers.retry_provider import IterListProvider
11
11
12
12
def convert_to_provider(provider: str) -> ProviderType:
13
13
if " " in provider:
14
14
provider_list = [ProviderUtils.convert[p] for p in provider.split() if p in ProviderUtils.convert]
15
15
if not provider_list:
16
16
raise ProviderNotFoundError(f'Providers not found: {provider}')
17
provider = IterProvider(provider_list)
17
provider = IterListProvider(provider_list, False)
18
18
elif provider in ProviderUtils.convert:
19
19
provider = ProviderUtils.convert[provider]
20
20
elif provider:
@@ -57,7 +57,9 @@ class AbstractProvider(BaseProvider):
57
57
loop = loop or asyncio.get_running_loop()
58
58
59
59
def create_func() -> str:
60
return "".join(cls.create_completion(model, messages, False, **kwargs))
60
chunks = [str(chunk) for chunk in cls.create_completion(model, messages, False, **kwargs) if chunk]
61
if chunks:
62
return "".join(chunks)
61
63
62
64
return await asyncio.wait_for(
63
65
loop.run_in_executor(executor, create_func),
@@ -205,7 +207,7 @@ class AsyncGeneratorProvider(AsyncProvider):
205
207
"""
206
208
return "".join([
207
209
str(chunk) async for chunk in cls.create_async_generator(model, messages, stream=False, **kwargs)
208
if not isinstance(chunk, (Exception, FinishReason, BaseConversation, SynthesizeData))
210
if chunk and not isinstance(chunk, (Exception, FinishReason, BaseConversation, SynthesizeData))
209
211
])
210
212
211
213
@staticmethod
@@ -8,6 +8,8 @@ from .types import BaseProvider, BaseRetryProvider, ProviderType
8
8
from .. import debug
9
9
from ..errors import RetryProviderError, RetryNoProviderError
10
10
11
DEFAULT_TIMEOUT = 60
12
11
13
class IterListProvider(BaseRetryProvider):
12
14
def __init__(
13
15
self,
@@ -50,12 +52,12 @@ class IterListProvider(BaseRetryProvider):
50
52
51
53
for provider in self.get_providers(stream):
52
54
self.last_provider = provider
55
debug.log(f"Using {provider.__name__} provider")
53
56
try:
54
if debug.logging:
55
print(f"Using {provider.__name__} provider")
56
for token in provider.create_completion(model, messages, stream, **kwargs):
57
yield token
58
started = True
57
for chunk in provider.create_completion(model, messages, stream, **kwargs):
58
if chunk:
59
yield chunk
60
started = True
59
61
if started:
60
62
return
61
63
except Exception as e:
@@ -87,13 +89,14 @@ class IterListProvider(BaseRetryProvider):
87
89
88
90
for provider in self.get_providers(False):
89
91
self.last_provider = provider
92
debug.log(f"Using {provider.__name__} provider")
90
93
try:
91
if debug.logging:
92
print(f"Using {provider.__name__} provider")
93
return await asyncio.wait_for(
94
chunk = await asyncio.wait_for(
94
95
provider.create_async(model, messages, **kwargs),
95
timeout=kwargs.get("timeout", 60),
96
timeout=kwargs.get("timeout", DEFAULT_TIMEOUT),
96
97
)
98
if chunk:
99
return chunk
97
100
except Exception as e:
98
101
exceptions[provider.__name__] = e
99
102
if debug.logging:
@@ -119,16 +122,21 @@ class IterListProvider(BaseRetryProvider):
119
122
120
123
for provider in self.get_providers(stream):
121
124
self.last_provider = provider
125
debug.log(f"Using {provider.__name__} provider")
122
126
try:
123
if debug.logging:
124
print(f"Using {provider.__name__} provider")
125
127
if not stream:
126
yield await provider.create_async(model, messages, **kwargs)
127
started = True
128
elif hasattr(provider, "create_async_generator"):
129
async for token in provider.create_async_generator(model, messages, stream=stream, **kwargs):
130
yield token
128
chunk = await asyncio.wait_for(
129
provider.create_async(model, messages, **kwargs),
130
timeout=kwargs.get("timeout", DEFAULT_TIMEOUT),
131
)
132
if chunk:
133
yield chunk
131
134
started = True
135
elif hasattr(provider, "create_async_generator"):
136
async for chunk in provider.create_async_generator(model, messages, stream=stream, **kwargs):
137
if chunk:
138
yield chunk
139
started = True
132
140
else:
133
141
for token in provider.create_completion(model, messages, stream, **kwargs):
134
142
yield token
@@ -137,8 +145,7 @@ class IterListProvider(BaseRetryProvider):
137
145
return
138
146
except Exception as e:
139
147
exceptions[provider.__name__] = e
140
if debug.logging:
141
print(f"{provider.__name__}: {e.__class__.__name__}: {e}")
148
debug.log(f"{provider.__name__}: {e.__class__.__name__}: {e}")
142
149
if started:
143
150
raise e
144
151
@@ -243,76 +250,48 @@ class RetryProvider(IterListProvider):
243
250
else:
244
251
return await super().create_async(model, messages, **kwargs)
245
252
246
class IterProvider(BaseRetryProvider):
247
__name__ = "IterProvider"
248
249
def __init__(
250
self,
251
providers: List[BaseProvider],
252
) -> None:
253
providers.reverse()
254
self.providers: List[BaseProvider] = providers
255
self.working: bool = True
256
self.last_provider: BaseProvider = None
257
258
def create_completion(
259
self,
260
model: str,
261
messages: Messages,
262
stream: bool = False,
263
**kwargs
264
) -> CreateResult:
265
exceptions: dict = {}
266
started: bool = False
267
for provider in self.iter_providers():
268
if stream and not provider.supports_stream:
269
continue
270
try:
271
for token in provider.create_completion(model, messages, stream, **kwargs):
272
yield token
273
started = True
274
if started:
275
return
276
except Exception as e:
277
exceptions[provider.__name__] = e
278
if debug.logging:
279
print(f"{provider.__name__}: {e.__class__.__name__}: {e}")
280
if started:
281
raise e
282
raise_exceptions(exceptions)
283
284
async def create_async(
253
async def create_async_generator(
285
254
self,
286
255
model: str,
287
256
messages: Messages,
257
stream: bool = True,
288
258
**kwargs
289
) -> str:
290
exceptions: dict = {}
291
for provider in self.iter_providers():
292
try:
293
return await asyncio.wait_for(
294
provider.create_async(model, messages, **kwargs),
295
timeout=kwargs.get("timeout", 60)
296
)
297
except Exception as e:
298
exceptions[provider.__name__] = e
299
if debug.logging:
300
print(f"{provider.__name__}: {e.__class__.__name__}: {e}")
301
raise_exceptions(exceptions)
259
) -> AsyncResult:
260
exceptions = {}
261
started = False
302
262
303
def iter_providers(self) -> Iterator[BaseProvider]:
304
used_provider = []
305
try:
306
while self.providers:
307
provider = self.providers.pop()
308
used_provider.append(provider)
309
self.last_provider = provider
310
if debug.logging:
311
print(f"Using {provider.__name__} provider")
312
yield provider
313
finally:
314
used_provider.reverse()
315
self.providers = [*used_provider, *self.providers]
263
if self.single_provider_retry:
264
provider = self.providers[0]
265
self.last_provider = provider
266
for attempt in range(self.max_retries):
267
try:
268
debug.log(f"Using {provider.__name__} provider (attempt {attempt + 1})")
269
if not stream:
270
chunk = await asyncio.wait_for(
271
provider.create_async(model, messages, **kwargs),
272
timeout=kwargs.get("timeout", DEFAULT_TIMEOUT),
273
)
274
if chunk:
275
started = True
276
elif hasattr(provider, "create_async_generator"):
277
async for chunk in provider.create_async_generator(model, messages, stream=stream, **kwargs):
278
if chunk:
279
yield chunk
280
started = True
281
else:
282
for token in provider.create_completion(model, messages, stream, **kwargs):
283
yield token
284
started = True
285
if started:
286
return
287
except Exception as e:
288
exceptions[provider.__name__] = e
289
if debug.logging:
290
print(f"{provider.__name__}: {e.__class__.__name__}: {e}")
291
raise_exceptions(exceptions)
292
else:
293
async for chunk in super().create_async_generator(model, messages, stream, **kwargs):
294
yield chunk
316
295
317
296
def raise_exceptions(exceptions: dict) -> None:
318
297
"""