Underground-Digital/Workflow-Engine
0
1import json2import logging3from typing import Optional4 5from httpx import get6 7from core.model_runtime.utils.encoders import jsonable_encoder8from core.tools.entities.api_entities import UserTool, UserToolProvider9from core.tools.entities.common_entities import I18nObject10from core.tools.entities.tool_bundle import ApiToolBundle11from core.tools.entities.tool_entities import (12 ApiProviderAuthType,13 ApiProviderSchemaType,14 ToolCredentialsOption,15 ToolProviderCredentials,16)17from core.tools.provider.api_tool_provider import ApiToolProviderController18from core.tools.tool_label_manager import ToolLabelManager19from core.tools.tool_manager import ToolManager20from core.tools.utils.configuration import ToolConfigurationManager21from core.tools.utils.parser import ApiBasedToolSchemaParser22from extensions.ext_database import db23from models.tools import ApiToolProvider24from services.tools.tools_transform_service import ToolTransformService25 26logger = logging.getLogger(__name__)27 28 29class ApiToolManageService:30 @staticmethod31 def parser_api_schema(schema: str) -> list[ApiToolBundle]:32 """33 parse api schema to tool bundle34 """35 try:36 warnings = {}37 try:38 tool_bundles, schema_type = ApiBasedToolSchemaParser.auto_parse_to_tool_bundle(schema, warning=warnings)39 except Exception as e:40 raise ValueError(f"invalid schema: {str(e)}")41 42 credentials_schema = [43 ToolProviderCredentials(44 name="auth_type",45 type=ToolProviderCredentials.CredentialsType.SELECT,46 required=True,47 default="none",48 options=[49 ToolCredentialsOption(value="none", label=I18nObject(en_US="None", zh_Hans="无")),50 ToolCredentialsOption(value="api_key", label=I18nObject(en_US="Api Key", zh_Hans="Api Key")),51 ],52 placeholder=I18nObject(en_US="Select auth type", zh_Hans="选择认证方式"),53 ),54 ToolProviderCredentials(55 name="api_key_header",56 type=ToolProviderCredentials.CredentialsType.TEXT_INPUT,57 required=False,58 placeholder=I18nObject(en_US="Enter api key header", zh_Hans="输入 api key header,如:X-API-KEY"),59 default="api_key",60 help=I18nObject(en_US="HTTP header name for api key", zh_Hans="HTTP 头部字段名,用于传递 api key"),61 ),62 ToolProviderCredentials(63 name="api_key_value",64 type=ToolProviderCredentials.CredentialsType.TEXT_INPUT,65 required=False,66 placeholder=I18nObject(en_US="Enter api key", zh_Hans="输入 api key"),67 default="",68 ),69 ]70 71 return jsonable_encoder(72 {73 "schema_type": schema_type,74 "parameters_schema": tool_bundles,75 "credentials_schema": credentials_schema,76 "warning": warnings,77 }78 )79 except Exception as e:80 raise ValueError(f"invalid schema: {str(e)}")81 82 @staticmethod83 def convert_schema_to_tool_bundles(84 schema: str, extra_info: Optional[dict] = None85 ) -> tuple[list[ApiToolBundle], str]:86 """87 convert schema to tool bundles88 89 :return: the list of tool bundles, description90 """91 try:92 tool_bundles = ApiBasedToolSchemaParser.auto_parse_to_tool_bundle(schema, extra_info=extra_info)93 return tool_bundles94 except Exception as e:95 raise ValueError(f"invalid schema: {str(e)}")96 97 @staticmethod98 def create_api_tool_provider(99 user_id: str,100 tenant_id: str,101 provider_name: str,102 icon: dict,103 credentials: dict,104 schema_type: str,105 schema: str,106 privacy_policy: str,107 custom_disclaimer: str,108 labels: list[str],109 ):110 """111 create api tool provider112 """113 if schema_type not in [member.value for member in ApiProviderSchemaType]:114 raise ValueError(f"invalid schema type {schema}")115 116 # check if the provider exists117 provider: ApiToolProvider = (118 db.session.query(ApiToolProvider)119 .filter(120 ApiToolProvider.tenant_id == tenant_id,121 ApiToolProvider.name == provider_name,122 )123 .first()124 )125 126 if provider is not None:127 raise ValueError(f"provider {provider_name} already exists")128 129 # parse openapi to tool bundle130 extra_info = {}131 # extra info like description will be set here132 tool_bundles, schema_type = ApiToolManageService.convert_schema_to_tool_bundles(schema, extra_info)133 134 if len(tool_bundles) > 100:135 raise ValueError("the number of apis should be less than 100")136 137 # create db provider138 db_provider = ApiToolProvider(139 tenant_id=tenant_id,140 user_id=user_id,141 name=provider_name,142 icon=json.dumps(icon),143 schema=schema,144 description=extra_info.get("description", ""),145 schema_type_str=schema_type,146 tools_str=json.dumps(jsonable_encoder(tool_bundles)),147 credentials_str={},148 privacy_policy=privacy_policy,149 custom_disclaimer=custom_disclaimer,150 )151 152 if "auth_type" not in credentials:153 raise ValueError("auth_type is required")154 155 # get auth type, none or api key156 auth_type = ApiProviderAuthType.value_of(credentials["auth_type"])157 158 # create provider entity159 provider_controller = ApiToolProviderController.from_db(db_provider, auth_type)160 # load tools into provider entity161 provider_controller.load_bundled_tools(tool_bundles)162 163 # encrypt credentials164 tool_configuration = ToolConfigurationManager(tenant_id=tenant_id, provider_controller=provider_controller)165 encrypted_credentials = tool_configuration.encrypt_tool_credentials(credentials)166 db_provider.credentials_str = json.dumps(encrypted_credentials)167 168 db.session.add(db_provider)169 db.session.commit()170 171 # update labels172 ToolLabelManager.update_tool_labels(provider_controller, labels)173 174 return {"result": "success"}175 176 @staticmethod177 def get_api_tool_provider_remote_schema(user_id: str, tenant_id: str, url: str):178 """179 get api tool provider remote schema180 """181 headers = {182 "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko)"183 " Chrome/120.0.0.0 Safari/537.36 Edg/120.0.0.0",184 "Accept": "*/*",185 }186 187 try:188 response = get(url, headers=headers, timeout=10)189 if response.status_code != 200:190 raise ValueError(f"Got status code {response.status_code}")191 schema = response.text192 193 # try to parse schema, avoid SSRF attack194 ApiToolManageService.parser_api_schema(schema)195 except Exception as e:196 logger.error(f"parse api schema error: {str(e)}")197 raise ValueError("invalid schema, please check the url you provided")198 199 return {"schema": schema}200 201 @staticmethod202 def list_api_tool_provider_tools(user_id: str, tenant_id: str, provider: str) -> list[UserTool]:203 """204 list api tool provider tools205 """206 provider: ApiToolProvider = (207 db.session.query(ApiToolProvider)208 .filter(209 ApiToolProvider.tenant_id == tenant_id,210 ApiToolProvider.name == provider,211 )212 .first()213 )214 215 if provider is None:216 raise ValueError(f"you have not added provider {provider}")217 218 controller = ToolTransformService.api_provider_to_controller(db_provider=provider)219 labels = ToolLabelManager.get_tool_labels(controller)220 221 return [222 ToolTransformService.tool_to_user_tool(223 tool_bundle,224 labels=labels,225 )226 for tool_bundle in provider.tools227 ]228 229 @staticmethod230 def update_api_tool_provider(231 user_id: str,232 tenant_id: str,233 provider_name: str,234 original_provider: str,235 icon: dict,236 credentials: dict,237 schema_type: str,238 schema: str,239 privacy_policy: str,240 custom_disclaimer: str,241 labels: list[str],242 ):243 """244 update api tool provider245 """246 if schema_type not in [member.value for member in ApiProviderSchemaType]:247 raise ValueError(f"invalid schema type {schema}")248 249 # check if the provider exists250 provider: ApiToolProvider = (251 db.session.query(ApiToolProvider)252 .filter(253 ApiToolProvider.tenant_id == tenant_id,254 ApiToolProvider.name == original_provider,255 )256 .first()257 )258 259 if provider is None:260 raise ValueError(f"api provider {provider_name} does not exists")261 262 # parse openapi to tool bundle263 extra_info = {}264 # extra info like description will be set here265 tool_bundles, schema_type = ApiToolManageService.convert_schema_to_tool_bundles(schema, extra_info)266 267 # update db provider268 provider.name = provider_name269 provider.icon = json.dumps(icon)270 provider.schema = schema271 provider.description = extra_info.get("description", "")272 provider.schema_type_str = ApiProviderSchemaType.OPENAPI.value273 provider.tools_str = json.dumps(jsonable_encoder(tool_bundles))274 provider.privacy_policy = privacy_policy275 provider.custom_disclaimer = custom_disclaimer276 277 if "auth_type" not in credentials:278 raise ValueError("auth_type is required")279 280 # get auth type, none or api key281 auth_type = ApiProviderAuthType.value_of(credentials["auth_type"])282 283 # create provider entity284 provider_controller = ApiToolProviderController.from_db(provider, auth_type)285 # load tools into provider entity286 provider_controller.load_bundled_tools(tool_bundles)287 288 # get original credentials if exists289 tool_configuration = ToolConfigurationManager(tenant_id=tenant_id, provider_controller=provider_controller)290 291 original_credentials = tool_configuration.decrypt_tool_credentials(provider.credentials)292 masked_credentials = tool_configuration.mask_tool_credentials(original_credentials)293 # check if the credential has changed, save the original credential294 for name, value in credentials.items():295 if name in masked_credentials and value == masked_credentials[name]:296 credentials[name] = original_credentials[name]297 298 credentials = tool_configuration.encrypt_tool_credentials(credentials)299 provider.credentials_str = json.dumps(credentials)300 301 db.session.add(provider)302 db.session.commit()303 304 # delete cache305 tool_configuration.delete_tool_credentials_cache()306 307 # update labels308 ToolLabelManager.update_tool_labels(provider_controller, labels)309 310 return {"result": "success"}311 312 @staticmethod313 def delete_api_tool_provider(user_id: str, tenant_id: str, provider_name: str):314 """315 delete tool provider316 """317 provider: ApiToolProvider = (318 db.session.query(ApiToolProvider)319 .filter(320 ApiToolProvider.tenant_id == tenant_id,321 ApiToolProvider.name == provider_name,322 )323 .first()324 )325 326 if provider is None:327 raise ValueError(f"you have not added provider {provider_name}")328 329 db.session.delete(provider)330 db.session.commit()331 332 return {"result": "success"}333 334 @staticmethod335 def get_api_tool_provider(user_id: str, tenant_id: str, provider: str):336 """337 get api tool provider338 """339 return ToolManager.user_get_api_provider(provider=provider, tenant_id=tenant_id)340 341 @staticmethod342 def test_api_tool_preview(343 tenant_id: str,344 provider_name: str,345 tool_name: str,346 credentials: dict,347 parameters: dict,348 schema_type: str,349 schema: str,350 ):351 """352 test api tool before adding api tool provider353 """354 if schema_type not in [member.value for member in ApiProviderSchemaType]:355 raise ValueError(f"invalid schema type {schema_type}")356 357 try:358 tool_bundles, _ = ApiBasedToolSchemaParser.auto_parse_to_tool_bundle(schema)359 except Exception as e:360 raise ValueError("invalid schema")361 362 # get tool bundle363 tool_bundle = next(filter(lambda tb: tb.operation_id == tool_name, tool_bundles), None)364 if tool_bundle is None:365 raise ValueError(f"invalid tool name {tool_name}")366 367 db_provider: ApiToolProvider = (368 db.session.query(ApiToolProvider)369 .filter(370 ApiToolProvider.tenant_id == tenant_id,371 ApiToolProvider.name == provider_name,372 )373 .first()374 )375 376 if not db_provider:377 # create a fake db provider378 db_provider = ApiToolProvider(379 tenant_id="",380 user_id="",381 name="",382 icon="",383 schema=schema,384 description="",385 schema_type_str=ApiProviderSchemaType.OPENAPI.value,386 tools_str=json.dumps(jsonable_encoder(tool_bundles)),387 credentials_str=json.dumps(credentials),388 )389 390 if "auth_type" not in credentials:391 raise ValueError("auth_type is required")392 393 # get auth type, none or api key394 auth_type = ApiProviderAuthType.value_of(credentials["auth_type"])395 396 # create provider entity397 provider_controller = ApiToolProviderController.from_db(db_provider, auth_type)398 # load tools into provider entity399 provider_controller.load_bundled_tools(tool_bundles)400 401 # decrypt credentials402 if db_provider.id:403 tool_configuration = ToolConfigurationManager(tenant_id=tenant_id, provider_controller=provider_controller)404 decrypted_credentials = tool_configuration.decrypt_tool_credentials(credentials)405 # check if the credential has changed, save the original credential406 masked_credentials = tool_configuration.mask_tool_credentials(decrypted_credentials)407 for name, value in credentials.items():408 if name in masked_credentials and value == masked_credentials[name]:409 credentials[name] = decrypted_credentials[name]410 411 try:412 provider_controller.validate_credentials_format(credentials)413 # get tool414 tool = provider_controller.get_tool(tool_name)415 tool = tool.fork_tool_runtime(416 runtime={417 "credentials": credentials,418 "tenant_id": tenant_id,419 }420 )421 result = tool.validate_credentials(credentials, parameters)422 except Exception as e:423 return {"error": str(e)}424 425 return {"result": result or "empty response"}426 427 @staticmethod428 def list_api_tools(user_id: str, tenant_id: str) -> list[UserToolProvider]:429 """430 list api tools431 """432 # get all api providers433 db_providers: list[ApiToolProvider] = (434 db.session.query(ApiToolProvider).filter(ApiToolProvider.tenant_id == tenant_id).all() or []435 )436 437 result: list[UserToolProvider] = []438 439 for provider in db_providers:440 # convert provider controller to user provider441 provider_controller = ToolTransformService.api_provider_to_controller(db_provider=provider)442 labels = ToolLabelManager.get_tool_labels(provider_controller)443 user_provider = ToolTransformService.api_provider_to_user_provider(444 provider_controller, db_provider=provider, decrypt_credentials=True445 )446 user_provider.labels = labels447 448 # add icon449 ToolTransformService.repack_provider(user_provider)450 451 tools = provider_controller.get_tools(user_id=user_id, tenant_id=tenant_id)452 453 for tool in tools:454 user_provider.tools.append(455 ToolTransformService.tool_to_user_tool(456 tenant_id=tenant_id, tool=tool, credentials=user_provider.original_credentials, labels=labels457 )458 )459 460 result.append(user_provider)461 462 return result463 