services.py 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393
  1. from core_shared import JSONValue
  2. from core_shared.secrets import EncryptedSecret, SecretCipher
  3. from app.db.models import ToolBinding, ToolCredential, ToolDefinition, ToolVersion
  4. from app.domain.repositories import (
  5. ToolBindingRepository,
  6. ToolCredentialRepository,
  7. ToolDefinitionRepository,
  8. ToolVersionRepository,
  9. )
  10. from app.schemas.tool import (
  11. McpConnectData,
  12. McpConnectRequestDto,
  13. McpToolDto,
  14. ToolBindingCreateRequest,
  15. ToolBindingCreateRequestDto,
  16. ToolBindingDeleteRequestDto,
  17. ToolBindingDetailRequestDto,
  18. ToolBindingUpdateRequestDto,
  19. ToolCreateRequest,
  20. ToolCreateRequestDto,
  21. ToolCredentialCreateRequest,
  22. ToolCredentialCreateRequestDto,
  23. ToolCredentialDeleteRequestDto,
  24. ToolCredentialDetailRequestDto,
  25. ToolCredentialDto,
  26. ToolCredentialRevealDto,
  27. ToolCredentialUpdateRequestDto,
  28. ToolDeleteRequestDto,
  29. ToolDetailRequestDto,
  30. ToolDto,
  31. ToolUpdateRequestDto,
  32. ToolVersionCreateRequest,
  33. ToolVersionCreateRequestDto,
  34. ToolVersionDetailRequestDto,
  35. ToolVersionDto,
  36. ToolVersionUpdateRequestDto,
  37. )
  38. class ToolApplicationService:
  39. def __init__(
  40. self,
  41. tool_definition_repository: ToolDefinitionRepository,
  42. tool_version_repository: ToolVersionRepository,
  43. tool_binding_repository: ToolBindingRepository,
  44. tool_credential_repository: ToolCredentialRepository,
  45. secret_cipher: SecretCipher) -> None:
  46. self.tool_definition_repository = tool_definition_repository
  47. self.tool_version_repository = tool_version_repository
  48. self.tool_binding_repository = tool_binding_repository
  49. self.tool_credential_repository = tool_credential_repository
  50. self.secret_cipher = secret_cipher
  51. def create_tool_definition(self, payload: ToolCreateRequest) -> ToolDefinition:
  52. code = payload.code or self._build_tool_code(payload.name)
  53. return self.tool_definition_repository.create(
  54. plugin_id=payload.plugin_id,
  55. code=code,
  56. name=payload.name,
  57. tool_type=payload.tool_type,
  58. description=payload.description)
  59. def list_tool_definitions(self) -> list[ToolDefinition]:
  60. return self.tool_definition_repository.list_all()
  61. def create_tool_definition_from_contract(
  62. self,
  63. payload: ToolCreateRequestDto) -> ToolDefinition:
  64. return self.create_tool_definition(
  65. ToolCreateRequest(
  66. plugin_id=payload.pluginId,
  67. name=payload.name,
  68. tool_type=payload.toolType,
  69. description=payload.description))
  70. def delete_tool_definition_from_contract(self, payload: ToolDeleteRequestDto) -> bool:
  71. entity = self.tool_definition_repository.get_by_id(tool_id=payload.toolId)
  72. if entity is None:
  73. return False
  74. self.tool_definition_repository.delete(entity)
  75. return True
  76. def get_tool_definition_from_contract(
  77. self,
  78. payload: ToolDetailRequestDto) -> ToolDefinition | None:
  79. return self.tool_definition_repository.get_by_id(tool_id=payload.toolId)
  80. def update_tool_definition_from_contract(
  81. self,
  82. payload: ToolUpdateRequestDto) -> ToolDefinition | None:
  83. entity = self.tool_definition_repository.get_by_id(tool_id=payload.toolId)
  84. if entity is None:
  85. return None
  86. if payload.name is not None:
  87. entity.name = payload.name
  88. entity.code = self._build_tool_code(payload.name)
  89. if payload.toolType is not None:
  90. entity.tool_type = payload.toolType
  91. if payload.description is not None:
  92. entity.description = payload.description
  93. if payload.pluginId is not None:
  94. entity.plugin_id = payload.pluginId
  95. return self.tool_definition_repository.save(entity)
  96. def create_tool_version(self, payload: ToolVersionCreateRequest) -> ToolVersion:
  97. return self.tool_version_repository.create(
  98. tool_id=payload.tool_id,
  99. input_schema_json=payload.input_schema_json,
  100. output_schema_json=payload.output_schema_json,
  101. invoke_config_json=payload.invoke_config_json,
  102. timeout_ms=payload.timeout_ms,
  103. retry_policy_json=payload.retry_policy_json)
  104. def list_tool_versions(self, tool_id: str | None = None) -> list[ToolVersion]:
  105. if tool_id is None:
  106. return self.tool_version_repository.list_all()
  107. return self.tool_version_repository.list_by_tool(tool_id=tool_id)
  108. def create_tool_version_from_contract(
  109. self,
  110. payload: ToolVersionCreateRequestDto) -> ToolVersion:
  111. return self.create_tool_version(
  112. ToolVersionCreateRequest(
  113. tool_id=payload.toolId,
  114. input_schema_json=payload.inputSchema,
  115. output_schema_json=payload.outputSchema,
  116. invoke_config_json=payload.invokeConfig,
  117. timeout_ms=payload.timeoutMs,
  118. retry_policy_json=payload.retryPolicy))
  119. def delete_tool_version(self, *, connection_id: str) -> bool:
  120. entity = self.tool_version_repository.get_by_id(tool_version_id=connection_id)
  121. if entity is None:
  122. return False
  123. self.tool_version_repository.delete(entity)
  124. return True
  125. def get_tool_version_from_contract(
  126. self,
  127. payload: ToolVersionDetailRequestDto) -> ToolVersion | None:
  128. return self.tool_version_repository.get_by_id(
  129. tool_version_id=payload.connectionId)
  130. def update_tool_version_from_contract(
  131. self,
  132. payload: ToolVersionUpdateRequestDto) -> ToolVersion | None:
  133. entity = self.tool_version_repository.get_by_id(
  134. tool_version_id=payload.connectionId)
  135. if entity is None:
  136. return None
  137. if payload.inputSchema is not None:
  138. entity.input_schema_json = payload.inputSchema
  139. if payload.outputSchema is not None:
  140. entity.output_schema_json = payload.outputSchema
  141. if payload.invokeConfig is not None:
  142. entity.invoke_config_json = payload.invokeConfig
  143. if payload.timeoutMs is not None:
  144. entity.timeout_ms = payload.timeoutMs
  145. if payload.retryPolicy is not None:
  146. entity.retry_policy_json = payload.retryPolicy
  147. return self.tool_version_repository.save(entity)
  148. def connect_mcp_server(self, payload: McpConnectRequestDto) -> McpConnectData:
  149. server_name, config = self._normalize_mcp_config(payload)
  150. discovered_tools = self._extract_mcp_tools(config)
  151. if discovered_tools:
  152. config["mcp_tools"] = [
  153. tool.model_dump(mode="json", by_alias=False)
  154. for tool in discovered_tools
  155. ]
  156. tool = self.create_tool_definition(
  157. ToolCreateRequest(
  158. name=payload.name or server_name,
  159. tool_type="mcp",
  160. description=f"MCP SSE server: {config.get('url', '')}"))
  161. version = self.create_tool_version(
  162. ToolVersionCreateRequest(
  163. tool_id=tool.id,
  164. input_schema_json={},
  165. output_schema_json={},
  166. invoke_config_json=config,
  167. timeout_ms=self._timeout_ms(config),
  168. retry_policy_json={"max_attempts": 1}))
  169. return McpConnectData(
  170. tool=ToolDto.from_entity(tool),
  171. connection=ToolVersionDto.from_entity(version),
  172. discoveredTools=discovered_tools)
  173. def create_tool_binding(self, payload: ToolBindingCreateRequest) -> ToolBinding:
  174. if payload.credential_id is not None:
  175. credential = self.tool_credential_repository.get_by_id(
  176. credential_id=payload.credential_id)
  177. if credential is None:
  178. raise ValueError(f"tool credential not found: {payload.credential_id}")
  179. return self.tool_binding_repository.create(
  180. app_id=payload.app_id,
  181. tool_version_id=payload.tool_version_id,
  182. credential_id=payload.credential_id,
  183. binding_scope=payload.binding_scope,
  184. enabled=payload.enabled,
  185. config_json=payload.config_json)
  186. def create_tool_binding_from_contract(
  187. self,
  188. payload: ToolBindingCreateRequestDto) -> ToolBinding:
  189. return self.create_tool_binding(
  190. ToolBindingCreateRequest(
  191. app_id=payload.appId,
  192. tool_version_id=payload.toolVersionId,
  193. credential_id=payload.credentialId,
  194. binding_scope=payload.bindingScope,
  195. enabled=True,
  196. config_json=payload.configJson))
  197. def list_tool_bindings(self, app_id: str | None = None) -> list[ToolBinding]:
  198. return self.tool_binding_repository.list_by_scope(app_id=app_id)
  199. def delete_tool_binding(self, payload: ToolBindingDeleteRequestDto) -> bool:
  200. entity = self.tool_binding_repository.get_by_id(binding_id=payload.bindingId)
  201. if entity is None:
  202. return False
  203. self.tool_binding_repository.delete(entity)
  204. return True
  205. def get_tool_binding_from_contract(
  206. self,
  207. payload: ToolBindingDetailRequestDto) -> ToolBinding | None:
  208. return self.tool_binding_repository.get_by_id(binding_id=payload.bindingId)
  209. def update_tool_binding_from_contract(
  210. self,
  211. payload: ToolBindingUpdateRequestDto) -> ToolBinding | None:
  212. entity = self.tool_binding_repository.get_by_id(binding_id=payload.bindingId)
  213. if entity is None:
  214. return None
  215. if payload.credentialId is not None:
  216. entity.credential_id = payload.credentialId
  217. if payload.bindingScope is not None:
  218. entity.binding_scope = payload.bindingScope
  219. if payload.configJson is not None:
  220. entity.config_json = payload.configJson
  221. return self.tool_binding_repository.save(entity)
  222. def create_tool_credential(self, payload: ToolCredentialCreateRequest) -> ToolCredential:
  223. encrypted = self.secret_cipher.encrypt_json(payload.secret_json)
  224. return self.tool_credential_repository.create(
  225. name=payload.name,
  226. credential_type=payload.credential_type,
  227. encrypted_payload_text=encrypted.ciphertext,
  228. secret_fingerprint=encrypted.fingerprint,
  229. encryption_algorithm=encrypted.algorithm,
  230. metadata_json=payload.metadata_json)
  231. def create_tool_credential_from_contract(
  232. self,
  233. payload: ToolCredentialCreateRequestDto) -> ToolCredential:
  234. return self.create_tool_credential(
  235. ToolCredentialCreateRequest(
  236. name=payload.name,
  237. credential_type=payload.credentialType,
  238. secret_json=payload.secretJson,
  239. metadata_json=payload.metadataJson))
  240. def list_tool_credentials(self) -> list[ToolCredential]:
  241. return self.tool_credential_repository.list_all()
  242. def delete_tool_credential(self, payload: ToolCredentialDeleteRequestDto) -> bool:
  243. entity = self.tool_credential_repository.get_by_id(
  244. credential_id=payload.credentialId)
  245. if entity is None:
  246. return False
  247. self.tool_credential_repository.delete(entity)
  248. return True
  249. def get_tool_credential_from_contract(
  250. self,
  251. payload: ToolCredentialDetailRequestDto) -> ToolCredential | None:
  252. return self.tool_credential_repository.get_by_id(
  253. credential_id=payload.credentialId)
  254. def update_tool_credential_from_contract(
  255. self,
  256. payload: ToolCredentialUpdateRequestDto) -> ToolCredential | None:
  257. entity = self.tool_credential_repository.get_by_id(
  258. credential_id=payload.credentialId)
  259. if entity is None:
  260. return None
  261. if payload.name is not None:
  262. entity.name = payload.name
  263. if payload.metadataJson is not None:
  264. entity.metadata_json = payload.metadataJson
  265. return self.tool_credential_repository.save(entity)
  266. def reveal_tool_credential(
  267. self,
  268. *,
  269. credential_id: str) -> tuple[ToolCredential, dict[str, JSONValue]] | None:
  270. credential = self.tool_credential_repository.get_by_id(
  271. credential_id=credential_id)
  272. if credential is None:
  273. return None
  274. secret_json = self.secret_cipher.decrypt_json(
  275. EncryptedSecret(
  276. ciphertext=credential.encrypted_payload_text,
  277. fingerprint=credential.secret_fingerprint,
  278. algorithm=credential.encryption_algorithm)
  279. )
  280. return credential, secret_json
  281. def reveal_tool_credential_from_contract(
  282. self,
  283. *,
  284. credential_id: str) -> ToolCredentialRevealDto | None:
  285. result = self.reveal_tool_credential(credential_id=credential_id)
  286. if result is None:
  287. return None
  288. credential, secret_json = result
  289. return ToolCredentialRevealDto(
  290. credential=ToolCredentialDto.from_entity(credential),
  291. secretJson=secret_json)
  292. def get_tool_binding_detail(
  293. self,
  294. *,
  295. binding_id: str) -> tuple[ToolBinding, ToolVersion, ToolDefinition] | None:
  296. binding = self.tool_binding_repository.get_by_id(binding_id=binding_id)
  297. if binding is None:
  298. return None
  299. tool_version = self.tool_version_repository.get_by_id(
  300. tool_version_id=binding.tool_version_id)
  301. if tool_version is None:
  302. return None
  303. tool_definition = self.tool_definition_repository.get_by_id(
  304. tool_id=tool_version.tool_id)
  305. if tool_definition is None:
  306. return None
  307. return binding, tool_version, tool_definition
  308. def _build_tool_code(self, name: str) -> str:
  309. base = "".join(
  310. char.lower() if char.isalnum() else "_"
  311. for char in name
  312. ).strip("_") or "tool"
  313. return base[:64]
  314. def _normalize_mcp_config(
  315. self,
  316. payload: McpConnectRequestDto) -> tuple[str, dict[str, JSONValue]]:
  317. config = payload.config
  318. if len(config) == 1:
  319. server_name, value = next(iter(config.items()))
  320. if isinstance(value, dict):
  321. normalized = {str(key): item for key, item in value.items()}
  322. normalized.setdefault("server_name", server_name)
  323. normalized.setdefault("transport", "sse")
  324. return payload.name or server_name, normalized
  325. server_name_value = config.get("server_name") or config.get("serverName")
  326. server_name = str(server_name_value or payload.name or "mcp_server")
  327. normalized = {str(key): value for key, value in config.items()}
  328. normalized.setdefault("server_name", server_name)
  329. normalized.setdefault("transport", "sse")
  330. return server_name, normalized
  331. def _extract_mcp_tools(self, config: dict[str, JSONValue]) -> list[McpToolDto]:
  332. raw_tools = config.get("mcp_tools") or config.get("tools")
  333. if not isinstance(raw_tools, list):
  334. return []
  335. tools: list[McpToolDto] = []
  336. for item in raw_tools:
  337. if not isinstance(item, dict):
  338. continue
  339. name = item.get("name")
  340. if not isinstance(name, str) or not name:
  341. continue
  342. description = item.get("description")
  343. input_schema = item.get("inputSchema") or item.get("input_schema")
  344. tools.append(
  345. McpToolDto(
  346. name=name,
  347. description=description if isinstance(description, str) else None,
  348. inputSchema=input_schema if isinstance(input_schema, dict) else None))
  349. return tools
  350. def _timeout_ms(self, config: dict[str, JSONValue]) -> int | None:
  351. timeout = config.get("timeout") or config.get("timeout_seconds")
  352. if isinstance(timeout, int | float):
  353. return int(timeout * 1000)
  354. return None