diff --git a/backend/pkg/response/abort.go b/backend/pkg/response/abort.go index 16d5902b..3d25361c 100644 --- a/backend/pkg/response/abort.go +++ b/backend/pkg/response/abort.go @@ -4,9 +4,11 @@ package response import ( + "errors" "net/http" "github.com/gin-gonic/gin" + "gorm.io/gorm" ) // AbortBadRequest 以 400 中断请求并将错误挂载到 Gin Error 链,供全局中间件统一记录 Trace 并响应。 @@ -43,3 +45,27 @@ func AbortTooManyRequests(c *gin.Context, msg string) { func AbortConflict(c *gin.Context, msg string) { AbortWithError(c, http.StatusConflict, msg) } + +// AbortNotFoundIfMissing 在 err 非空时中断请求:gorm.ErrRecordNotFound 映射为 404(文案 notFoundMsg),其余映射为 400(文案 err.Error())。 +// 返回是否已中断。err 为 nil 时不写响应并返回 false。 +func AbortNotFoundIfMissing(c *gin.Context, err error, notFoundMsg string) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + AbortNotFound(c, notFoundMsg) + return true + } + AbortBadRequest(c, err.Error()) + return true +} + +// AbortBadRequestOnError 在 err 非空时以 400 中断请求,文案为 err.Error()。 +// 返回是否已中断。err 为 nil 时不写响应并返回 false。 +func AbortBadRequestOnError(c *gin.Context, err error) bool { + if err == nil { + return false + } + AbortBadRequest(c, err.Error()) + return true +} diff --git a/backend/pkg/response/abort_helpers_test.go b/backend/pkg/response/abort_helpers_test.go new file mode 100644 index 00000000..d05f1235 --- /dev/null +++ b/backend/pkg/response/abort_helpers_test.go @@ -0,0 +1,68 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package response + +import ( + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func TestAbortNotFoundIfMissing(t *testing.T) { + t.Run("nil", func(t *testing.T) { + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + assert.False(t, AbortNotFoundIfMissing(c, nil, "gone")) + assert.False(t, c.IsAborted()) + assert.Empty(t, c.Errors) + }) + + t.Run("record not found", func(t *testing.T) { + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + assert.True(t, AbortNotFoundIfMissing(c, gorm.ErrRecordNotFound, "记录不存在")) + assert.True(t, c.IsAborted()) + var apiErr *APIError + require.True(t, errors.As(c.Errors.Last().Err, &apiErr)) + assert.Equal(t, http.StatusNotFound, apiErr.Code) + assert.Equal(t, "记录不存在", apiErr.Msg) + }) + + t.Run("other error", func(t *testing.T) { + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + assert.True(t, AbortNotFoundIfMissing(c, errors.New("boom"), "记录不存在")) + assert.True(t, c.IsAborted()) + var apiErr *APIError + require.True(t, errors.As(c.Errors.Last().Err, &apiErr)) + assert.Equal(t, http.StatusBadRequest, apiErr.Code) + assert.Equal(t, "boom", apiErr.Msg) + }) +} + +func TestAbortBadRequestOnError(t *testing.T) { + t.Run("nil", func(t *testing.T) { + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + assert.False(t, AbortBadRequestOnError(c, nil)) + assert.False(t, c.IsAborted()) + }) + + t.Run("error", func(t *testing.T) { + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + assert.True(t, AbortBadRequestOnError(c, errors.New("bad"))) + assert.True(t, c.IsAborted()) + var apiErr *APIError + require.True(t, errors.As(c.Errors.Last().Err, &apiErr)) + assert.Equal(t, http.StatusBadRequest, apiErr.Code) + assert.Equal(t, "bad", apiErr.Msg) + }) +}