mirror of
https://github.com/navidrome/navidrome.git
synced 2026-10-08 02:17:25 +02:00
feat(api): enforce API v1 security from the spec and harden problem responses
This commit is contained in:
parent
71c379ff44
commit
3e645959f8
13 changed files with 918 additions and 34 deletions
|
|
@ -29,6 +29,7 @@
|
|||
"tags": [
|
||||
"server"
|
||||
],
|
||||
"security": [],
|
||||
"summary": "Describe the server",
|
||||
"description": "Returns the public server description. No authentication required.\nAuthenticated requests will additionally receive the implemented capability modules\nonce authentication is available.\n",
|
||||
"responses": {
|
||||
|
|
@ -56,6 +57,7 @@
|
|||
"tags": [
|
||||
"server"
|
||||
],
|
||||
"security": [],
|
||||
"summary": "Get the OpenAPI document (JSON)",
|
||||
"description": "The bundled OpenAPI document of the running server version. Supports ETag revalidation.",
|
||||
"responses": {
|
||||
|
|
@ -89,6 +91,7 @@
|
|||
"tags": [
|
||||
"server"
|
||||
],
|
||||
"security": [],
|
||||
"summary": "Get the OpenAPI document (YAML)",
|
||||
"description": "The bundled OpenAPI document of the running server version. Supports ETag revalidation.",
|
||||
"responses": {
|
||||
|
|
@ -195,13 +198,23 @@
|
|||
"enum": [
|
||||
"validation",
|
||||
"unauthorized",
|
||||
"token_expired",
|
||||
"forbidden",
|
||||
"insufficient_scope",
|
||||
"not_found",
|
||||
"method_not_allowed",
|
||||
"setup_complete",
|
||||
"password_managed_externally",
|
||||
"payload_too_large",
|
||||
"rate_limited",
|
||||
"unavailable",
|
||||
"internal"
|
||||
]
|
||||
},
|
||||
"referenceId": {
|
||||
"type": "string",
|
||||
"description": "Present on internal errors. Quote it when reporting a problem; it tags the server's log lines for this request."
|
||||
},
|
||||
"errors": {
|
||||
"type": "array",
|
||||
"description": "Per-field failures. Present only when `code` is `validation`.",
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ paths:
|
|||
x-module: core
|
||||
x-stability-level: alpha
|
||||
tags: [server]
|
||||
security: []
|
||||
summary: Describe the server
|
||||
description: |
|
||||
Returns the public server description. No authentication required.
|
||||
|
|
@ -50,6 +51,7 @@ paths:
|
|||
x-module: core
|
||||
x-stability-level: alpha
|
||||
tags: [server]
|
||||
security: []
|
||||
summary: Get the OpenAPI document (JSON)
|
||||
description: The bundled OpenAPI document of the running server version. Supports ETag revalidation.
|
||||
responses:
|
||||
|
|
@ -71,6 +73,7 @@ paths:
|
|||
x-module: core
|
||||
x-stability-level: alpha
|
||||
tags: [server]
|
||||
security: []
|
||||
summary: Get the OpenAPI document (YAML)
|
||||
description: The bundled OpenAPI document of the running server version. Supports ETag revalidation.
|
||||
responses:
|
||||
|
|
@ -152,11 +155,20 @@ components:
|
|||
enum:
|
||||
- validation
|
||||
- unauthorized
|
||||
- token_expired
|
||||
- forbidden
|
||||
- insufficient_scope
|
||||
- not_found
|
||||
- method_not_allowed
|
||||
- setup_complete
|
||||
- password_managed_externally
|
||||
- payload_too_large
|
||||
- rate_limited
|
||||
- unavailable
|
||||
- internal
|
||||
referenceId:
|
||||
type: string
|
||||
description: Present on internal errors. Quote it when reporting a problem; it tags the server's log lines for this request.
|
||||
errors:
|
||||
type: array
|
||||
description: Per-field failures. Present only when `code` is `validation`.
|
||||
|
|
|
|||
|
|
@ -23,11 +23,20 @@ properties:
|
|||
enum:
|
||||
- validation
|
||||
- unauthorized
|
||||
- token_expired
|
||||
- forbidden
|
||||
- insufficient_scope
|
||||
- not_found
|
||||
- method_not_allowed
|
||||
- setup_complete
|
||||
- password_managed_externally
|
||||
- payload_too_large
|
||||
- rate_limited
|
||||
- unavailable
|
||||
- internal
|
||||
referenceId:
|
||||
type: string
|
||||
description: Present on internal errors. Quote it when reporting a problem; it tags the server's log lines for this request.
|
||||
errors:
|
||||
type: array
|
||||
description: Per-field failures. Present only when `code` is `validation`.
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ json:
|
|||
x-module: core
|
||||
x-stability-level: alpha
|
||||
tags: [server]
|
||||
security: []
|
||||
summary: Get the OpenAPI document (JSON)
|
||||
description: The bundled OpenAPI document of the running server version. Supports ETag revalidation.
|
||||
responses:
|
||||
|
|
@ -25,6 +26,7 @@ yaml:
|
|||
x-module: core
|
||||
x-stability-level: alpha
|
||||
tags: [server]
|
||||
security: []
|
||||
summary: Get the OpenAPI document (YAML)
|
||||
description: The bundled OpenAPI document of the running server version. Supports ETag revalidation.
|
||||
responses:
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ get:
|
|||
x-module: core
|
||||
x-stability-level: alpha
|
||||
tags: [server]
|
||||
security: []
|
||||
summary: Describe the server
|
||||
description: |
|
||||
Returns the public server description. No authentication required.
|
||||
|
|
|
|||
|
|
@ -7,26 +7,54 @@ import (
|
|||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/getkin/kin-openapi/openapi3"
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/navidrome/navidrome/api"
|
||||
"github.com/navidrome/navidrome/core/apiauth"
|
||||
"github.com/navidrome/navidrome/log"
|
||||
"github.com/navidrome/navidrome/model"
|
||||
)
|
||||
|
||||
const maxBodyBytes = 1 << 20
|
||||
|
||||
type Router struct {
|
||||
http.Handler
|
||||
ds model.DataStore
|
||||
ds model.DataStore
|
||||
auth *apiauth.Service
|
||||
}
|
||||
|
||||
func New(ds model.DataStore) *Router {
|
||||
rt := &Router{ds: ds}
|
||||
rt := &Router{ds: ds, auth: apiauth.New(ds)}
|
||||
rt.Handler = rt.routes()
|
||||
return rt
|
||||
}
|
||||
|
||||
var gateRulesV1 = gateRules{
|
||||
limited: map[string]bool{"login": true, "setupFirstAdmin": true, "changePassword": true},
|
||||
noScope: map[string]bool{"getCapabilities": true},
|
||||
grantOps: map[string]bool{"createAccessToken": true},
|
||||
}
|
||||
|
||||
func limitBody(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Body != nil {
|
||||
r.Body = http.MaxBytesReader(w, r.Body, maxBodyBytes)
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func (rt *Router) routes() http.Handler {
|
||||
r := chi.NewRouter()
|
||||
r.Use(problemRecoverer, headAsGet(r))
|
||||
doc, err := openapi3.NewLoader().LoadFromData(api.SpecJSON())
|
||||
if err != nil {
|
||||
log.Fatal("API v1: cannot load the embedded OpenAPI spec", err)
|
||||
}
|
||||
g, err := newGate(doc, r, rt.auth, gateRulesV1)
|
||||
if err != nil {
|
||||
log.Fatal("API v1: the embedded OpenAPI spec breaks the security rules", err)
|
||||
}
|
||||
r.Use(referenceIDMiddleware, problemRecoverer, headAsGet(r), limitBody, g.handler)
|
||||
r.NotFound(func(w http.ResponseWriter, req *http.Request) {
|
||||
writeProblemStatus(w, req, http.StatusNotFound, ProblemCodeNotFound, "no such endpoint")
|
||||
})
|
||||
|
|
@ -40,11 +68,19 @@ func (rt *Router) routes() http.Handler {
|
|||
|
||||
strict := NewStrictHandlerWithOptions(rt, nil, StrictHTTPServerOptions{
|
||||
RequestErrorHandlerFunc: func(w http.ResponseWriter, req *http.Request, err error) {
|
||||
writeProblemStatus(w, req, http.StatusBadRequest, "validation", err.Error())
|
||||
var tooLarge *http.MaxBytesError
|
||||
if errors.As(err, &tooLarge) {
|
||||
writeProblemStatus(w, req, http.StatusRequestEntityTooLarge, ProblemCodePayloadTooLarge, "request body too large")
|
||||
return
|
||||
}
|
||||
writeProblemStatus(w, req, http.StatusBadRequest, ProblemCodeValidation, "request body is not valid JSON")
|
||||
},
|
||||
ResponseErrorHandlerFunc: writeProblem,
|
||||
})
|
||||
HandlerWithOptions(strict, ChiServerOptions{BaseRouter: r, ErrorHandlerFunc: bindingErrorHandler})
|
||||
if err := g.checkRoutes(); err != nil {
|
||||
log.Fatal("API v1: routes and the embedded OpenAPI spec disagree", err)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -15,13 +15,19 @@ import (
|
|||
|
||||
// Defines values for ProblemCode.
|
||||
const (
|
||||
ProblemCodeForbidden ProblemCode = "forbidden"
|
||||
ProblemCodeInternal ProblemCode = "internal"
|
||||
ProblemCodeMethodNotAllowed ProblemCode = "method_not_allowed"
|
||||
ProblemCodeNotFound ProblemCode = "not_found"
|
||||
ProblemCodeUnauthorized ProblemCode = "unauthorized"
|
||||
ProblemCodeUnavailable ProblemCode = "unavailable"
|
||||
ProblemCodeValidation ProblemCode = "validation"
|
||||
ProblemCodeForbidden ProblemCode = "forbidden"
|
||||
ProblemCodeInsufficientScope ProblemCode = "insufficient_scope"
|
||||
ProblemCodeInternal ProblemCode = "internal"
|
||||
ProblemCodeMethodNotAllowed ProblemCode = "method_not_allowed"
|
||||
ProblemCodeNotFound ProblemCode = "not_found"
|
||||
ProblemCodePasswordManagedExternally ProblemCode = "password_managed_externally"
|
||||
ProblemCodePayloadTooLarge ProblemCode = "payload_too_large"
|
||||
ProblemCodeRateLimited ProblemCode = "rate_limited"
|
||||
ProblemCodeSetupComplete ProblemCode = "setup_complete"
|
||||
ProblemCodeTokenExpired ProblemCode = "token_expired"
|
||||
ProblemCodeUnauthorized ProblemCode = "unauthorized"
|
||||
ProblemCodeUnavailable ProblemCode = "unavailable"
|
||||
ProblemCodeValidation ProblemCode = "validation"
|
||||
)
|
||||
|
||||
// Valid indicates whether the value is a known member of the ProblemCode enum.
|
||||
|
|
@ -29,12 +35,24 @@ func (e ProblemCode) Valid() bool {
|
|||
switch e {
|
||||
case ProblemCodeForbidden:
|
||||
return true
|
||||
case ProblemCodeInsufficientScope:
|
||||
return true
|
||||
case ProblemCodeInternal:
|
||||
return true
|
||||
case ProblemCodeMethodNotAllowed:
|
||||
return true
|
||||
case ProblemCodeNotFound:
|
||||
return true
|
||||
case ProblemCodePasswordManagedExternally:
|
||||
return true
|
||||
case ProblemCodePayloadTooLarge:
|
||||
return true
|
||||
case ProblemCodeRateLimited:
|
||||
return true
|
||||
case ProblemCodeSetupComplete:
|
||||
return true
|
||||
case ProblemCodeTokenExpired:
|
||||
return true
|
||||
case ProblemCodeUnauthorized:
|
||||
return true
|
||||
case ProblemCodeUnavailable:
|
||||
|
|
@ -72,6 +90,9 @@ type Problem struct {
|
|||
// Errors Per-field failures. Present only when `code` is `validation`.
|
||||
Errors *[]ValidationError `json:"errors,omitempty"`
|
||||
|
||||
// ReferenceId Present on internal errors. Quote it when reporting a problem; it tags the server's log lines for this request.
|
||||
ReferenceId *string `json:"referenceId,omitempty"`
|
||||
|
||||
// Status HTTP status code of this response.
|
||||
Status int `json:"status"`
|
||||
|
||||
|
|
|
|||
|
|
@ -67,6 +67,19 @@ var _ = Describe("Router", func() {
|
|||
Expect(p.Detail).To(BeNil())
|
||||
})
|
||||
|
||||
It("tags internal errors with a referenceId that is also on the request's log lines", func() {
|
||||
logs := captureLogs()
|
||||
w := httptest.NewRecorder()
|
||||
h := referenceIDMiddleware(problemRecoverer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
panic("kaboom")
|
||||
})))
|
||||
h.ServeHTTP(w, httptest.NewRequestWithContext(GinkgoT().Context(), http.MethodGet, "/boom", nil))
|
||||
p := decodeProblem(w)
|
||||
Expect(p.ReferenceId).ToNot(BeNil())
|
||||
Expect(*p.ReferenceId).To(MatchRegexp(`^[0-9A-Za-z]{22}$`))
|
||||
Expect(logs.String()).To(ContainSubstring(*p.ReferenceId))
|
||||
})
|
||||
|
||||
It("re-panics http.ErrAbortHandler so the server can drop the connection", func() {
|
||||
Expect(func() {
|
||||
panicking(http.ErrAbortHandler).ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/abort", nil))
|
||||
|
|
|
|||
341
server/apiv1/gate.go
Normal file
341
server/apiv1/gate.go
Normal file
|
|
@ -0,0 +1,341 @@
|
|||
package apiv1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
"github.com/getkin/kin-openapi/openapi3"
|
||||
"github.com/getkin/kin-openapi/openapi3filter"
|
||||
"github.com/getkin/kin-openapi/routers"
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/go-chi/httprate"
|
||||
"github.com/navidrome/navidrome/conf"
|
||||
"github.com/navidrome/navidrome/core/apiauth"
|
||||
"github.com/navidrome/navidrome/log"
|
||||
"github.com/navidrome/navidrome/model"
|
||||
"github.com/navidrome/navidrome/model/request"
|
||||
"github.com/navidrome/navidrome/server"
|
||||
)
|
||||
|
||||
type authenticator interface {
|
||||
Authenticate(ctx context.Context, token, ip string) (*apiauth.Principal, error)
|
||||
ResolveGrant(ctx context.Context, secret, ip string) (*apiauth.Principal, error)
|
||||
}
|
||||
|
||||
type authKind int
|
||||
|
||||
const (
|
||||
authPublic authKind = iota
|
||||
authToken
|
||||
authGrant
|
||||
)
|
||||
|
||||
type gateOp struct {
|
||||
id string
|
||||
route *routers.Route
|
||||
kind authKind
|
||||
scope string
|
||||
limited bool
|
||||
}
|
||||
|
||||
type gate struct {
|
||||
mux chi.Routes
|
||||
ops map[string]*gateOp
|
||||
auth authenticator
|
||||
limiter func(http.Handler) http.Handler
|
||||
}
|
||||
|
||||
// Modules that ride another module's scope; every other module's scope is its own name.
|
||||
var moduleScope = map[string]string{
|
||||
"core": apiauth.ScopeRead,
|
||||
"transcoding": "streaming",
|
||||
"custom-tags": apiauth.ScopeRead,
|
||||
"grouping": apiauth.ScopeRead,
|
||||
"smart-playlists": "playlists:write",
|
||||
}
|
||||
|
||||
type gateRules struct {
|
||||
limited map[string]bool // login-type operations, throttled per client IP
|
||||
noScope map[string]bool // the only token operations allowed without x-scope
|
||||
grantOps map[string]bool // the only operations allowed to use grantAuth
|
||||
}
|
||||
|
||||
func newGate(doc *openapi3.T, mux chi.Routes, auth authenticator, rules gateRules) (*gate, error) {
|
||||
g := &gate{mux: mux, ops: map[string]*gateOp{}, auth: auth}
|
||||
for path, item := range doc.Paths.Map() {
|
||||
for method, op := range item.Operations() {
|
||||
gop, err := buildGateOp(doc, path, item, method, op, rules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
gop.limited = rules.limited[op.OperationID]
|
||||
g.ops[method+" "+path] = gop
|
||||
}
|
||||
}
|
||||
if conf.Server.AuthRequestLimit > 0 {
|
||||
g.limiter = httprate.LimitBy(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength,
|
||||
func(r *http.Request) (string, error) { return server.ClientIP(r), nil },
|
||||
httprate.WithLimitHandler(func(w http.ResponseWriter, r *http.Request) {
|
||||
writeProblemStatus(w, r, http.StatusTooManyRequests, ProblemCodeRateLimited, "too many requests")
|
||||
}))
|
||||
}
|
||||
return g, nil
|
||||
}
|
||||
|
||||
// buildGateOp enforces the allowed security forms, so a spec edit cannot silently drop a requirement.
|
||||
func buildGateOp(doc *openapi3.T, path string, item *openapi3.PathItem, method string, op *openapi3.Operation, rules gateRules) (*gateOp, error) {
|
||||
id := op.OperationID
|
||||
gop := &gateOp{id: id, route: &routers.Route{Spec: doc, Path: path, PathItem: item, Method: method, Operation: op}}
|
||||
if op.Security == nil {
|
||||
return nil, fmt.Errorf("operation %s must declare security explicitly", id)
|
||||
}
|
||||
rawScope, hasScope := op.Extensions["x-scope"]
|
||||
scope, isString := rawScope.(string)
|
||||
if hasScope && (!isString || scope == "") {
|
||||
return nil, fmt.Errorf("operation %s: x-scope must be a non-empty string", id)
|
||||
}
|
||||
module, _ := op.Extensions["x-module"].(string)
|
||||
switch reqs := *op.Security; {
|
||||
case len(reqs) == 0:
|
||||
gop.kind = authPublic
|
||||
case len(reqs) == 1 && isScheme(reqs[0], "bearerAuth"):
|
||||
gop.kind = authToken
|
||||
case len(reqs) == 1 && isScheme(reqs[0], "grantAuth") && rules.grantOps[id]:
|
||||
gop.kind = authGrant
|
||||
default:
|
||||
return nil, fmt.Errorf("operation %s has a security requirement outside the allowed forms", id)
|
||||
}
|
||||
if gop.kind == authToken && scope == "" && !rules.noScope[id] {
|
||||
return nil, fmt.Errorf("operation %s: bearerAuth needs x-scope", id)
|
||||
}
|
||||
if scope != "" {
|
||||
if gop.kind != authToken {
|
||||
return nil, fmt.Errorf("operation %s: x-scope needs bearerAuth", id)
|
||||
}
|
||||
base := module
|
||||
if s, ok := moduleScope[module]; ok {
|
||||
base = s
|
||||
}
|
||||
if scope != base && scope != base+":write" {
|
||||
return nil, fmt.Errorf("operation %s: x-scope %q does not match module %q", op.OperationID, scope, module)
|
||||
}
|
||||
if !slices.Contains(apiauth.KnownScopes, scope) && scope != apiauth.ScopeAdmin {
|
||||
return nil, fmt.Errorf("operation %s: unknown x-scope %q", op.OperationID, scope)
|
||||
}
|
||||
}
|
||||
gop.scope = scope
|
||||
return gop, nil
|
||||
}
|
||||
|
||||
// isScheme requires the scheme alone with an empty scope list, as OpenAPI 3.0.3 demands for http schemes.
|
||||
func isScheme(req openapi3.SecurityRequirement, name string) bool {
|
||||
scopes, ok := req[name]
|
||||
return ok && len(req) == 1 && len(scopes) == 0
|
||||
}
|
||||
|
||||
// checkRoutes fails when a routed pattern has no spec operation or a spec operation has no route.
|
||||
func (g *gate) checkRoutes() error {
|
||||
err := chi.Walk(g.mux, func(method, route string, _ http.Handler, _ ...func(http.Handler) http.Handler) error {
|
||||
if _, ok := g.ops[method+" "+route]; !ok {
|
||||
return fmt.Errorf("route %s %s is not in the spec", method, route)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, op := range g.ops {
|
||||
if g.mux.Find(chi.NewRouteContext(), op.route.Method, op.route.Path) != op.route.Path {
|
||||
return fmt.Errorf("spec operation %s (%s %s) has no route", op.id, op.route.Method, op.route.Path)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (g *gate) handler(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
path := routePath(r)
|
||||
method := r.Method
|
||||
if method == http.MethodHead && !g.mux.Match(chi.NewRouteContext(), http.MethodHead, path) {
|
||||
method = http.MethodGet
|
||||
}
|
||||
rctx := chi.NewRouteContext()
|
||||
pattern := g.mux.Find(rctx, method, path)
|
||||
if pattern == "" {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
op, ok := g.ops[method+" "+pattern]
|
||||
if !ok {
|
||||
log.Error(r.Context(), "API v1: routed pattern missing from the spec", "method", method, "pattern", pattern)
|
||||
writeProblemStatus(w, r, http.StatusInternalServerError, ProblemCodeInternal, "")
|
||||
return
|
||||
}
|
||||
serve := func(w http.ResponseWriter, r *http.Request) {
|
||||
r, ok := g.authorize(w, r, op)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if !g.validate(w, r, op, rctx) {
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
}
|
||||
if op.limited && g.limiter != nil {
|
||||
g.limiter(http.HandlerFunc(serve)).ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
serve(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func (g *gate) authorize(w http.ResponseWriter, r *http.Request, op *gateOp) (*http.Request, bool) {
|
||||
if op.kind == authPublic {
|
||||
return r, true
|
||||
}
|
||||
token, ok := bearerToken(r)
|
||||
if !ok {
|
||||
w.Header().Set("WWW-Authenticate", "Bearer")
|
||||
writeProblemStatus(w, r, http.StatusUnauthorized, ProblemCodeUnauthorized, "")
|
||||
return r, false
|
||||
}
|
||||
ip := server.ClientIP(r)
|
||||
var p *apiauth.Principal
|
||||
var err error
|
||||
if op.kind == authGrant {
|
||||
p, err = g.auth.ResolveGrant(r.Context(), token, ip)
|
||||
} else {
|
||||
p, err = g.auth.Authenticate(r.Context(), token, ip)
|
||||
}
|
||||
switch {
|
||||
case errors.Is(err, apiauth.ErrTokenExpired), errors.Is(err, model.ErrInvalidAuth):
|
||||
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
|
||||
writeProblem(w, r, err)
|
||||
return r, false
|
||||
case errors.Is(err, apiauth.ErrInsufficientScope):
|
||||
w.Header().Set("WWW-Authenticate", `Bearer error="insufficient_scope"`)
|
||||
writeProblem(w, r, err)
|
||||
return r, false
|
||||
case err != nil:
|
||||
writeProblem(w, r, err)
|
||||
return r, false
|
||||
}
|
||||
if op.scope != "" && !apiauth.Satisfies(p.Scopes, op.scope) {
|
||||
w.Header().Set("WWW-Authenticate", fmt.Sprintf(`Bearer error="insufficient_scope", scope=%q`, op.scope))
|
||||
writeProblem(w, r, apiauth.ErrInsufficientScope)
|
||||
return r, false
|
||||
}
|
||||
ctx := apiauth.WithPrincipal(request.WithUser(r.Context(), p.User), p)
|
||||
return r.WithContext(ctx), true
|
||||
}
|
||||
|
||||
func bearerToken(r *http.Request) (string, bool) {
|
||||
scheme, token, ok := strings.Cut(strings.TrimSpace(r.Header.Get("Authorization")), " ")
|
||||
token = strings.TrimSpace(token)
|
||||
if !ok || !strings.EqualFold(scheme, "Bearer") || token == "" {
|
||||
return "", false
|
||||
}
|
||||
return token, true
|
||||
}
|
||||
|
||||
func (g *gate) validate(w http.ResponseWriter, r *http.Request, op *gateOp, rctx *chi.Context) bool {
|
||||
params := map[string]string{}
|
||||
for i, k := range rctx.URLParams.Keys {
|
||||
params[k] = rctx.URLParams.Values[i]
|
||||
}
|
||||
err := openapi3filter.ValidateRequest(r.Context(), &openapi3filter.RequestValidationInput{
|
||||
Request: r, PathParams: params, Route: op.route,
|
||||
Options: &openapi3filter.Options{AuthenticationFunc: openapi3filter.NoopAuthenticationFunc, MultiError: true},
|
||||
})
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
var tooLarge *http.MaxBytesError
|
||||
if errors.As(err, &tooLarge) {
|
||||
writeProblemStatus(w, r, http.StatusRequestEntityTooLarge, ProblemCodePayloadTooLarge, "request body too large")
|
||||
return false
|
||||
}
|
||||
fields := sanitizeValidation(err)
|
||||
log.Debug(r.Context(), "API v1: request failed validation", "operation", op.id, "errors", fields)
|
||||
writeProblemStatus(w, r, http.StatusBadRequest, ProblemCodeValidation, "the request does not match the API schema", fields...)
|
||||
return false
|
||||
}
|
||||
|
||||
var missingProperty = regexp.MustCompile(`property "([^"]+)" is missing`)
|
||||
|
||||
// sanitizeValidation keeps only field paths and fixed messages: kin-openapi errors can embed the submitted value.
|
||||
// It walks wrappers by concrete type, not errors.As, because MultiError.As would skip the RequestError that names the parameter.
|
||||
func sanitizeValidation(err error) []ValidationError {
|
||||
var out []ValidationError
|
||||
var walk func(err error, param string)
|
||||
walk = func(err error, param string) {
|
||||
switch e := err.(type) { //nolint:errorlint
|
||||
case openapi3.MultiError:
|
||||
for _, child := range e {
|
||||
walk(child, param)
|
||||
}
|
||||
case *openapi3filter.RequestError:
|
||||
if e.Parameter != nil {
|
||||
param = e.Parameter.Name
|
||||
}
|
||||
switch {
|
||||
case errors.Is(e.Err, openapi3filter.ErrInvalidRequired), errors.Is(e.Err, openapi3filter.ErrInvalidEmptyValue):
|
||||
out = append(out, ValidationError{Field: param, Message: "is required"})
|
||||
case e.Err != nil:
|
||||
walk(e.Err, param)
|
||||
default:
|
||||
out = append(out, ValidationError{Field: param, Message: "is invalid"})
|
||||
}
|
||||
case *openapi3.SchemaError:
|
||||
field := strings.Join(e.JSONPointer(), ".")
|
||||
if field == "" && e.SchemaField == "required" {
|
||||
if m := missingProperty.FindStringSubmatch(e.Reason); m != nil {
|
||||
field = m[1]
|
||||
}
|
||||
}
|
||||
switch {
|
||||
case param != "" && field != "":
|
||||
field = param + "." + field
|
||||
case field == "":
|
||||
field = param
|
||||
}
|
||||
out = append(out, ValidationError{Field: field, Message: schemaMessage(e.SchemaField)})
|
||||
default:
|
||||
if inner := errors.Unwrap(err); inner != nil {
|
||||
walk(inner, param)
|
||||
return
|
||||
}
|
||||
out = append(out, ValidationError{Field: param, Message: "is invalid"})
|
||||
}
|
||||
}
|
||||
walk(err, "")
|
||||
return out
|
||||
}
|
||||
|
||||
func schemaMessage(keyword string) string {
|
||||
switch keyword {
|
||||
case "required":
|
||||
return "is required"
|
||||
case "maxLength", "maxItems":
|
||||
return "is too long"
|
||||
case "minLength", "minItems":
|
||||
return "is too short"
|
||||
case "maximum", "exclusiveMaximum":
|
||||
return "is too large"
|
||||
case "minimum", "exclusiveMinimum":
|
||||
return "is too small"
|
||||
case "pattern", "format":
|
||||
return "has an invalid format"
|
||||
case "enum":
|
||||
return "is not an allowed value"
|
||||
case "type", "nullable":
|
||||
return "has the wrong type"
|
||||
default:
|
||||
return "is invalid"
|
||||
}
|
||||
}
|
||||
315
server/apiv1/gate_test.go
Normal file
315
server/apiv1/gate_test.go
Normal file
|
|
@ -0,0 +1,315 @@
|
|||
package apiv1
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/getkin/kin-openapi/openapi3"
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/navidrome/navidrome/conf"
|
||||
"github.com/navidrome/navidrome/conf/configtest"
|
||||
"github.com/navidrome/navidrome/core/apiauth"
|
||||
"github.com/navidrome/navidrome/log"
|
||||
"github.com/navidrome/navidrome/model"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
const gateSpec = `
|
||||
openapi: 3.0.3
|
||||
info: {title: t, version: "1"}
|
||||
paths:
|
||||
/open:
|
||||
get: {operationId: open, x-module: core, security: [], responses: {'200': {description: ok}}}
|
||||
/things/{id}:
|
||||
get:
|
||||
operationId: getThing
|
||||
x-module: core
|
||||
x-scope: read
|
||||
security: [{bearerAuth: []}]
|
||||
parameters: [{name: id, in: path, required: true, schema: {type: string, maxLength: 3}}]
|
||||
responses: {'200': {description: ok}}
|
||||
/things:
|
||||
post:
|
||||
operationId: createThing
|
||||
x-module: password
|
||||
x-scope: password
|
||||
security: [{bearerAuth: []}]
|
||||
requestBody:
|
||||
required: true
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
type: object
|
||||
required: [name]
|
||||
properties: {name: {type: string, maxLength: 5}}
|
||||
responses: {'200': {description: ok}}
|
||||
/caps:
|
||||
get: {operationId: caps, x-module: core, security: [{bearerAuth: []}], responses: {'200': {description: ok}}}
|
||||
/mint:
|
||||
post: {operationId: mint, x-module: core, security: [{grantAuth: []}], responses: {'200': {description: ok}}}
|
||||
/limited:
|
||||
post: {operationId: limited, x-module: core, security: [], responses: {'200': {description: ok}}}
|
||||
components:
|
||||
securitySchemes:
|
||||
bearerAuth: {type: http, scheme: bearer}
|
||||
grantAuth: {type: http, scheme: bearer}
|
||||
`
|
||||
|
||||
type fakeAuth struct {
|
||||
principal *apiauth.Principal
|
||||
err error
|
||||
gotToken string
|
||||
gotSecret string
|
||||
}
|
||||
|
||||
func (f *fakeAuth) Authenticate(_ context.Context, token, _ string) (*apiauth.Principal, error) {
|
||||
f.gotToken = token
|
||||
return f.principal, f.err
|
||||
}
|
||||
|
||||
func (f *fakeAuth) ResolveGrant(_ context.Context, secret, _ string) (*apiauth.Principal, error) {
|
||||
f.gotSecret = secret
|
||||
return f.principal, f.err
|
||||
}
|
||||
|
||||
var testGateRules = gateRules{
|
||||
limited: map[string]bool{"limited": true},
|
||||
noScope: map[string]bool{"caps": true},
|
||||
grantOps: map[string]bool{"mint": true},
|
||||
}
|
||||
|
||||
var _ = Describe("spec gate", func() {
|
||||
var ctx context.Context
|
||||
var fa *fakeAuth
|
||||
var mux *chi.Mux
|
||||
var g *gate
|
||||
var reached string
|
||||
|
||||
build := func(spec string) (*chi.Mux, error) {
|
||||
doc, err := openapi3.NewLoader().LoadFromData([]byte(spec))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
m := chi.NewRouter()
|
||||
g, err = newGate(doc, m, fa, testGateRules)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m.Use(headAsGet(m), g.handler)
|
||||
ok := func(name string) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
reached = name
|
||||
if p, found := apiauth.PrincipalFrom(r.Context()); found {
|
||||
w.Header().Set("X-User", p.User.ID)
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}
|
||||
}
|
||||
m.Get("/open", ok("open"))
|
||||
m.Get("/things/{id}", ok("getThing"))
|
||||
m.Post("/things", ok("createThing"))
|
||||
m.Get("/caps", ok("caps"))
|
||||
m.Post("/mint", ok("mint"))
|
||||
m.Post("/limited", ok("limited"))
|
||||
return m, nil
|
||||
}
|
||||
|
||||
do := func(method, path, auth, body string) *httptest.ResponseRecorder {
|
||||
var req *http.Request
|
||||
if body != "" {
|
||||
req = httptest.NewRequestWithContext(ctx, method, path, strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
} else {
|
||||
req = httptest.NewRequestWithContext(ctx, method, path, nil)
|
||||
}
|
||||
if auth != "" {
|
||||
req.Header.Set("Authorization", auth)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
mux.ServeHTTP(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
BeforeEach(func() {
|
||||
ctx = GinkgoT().Context()
|
||||
reached = ""
|
||||
fa = &fakeAuth{principal: &apiauth.Principal{User: model.User{ID: "u1"}, GrantID: "g1", Scopes: []string{"read"}}}
|
||||
var err error
|
||||
mux, err = build(gateSpec)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
})
|
||||
|
||||
It("lets public operations through without a token", func() {
|
||||
Expect(do(http.MethodGet, "/open", "", "").Code).To(Equal(http.StatusOK))
|
||||
Expect(reached).To(Equal("open"))
|
||||
})
|
||||
|
||||
It("requires a token, with a Bearer challenge", func() {
|
||||
w := do(http.MethodGet, "/things/1", "", "")
|
||||
Expect(w.Code).To(Equal(http.StatusUnauthorized))
|
||||
Expect(w.Header().Get("WWW-Authenticate")).To(Equal("Bearer"))
|
||||
Expect(decodeProblem(w).Code).To(Equal(ProblemCodeUnauthorized))
|
||||
Expect(reached).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("accepts the Bearer scheme in any case and trims spaces", func() {
|
||||
w := do(http.MethodGet, "/things/1", "bearer tok-1 ", "")
|
||||
Expect(w.Code).To(Equal(http.StatusOK))
|
||||
Expect(fa.gotToken).To(Equal("tok-1"))
|
||||
Expect(w.Header().Get("X-User")).To(Equal("u1"))
|
||||
})
|
||||
|
||||
It("maps an expired token to token_expired", func() {
|
||||
fa.err = apiauth.ErrTokenExpired
|
||||
w := do(http.MethodGet, "/things/1", "Bearer x", "")
|
||||
Expect(w.Code).To(Equal(http.StatusUnauthorized))
|
||||
Expect(w.Header().Get("WWW-Authenticate")).To(Equal(`Bearer error="invalid_token"`))
|
||||
Expect(decodeProblem(w).Code).To(Equal(ProblemCodeTokenExpired))
|
||||
})
|
||||
|
||||
It("maps other auth failures to unauthorized with invalid_token", func() {
|
||||
fa.err = model.ErrInvalidAuth
|
||||
w := do(http.MethodGet, "/things/1", "Bearer x", "")
|
||||
Expect(w.Code).To(Equal(http.StatusUnauthorized))
|
||||
Expect(w.Header().Get("WWW-Authenticate")).To(Equal(`Bearer error="invalid_token"`))
|
||||
Expect(decodeProblem(w).Code).To(Equal(ProblemCodeUnauthorized))
|
||||
})
|
||||
|
||||
It("rejects a token without the operation's scope", func() {
|
||||
w := do(http.MethodPost, "/things", "Bearer x", `{"name":"a"}`)
|
||||
Expect(w.Code).To(Equal(http.StatusForbidden))
|
||||
Expect(w.Header().Get("WWW-Authenticate")).To(Equal(`Bearer error="insufficient_scope", scope="password"`))
|
||||
Expect(decodeProblem(w).Code).To(Equal(ProblemCodeInsufficientScope))
|
||||
})
|
||||
|
||||
It("lets any valid token through an operation with no x-scope", func() {
|
||||
fa.principal.Scopes = nil
|
||||
Expect(do(http.MethodGet, "/caps", "Bearer x", "").Code).To(Equal(http.StatusOK))
|
||||
})
|
||||
|
||||
It("uses ResolveGrant for grantAuth operations", func() {
|
||||
Expect(do(http.MethodPost, "/mint", "Bearer ndg_secret", "").Code).To(Equal(http.StatusOK))
|
||||
Expect(fa.gotSecret).To(Equal("ndg_secret"))
|
||||
Expect(fa.gotToken).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("checks HEAD on a protected GET", func() {
|
||||
w := do(http.MethodHead, "/things/1", "", "")
|
||||
Expect(w.Code).To(Equal(http.StatusUnauthorized))
|
||||
})
|
||||
|
||||
It("turns an insufficient-scope error from Authenticate into a 403 challenge", func() {
|
||||
fa.err = apiauth.ErrInsufficientScope // e.g. a token carrying admin after demotion
|
||||
w := do(http.MethodGet, "/caps", "Bearer x", "")
|
||||
Expect(w.Code).To(Equal(http.StatusForbidden))
|
||||
Expect(w.Header().Get("WWW-Authenticate")).To(Equal(`Bearer error="insufficient_scope"`))
|
||||
Expect(decodeProblem(w).Code).To(Equal(ProblemCodeInsufficientScope))
|
||||
})
|
||||
|
||||
It("works when mounted under a base path", func() {
|
||||
root := chi.NewRouter()
|
||||
root.Mount("/music/api/v1", mux)
|
||||
req := httptest.NewRequestWithContext(ctx, http.MethodGet, "/music/api/v1/things/1", nil)
|
||||
w := httptest.NewRecorder()
|
||||
root.ServeHTTP(w, req)
|
||||
Expect(w.Code).To(Equal(http.StatusUnauthorized))
|
||||
})
|
||||
|
||||
It("authenticates before validating", func() {
|
||||
w := do(http.MethodPost, "/things", "", `{"name":"far-too-long"}`)
|
||||
Expect(w.Code).To(Equal(http.StatusUnauthorized))
|
||||
})
|
||||
|
||||
It("returns and logs sanitised validation errors that never echo the value", func() {
|
||||
logs := captureLogs()
|
||||
fa.principal.Scopes = []string{"password"}
|
||||
w := do(http.MethodPost, "/things", "Bearer x", `{"name":"hunter2-secret"}`)
|
||||
Expect(w.Code).To(Equal(http.StatusBadRequest))
|
||||
p := decodeProblem(w)
|
||||
Expect(p.Code).To(Equal(ProblemCodeValidation))
|
||||
Expect(*p.Errors).To(ConsistOf(ValidationError{Field: "name", Message: "is too long"}))
|
||||
Expect(w.Body.String()).ToNot(ContainSubstring("hunter2"))
|
||||
Expect(logs.String()).To(ContainSubstring("failed validation"))
|
||||
Expect(logs.String()).ToNot(ContainSubstring("hunter2"))
|
||||
})
|
||||
|
||||
It("reports a missing required body field by name", func() {
|
||||
fa.principal.Scopes = []string{"password"}
|
||||
w := do(http.MethodPost, "/things", "Bearer x", `{}`)
|
||||
Expect(*decodeProblem(w).Errors).To(ConsistOf(ValidationError{Field: "name", Message: "is required"}))
|
||||
})
|
||||
|
||||
It("validates path parameters", func() {
|
||||
w := do(http.MethodGet, "/things/toolong", "Bearer x", "")
|
||||
Expect(w.Code).To(Equal(http.StatusBadRequest))
|
||||
Expect(*decodeProblem(w).Errors).To(ConsistOf(ValidationError{Field: "id", Message: "is too long"}))
|
||||
})
|
||||
|
||||
It("passes unknown paths through to the router's 404", func() {
|
||||
Expect(do(http.MethodGet, "/nope", "", "").Code).To(Equal(http.StatusNotFound))
|
||||
})
|
||||
|
||||
It("fails closed for a routed pattern the spec does not know", func() {
|
||||
mux.Get("/extra", func(w http.ResponseWriter, r *http.Request) { reached = "extra" })
|
||||
w := do(http.MethodGet, "/extra", "", "")
|
||||
Expect(w.Code).To(Equal(http.StatusInternalServerError))
|
||||
Expect(reached).To(BeEmpty())
|
||||
})
|
||||
|
||||
It("rate-limits the listed operations with a 429 problem", func() {
|
||||
DeferCleanup(configtest.SetupConfig())
|
||||
conf.Server.AuthRequestLimit = 1
|
||||
var err error
|
||||
mux, err = build(gateSpec)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(do(http.MethodPost, "/limited", "", "").Code).To(Equal(http.StatusOK))
|
||||
w := do(http.MethodPost, "/limited", "", "")
|
||||
Expect(w.Code).To(Equal(http.StatusTooManyRequests))
|
||||
Expect(w.Header().Get("Retry-After")).ToNot(BeEmpty())
|
||||
Expect(decodeProblem(w).Code).To(Equal(ProblemCodeRateLimited))
|
||||
})
|
||||
|
||||
DescribeTable("refuses specs that break the security rules",
|
||||
func(bad string) {
|
||||
_, err := build(bad)
|
||||
Expect(err).To(HaveOccurred())
|
||||
},
|
||||
Entry("missing security", strings.Replace(gateSpec, "operationId: open, x-module: core, security: [],", "operationId: open, x-module: core,", 1)),
|
||||
Entry("scope not matching module", strings.Replace(gateSpec, "x-scope: read", "x-scope: password", 1)),
|
||||
Entry("unknown scope", strings.Replace(gateSpec, "x-scope: read", "x-scope: bogus", 1)),
|
||||
Entry("bearer without x-scope outside the allowlist", strings.Replace(gateSpec, " x-scope: read\n", "", 1)),
|
||||
Entry("grantAuth outside the allowlist", strings.Replace(gateSpec, "operationId: limited, x-module: core, security: []", "operationId: limited, x-module: core, security: [{grantAuth: []}]", 1)),
|
||||
Entry("non-empty scope list on a bearer scheme", strings.Replace(gateSpec, "operationId: caps, x-module: core, security: [{bearerAuth: []}]", "operationId: caps, x-module: core, security: [{bearerAuth: [read]}]", 1)),
|
||||
Entry("x-scope on a public operation", strings.Replace(gateSpec, "operationId: open, x-module: core, security: [],", "operationId: open, x-module: core, x-scope: read, security: [],", 1)),
|
||||
Entry("x-scope that is not a string", strings.Replace(gateSpec, "x-scope: read", "x-scope: [read]", 1)),
|
||||
)
|
||||
|
||||
It("checks routes against the spec in both directions", func() {
|
||||
Expect(g.checkRoutes()).To(Succeed())
|
||||
|
||||
mux.Get("/extra", func(http.ResponseWriter, *http.Request) {})
|
||||
Expect(g.checkRoutes()).To(MatchError(ContainSubstring("GET /extra is not in the spec")))
|
||||
|
||||
extraOp := strings.Replace(gateSpec, "components:", ` /unrouted:
|
||||
get: {operationId: unrouted, x-module: core, security: [], responses: {'200': {description: ok}}}
|
||||
components:`, 1)
|
||||
_, err := build(extraOp)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(g.checkRoutes()).To(MatchError(ContainSubstring("unrouted")))
|
||||
})
|
||||
})
|
||||
|
||||
// captureLogs sends debug logs to a buffer for the rest of the spec.
|
||||
func captureLogs() *bytes.Buffer {
|
||||
buf := &bytes.Buffer{}
|
||||
log.SetOutput(buf)
|
||||
log.SetLevel(log.LevelDebug)
|
||||
DeferCleanup(func() {
|
||||
log.SetOutput(os.Stderr)
|
||||
log.SetLevel(log.LevelFatal)
|
||||
})
|
||||
return buf
|
||||
}
|
||||
|
|
@ -5,24 +5,69 @@ import (
|
|||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/navidrome/navidrome/core/apiauth"
|
||||
"github.com/navidrome/navidrome/core/auth"
|
||||
"github.com/navidrome/navidrome/log"
|
||||
"github.com/navidrome/navidrome/model"
|
||||
)
|
||||
|
||||
const problemContentType = "application/problem+json"
|
||||
|
||||
type clientError struct {
|
||||
err error
|
||||
detail string
|
||||
}
|
||||
|
||||
func (e *clientError) Error() string { return e.detail }
|
||||
func (e *clientError) Unwrap() error { return e.err }
|
||||
|
||||
// ClientError marks detail as safe to show clients; err still decides the status and code.
|
||||
func ClientError(err error, detail string) error {
|
||||
return &clientError{err: err, detail: detail}
|
||||
}
|
||||
|
||||
type fieldErrors struct {
|
||||
fields []ValidationError
|
||||
}
|
||||
|
||||
func (e *fieldErrors) Error() string { return "validation failed" }
|
||||
func (e *fieldErrors) Unwrap() error { return model.ErrValidation }
|
||||
|
||||
func validationFailed(fields ...ValidationError) error {
|
||||
return &fieldErrors{fields: fields}
|
||||
}
|
||||
|
||||
func writeProblem(w http.ResponseWriter, r *http.Request, err error) {
|
||||
status, code := classifyError(err)
|
||||
detail := err.Error()
|
||||
if status == http.StatusInternalServerError {
|
||||
log.Error(r.Context(), "API v1: unexpected error", "path", r.URL.Path, err)
|
||||
detail = ""
|
||||
writeProblemStatus(w, r, status, code, "")
|
||||
return
|
||||
}
|
||||
log.Debug(r.Context(), "API v1: request failed", "path", r.URL.Path, "status", status, "code", code, err)
|
||||
var detail string
|
||||
var ce *clientError
|
||||
if errors.As(err, &ce) {
|
||||
detail = ce.detail
|
||||
}
|
||||
var fe *fieldErrors
|
||||
if errors.As(err, &fe) {
|
||||
writeProblemStatus(w, r, status, code, detail, fe.fields...)
|
||||
return
|
||||
}
|
||||
writeProblemStatus(w, r, status, code, detail)
|
||||
}
|
||||
|
||||
func classifyError(err error) (int, ProblemCode) {
|
||||
switch {
|
||||
case errors.Is(err, apiauth.ErrTokenExpired):
|
||||
return http.StatusUnauthorized, ProblemCodeTokenExpired
|
||||
case errors.Is(err, apiauth.ErrInsufficientScope):
|
||||
return http.StatusForbidden, ProblemCodeInsufficientScope
|
||||
case errors.Is(err, auth.ErrSetupComplete):
|
||||
return http.StatusConflict, ProblemCodeSetupComplete
|
||||
case errors.Is(err, apiauth.ErrPasswordManagedExternally):
|
||||
return http.StatusConflict, ProblemCodePasswordManagedExternally
|
||||
case errors.Is(err, model.ErrNotFound):
|
||||
return http.StatusNotFound, ProblemCodeNotFound
|
||||
case errors.Is(err, model.ErrNotAuthorized):
|
||||
|
|
@ -45,6 +90,15 @@ func writeProblemStatus(w http.ResponseWriter, r *http.Request, status int, code
|
|||
if len(fieldErrors) > 0 {
|
||||
p.Errors = &fieldErrors
|
||||
}
|
||||
if status == http.StatusInternalServerError {
|
||||
if ref := referenceIDFrom(r.Context()); ref != "" {
|
||||
p.ReferenceId = &ref
|
||||
}
|
||||
}
|
||||
// Every 401 carries a Bearer challenge; callers may set a more specific one first.
|
||||
if status == http.StatusUnauthorized && w.Header().Get("WWW-Authenticate") == "" {
|
||||
w.Header().Set("WWW-Authenticate", "Bearer")
|
||||
}
|
||||
w.Header().Set("Content-Type", problemContentType)
|
||||
w.WriteHeader(status)
|
||||
if err := json.NewEncoder(w).Encode(p); err != nil {
|
||||
|
|
@ -53,20 +107,20 @@ func writeProblemStatus(w http.ResponseWriter, r *http.Request, status int, code
|
|||
}
|
||||
|
||||
func bindingErrorHandler(w http.ResponseWriter, r *http.Request, err error) {
|
||||
var fieldErrors []ValidationError
|
||||
var fieldErrs []ValidationError
|
||||
var required *RequiredParamError
|
||||
var invalid *InvalidParamFormatError
|
||||
var tooMany *TooManyValuesForParamError
|
||||
var unmarshal *UnmarshalingParamError
|
||||
switch {
|
||||
case errors.As(err, &required):
|
||||
fieldErrors = append(fieldErrors, ValidationError{Field: required.ParamName, Message: "is required"})
|
||||
fieldErrs = append(fieldErrs, ValidationError{Field: required.ParamName, Message: "is required"})
|
||||
case errors.As(err, &invalid):
|
||||
fieldErrors = append(fieldErrors, ValidationError{Field: invalid.ParamName, Message: invalid.Err.Error()})
|
||||
fieldErrs = append(fieldErrs, ValidationError{Field: invalid.ParamName, Message: "has an invalid value"})
|
||||
case errors.As(err, &tooMany):
|
||||
fieldErrors = append(fieldErrors, ValidationError{Field: tooMany.ParamName, Message: "expected a single value"})
|
||||
fieldErrs = append(fieldErrs, ValidationError{Field: tooMany.ParamName, Message: "expected a single value"})
|
||||
case errors.As(err, &unmarshal):
|
||||
fieldErrors = append(fieldErrors, ValidationError{Field: unmarshal.ParamName, Message: unmarshal.Err.Error()})
|
||||
fieldErrs = append(fieldErrs, ValidationError{Field: unmarshal.ParamName, Message: "has an invalid value"})
|
||||
}
|
||||
writeProblemStatus(w, r, http.StatusBadRequest, ProblemCodeValidation, err.Error(), fieldErrors...)
|
||||
writeProblemStatus(w, r, http.StatusBadRequest, ProblemCodeValidation, "invalid request parameters", fieldErrs...)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,8 @@ import (
|
|||
"net/http"
|
||||
"net/http/httptest"
|
||||
|
||||
"github.com/navidrome/navidrome/core/apiauth"
|
||||
"github.com/navidrome/navidrome/core/auth"
|
||||
"github.com/navidrome/navidrome/model"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
|
@ -46,20 +48,51 @@ var _ = Describe("problem", func() {
|
|||
Entry("expired", model.ErrExpired, http.StatusUnauthorized, ProblemCodeUnauthorized),
|
||||
Entry("validation", model.ErrValidation, http.StatusBadRequest, ProblemCodeValidation),
|
||||
Entry("not available", model.ErrNotAvailable, http.StatusServiceUnavailable, ProblemCodeUnavailable),
|
||||
Entry("token expired", apiauth.ErrTokenExpired, http.StatusUnauthorized, ProblemCodeTokenExpired),
|
||||
Entry("insufficient scope", apiauth.ErrInsufficientScope, http.StatusForbidden, ProblemCodeInsufficientScope),
|
||||
Entry("setup complete", auth.ErrSetupComplete, http.StatusConflict, ProblemCodeSetupComplete),
|
||||
Entry("password managed externally", apiauth.ErrPasswordManagedExternally, http.StatusConflict, ProblemCodePasswordManagedExternally),
|
||||
Entry("unknown", errors.New("boom"), http.StatusInternalServerError, ProblemCodeInternal),
|
||||
)
|
||||
|
||||
DescribeTable("keeps the wrapping context as detail for client errors",
|
||||
func(err error) {
|
||||
writeProblem(w, r, err)
|
||||
p := decodeProblem(w)
|
||||
Expect(p.Status).To(Equal(http.StatusNotFound))
|
||||
Expect(p.Detail).ToNot(BeNil())
|
||||
Expect(*p.Detail).To(ContainSubstring("album 123"))
|
||||
},
|
||||
Entry("fmt.Errorf %w", fmt.Errorf("album 123: %w", model.ErrNotFound)),
|
||||
Entry("errors.Join", errors.Join(errors.New("album 123"), model.ErrNotFound)),
|
||||
)
|
||||
It("shows detail only for errors marked as client-facing", func() {
|
||||
writeProblem(w, r, fmt.Errorf("album 123: %w", model.ErrNotFound))
|
||||
Expect(decodeProblem(w).Detail).To(BeNil())
|
||||
|
||||
w = httptest.NewRecorder()
|
||||
writeProblem(w, r, ClientError(model.ErrNotFound, "album not found"))
|
||||
p := decodeProblem(w)
|
||||
Expect(p.Code).To(Equal(ProblemCodeNotFound))
|
||||
Expect(*p.Detail).To(Equal("album not found"))
|
||||
})
|
||||
|
||||
It("writes field errors from validationFailed", func() {
|
||||
writeProblem(w, r, validationFailed(ValidationError{Field: "currentPassword", Message: "is incorrect"}))
|
||||
p := decodeProblem(w)
|
||||
Expect(w.Code).To(Equal(http.StatusBadRequest))
|
||||
Expect(p.Code).To(Equal(ProblemCodeValidation))
|
||||
Expect(*p.Errors).To(ConsistOf(ValidationError{Field: "currentPassword", Message: "is incorrect"}))
|
||||
})
|
||||
|
||||
It("adds a Bearer challenge to every 401 unless one is already set", func() {
|
||||
writeProblem(w, r, model.ErrInvalidAuth)
|
||||
Expect(w.Header().Get("WWW-Authenticate")).To(Equal("Bearer"))
|
||||
|
||||
w = httptest.NewRecorder()
|
||||
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token"`)
|
||||
writeProblem(w, r, model.ErrInvalidAuth)
|
||||
Expect(w.Header().Get("WWW-Authenticate")).To(Equal(`Bearer error="invalid_token"`))
|
||||
})
|
||||
|
||||
It("adds the request's referenceId to internal errors only", func() {
|
||||
r = r.WithContext(withReferenceID(r.Context(), "ref-123"))
|
||||
writeProblem(w, r, errors.New("boom"))
|
||||
Expect(*decodeProblem(w).ReferenceId).To(Equal("ref-123"))
|
||||
|
||||
w = httptest.NewRecorder()
|
||||
writeProblem(w, r, model.ErrNotFound)
|
||||
Expect(decodeProblem(w).ReferenceId).To(BeNil())
|
||||
})
|
||||
|
||||
It("hides details for internal errors", func() {
|
||||
writeProblem(w, r, errors.New("db password is hunter2"))
|
||||
|
|
@ -95,14 +128,19 @@ var _ = Describe("problem", func() {
|
|||
Expect(p.Errors).ToNot(BeNil())
|
||||
Expect(*p.Errors).To(HaveLen(1))
|
||||
Expect((*p.Errors)[0].Field).To(Equal(field))
|
||||
Expect((*p.Errors)[0].Message).To(ContainSubstring(message))
|
||||
Expect((*p.Errors)[0].Message).To(Equal(message))
|
||||
},
|
||||
Entry("required", &RequiredParamError{ParamName: "limit"}, "limit", "is required"),
|
||||
Entry("invalid format", &InvalidParamFormatError{ParamName: "offset", Err: errors.New("not a number")}, "offset", "not a number"),
|
||||
Entry("too many values", &TooManyValuesForParamError{ParamName: "sort", Count: 2}, "sort", "single value"),
|
||||
Entry("unmarshaling", &UnmarshalingParamError{ParamName: "ids", Err: errors.New("bad json")}, "ids", "bad json"),
|
||||
Entry("invalid format", &InvalidParamFormatError{ParamName: "offset", Err: errors.New(`parsing "abc": invalid syntax`)}, "offset", "has an invalid value"),
|
||||
Entry("too many values", &TooManyValuesForParamError{ParamName: "sort", Count: 2}, "sort", "expected a single value"),
|
||||
Entry("unmarshaling", &UnmarshalingParamError{ParamName: "ids", Err: errors.New("bad json")}, "ids", "has an invalid value"),
|
||||
)
|
||||
|
||||
It("never echoes the submitted value", func() {
|
||||
bindingErrorHandler(w, r, &InvalidParamFormatError{ParamName: "offset", Err: errors.New(`parsing "hunter2": invalid syntax`)})
|
||||
Expect(w.Body.String()).ToNot(ContainSubstring("hunter2"))
|
||||
})
|
||||
|
||||
It("still returns a validation problem for unknown binding errors", func() {
|
||||
bindingErrorHandler(w, r, errors.New("weird"))
|
||||
p := decodeProblem(w)
|
||||
|
|
|
|||
29
server/apiv1/reference.go
Normal file
29
server/apiv1/reference.go
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
package apiv1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/navidrome/navidrome/log"
|
||||
"github.com/navidrome/navidrome/model/id"
|
||||
)
|
||||
|
||||
type referenceIDKey struct{}
|
||||
|
||||
func withReferenceID(ctx context.Context, ref string) context.Context {
|
||||
return context.WithValue(ctx, referenceIDKey{}, ref)
|
||||
}
|
||||
|
||||
func referenceIDFrom(ctx context.Context) string {
|
||||
ref, _ := ctx.Value(referenceIDKey{}).(string)
|
||||
return ref
|
||||
}
|
||||
|
||||
// referenceIDMiddleware tags every log line of the request with an id that 500 problems also carry.
|
||||
func referenceIDMiddleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
ref := id.NewRandom()
|
||||
ctx := log.NewContext(withReferenceID(r.Context(), ref), "referenceId", ref)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue