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

XFEstudio/gpt4free

Fix api with default providers, add unittests for RetryProvider

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

代码差异

7 个文件 +161 -91
Modified etc/unittest/__main__.py +2 -1
@@ -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()
Modified etc/unittest/mocks.py +30 -2
@@ -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
Added etc/unittest/retry_provider.py +60 -0
@@ -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")
Modified g4f/client/__init__.py +1 -1
@@ -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)
Modified g4f/client/service.py +2 -2
@@ -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:
Modified g4f/providers/base_provider.py +4 -2
@@ -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
Modified g4f/providers/retry_provider.py +62 -83
@@ -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 """