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

XFEstudio/gpt4free

feat: implement quota retrieval for providers and update related methods

a9395404
hlohaus <hlohaus@users.noreply.github.com>
提交于

代码差异

8 个文件 +82 -25
Added etc/unittest/test_api_quota.py +31 -0
@@ -0,0 +1,31 @@
1 from __future__ import annotations
2
3 import unittest
4 from fastapi.testclient import TestClient
5
6 import g4f.api
7 from g4f import Provider
8
9 class TestApiQuota(unittest.TestCase):
10 def setUp(self):
11 # create fresh FastAPI app instance for each test
12 self.app = g4f.api.create_app()
13 self.client = TestClient(self.app)
14
15 def test_nonexistent_provider_returns_404(self):
16 resp = self.client.get("/api/NoSuchProvider/quota")
17 self.assertEqual(resp.status_code, 404)
18
19 def test_dummy_provider_quota_route(self):
20 # monkeypatch a fake provider with async get_quota method
21 class DummyProvider:
22 async def get_quota(self, api_key=None):
23 return {"foo": "bar"}
24
25 Provider.__map__["dummy"] = DummyProvider()
26 try:
27 resp = self.client.get("/api/dummy/quota")
28 self.assertEqual(resp.status_code, 200)
29 self.assertEqual(resp.json(), {"foo": "bar"})
30 finally:
31 Provider.__map__.pop("dummy", None)
Modified g4f/Provider/github/GithubCopilot.py +1 -1
@@ -237,7 +237,7 @@ class GithubCopilot(OpenaiTemplate):
237 237 return None
238 238
239 239 @classmethod
240 async def get_usage(cls) -> dict:
240 async def get_quota(cls) -> dict:
241 241 """
242 242 Fetch and summarize current GitHub Copilot usage/quota information.
243 243 Returns a dictionary with usage details or raises an exception on failure.
Modified g4f/Provider/needs_auth/Antigravity.py +1 -1
@@ -1295,7 +1295,7 @@ class Antigravity(AsyncGeneratorProvider, ProviderModelMixin):
1295 1295 return []
1296 1296
1297 1297 @classmethod
1298 async def get_usage(cls) -> dict:
1298 async def get_quota(cls) -> dict:
1299 1299 """
1300 1300 Fetch and summarize quota usage for Antigravity account.
1301 1301 Returns a dict with OpenAI Usage keys if possible, or quota info.
Modified g4f/Provider/needs_auth/GeminiCLI.py +1 -1
@@ -855,7 +855,7 @@ class GeminiCLI(AsyncGeneratorProvider, ProviderModelMixin):
855 855 return cls.models if cls.models else cls.fallback_models
856 856
857 857 @classmethod
858 async def get_usage(cls) -> dict:
858 async def get_quota(cls) -> dict:
859 859 if cls.auth_manager is None:
860 860 cls.auth_manager = AuthManager(env=os.environ)
861 861 provider = GeminiCLIProvider(env=os.environ, auth_manager=cls.auth_manager)
Modified g4f/Provider/template/OpenaiTemplate.py +5 -2
@@ -59,9 +59,12 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
59 59 cls.image_models = [model.get("name") if cls.use_model_names else model.get("id", model.get("name")) for model in data if model.get("image") or model.get("type") == "image" or model.get("supports_images")]
60 60 cls.vision_models = cls.vision_models.copy()
61 61 cls.vision_models += [model.get("name") if cls.use_model_names else model.get("id", model.get("name")) for model in data if model.get("vision")]
62 cls.models = [model.get("name") if cls.use_model_names else model.get("id", model.get("name")) for model in data]
62 cls.models = {model.get("name") if cls.use_model_names else model.get("id", model.get("name")): model for model in data}
63 for key, value in cls.models.items():
64 value.pop("id")
65 cls.models[key] = {"id": key, **value}
63 66 cls.models_count = {model.get("name") if cls.use_model_names else model.get("id", model.get("name")): len(model.get("providers", [])) for model in data if len(model.get("providers", [])) > 1}
64 if cls.sort_models:
67 if cls.sort_models and isinstance(cls.models, list):
65 68 cls.models.sort()
66 69 except MissingAuthError:
67 70 raise
Modified g4f/api/__init__.py +28 -7
@@ -376,18 +376,39 @@ class Api:
376 376 return {
377 377 "object": "list",
378 378 "data": [{
379 "id": model,
379 "id": model.get("id") if isinstance(model, dict) else model,
380 380 "object": "model",
381 381 "created": 0,
382 382 "owned_by": getattr(provider, "label", provider.__name__),
383 "image": model in getattr(provider, "image_models", []),
384 "vision": model in getattr(provider, "vision_models", []),
385 "audio": model in getattr(provider, "audio_models", []),
386 "video": model in getattr(provider, "video_models", []),
387 "type": "image" if model in getattr(provider, "image_models", []) else "chat",
388 } for model in models]
383 "image": (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "image_models", []),
384 "vision": (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "vision_models", []),
385 "audio": (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "audio_models", []),
386 "video": (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "video_models", []),
387 "type": "image" if (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "image_models", []) else "chat",
388 **(model if isinstance(model, dict) else {})
389 } for model in (models.values() if isinstance(models, dict) else models)]
389 390 }
390 391
392 # quota endpoint mimics backend-api/v2/quota but exposed on public API
393 @self.app.get("/api/{provider}/quota")
394 async def provider_quota(provider: str, credentials: Annotated[HTTPAuthorizationCredentials, Depends(Api.security)] = None):
395 # provider must exist
396 if provider not in Provider.__map__:
397 return ErrorResponse.from_message("The provider does not exist.", 404)
398 provider_obj: ProviderType = Provider.__map__[provider]
399 if not hasattr(provider_obj, "get_quota"):
400 return ErrorResponse.from_message("Provider doesn't support get_quota", HTTP_500_INTERNAL_SERVER_ERROR)
401 try:
402 if credentials is not None and credentials.credentials != "secret":
403 usage = await provider_obj.get_quota(api_key=credentials.credentials)
404 else:
405 usage = await provider_obj.get_quota()
406 return usage
407 except MissingAuthError as e:
408 return ErrorResponse.from_message(f"{type(e).__name__}: {e}", HTTP_401_UNAUTHORIZED)
409 except Exception as e:
410 return ErrorResponse.from_message(f"{type(e).__name__}: {e}", HTTP_500_INTERNAL_SERVER_ERROR)
411
391 412 @self.app.get("/v1/models/{model_name}", responses={
392 413 HTTP_200_OK: {"model": ModelResponseModel},
393 414 HTTP_404_NOT_FOUND: {"model": ErrorResponseModel},
Modified g4f/gui/server/api.py +12 -10
@@ -50,16 +50,18 @@ class Api:
50 50 @staticmethod
51 51 def get_provider_models(provider: str, api_key: str = None, base_url: str = None, ignored: list = None):
52 52 def get_model_data(provider: ProviderModelMixin, model: str, default: bool = False) -> dict:
53 model_id = model.get("id") if isinstance(model, dict) else model
53 54 return {
54 "model": model,
55 "label": model.split(":")[-1] if provider.__name__ == "AnyProvider" and not model.startswith("openrouter:") else model,
56 "default": default or model == provider.default_model,
57 "vision": model in provider.vision_models,
58 "audio": False if provider.audio_models is None else model in provider.audio_models,
59 "video": model in provider.video_models,
60 "image": model in provider.image_models,
61 "count": False if provider.models_count is None else provider.models_count.get(model),
62 "tags": [] if provider.models_tags is None else provider.models_tags.get(model, []),
55 "model": model_id,
56 "label": model_id.split(":")[-1] if provider.__name__ == "AnyProvider" and not model_id.startswith("openrouter:") else model_id,
57 "default": default or model_id == provider.default_model,
58 "vision": model_id in provider.vision_models,
59 "audio": False if provider.audio_models is None else model_id in provider.audio_models,
60 "video": model_id in provider.video_models,
61 "image": model_id in provider.image_models,
62 "count": False if provider.models_count is None else provider.models_count.get(model_id),
63 "tags": [] if provider.models_tags is None else provider.models_tags.get(model_id, []),
64 **(model if isinstance(model, dict) else {})
63 65 }
64 66 if provider in Provider.__map__:
65 67 provider = Provider.__map__[provider]
@@ -79,7 +81,7 @@ class Api:
79 81 } for model in models]
80 82 return [
81 83 get_model_data(provider, model)
82 for model in models
84 for model in (models.values() if isinstance(models, dict) else models)
83 85 ]
84 86 elif provider in model_map:
85 87 return [get_model_data(AnyProvider, provider, True)]
Modified g4f/gui/server/backend_api.py +3 -3
@@ -259,10 +259,10 @@ class Backend_Api(Api):
259 259 provider_handler = convert_to_provider(provider)
260 260 except ProviderNotFoundError:
261 261 return "Provider not found", 404
262 if not hasattr(provider_handler, "get_usage"):
263 return "Provider doesn't support get_usage", 500
262 if not hasattr(provider_handler, "get_quota"):
263 return "Provider doesn't support get_quota", 500
264 264 try:
265 response_data = await provider_handler.get_usage()
265 response_data = await provider_handler.get_quota()
266 266 return jsonify(response_data)
267 267 except MissingAuthError as e:
268 268 return jsonify({"error": {"message": f"{type(e).__name__}: {e}"}}), 401