Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
model_provider_service.py569 linesDownload Raw Back to services
1import logging2import mimetypes3import os4from pathlib import Path5from typing import Optional, cast6 7import requests8from flask import current_app9 10from core.entities.model_entities import ModelStatus, ProviderModelWithStatusEntity11from core.model_runtime.entities.model_entities import ModelType, ParameterRule12from core.model_runtime.model_providers import model_provider_factory13from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel14from core.provider_manager import ProviderManager15from models.provider import ProviderType16from services.entities.model_provider_entities import (17    CustomConfigurationResponse,18    CustomConfigurationStatus,19    DefaultModelResponse,20    ModelWithProviderEntityResponse,21    ProviderResponse,22    ProviderWithModelsResponse,23    SimpleProviderEntityResponse,24    SystemConfigurationResponse,25)26 27logger = logging.getLogger(__name__)28 29 30class ModelProviderService:31    """32    Model Provider Service33    """34 35    def __init__(self) -> None:36        self.provider_manager = ProviderManager()37 38    def get_provider_list(self, tenant_id: str, model_type: Optional[str] = None) -> list[ProviderResponse]:39        """40        get provider list.41 42        :param tenant_id: workspace id43        :param model_type: model type44        :return:45        """46        # Get all provider configurations of the current workspace47        provider_configurations = self.provider_manager.get_configurations(tenant_id)48 49        provider_responses = []50        for provider_configuration in provider_configurations.values():51            if model_type:52                model_type_entity = ModelType.value_of(model_type)53                if model_type_entity not in provider_configuration.provider.supported_model_types:54                    continue55 56            provider_response = ProviderResponse(57                provider=provider_configuration.provider.provider,58                label=provider_configuration.provider.label,59                description=provider_configuration.provider.description,60                icon_small=provider_configuration.provider.icon_small,61                icon_large=provider_configuration.provider.icon_large,62                background=provider_configuration.provider.background,63                help=provider_configuration.provider.help,64                supported_model_types=provider_configuration.provider.supported_model_types,65                configurate_methods=provider_configuration.provider.configurate_methods,66                provider_credential_schema=provider_configuration.provider.provider_credential_schema,67                model_credential_schema=provider_configuration.provider.model_credential_schema,68                preferred_provider_type=provider_configuration.preferred_provider_type,69                custom_configuration=CustomConfigurationResponse(70                    status=CustomConfigurationStatus.ACTIVE71                    if provider_configuration.is_custom_configuration_available()72                    else CustomConfigurationStatus.NO_CONFIGURE73                ),74                system_configuration=SystemConfigurationResponse(75                    enabled=provider_configuration.system_configuration.enabled,76                    current_quota_type=provider_configuration.system_configuration.current_quota_type,77                    quota_configurations=provider_configuration.system_configuration.quota_configurations,78                ),79            )80 81            provider_responses.append(provider_response)82 83        return provider_responses84 85    def get_models_by_provider(self, tenant_id: str, provider: str) -> list[ModelWithProviderEntityResponse]:86        """87        get provider models.88        For the model provider page,89        only supports passing in a single provider to query the list of supported models.90 91        :param tenant_id:92        :param provider:93        :return:94        """95        # Get all provider configurations of the current workspace96        provider_configurations = self.provider_manager.get_configurations(tenant_id)97 98        # Get provider available models99        return [100            ModelWithProviderEntityResponse(model) for model in provider_configurations.get_models(provider=provider)101        ]102 103    def get_provider_credentials(self, tenant_id: str, provider: str) -> dict:104        """105        get provider credentials.106 107        :param tenant_id:108        :param provider:109        :return:110        """111        # Get all provider configurations of the current workspace112        provider_configurations = self.provider_manager.get_configurations(tenant_id)113 114        # Get provider configuration115        provider_configuration = provider_configurations.get(provider)116        if not provider_configuration:117            raise ValueError(f"Provider {provider} does not exist.")118 119        # Get provider custom credentials from workspace120        return provider_configuration.get_custom_credentials(obfuscated=True)121 122    def provider_credentials_validate(self, tenant_id: str, provider: str, credentials: dict) -> None:123        """124        validate provider credentials.125 126        :param tenant_id:127        :param provider:128        :param credentials:129        """130        # Get all provider configurations of the current workspace131        provider_configurations = self.provider_manager.get_configurations(tenant_id)132 133        # Get provider configuration134        provider_configuration = provider_configurations.get(provider)135        if not provider_configuration:136            raise ValueError(f"Provider {provider} does not exist.")137 138        provider_configuration.custom_credentials_validate(credentials)139 140    def save_provider_credentials(self, tenant_id: str, provider: str, credentials: dict) -> None:141        """142        save custom provider config.143 144        :param tenant_id: workspace id145        :param provider: provider name146        :param credentials: provider credentials147        :return:148        """149        # Get all provider configurations of the current workspace150        provider_configurations = self.provider_manager.get_configurations(tenant_id)151 152        # Get provider configuration153        provider_configuration = provider_configurations.get(provider)154        if not provider_configuration:155            raise ValueError(f"Provider {provider} does not exist.")156 157        # Add or update custom provider credentials.158        provider_configuration.add_or_update_custom_credentials(credentials)159 160    def remove_provider_credentials(self, tenant_id: str, provider: str) -> None:161        """162        remove custom provider config.163 164        :param tenant_id: workspace id165        :param provider: provider name166        :return:167        """168        # Get all provider configurations of the current workspace169        provider_configurations = self.provider_manager.get_configurations(tenant_id)170 171        # Get provider configuration172        provider_configuration = provider_configurations.get(provider)173        if not provider_configuration:174            raise ValueError(f"Provider {provider} does not exist.")175 176        # Remove custom provider credentials.177        provider_configuration.delete_custom_credentials()178 179    def get_model_credentials(self, tenant_id: str, provider: str, model_type: str, model: str) -> dict:180        """181        get model credentials.182 183        :param tenant_id: workspace id184        :param provider: provider name185        :param model_type: model type186        :param model: model name187        :return:188        """189        # Get all provider configurations of the current workspace190        provider_configurations = self.provider_manager.get_configurations(tenant_id)191 192        # Get provider configuration193        provider_configuration = provider_configurations.get(provider)194        if not provider_configuration:195            raise ValueError(f"Provider {provider} does not exist.")196 197        # Get model custom credentials from ProviderModel if exists198        return provider_configuration.get_custom_model_credentials(199            model_type=ModelType.value_of(model_type), model=model, obfuscated=True200        )201 202    def model_credentials_validate(203        self, tenant_id: str, provider: str, model_type: str, model: str, credentials: dict204    ) -> None:205        """206        validate model credentials.207 208        :param tenant_id: workspace id209        :param provider: provider name210        :param model_type: model type211        :param model: model name212        :param credentials: model credentials213        :return:214        """215        # Get all provider configurations of the current workspace216        provider_configurations = self.provider_manager.get_configurations(tenant_id)217 218        # Get provider configuration219        provider_configuration = provider_configurations.get(provider)220        if not provider_configuration:221            raise ValueError(f"Provider {provider} does not exist.")222 223        # Validate model credentials224        provider_configuration.custom_model_credentials_validate(225            model_type=ModelType.value_of(model_type), model=model, credentials=credentials226        )227 228    def save_model_credentials(229        self, tenant_id: str, provider: str, model_type: str, model: str, credentials: dict230    ) -> None:231        """232        save model credentials.233 234        :param tenant_id: workspace id235        :param provider: provider name236        :param model_type: model type237        :param model: model name238        :param credentials: model credentials239        :return:240        """241        # Get all provider configurations of the current workspace242        provider_configurations = self.provider_manager.get_configurations(tenant_id)243 244        # Get provider configuration245        provider_configuration = provider_configurations.get(provider)246        if not provider_configuration:247            raise ValueError(f"Provider {provider} does not exist.")248 249        # Add or update custom model credentials250        provider_configuration.add_or_update_custom_model_credentials(251            model_type=ModelType.value_of(model_type), model=model, credentials=credentials252        )253 254    def remove_model_credentials(self, tenant_id: str, provider: str, model_type: str, model: str) -> None:255        """256        remove model credentials.257 258        :param tenant_id: workspace id259        :param provider: provider name260        :param model_type: model type261        :param model: model name262        :return:263        """264        # Get all provider configurations of the current workspace265        provider_configurations = self.provider_manager.get_configurations(tenant_id)266 267        # Get provider configuration268        provider_configuration = provider_configurations.get(provider)269        if not provider_configuration:270            raise ValueError(f"Provider {provider} does not exist.")271 272        # Remove custom model credentials273        provider_configuration.delete_custom_model_credentials(model_type=ModelType.value_of(model_type), model=model)274 275    def get_models_by_model_type(self, tenant_id: str, model_type: str) -> list[ProviderWithModelsResponse]:276        """277        get models by model type.278 279        :param tenant_id: workspace id280        :param model_type: model type281        :return:282        """283        # Get all provider configurations of the current workspace284        provider_configurations = self.provider_manager.get_configurations(tenant_id)285 286        # Get provider available models287        models = provider_configurations.get_models(model_type=ModelType.value_of(model_type))288 289        # Group models by provider290        provider_models = {}291        for model in models:292            if model.provider.provider not in provider_models:293                provider_models[model.provider.provider] = []294 295            if model.deprecated:296                continue297 298            if model.status != ModelStatus.ACTIVE:299                continue300 301            provider_models[model.provider.provider].append(model)302 303        # convert to ProviderWithModelsResponse list304        providers_with_models: list[ProviderWithModelsResponse] = []305        for provider, models in provider_models.items():306            if not models:307                continue308 309            first_model = models[0]310 311            providers_with_models.append(312                ProviderWithModelsResponse(313                    provider=provider,314                    label=first_model.provider.label,315                    icon_small=first_model.provider.icon_small,316                    icon_large=first_model.provider.icon_large,317                    status=CustomConfigurationStatus.ACTIVE,318                    models=[319                        ProviderModelWithStatusEntity(320                            model=model.model,321                            label=model.label,322                            model_type=model.model_type,323                            features=model.features,324                            fetch_from=model.fetch_from,325                            model_properties=model.model_properties,326                            status=model.status,327                            load_balancing_enabled=model.load_balancing_enabled,328                        )329                        for model in models330                    ],331                )332            )333 334        return providers_with_models335 336    def get_model_parameter_rules(self, tenant_id: str, provider: str, model: str) -> list[ParameterRule]:337        """338        get model parameter rules.339        Only supports LLM.340 341        :param tenant_id: workspace id342        :param provider: provider name343        :param model: model name344        :return:345        """346        # Get all provider configurations of the current workspace347        provider_configurations = self.provider_manager.get_configurations(tenant_id)348 349        # Get provider configuration350        provider_configuration = provider_configurations.get(provider)351        if not provider_configuration:352            raise ValueError(f"Provider {provider} does not exist.")353 354        # Get model instance of LLM355        model_type_instance = provider_configuration.get_model_type_instance(ModelType.LLM)356        model_type_instance = cast(LargeLanguageModel, model_type_instance)357 358        # fetch credentials359        credentials = provider_configuration.get_current_credentials(model_type=ModelType.LLM, model=model)360 361        if not credentials:362            return []363 364        # Call get_parameter_rules method of model instance to get model parameter rules365        return model_type_instance.get_parameter_rules(model=model, credentials=credentials)366 367    def get_default_model_of_model_type(self, tenant_id: str, model_type: str) -> Optional[DefaultModelResponse]:368        """369        get default model of model type.370 371        :param tenant_id: workspace id372        :param model_type: model type373        :return:374        """375        model_type_enum = ModelType.value_of(model_type)376        result = self.provider_manager.get_default_model(tenant_id=tenant_id, model_type=model_type_enum)377        try:378            return (379                DefaultModelResponse(380                    model=result.model,381                    model_type=result.model_type,382                    provider=SimpleProviderEntityResponse(383                        provider=result.provider.provider,384                        label=result.provider.label,385                        icon_small=result.provider.icon_small,386                        icon_large=result.provider.icon_large,387                        supported_model_types=result.provider.supported_model_types,388                    ),389                )390                if result391                else None392            )393        except Exception as e:394            logger.info(f"get_default_model_of_model_type error: {e}")395            return None396 397    def update_default_model_of_model_type(self, tenant_id: str, model_type: str, provider: str, model: str) -> None:398        """399        update default model of model type.400 401        :param tenant_id: workspace id402        :param model_type: model type403        :param provider: provider name404        :param model: model name405        :return:406        """407        model_type_enum = ModelType.value_of(model_type)408        self.provider_manager.update_default_model_record(409            tenant_id=tenant_id, model_type=model_type_enum, provider=provider, model=model410        )411 412    def get_model_provider_icon(413        self, provider: str, icon_type: str, lang: str414    ) -> tuple[Optional[bytes], Optional[str]]:415        """416        get model provider icon.417 418        :param provider: provider name419        :param icon_type: icon type (icon_small or icon_large)420        :param lang: language (zh_Hans or en_US)421        :return:422        """423        provider_instance = model_provider_factory.get_provider_instance(provider)424        provider_schema = provider_instance.get_provider_schema()425 426        if icon_type.lower() == "icon_small":427            if not provider_schema.icon_small:428                raise ValueError(f"Provider {provider} does not have small icon.")429 430            if lang.lower() == "zh_hans":431                file_name = provider_schema.icon_small.zh_Hans432            else:433                file_name = provider_schema.icon_small.en_US434        else:435            if not provider_schema.icon_large:436                raise ValueError(f"Provider {provider} does not have large icon.")437 438            if lang.lower() == "zh_hans":439                file_name = provider_schema.icon_large.zh_Hans440            else:441                file_name = provider_schema.icon_large.en_US442 443        root_path = current_app.root_path444        provider_instance_path = os.path.dirname(445            os.path.join(root_path, provider_instance.__class__.__module__.replace(".", "/"))446        )447        file_path = os.path.join(provider_instance_path, "_assets")448        file_path = os.path.join(file_path, file_name)449 450        if not os.path.exists(file_path):451            return None, None452 453        mimetype, _ = mimetypes.guess_type(file_path)454        mimetype = mimetype or "application/octet-stream"455 456        # read binary from file457        byte_data = Path(file_path).read_bytes()458        return byte_data, mimetype459 460    def switch_preferred_provider(self, tenant_id: str, provider: str, preferred_provider_type: str) -> None:461        """462        switch preferred provider.463 464        :param tenant_id: workspace id465        :param provider: provider name466        :param preferred_provider_type: preferred provider type467        :return:468        """469        # Get all provider configurations of the current workspace470        provider_configurations = self.provider_manager.get_configurations(tenant_id)471 472        # Convert preferred_provider_type to ProviderType473        preferred_provider_type_enum = ProviderType.value_of(preferred_provider_type)474 475        # Get provider configuration476        provider_configuration = provider_configurations.get(provider)477        if not provider_configuration:478            raise ValueError(f"Provider {provider} does not exist.")479 480        # Switch preferred provider type481        provider_configuration.switch_preferred_provider_type(preferred_provider_type_enum)482 483    def enable_model(self, tenant_id: str, provider: str, model: str, model_type: str) -> None:484        """485        enable model.486 487        :param tenant_id: workspace id488        :param provider: provider name489        :param model: model name490        :param model_type: model type491        :return:492        """493        # Get all provider configurations of the current workspace494        provider_configurations = self.provider_manager.get_configurations(tenant_id)495 496        # Get provider configuration497        provider_configuration = provider_configurations.get(provider)498        if not provider_configuration:499            raise ValueError(f"Provider {provider} does not exist.")500 501        # Enable model502        provider_configuration.enable_model(model=model, model_type=ModelType.value_of(model_type))503 504    def disable_model(self, tenant_id: str, provider: str, model: str, model_type: str) -> None:505        """506        disable model.507 508        :param tenant_id: workspace id509        :param provider: provider name510        :param model: model name511        :param model_type: model type512        :return:513        """514        # Get all provider configurations of the current workspace515        provider_configurations = self.provider_manager.get_configurations(tenant_id)516 517        # Get provider configuration518        provider_configuration = provider_configurations.get(provider)519        if not provider_configuration:520            raise ValueError(f"Provider {provider} does not exist.")521 522        # Enable model523        provider_configuration.disable_model(model=model, model_type=ModelType.value_of(model_type))524 525    def free_quota_submit(self, tenant_id: str, provider: str):526        api_key = os.environ.get("FREE_QUOTA_APPLY_API_KEY")527        api_base_url = os.environ.get("FREE_QUOTA_APPLY_BASE_URL")528        api_url = api_base_url + "/api/v1/providers/apply"529 530        headers = {"Content-Type": "application/json", "Authorization": f"Bearer {api_key}"}531        response = requests.post(api_url, headers=headers, json={"workspace_id": tenant_id, "provider_name": provider})532        if not response.ok:533            logger.error(f"Request FREE QUOTA APPLY SERVER Error: {response.status_code} ")534            raise ValueError(f"Error: {response.status_code} ")535 536        if response.json()["code"] != "success":537            raise ValueError(f"error: {response.json()['message']}")538 539        rst = response.json()540 541        if rst["type"] == "redirect":542            return {"type": rst["type"], "redirect_url": rst["redirect_url"]}543        else:544            return {"type": rst["type"], "result": "success"}545 546    def free_quota_qualification_verify(self, tenant_id: str, provider: str, token: Optional[str]):547        api_key = os.environ.get("FREE_QUOTA_APPLY_API_KEY")548        api_base_url = os.environ.get("FREE_QUOTA_APPLY_BASE_URL")549        api_url = api_base_url + "/api/v1/providers/qualification-verify"550 551        headers = {"Content-Type": "application/json", "Authorization": f"Bearer {api_key}"}552        json_data = {"workspace_id": tenant_id, "provider_name": provider}553        if token:554            json_data["token"] = token555        response = requests.post(api_url, headers=headers, json=json_data)556        if not response.ok:557            logger.error(f"Request FREE QUOTA APPLY SERVER Error: {response.status_code} ")558            raise ValueError(f"Error: {response.status_code} ")559 560        rst = response.json()561        if rst["code"] != "success":562            raise ValueError(f"error: {rst['message']}")563 564        data = rst["data"]565        if data["qualified"] is True:566            return {"result": "success", "provider_name": provider, "flag": True}567        else:568            return {"result": "success", "provider_name": provider, "flag": False, "reason": data["reason"]}569