tooling.py 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127
  1. """Tool 注册、上下文注入与执行审计。"""
  2. from collections.abc import Callable
  3. from dataclasses import dataclass, field
  4. from typing import Any
  5. from langchain_core.tools import BaseTool, StructuredTool
  6. from pydantic import BaseModel
  7. from app.core.errors import AppError
  8. from app.domains.identity.models import AdminUser, H5User
  9. from app.harness.policy import HarnessPolicyEngine, ToolPolicy
  10. from app.harness.schemas import HarnessToolResult
  11. ToolHandler = Callable[["ToolExecutionContext", BaseModel], HarnessToolResult]
  12. @dataclass(frozen=True)
  13. class ToolDefinition:
  14. name: str
  15. description: str
  16. arguments: type[BaseModel]
  17. policy: ToolPolicy
  18. handler: ToolHandler
  19. @dataclass(frozen=True)
  20. class ToolExecution:
  21. name: str
  22. result: HarnessToolResult
  23. @dataclass
  24. class ToolExecutionContext:
  25. persona: str
  26. h5_user: H5User | None = None
  27. admin_user: AdminUser | None = None
  28. executions: list[ToolExecution] = field(default_factory=list)
  29. @property
  30. def principal_id(self) -> str:
  31. principal = self.h5_user or self.admin_user
  32. if principal is None:
  33. raise AppError("AGENT_PRINCIPAL_REQUIRED", "Agent缺少登录身份", 401)
  34. return principal.id
  35. @property
  36. def permissions(self) -> tuple[str, ...]:
  37. return self.admin_user.permissions if self.admin_user else ()
  38. class ToolRegistry:
  39. def __init__(self) -> None:
  40. self._items: dict[str, ToolDefinition] = {}
  41. def register(self, definition: ToolDefinition) -> None:
  42. if definition.name in self._items:
  43. raise AppError(
  44. "HARNESS_TOOL_DUPLICATED",
  45. f"Tool标识重复:{definition.name}",
  46. 500,
  47. )
  48. self._items[definition.name] = definition
  49. def names(self) -> tuple[str, ...]:
  50. return tuple(sorted(self._items))
  51. def build(
  52. self,
  53. *,
  54. names: tuple[str, ...],
  55. context: ToolExecutionContext,
  56. policy_engine: HarnessPolicyEngine,
  57. ) -> list[BaseTool]:
  58. tools: list[BaseTool] = []
  59. for name in names:
  60. definition = self._items.get(name)
  61. if definition is None:
  62. raise AppError(
  63. "HARNESS_TOOL_NOT_FOUND",
  64. f"未找到Tool:{name}",
  65. 500,
  66. )
  67. if not policy_engine.can_use(
  68. persona=context.persona,
  69. allowed_by_persona=names,
  70. tool_name=name,
  71. policy=definition.policy,
  72. permissions=context.permissions,
  73. ):
  74. continue
  75. tools.append(
  76. self._to_langchain_tool(
  77. definition,
  78. names,
  79. context,
  80. policy_engine,
  81. )
  82. )
  83. return tools
  84. @staticmethod
  85. def _to_langchain_tool(
  86. definition: ToolDefinition,
  87. allowed_names: tuple[str, ...],
  88. context: ToolExecutionContext,
  89. policy_engine: HarnessPolicyEngine,
  90. ) -> BaseTool:
  91. def invoke(**kwargs: Any) -> str:
  92. policy_engine.authorize(
  93. persona=context.persona,
  94. allowed_by_persona=allowed_names,
  95. tool_name=definition.name,
  96. policy=definition.policy,
  97. permissions=context.permissions,
  98. )
  99. arguments = definition.arguments.model_validate(kwargs)
  100. result = definition.handler(context, arguments)
  101. context.executions.append(ToolExecution(name=definition.name, result=result))
  102. return result.model_dump_json(exclude_none=True)
  103. return StructuredTool.from_function(
  104. func=invoke,
  105. name=definition.name,
  106. description=definition.description,
  107. args_schema=definition.arguments,
  108. )