feat(captain): validate custom tools

This commit is contained in:
2026-06-07 16:20:39 +08:00
parent dd7bcf4118
commit 8f9adecd04
5 changed files with 213 additions and 12 deletions
@@ -246,6 +246,67 @@ func (s *CaptainCustomToolCRUDTestSuite) TestCreate_超过每账户15个工具
assert.Equal(s.T(), "You can create a maximum of 15 custom tools per account", resp["error"])
}
func (s *CaptainCustomToolCRUDTestSuite) TestCreate_无效枚举返回RecordInvalid形态422() {
body := map[string]interface{}{
"custom_tool": map[string]interface{}{
"title": "Invalid enums",
"endpoint_url": "https://example.com/invalid",
"http_method": "DELETE",
"auth_type": "oauth",
},
}
w := s.makeRequest("POST", s.accountPath()+"/captain/custom_tools/", body)
assert.Equal(s.T(), http.StatusUnprocessableEntity, w.Code)
var resp map[string]interface{}
s.Require().NoError(json.Unmarshal(w.Body.Bytes(), &resp))
assert.Contains(s.T(), resp["message"], "Http method is not included in the list")
assert.Contains(s.T(), resp["message"], "Auth type is not included in the list")
assert.ElementsMatch(s.T(), []interface{}{"http_method", "auth_type"}, resp["attributes"].([]interface{}))
}
func (s *CaptainCustomToolCRUDTestSuite) TestCreate_无效ParamSchema返回RecordInvalid形态422() {
body := map[string]interface{}{
"custom_tool": map[string]interface{}{
"title": "Invalid schema",
"endpoint_url": "https://example.com/schema",
"param_schema": []map[string]interface{}{{
"name": "order_id",
"type": "string",
"extra": "not allowed",
}},
},
}
w := s.makeRequest("POST", s.accountPath()+"/captain/custom_tools/", body)
assert.Equal(s.T(), http.StatusUnprocessableEntity, w.Code)
var resp map[string]interface{}
s.Require().NoError(json.Unmarshal(w.Body.Bytes(), &resp))
assert.Contains(s.T(), resp["message"], "Description is required")
assert.Contains(s.T(), resp["message"], "Extra is not permitted")
assert.ElementsMatch(s.T(), []interface{}{"description", "extra"}, resp["attributes"].([]interface{}))
}
func (s *CaptainCustomToolCRUDTestSuite) TestCreate_重复显式Slug返回422() {
_ = s.createToolAndGetID("First duplicate", "explicit-duplicate", "https://example.com/first")
w := s.makeRequest("POST", s.accountPath()+"/captain/custom_tools/", map[string]interface{}{
"custom_tool": map[string]interface{}{
"title": "Second duplicate",
"slug": "explicit-duplicate",
"endpoint_url": "https://example.com/second",
},
})
assert.Equal(s.T(), http.StatusUnprocessableEntity, w.Code)
var resp map[string]interface{}
s.Require().NoError(json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(s.T(), "Slug has already been taken", resp["message"])
assert.ElementsMatch(s.T(), []interface{}{"slug"}, resp["attributes"].([]interface{}))
}
func (s *CaptainCustomToolCRUDTestSuite) TestCreate_无效JSON返回400() {
// ShouldBindJSON 在 JSON 解析失败时返回 400
w := httptest.NewRecorder()
@@ -47,6 +47,9 @@ func (h *CaptainCustomToolHandler) Create(c *gin.Context) {
c.JSON(http.StatusUnprocessableEntity, gin.H{"error": service.ErrCaptainCustomToolLimitExceeded.Error()})
return
}
if renderCaptainCustomToolValidationError(c, err) {
return
}
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to create custom tool")
return
}
@@ -107,6 +110,9 @@ func (h *CaptainCustomToolHandler) Update(c *gin.Context) {
tool, err := h.svc.UpdateByAccount(c.Request.Context(), accountID, id, &req)
if err != nil {
applogger.L().Errorf("Update captain custom tool: %v", err)
if renderCaptainCustomToolValidationError(c, err) {
return
}
handleServiceError(c, err)
return
}
@@ -227,6 +233,18 @@ func (h *CaptainCustomToolHandler) ensureCustomToolsEnabled(c *gin.Context, acco
return false
}
func renderCaptainCustomToolValidationError(c *gin.Context, err error) bool {
var validationErr *service.CaptainCustomToolValidationError
if !errors.As(err, &validationErr) {
return false
}
c.JSON(http.StatusUnprocessableEntity, gin.H{
"message": validationErr.Message,
"attributes": validationErr.Attributes,
})
return true
}
func captainCustomToolPayload(tool *model.CaptainCustomTool) gin.H {
return gin.H{
"id": tool.ID,
@@ -187,7 +187,7 @@ func TestCaptainCustomToolHandler_ChatwootToolPayloadsAndScope(t *testing.T) {
"endpoint_url": "https://example.com/orders",
"http_method": "POST",
"auth_type": "none",
"param_schema": []map[string]any{{"name": "order_id", "type": "string", "required": true}},
"param_schema": []map[string]any{{"name": "order_id", "type": "string", "description": "Order ID", "required": true}},
"request_template": "{\"id\":\"{{.order_id}}\"}",
}}
w := captainResourceJSONRequest(t, router, http.MethodPost, basePath+"/", body)
+126 -7
View File
@@ -36,6 +36,15 @@ const (
var ErrCaptainCustomToolLimitExceeded = errors.New("You can create a maximum of 15 custom tools per account")
type CaptainCustomToolValidationError struct {
Message string
Attributes []string
}
func (e *CaptainCustomToolValidationError) Error() string {
return e.Message
}
type HTTPDoer interface {
Do(req *http.Request) (*http.Response, error)
}
@@ -100,13 +109,7 @@ type UpdateCustomToolRequest struct {
// Create creates a new CaptainCustomTool.
func (s *CaptainCustomToolService) Create(ctx context.Context, accountID uint, req *CreateCustomToolRequest) (*model.CaptainCustomTool, error) {
count, err := s.toolRepo.CountByAccount(ctx, accountID)
if err != nil {
return nil, fmt.Errorf("count custom tools: %w", err)
}
if count >= maxCaptainCustomToolsPerAccount {
return nil, ErrCaptainCustomToolLimitExceeded
}
var err error
// Default values
httpMethod := req.HTTPMethod
@@ -139,6 +142,19 @@ func (s *CaptainCustomToolService) Create(ctx context.Context, accountID uint, r
ResponseTemplate: req.ResponseTemplate,
Enabled: true,
}
if err := validateCaptainCustomTool(tool); err != nil {
return nil, err
}
if req.Slug != "" && s.customToolSlugExists(ctx, accountID, req.Slug) {
return nil, newCaptainCustomToolValidationError("Slug has already been taken", "slug")
}
count, err := s.toolRepo.CountByAccount(ctx, accountID)
if err != nil {
return nil, fmt.Errorf("count custom tools: %w", err)
}
if count >= maxCaptainCustomToolsPerAccount {
return nil, ErrCaptainCustomToolLimitExceeded
}
if err := s.toolRepo.Create(ctx, tool); err != nil {
applogger.L().Errorf("Create captain custom tool: %v", err)
@@ -217,6 +233,9 @@ func (s *CaptainCustomToolService) UpdateByAccount(ctx context.Context, accountI
return nil, fmt.Errorf("custom tool not found: %w", err)
}
applyCustomToolUpdate(tool, req)
if err := validateCaptainCustomTool(tool); err != nil {
return nil, err
}
if err := s.toolRepo.Update(ctx, tool); err != nil {
applogger.L().Errorf("Update captain custom tool: %v", err)
return nil, fmt.Errorf("update custom tool: %w", err)
@@ -308,6 +327,106 @@ func (s *CaptainCustomToolService) customToolSlugExists(ctx context.Context, acc
return err == nil
}
func validateCaptainCustomTool(tool *model.CaptainCustomTool) error {
var messages []string
var attributes []string
add := func(attr, message string) {
messages = append(messages, message)
attributes = append(attributes, attr)
}
if strings.TrimSpace(tool.Title) == "" {
add("title", "Title can't be blank")
}
if strings.TrimSpace(tool.EndpointURL) == "" {
add("endpoint_url", "Endpoint url can't be blank")
}
if len(tool.Slug) > maxCaptainCustomToolSlugLength {
add("slug", "Slug is too long (maximum is 64 characters)")
}
if tool.HTTPMethod != "GET" && tool.HTTPMethod != "POST" {
add("http_method", "Http method is not included in the list")
}
switch tool.AuthType {
case model.ToolAuthTypeNone, model.ToolAuthTypeBearer, model.ToolAuthTypeBasic, model.ToolAuthTypeApiKey:
default:
add("auth_type", "Auth type is not included in the list")
}
for _, validation := range validateCaptainCustomToolParamSchema(tool.ParamSchema) {
add(validation.attribute, validation.message)
}
if len(messages) == 0 {
return nil
}
return &CaptainCustomToolValidationError{Message: strings.Join(messages, ", "), Attributes: uniqueStrings(attributes)}
}
type customToolParamSchemaValidation struct {
attribute string
message string
}
func validateCaptainCustomToolParamSchema(raw json.RawMessage) []customToolParamSchemaValidation {
if len(raw) == 0 || string(raw) == "null" {
return nil
}
var items []map[string]any
if err := json.Unmarshal(raw, &items); err != nil {
return []customToolParamSchemaValidation{{attribute: "param_schema", message: "Param schema must be of type array"}}
}
allowed := map[string]bool{"name": true, "type": true, "description": true, "required": true}
var validations []customToolParamSchemaValidation
for _, item := range items {
for _, field := range []string{"name", "type", "description"} {
value, ok := item[field]
if !ok {
validations = append(validations, customToolParamSchemaValidation{attribute: field, message: customToolFieldLabel(field) + " is required"})
continue
}
if _, ok := value.(string); !ok {
validations = append(validations, customToolParamSchemaValidation{attribute: field, message: customToolFieldLabel(field) + " must be of type string"})
}
}
if value, ok := item["required"]; ok {
if _, ok := value.(bool); !ok {
validations = append(validations, customToolParamSchemaValidation{attribute: "required", message: "Required must be of type boolean"})
}
}
for field := range item {
if !allowed[field] {
validations = append(validations, customToolParamSchemaValidation{attribute: field, message: customToolFieldLabel(field) + " is not permitted"})
}
}
}
return validations
}
func customToolFieldLabel(field string) string {
if field == "" {
return field
}
return strings.ToUpper(field[:1]) + field[1:]
}
func newCaptainCustomToolValidationError(message, attribute string) error {
return &CaptainCustomToolValidationError{Message: message, Attributes: []string{attribute}}
}
func uniqueStrings(values []string) []string {
seen := make(map[string]bool, len(values))
unique := make([]string, 0, len(values))
for _, value := range values {
if seen[value] {
continue
}
seen[value] = true
unique = append(unique, value)
}
return unique
}
func customToolSlug(title string) string {
slug := strings.ToLower(strings.TrimSpace(title))
slug = regexp.MustCompile(`[^a-z0-9]+`).ReplaceAllString(slug, "_")