返回提交历史
Added
etc/unittest/test_api_quota.py
+31
-0
Modified
g4f/Provider/github/GithubCopilot.py
+1
-1
Modified
g4f/Provider/needs_auth/Antigravity.py
+1
-1
Modified
g4f/Provider/needs_auth/GeminiCLI.py
+1
-1
Modified
g4f/Provider/template/OpenaiTemplate.py
+5
-2
Modified
g4f/api/__init__.py
+28
-7
Modified
g4f/gui/server/api.py
+12
-10
Modified
g4f/gui/server/backend_api.py
+3
-3
XFEstudio/gpt4free
feat: implement quota retrieval for providers and update related methods
a9395404
代码差异
8 个文件
+82
-25
@@ -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)
@@ -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.
@@ -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.
@@ -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)
@@ -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
@@ -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},
@@ -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)]
@@ -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