mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-07 18:29:04 +02:00
refactor: standardize API error handling (#1635)
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
)
|
||||
|
||||
@@ -29,22 +30,20 @@ func newHandler(service *Service) *handler {
|
||||
// @Param sort[direction] query string false "Sort direction (asc or desc)" default("asc")
|
||||
// @Success 200 {object} dto.Paginated[apiResponseDto]
|
||||
// @Router /api/apis [get]
|
||||
func (h *handler) list(c *gin.Context) {
|
||||
func (h *handler) list(c *gin.Context) error {
|
||||
search := c.Query("search")
|
||||
listRequestOptions := utils.ParseListRequestOptions(c)
|
||||
|
||||
apis, pagination, err := h.service.List(c.Request.Context(), search, listRequestOptions)
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
items := make([]apiResponseDto, len(apis))
|
||||
for i, api := range apis {
|
||||
var item apiResponseDto
|
||||
if err := dto.MapStruct(api, &item); err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
item.Resource = api.Audience
|
||||
items[i] = item
|
||||
@@ -54,6 +53,7 @@ func (h *handler) list(c *gin.Context) {
|
||||
Data: items,
|
||||
Pagination: pagination,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// get godoc
|
||||
@@ -64,14 +64,13 @@ func (h *handler) list(c *gin.Context) {
|
||||
// @Param id path string true "API ID"
|
||||
// @Success 200 {object} apiResponseDto
|
||||
// @Router /api/apis/{id} [get]
|
||||
func (h *handler) get(c *gin.Context) {
|
||||
func (h *handler) get(c *gin.Context) error {
|
||||
api, err := h.service.Get(c.Request.Context(), nil, c.Param("id"))
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
h.respond(c, http.StatusOK, api)
|
||||
return h.respond(c, http.StatusOK, api)
|
||||
}
|
||||
|
||||
// create godoc
|
||||
@@ -83,20 +82,18 @@ func (h *handler) get(c *gin.Context) {
|
||||
// @Param api body apiCreateDto true "API information"
|
||||
// @Success 201 {object} apiResponseDto "Created API"
|
||||
// @Router /api/apis [post]
|
||||
func (h *handler) create(c *gin.Context) {
|
||||
func (h *handler) create(c *gin.Context) error {
|
||||
var input apiCreateDto
|
||||
if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
if err := httpserver.BindJSON(c, &input); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
api, err := h.service.Create(c.Request.Context(), input)
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
h.respond(c, http.StatusCreated, api)
|
||||
return h.respond(c, http.StatusCreated, api)
|
||||
}
|
||||
|
||||
// update godoc
|
||||
@@ -109,20 +106,18 @@ func (h *handler) create(c *gin.Context) {
|
||||
// @Param api body apiUpdateDto true "API information"
|
||||
// @Success 200 {object} apiResponseDto "Updated API"
|
||||
// @Router /api/apis/{id} [put]
|
||||
func (h *handler) update(c *gin.Context) {
|
||||
func (h *handler) update(c *gin.Context) error {
|
||||
var input apiUpdateDto
|
||||
if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
if err := httpserver.BindJSON(c, &input); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
api, err := h.service.Update(c.Request.Context(), c.Param("id"), input)
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
h.respond(c, http.StatusOK, api)
|
||||
return h.respond(c, http.StatusOK, api)
|
||||
}
|
||||
|
||||
// delete godoc
|
||||
@@ -132,13 +127,13 @@ func (h *handler) update(c *gin.Context) {
|
||||
// @Param id path string true "API ID"
|
||||
// @Success 204 "No Content"
|
||||
// @Router /api/apis/{id} [delete]
|
||||
func (h *handler) delete(c *gin.Context) {
|
||||
func (h *handler) delete(c *gin.Context) error {
|
||||
if err := h.service.Delete(c.Request.Context(), c.Param("id")); err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
c.Status(http.StatusNoContent)
|
||||
return nil
|
||||
}
|
||||
|
||||
// updatePermissions godoc
|
||||
@@ -151,20 +146,18 @@ func (h *handler) delete(c *gin.Context) {
|
||||
// @Param permissions body apiPermissionsUpdateDto true "Permissions to set"
|
||||
// @Success 200 {object} apiResponseDto "Updated API"
|
||||
// @Router /api/apis/{id}/permissions [put]
|
||||
func (h *handler) updatePermissions(c *gin.Context) {
|
||||
func (h *handler) updatePermissions(c *gin.Context) error {
|
||||
var input apiPermissionsUpdateDto
|
||||
if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
if err := httpserver.BindJSON(c, &input); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
api, err := h.service.UpdatePermissions(c.Request.Context(), c.Param("id"), input)
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
h.respond(c, http.StatusOK, api)
|
||||
return h.respond(c, http.StatusOK, api)
|
||||
}
|
||||
|
||||
// getClientAccess godoc
|
||||
@@ -175,14 +168,14 @@ func (h *handler) updatePermissions(c *gin.Context) {
|
||||
// @Param clientId path string true "OIDC Client ID"
|
||||
// @Success 200 {object} clientApiAccessDto
|
||||
// @Router /api/api-access/{clientId} [get]
|
||||
func (h *handler) getClientAccess(c *gin.Context) {
|
||||
func (h *handler) getClientAccess(c *gin.Context) error {
|
||||
access, err := h.service.GetClientAPIAccess(c.Request.Context(), c.Param("clientId"))
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, newClientApiAccessDto(access))
|
||||
return nil
|
||||
}
|
||||
|
||||
// updateClientAccess godoc
|
||||
@@ -195,21 +188,20 @@ func (h *handler) getClientAccess(c *gin.Context) {
|
||||
// @Param access body clientApiAccessUpdateDto true "Allowed permission IDs per subject type"
|
||||
// @Success 200 {object} clientApiAccessDto
|
||||
// @Router /api/api-access/{clientId} [put]
|
||||
func (h *handler) updateClientAccess(c *gin.Context) {
|
||||
func (h *handler) updateClientAccess(c *gin.Context) error {
|
||||
var input clientApiAccessUpdateDto
|
||||
err := c.ShouldBindJSON(&input)
|
||||
err := httpserver.BindJSON(c, &input)
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
applied, err := h.service.SetClientAPIAccess(c.Request.Context(), c.Param("clientId"), ClientAPIAccess(input))
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, newClientApiAccessDto(applied))
|
||||
return nil
|
||||
}
|
||||
|
||||
// newClientApiAccessDto always serializes both permission lists as arrays rather than null
|
||||
@@ -224,12 +216,12 @@ func newClientApiAccessDto(access ClientAPIAccess) clientApiAccessDto {
|
||||
return dto
|
||||
}
|
||||
|
||||
func (h *handler) respond(c *gin.Context, status int, api API) {
|
||||
func (h *handler) respond(c *gin.Context, status int, api API) error {
|
||||
var responseDto apiResponseDto
|
||||
if err := dto.MapStruct(api, &responseDto); err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
responseDto.Resource = api.Audience
|
||||
c.JSON(status, responseDto)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/oidc"
|
||||
)
|
||||
|
||||
@@ -63,16 +64,16 @@ func (m *Module) DescribePermissions(ctx context.Context, audience string, keys
|
||||
func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, adminAuth gin.HandlerFunc) {
|
||||
apis := apiGroup.Group("/apis")
|
||||
apis.Use(adminAuth)
|
||||
apis.GET("", m.handler.list)
|
||||
apis.POST("", m.handler.create)
|
||||
apis.GET("/:id", m.handler.get)
|
||||
apis.PUT("/:id", m.handler.update)
|
||||
apis.DELETE("/:id", m.handler.delete)
|
||||
apis.PUT("/:id/permissions", m.handler.updatePermissions)
|
||||
apis.GET("", httpserver.Handle(m.handler.list))
|
||||
apis.POST("", httpserver.Handle(m.handler.create))
|
||||
apis.GET("/:id", httpserver.Handle(m.handler.get))
|
||||
apis.PUT("/:id", httpserver.Handle(m.handler.update))
|
||||
apis.DELETE("/:id", httpserver.Handle(m.handler.delete))
|
||||
apis.PUT("/:id/permissions", httpserver.Handle(m.handler.updatePermissions))
|
||||
|
||||
// The per-client API-access allow-list lives on a separate path so it does not collide with the /apis/:id wildcard
|
||||
access := apiGroup.Group("/api-access")
|
||||
access.Use(adminAuth)
|
||||
access.GET("/:clientId", m.handler.getClientAccess)
|
||||
access.PUT("/:clientId", m.handler.updateClientAccess)
|
||||
access.GET("/:clientId", httpserver.Handle(m.handler.getClientAccess))
|
||||
access.PUT("/:clientId", httpserver.Handle(m.handler.updateClientAccess))
|
||||
}
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/ory/fosite"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/common"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/oidc"
|
||||
@@ -81,13 +81,16 @@ func (s *Service) Get(ctx context.Context, tx *gorm.DB, id string) (api API, err
|
||||
Where("id = ?", id).
|
||||
First(&api).
|
||||
Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return API{}, apperror.NotFound("API")
|
||||
}
|
||||
return api, err
|
||||
}
|
||||
|
||||
func (s *Service) Create(ctx context.Context, input apiCreateDto) (api API, err error) {
|
||||
// Reject the issuer as an audience so a custom API cannot impersonate Pocket ID's own identity tokens
|
||||
if isIssuerAudience(input.Resource, s.issuer) {
|
||||
return API{}, &common.ValidationError{Message: "the resource is reserved by Pocket ID and cannot be used for a custom API"}
|
||||
return API{}, apperror.InvalidField("resource", "reserved", "is reserved by Pocket ID and cannot be used for a custom API")
|
||||
}
|
||||
|
||||
api = API{
|
||||
@@ -98,7 +101,7 @@ func (s *Service) Create(ctx context.Context, input apiCreateDto) (api API, err
|
||||
err = s.db.WithContext(ctx).Create(&api).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrDuplicatedKey) {
|
||||
return API{}, &common.AlreadyInUseError{Property: "resource"}
|
||||
return API{}, apperror.AlreadyInUse("resource")
|
||||
}
|
||||
return API{}, err
|
||||
}
|
||||
@@ -170,16 +173,17 @@ func (s *Service) UpdatePermissions(ctx context.Context, id string, input apiPer
|
||||
// Reject keys with invalid characters, that collide with Pocket ID's reserved scopes and claims, or that repeat within the request before persisting anything
|
||||
// A duplicate key would otherwise be silently coalesced last-wins into the map below, dropping a row behind a 200
|
||||
seen := make(map[string]struct{}, len(input.Permissions))
|
||||
for _, permission := range input.Permissions {
|
||||
for index, permission := range input.Permissions {
|
||||
field := fmt.Sprintf("permissions[%d].key", index)
|
||||
if !isValidPermissionKey(permission.Key) {
|
||||
return API{}, &common.ValidationError{Message: fmt.Sprintf("the permission key %q contains invalid characters", permission.Key)}
|
||||
return API{}, apperror.InvalidField(field, "invalid_format", "contains characters that are not valid in an OAuth scope")
|
||||
}
|
||||
if isPermissionKeyReserved(permission.Key) {
|
||||
return API{}, &common.ValidationError{Message: fmt.Sprintf("the permission key %q is reserved by Pocket ID", permission.Key)}
|
||||
return API{}, apperror.InvalidField(field, "reserved", "is reserved by Pocket ID")
|
||||
}
|
||||
_, ok := seen[permission.Key]
|
||||
if ok {
|
||||
return API{}, &common.ValidationError{Message: fmt.Sprintf("the permission key %q is listed more than once", permission.Key)}
|
||||
return API{}, apperror.InvalidField(field, "duplicate", "is listed more than once")
|
||||
}
|
||||
seen[permission.Key] = struct{}{}
|
||||
}
|
||||
@@ -261,8 +265,17 @@ type ClientAPIAccess struct {
|
||||
// GetClientAPIAccess returns the API permissions a client is allowed to request, split by subject type
|
||||
// Only custom-API permissions are tracked here because the identity scopes are freely requestable by every client
|
||||
func (s *Service) GetClientAPIAccess(ctx context.Context, clientID string) (access ClientAPIAccess, err error) {
|
||||
tx := s.db.Begin()
|
||||
defer func() {
|
||||
tx.Rollback()
|
||||
}()
|
||||
|
||||
if err := ensureOIDCClientExists(ctx, tx, clientID); err != nil {
|
||||
return ClientAPIAccess{}, err
|
||||
}
|
||||
|
||||
var rows []OidcClientAllowedAPIPermission
|
||||
err = s.db.WithContext(ctx).
|
||||
err = tx.WithContext(ctx).
|
||||
Where("oidc_client_id = ?", clientID).
|
||||
Find(&rows).
|
||||
Error
|
||||
@@ -293,9 +306,7 @@ func (s *Service) SetClientAPIAccess(ctx context.Context, clientID string, acces
|
||||
tx.Rollback()
|
||||
}()
|
||||
|
||||
// Ensure the client exists so callers get a 404 for an unknown client
|
||||
var client model.OidcClient
|
||||
if err = tx.WithContext(ctx).Select("id").Where("id = ?", clientID).First(&client).Error; err != nil {
|
||||
if err = ensureOIDCClientExists(ctx, tx, clientID); err != nil {
|
||||
return ClientAPIAccess{}, err
|
||||
}
|
||||
|
||||
@@ -337,6 +348,20 @@ func (s *Service) SetClientAPIAccess(ctx context.Context, clientID string, acces
|
||||
return applied, nil
|
||||
}
|
||||
|
||||
func ensureOIDCClientExists(ctx context.Context, db *gorm.DB, clientID string) error {
|
||||
var client model.OidcClient
|
||||
err := db.WithContext(ctx).
|
||||
Select("id").
|
||||
Where("id = ?", clientID).
|
||||
First(&client).
|
||||
Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return apperror.NotFound("OIDC client")
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// ClientAPIScopesAndAudiences returns the permission keys a client may request and the distinct audiences of the custom APIs those permissions belong to, across both subject types
|
||||
// The OIDC module uses this to widen fosite's scope and audience validation for the client; the per-flow subject-type enforcement happens when the resource is resolved
|
||||
func (s *Service) ClientAPIScopesAndAudiences(ctx context.Context, tx *gorm.DB, clientID string) (scopes []string, audiences []string, err error) {
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/common"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/oidc"
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
@@ -22,7 +22,7 @@ func TestAPICrudAndPermissionDiff(t *testing.T) {
|
||||
|
||||
// The resource is unique.
|
||||
_, err = svc.Create(t.Context(), apiCreateDto{Name: "Dup", Resource: "https://api.orders.example.com"})
|
||||
require.ErrorIs(t, err, &common.AlreadyInUseError{})
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeAlreadyInUse))
|
||||
|
||||
desc := "Read orders"
|
||||
updated, err := svc.UpdatePermissions(t.Context(), created.ID, apiPermissionsUpdateDto{Permissions: []apiPermissionInputDto{
|
||||
@@ -59,7 +59,7 @@ func TestAPICrudAndPermissionDiff(t *testing.T) {
|
||||
|
||||
require.NoError(t, svc.Delete(t.Context(), created.ID))
|
||||
_, err = svc.Get(t.Context(), nil, created.ID)
|
||||
require.Error(t, err)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeNotFound))
|
||||
}
|
||||
|
||||
func TestClientApiAccessAllowList(t *testing.T) {
|
||||
@@ -121,7 +121,10 @@ func TestClientApiAccessAllowList(t *testing.T) {
|
||||
|
||||
// An unknown client is rejected (surfaces as 404 at the HTTP layer).
|
||||
_, err = svc.SetClientAPIAccess(t.Context(), "nope", ClientAPIAccess{UserDelegatedPermissionIDs: []string{readID}})
|
||||
require.Error(t, err)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeNotFound))
|
||||
|
||||
_, err = svc.GetClientAPIAccess(t.Context(), "nope")
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeNotFound))
|
||||
}
|
||||
|
||||
// TestAllowedScopesForAudienceFiltersBySubjectType guards that the scopes resolved for a flow
|
||||
@@ -177,8 +180,7 @@ func TestUpdatePermissionsRejectsReservedKeys(t *testing.T) {
|
||||
{Key: key, Name: "Reserved"},
|
||||
}})
|
||||
require.Error(t, err, "key %q must be rejected", key)
|
||||
var validationErr *common.ValidationError
|
||||
require.ErrorAs(t, err, &validationErr)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeValidationFailed))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -195,8 +197,7 @@ func TestUpdatePermissionsRejectsDuplicateKeys(t *testing.T) {
|
||||
{Key: "read:orders", Name: "Read again"},
|
||||
}})
|
||||
require.Error(t, err)
|
||||
var validationErr *common.ValidationError
|
||||
require.ErrorAs(t, err, &validationErr)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeValidationFailed))
|
||||
}
|
||||
|
||||
func TestUpdatePermissionsRejectsInvalidKeyCharacters(t *testing.T) {
|
||||
@@ -212,8 +213,7 @@ func TestUpdatePermissionsRejectsInvalidKeyCharacters(t *testing.T) {
|
||||
{Key: key, Name: "Invalid"},
|
||||
}})
|
||||
require.Error(t, err, "key %q must be rejected", key)
|
||||
var validationErr *common.ValidationError
|
||||
require.ErrorAs(t, err, &validationErr)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeValidationFailed))
|
||||
}
|
||||
|
||||
// A valid scope-token key is accepted
|
||||
@@ -232,8 +232,7 @@ func TestCreateRejectsIssuerResource(t *testing.T) {
|
||||
for _, resource := range []string{issuer, issuer + "/", "https://ID.example.com"} {
|
||||
_, err := svc.Create(t.Context(), apiCreateDto{Name: "Reserved", Resource: resource})
|
||||
require.Error(t, err, "resource %q must be rejected", resource)
|
||||
var validationErr *common.ValidationError
|
||||
require.ErrorAs(t, err, &validationErr)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeValidationFailed))
|
||||
}
|
||||
|
||||
// A normal resource is accepted
|
||||
|
||||
Reference in New Issue
Block a user