feat(captain): validate custom tools
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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, "_")
|
||||
|
||||
Reference in New Issue
Block a user