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

XFEstudio/gpt4free

Refactor provider handling in Api class and update ClientFactory usage

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

代码差异

4 个文件 +26 -32
Modified g4f/api/__init__.py +16 -25
@@ -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)
Modified g4f/client/__init__.py +1 -1
@@ -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(
Modified g4f/client/factory.py +4 -3
@@ -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:
Modified g4f/mcp/pa_provider.py +5 -3
@@ -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()