返回提交历史
Modified
g4f/api/__init__.py
+16
-25
Modified
g4f/client/__init__.py
+1
-1
Modified
g4f/client/factory.py
+4
-3
Modified
g4f/mcp/pa_provider.py
+5
-3
XFEstudio/gpt4free
Refactor provider handling in Api class and update ClientFactory usage
3645dc39
代码差异
4 个文件
+26
-32
@@ -74,7 +74,7 @@ from g4f.providers.response import AudioResponse
74
74
from g4f.providers.any_provider import AnyProvider
75
75
from g4f.providers.any_model_map import model_map, vision_models, image_models, audio_models, video_models
76
76
from g4f.config import AppConfig
77
from g4f.client import ClientFactory
77
from g4f.client.factory import AbstractClientFactory
78
78
from g4f import Provider
79
79
from g4f.Provider import ProviderUtils
80
80
@@ -432,15 +432,6 @@ def update_headers(request: Request, new_api_key: str = None, user: str = None)
432
432
delattr(request, "_headers")
433
433
return request
434
434
435
def get_provider_by_label(provider: str) -> ProviderType:
436
try:
437
return ProviderUtils.get_by_label(provider)
438
except ValueError as e:
439
try:
440
return ClientFactory.create_provider(None, provider)
441
except ProviderNotFoundError:
442
raise e
443
444
435
class Api:
445
436
def __init__(self, app: FastAPI) -> None:
446
437
self.app = app
@@ -620,8 +611,8 @@ class Api:
620
611
})
621
612
async def models(provider: str, credentials: Annotated[HTTPAuthorizationCredentials, Depends(Api.security)] = None):
622
613
try:
623
provider = get_provider_by_label(provider)
624
except ValueError as e:
614
provider = AbstractClientFactory.create_provider(None, provider)
615
except ProviderNotFoundError as e:
625
616
return ErrorResponse.from_message(str(e), 404)
626
617
if not hasattr(provider, "get_models"):
627
618
models = getattr(provider, "models", [])
@@ -641,7 +632,7 @@ class Api:
641
632
"audio": (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "audio_models", []),
642
633
"video": (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "video_models", []),
643
634
"type": "image" if (model.get("id") if isinstance(model, dict) else model) in getattr(provider, "image_models", []) else "chat",
644
"count": provider.models_count.get(model.get("id"), 0) if isinstance(model, dict) else 0,
635
"count": getattr(provider, "models_count", {}).get(model.get("id") if isinstance(model, dict) else model, 0),
645
636
**(model if isinstance(model, dict) else {})
646
637
} for model in (models.values() if isinstance(models, dict) else models)]
647
638
}
@@ -650,8 +641,8 @@ class Api:
650
641
@self.app.get("/api/{provider}/quota")
651
642
async def provider_quota(provider: str, credentials: Annotated[HTTPAuthorizationCredentials, Depends(Api.security)] = None):
652
643
try:
653
provider = get_provider_by_label(provider)
654
except ValueError as e:
644
provider = AbstractClientFactory.create_provider(None, provider)
645
except ProviderNotFoundError as e:
655
646
return ErrorResponse.from_message(str(e), 404)
656
647
if not hasattr(provider, "get_quota"):
657
648
return ErrorResponse.from_message("Provider doesn't support get_quota", HTTP_500_INTERNAL_SERVER_ERROR)
@@ -708,8 +699,8 @@ class Api:
708
699
if config.provider is None:
709
700
config.provider = AppConfig.provider
710
701
try:
711
provider = get_provider_by_label(config.provider)
712
except ValueError as e:
702
provider = AbstractClientFactory.create_provider(None, config.provider)
703
except ProviderNotFoundError as e:
713
704
return ErrorResponse.from_message(str(e), 404)
714
705
try:
715
706
if config.conversation_id is None:
@@ -822,8 +813,8 @@ class Api:
822
813
if provider is None:
823
814
provider = AppConfig.provider
824
815
try:
825
provider = get_provider_by_label(provider)
826
except ValueError as e:
816
provider = AbstractClientFactory.create_provider(None, provider)
817
except ProviderNotFoundError as e:
827
818
return ErrorResponse.from_message(str(e), 404)
828
819
if config.api_key is None and credentials is not None and credentials.credentials != "secret":
829
820
config.api_key = credentials.credentials
@@ -866,8 +857,8 @@ class Api:
866
857
})
867
858
async def providers_info(provider: str):
868
859
try:
869
provider = get_provider_by_label(provider)
870
except ValueError as e:
860
provider = AbstractClientFactory.create_provider(None, provider)
861
except ProviderNotFoundError as e:
871
862
return ErrorResponse.from_message(str(e), 404)
872
863
def safe_get_models(provider: ProviderType) -> list[str]:
873
864
try:
@@ -1173,8 +1164,8 @@ class Api:
1173
1164
if provider is None:
1174
1165
provider = "MarkItDown"
1175
1166
try:
1176
provider = get_provider_by_label(provider)
1177
except ValueError as e:
1167
provider = AbstractClientFactory.create_provider(None, provider)
1168
except ProviderNotFoundError as e:
1178
1169
return ErrorResponse.from_message(str(e), 404)
1179
1170
kwargs = {"modalities": ["text"]}
1180
1171
try:
@@ -1217,8 +1208,8 @@ class Api:
1217
1208
if provider is None:
1218
1209
provider = AppConfig.media_provider
1219
1210
try:
1220
provider = get_provider_by_label(provider)
1221
except ValueError as e:
1211
provider = AbstractClientFactory.create_provider(None, provider)
1212
except ProviderNotFoundError as e:
1222
1213
return ErrorResponse.from_message(str(e), 404)
1223
1214
try:
1224
1215
audio = filter_none(voice=config.voice, format=config.response_format, language=config.language)
@@ -732,7 +732,7 @@ class ClientFactory(AbstractClientFactory):
732
732
733
733
Example usage:
734
734
# Create client with a named provider
735
client = AbstractClientFactory.create_client("PollinationsAI")
735
client = ClientFactory.create_client("PollinationsAI")
736
736
737
737
# Create client with custom provider
738
738
client = ClientFactory.create_client(
@@ -103,9 +103,10 @@ class AbstractClientFactory:
103
103
if path.exists():
104
104
with open(path, "r", encoding="utf-8") as f:
105
105
cls._live_providers = json.load(f)
106
cls._live_providers = requests.get(cls._live_providers_url).json()
107
with open(path, "w", encoding="utf-8") as f:
108
json.dump(cls._live_providers, f, indent=4)
106
if not cls._live_providers:
107
cls._live_providers = requests.get(cls._live_providers_url).json()
108
with open(path, "w", encoding="utf-8") as f:
109
json.dump(cls._live_providers, f, indent=4)
109
110
if provider in cls._live_providers.get("providers", {}):
110
111
config = cls._live_providers["providers"][provider]
111
112
if "provider" in config and config.get("provider") in ProviderUtils.convert:
@@ -67,7 +67,7 @@ import traceback
67
67
import builtins as _builtins
68
68
from pathlib import Path
69
69
from typing import Any, Dict, FrozenSet, List, Optional, Type
70
70
from .. import debug
71
71
# ---------------------------------------------------------------------------
72
72
# Workspace directory
73
73
# ---------------------------------------------------------------------------
@@ -227,6 +227,7 @@ def _make_restricted_import(allowed: FrozenSet[str]):
227
227
_ALLOWED_G4F_SUBPATHS: FrozenSet[str] = frozenset({
228
228
"g4f.Provider.helper",
229
229
"g4f.Provider.base_provider",
230
"g4f.Provider.template",
230
231
"g4f.typing",
231
232
})
232
233
@@ -605,9 +606,10 @@ class PaProviderRegistry:
605
606
bool(getattr(cls, "working", True)),
606
607
getattr(cls, "url", None),
607
608
cls,
608
str(pa_path)[len(str(directory)):]
609
str(pa_path)[len(str(directory))+1:]
609
610
))
610
except Exception:
611
except Exception as e:
612
debug.error(f"Failed to load PA provider from {pa_path}:", e)
611
613
pass
612
614
self._entries = entries
613
615
self._loaded_at = _time_module.monotonic()