XFE Git
XFE Studio Git
Git 首页 全局搜索
XFE 主站 文档 NuGet
公开
关注 0 Fork 0 Star 1
UTF-8
from __future__ import annotations

import inspect
import logging
import os
import asyncio
import threading
from typing import Iterator
from flask import send_from_directory, request
from inspect import signature
import concurrent.futures

try:
    from PIL import Image

    has_pillow = True
except ImportError:
    has_pillow = False

from ...errors import VersionNotFoundError, MissingAuthError
from ...image.copy_images import copy_media, ensure_media_dir, get_media_dir
from ...image import get_width_height
from ...tools.run_tools import iter_run_tools
from ...providers.base_provider import ProviderModelMixin
from ...providers.retry_provider import BaseRetryProvider
from ...providers.helper import format_media_prompt
from ...providers.response import *
from ...providers.any_model_map import model_map
from ...providers.any_provider import AnyProvider
from ...providers.cache import FileStorage
from ...version import utils as version_utils
from ... import Provider
from ... import debug

logger = logging.getLogger(__name__)
storage = FileStorage()


class Api:
    models_lock = threading.Lock()

    @staticmethod
    async def get_provider_models(
        provider: str, api_key: "str | None" = None, base_url: "str | None" = None, ignored: "list | None" = None
    ):
        def get_model_data(
            provider: ProviderModelMixin, model: str, default: bool = False
        ) -> dict:
            model_id = model.get("id") if isinstance(model, dict) else model
            return {
                "id": model_id,
                "label": model_id,
                "default": default or model_id == provider.default_model,
                "vision": model_id in provider.vision_models,
                "audio": False
                if provider.audio_models is None
                else model_id in provider.audio_models,
                "video": model_id in provider.video_models,
                "image": model_id in provider.image_models,
                "count": False
                if provider.models_count is None
                else provider.models_count.get(model_id),
                "tags": []
                if provider.models_tags is None
                else provider.models_tags.get(model_id, []),
                **(model if isinstance(model, dict) else {}),
            }

        if provider in Provider.__map__:
            provider = Provider.__map__[provider]
            if issubclass(provider, ProviderModelMixin):
                has_grouped_models = hasattr(provider, "get_grouped_models")
                method = (
                    provider.get_grouped_models
                    if has_grouped_models
                    else provider.get_models
                )
                if "api_key" in signature(provider.get_models).parameters:
                    models = method(api_key=api_key, base_url=base_url)
                elif "ignored" in signature(provider.get_models).parameters:
                    models = method(ignored=ignored)
                else:
                    models = method()
                if inspect.isawaitable(models):
                    models = await models
                if has_grouped_models:
                    return [
                        {
                            "group": model.get("group"),
                            "models": [
                                get_model_data(provider, name)
                                for name in (
                                    model.get("models", {}).values()
                                    if isinstance(model.get("models"), dict)
                                    else model.get("models", [])
                                )
                            ],
                        }
                        if model.get("models")
                        else model
                        for model in models
                    ]
                return [
                    get_model_data(provider, model)
                    for model in (
                        models.values() if isinstance(models, dict) else models
                    )
                ]
        elif provider in model_map:
            return [get_model_data(AnyProvider, provider, True)]

        return []

    @staticmethod
    def get_providers() -> dict[str, str]:
        saved = storage.get(f"{version_utils.current_version}/providers")
        if saved is not None:
            return saved
        result = [
            {
                "name": provider.__name__,
                "label": getattr(provider, "label", provider.__name__),
                "parent": getattr(provider, "parent", None),
                "image": len(getattr(provider, "image_models", [])),
                "audio": len(getattr(provider, "audio_models", [])),
                "video": len(getattr(provider, "video_models", [])),
                "vision": getattr(provider, "default_vision_model", None) is not None,
                "nodriver": getattr(provider, "use_nodriver", False),
                "hf_space": getattr(provider, "hf_space", False),
                "active_by_default": False
                if provider.active_by_default is None
                else provider.active_by_default,
                "auth": provider.needs_auth,
                "login_url": getattr(provider, "login_url", None),
                "live": provider.live,
                "login": hasattr(provider, "login"),
            }
            for provider in Provider.__providers__
            if provider.working
        ]
        storage.set(f"{version_utils.current_version}/providers", result)
        return result

    def get_all_models(self) -> dict[str, list]:
        storage_key = f"{version_utils.current_version}/all_models"
        saved = storage.get(storage_key)
        if saved is not None:
            return saved
        with self.models_lock:

            def safe_get_provider_models(provider) -> tuple[str, list[str]]:
                try:
                    models = provider.get_models(timeout=10)
                    if inspect.isawaitable(models):
                        models = asyncio.run(models)
                    return provider.__name__, list(models)
                except MissingAuthError as e:
                    return provider.__name__, []
                except Exception as e:
                    debug.error(f"{provider.__name__}: get_models error:", e)
                    return provider.__name__, []

            providers = [
                p
                for p in Provider.__providers__
                if p.working and hasattr(p, "get_models")
            ]
            results = {}

            with concurrent.futures.ThreadPoolExecutor(max_workers=20) as executor:
                futures = {
                    executor.submit(safe_get_provider_models, p): p for p in providers
                }
                for future in concurrent.futures.as_completed(futures):
                    name, models = future.result()
                    results[name] = models

            storage.set(storage_key, results)
            return results

    @staticmethod
    def get_version() -> dict:
        current_version = None
        latest_version = None
        try:
            current_version = version_utils.current_version
            try:
                if request.args.get("cache"):
                    latest_version = version_utils.latest_version_cached
            except RuntimeError:
                pass
            if latest_version is None:
                latest_version = version_utils.latest_version
        except VersionNotFoundError:
            pass
        return {
            "version": current_version,
            "latest_version": latest_version,
        }

    def serve_images(self, name):
        ensure_media_dir()
        return send_from_directory(os.path.abspath(get_media_dir()), name)

    def _prepare_conversation_kwargs(self, json_data: dict):
        kwargs = {**json_data}
        model = kwargs.pop("model", None)
        provider = kwargs.pop("provider", None)
        messages = kwargs.pop("messages", None)
        action = kwargs.get("action")
        if action == "continue":
            if "tool_calls" not in kwargs:
                kwargs["tool_calls"] = []
            kwargs["tool_calls"].append(
                {"function": {"name": "continue_tool"}, "type": "function"}
            )
        conversation = kwargs.pop("conversation", None)
        if isinstance(conversation, dict):
            kwargs["conversation"] = JsonConversation(**conversation)
        return {
            "model": model,
            "provider": provider,
            "messages": messages,
            "ignore_stream": True,
            **kwargs,
        }

    def _create_response_stream(
        self,
        kwargs: dict,
        provider: str,
        download_media: bool = True,
        tempfiles: list[str] = [],
    ) -> Iterator:
        def decorated_log(*values: str, file=None):
            debug.logs.append(" ".join([str(value) for value in values]))
            if debug.logging:
                debug.log_handler(*values, file=file)

        debug.log = decorated_log
        proxy = os.environ.get("G4F_PROXY")
        try:
            model = kwargs.get("model")
            provider_handler = provider or AnyProvider
            provider = provider_handler.__name__ if provider_handler else provider
            if "user" in kwargs:
                debug.error("User:", kwargs.get("user", "Unknown"))
                debug.error("Referrer:", kwargs.get("referer", ""))
                debug.error("User-Agent:", kwargs.get("user-agent", ""))
        except Exception as e:
            logger.exception(e)
            yield self._format_json(
                "error", type(e).__name__, message=get_error_message(e)
            )
            return
        if not isinstance(provider_handler, BaseRetryProvider):
            yield self.handle_provider(provider_handler, model)
            if hasattr(provider_handler, "get_parameters"):
                yield self._format_json(
                    "parameters", provider_handler.get_parameters(as_json=True)
                )
        try:
            result = iter_run_tools(
                provider_handler,
                **{
                    **kwargs,
                    "model": model,
                    "download_media": download_media,
                    "proxy": proxy,
                },
            )
            for chunk in result:
                if isinstance(chunk, ProviderInfo):
                    model = getattr(chunk, "model", model)
                    provider = getattr(chunk, "provider", provider)
                    yield self.handle_provider(chunk, model)
                elif isinstance(chunk, JsonConversation):
                    if provider is not None:
                        yield self._format_json(
                            "conversation",
                            chunk.get_dict()
                            if provider == "AnyProvider"
                            else {provider: chunk.get_dict()},
                        )
                elif isinstance(chunk, Exception):
                    logger.exception(chunk)
                    yield self._format_json(
                        "message", get_error_message(chunk), error=type(chunk).__name__
                    )
                elif isinstance(chunk, RequestLogin):
                    yield self._format_json("preview", chunk.to_string())
                elif isinstance(chunk, PreviewResponse):
                    yield self._format_json("preview", chunk.to_string())
                elif isinstance(chunk, ImagePreview):
                    yield self._format_json(
                        "preview", chunk.to_string(), urls=chunk.urls, alt=chunk.alt
                    )
                elif isinstance(chunk, MediaResponse):
                    media = chunk
                    if download_media or chunk.get("cookies") or chunk.get("headers"):
                        chunk.alt = format_media_prompt(
                            kwargs.get("messages"), chunk.alt
                        )
                        width, height = get_width_height(
                            chunk.get("width"), chunk.get("height")
                        )
                        tags = [
                            model,
                            kwargs.get("aspect_ratio"),
                            kwargs.get("resolution"),
                        ]
                        media = asyncio.run(
                            copy_media(
                                chunk.get_list(),
                                chunk.get("cookies"),
                                chunk.get("headers"),
                                proxy=proxy,
                                alt=chunk.alt,
                                tags=tags,
                                add_url=True,
                                timeout=kwargs.get("timeout"),
                                return_target=True
                                if isinstance(chunk, ImageResponse)
                                else False,
                            )
                        )
                        options = {}
                        target_paths, urls = get_target_paths_and_urls(media)
                        if target_paths:
                            if has_pillow:
                                try:
                                    with Image.open(target_paths[0]) as img:
                                        width, height = img.size
                                        options = {"width": width, "height": height}
                                except Exception as e:
                                    logger.exception(e)
                            options["target_paths"] = target_paths
                        media = (
                            ImageResponse(urls, chunk.alt, options)
                            if isinstance(chunk, ImageResponse)
                            else VideoResponse(media, chunk.alt)
                        )
                    yield self._format_json(
                        "content", str(media), urls=media.urls, alt=media.alt
                    )
                elif isinstance(chunk, SynthesizeData):
                    yield self._format_json("synthesize", chunk.get_dict())
                elif isinstance(chunk, TitleGeneration):
                    yield self._format_json("title", chunk.title)
                elif isinstance(chunk, Parameters):
                    yield self._format_json("parameters", chunk.get_dict())
                elif isinstance(chunk, FinishReason):
                    yield self._format_json("finish", chunk.get_dict())
                elif isinstance(chunk, Usage):
                    yield self._format_json(
                        "usage", chunk.get_dict(), model=model, provider=provider
                    )
                elif isinstance(chunk, Reasoning):
                    yield self._format_json("reasoning", **chunk.get_dict())
                elif isinstance(chunk, YouTubeResponse):
                    yield self._format_json("content", chunk.to_string())
                elif isinstance(chunk, AudioResponse):
                    yield self._format_json("content", str(chunk), data=chunk.data)
                elif isinstance(chunk, SuggestedFollowups):
                    yield self._format_json("suggestions", chunk.suggestions)
                elif isinstance(chunk, DebugResponse):
                    yield self._format_json("log", chunk.log)
                elif isinstance(chunk, ContinueResponse):
                    yield self._format_json("continue", chunk.text)
                elif isinstance(chunk, VariantResponse):
                    yield self._format_json("variant", chunk.text)
                elif isinstance(chunk, ToolCalls):
                    yield self._format_json("tool_calls", chunk.list)
                elif isinstance(chunk, RawResponse):
                    yield self._format_json(chunk.type, **chunk.get_dict())
                elif isinstance(chunk, JsonRequest):
                    yield self._format_json("request", chunk.get_dict())
                elif isinstance(chunk, JsonResponse):
                    yield self._format_json("response", chunk.get_dict())
                elif isinstance(chunk, PlainTextResponse):
                    yield self._format_json("response", chunk.text)
                elif isinstance(chunk, HeadersResponse):
                    yield self._format_json("headers", chunk.get_dict())
                else:
                    yield self._format_json("content", str(chunk))
        except MissingAuthError as e:
            yield self._format_json(
                "auth", type(e).__name__, message=get_error_message(e)
            )
        except (TimeoutError, asyncio.exceptions.CancelledError) as e:
            if "user" in kwargs:
                debug.error(e, "User:", kwargs.get("user", "Unknown"))
            yield self._format_json(
                "error", type(e).__name__, message=get_error_message(e)
            )
        except Exception as e:
            if "user" in kwargs:
                debug.error(e, "User:", kwargs.get("user", "Unknown"))
            logger.exception(e)
            yield self._format_json(
                "error", type(e).__name__, message=get_error_message(e)
            )
        finally:
            yield from self._yield_logs()
            for tempfile in tempfiles:
                try:
                    os.remove(tempfile)
                except Exception as e:
                    logger.exception(e)

    def _yield_logs(self):
        if debug.logs:
            for log in debug.logs:
                yield self._format_json("log", log)
            debug.logs = []

    def _format_json(self, response_type: str, content=None, **kwargs):
        if content is not None and isinstance(response_type, str):
            return {"type": response_type, response_type: content, **kwargs}
        return {"type": response_type, **kwargs}

    def handle_provider(self, provider_handler, model):
        if not getattr(provider_handler, "model", False):
            return self._format_json(
                "provider", {**provider_handler.get_dict(), "model": model}
            )
        return self._format_json("provider", provider_handler.get_dict())


def get_error_message(exception: Exception) -> str:
    return f"{type(exception).__name__}: {exception}"


def get_target_paths_and_urls(
    media: list[Union[str, tuple[str, str]]]
) -> tuple[list[str], list[str]]:
    target_paths = []
    urls = []
    for item in media:
        if isinstance(item, tuple):
            item, target_path = item
            target_paths.append(target_path)
        urls.append(item)
    return target_paths, urls