mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-09-18 17:09:05 +02:00
201 lines
7.4 KiB
Go
201 lines
7.4 KiB
Go
package oidc
|
|
|
|
import (
|
|
"bytes"
|
|
"crypto/ecdsa"
|
|
"crypto/elliptic"
|
|
"crypto/rand"
|
|
"encoding/base64"
|
|
"mime/multipart"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/pocket-id/pocket-id/backend/internal/model"
|
|
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
|
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
type authorizationPARTestCase struct {
|
|
name string
|
|
method string
|
|
queryURI string
|
|
bodyURI []string
|
|
createPAR bool
|
|
wantCode bool
|
|
multipart bool
|
|
requestObject bool
|
|
malformed bool
|
|
}
|
|
|
|
func TestAuthorizationHandlerRequiresPAR(t *testing.T) {
|
|
gin.SetMode(gin.TestMode)
|
|
const fakeURI = parRequestURIPrefix + "missing"
|
|
for _, clientType := range []string{"public", "confidential"} {
|
|
t.Run(clientType, func(t *testing.T) {
|
|
for _, tt := range []authorizationPARTestCase{
|
|
{name: "missing request URI", method: http.MethodPost},
|
|
{name: "unknown PAR URI", method: http.MethodGet, queryURI: fakeURI},
|
|
{name: "empty body shadows PAR query", method: http.MethodPost, queryURI: fakeURI, bodyURI: []string{""}},
|
|
{name: "non-PAR body shadows PAR query", method: http.MethodPost, queryURI: fakeURI, bodyURI: []string{"not-a-par-uri"}},
|
|
{name: "valid PAR in query", method: http.MethodGet, queryURI: "stored", createPAR: true, wantCode: true},
|
|
{name: "valid PAR in body", method: http.MethodPost, bodyURI: []string{"stored"}, createPAR: true, wantCode: true},
|
|
{name: "valid body overrides invalid query", method: http.MethodPost, queryURI: fakeURI, bodyURI: []string{"stored"}, createPAR: true, wantCode: true},
|
|
{name: "empty body shadows stored PAR", method: http.MethodPost, queryURI: "stored", bodyURI: []string{""}, createPAR: true},
|
|
{name: "valid PAR in multipart body", method: http.MethodPost, bodyURI: []string{"stored"}, createPAR: true, wantCode: true, multipart: true},
|
|
{name: "request object cannot introduce PAR URI", method: http.MethodPost, requestObject: true},
|
|
{name: "malformed body", method: http.MethodPost, queryURI: fakeURI, malformed: true},
|
|
} {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
testAuthorizationHandlerPAR(t, clientType, tt)
|
|
})
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func testAuthorizationHandlerPAR(t *testing.T, clientType string, tt authorizationPARTestCase) {
|
|
t.Helper()
|
|
const (
|
|
baseURL = "https://issuer.example.com"
|
|
callbackURL = "https://client.example.com/callback"
|
|
clientID = "par-client"
|
|
userID = "test-user"
|
|
clientSecret = "test-client-secret"
|
|
fakeURI = parRequestURIPrefix + "missing"
|
|
)
|
|
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
require.NoError(t, err)
|
|
|
|
// Use the real provider and an authenticated user so a bypass would issue an authorization code
|
|
db := testutils.NewDatabaseForTest(t)
|
|
require.NoError(t, db.Create(&model.User{Base: model.Base{ID: userID}}).Error)
|
|
require.NoError(t, db.Create(&model.OidcClient{
|
|
Base: model.Base{ID: clientID},
|
|
Name: "PAR client",
|
|
CallbackURLs: datatype.StringList{callbackURL},
|
|
IsPublic: clientType == "public",
|
|
PkceEnabled: true,
|
|
SkipConsent: true,
|
|
RequiresPushedAuthorizationRequests: true,
|
|
Credentials: testClientCredentials(clientSecret),
|
|
}).Error)
|
|
store := NewStore(db, nil)
|
|
provider, err := newProvider(store, nil, testTokenSigner{key: key}, Config{
|
|
BaseURL: baseURL, TokenBaseURL: baseURL, Secret: []byte("test-secret"),
|
|
}, nil)
|
|
require.NoError(t, err)
|
|
service := newAuthorizationService(db, newInteractionSessionService(db), newClaimsService(db, nil, baseURL, nil), nil, nil, nil)
|
|
handler := newAuthorizationHandler(provider, service)
|
|
params := url.Values{
|
|
"client_id": {clientID},
|
|
"response_type": {"code"},
|
|
"redirect_uri": {callbackURL},
|
|
"scope": {"openid"},
|
|
"state": {"state-with-enough-entropy"},
|
|
"prompt": {"none"},
|
|
"code_challenge": {"E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"},
|
|
"code_challenge_method": {"S256"},
|
|
}
|
|
if tt.requestObject {
|
|
header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"none"}`))
|
|
payload := base64.RawURLEncoding.EncodeToString([]byte(`{"request_uri":"` + fakeURI + `"}`))
|
|
params.Set("request", header+"."+payload+".")
|
|
}
|
|
|
|
// Create valid PAR sessions through Fosite to cover both accepted transports and consumption
|
|
var storedURI string
|
|
if tt.createPAR {
|
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/api/oidc/par", strings.NewReader(params.Encode()))
|
|
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
if clientType == "confidential" {
|
|
req.SetBasicAuth(clientID, clientSecret)
|
|
}
|
|
ar, err := provider.NewPushedAuthorizeRequest(t.Context(), req)
|
|
require.NoError(t, err)
|
|
response, err := provider.NewPushedAuthorizeResponse(t.Context(), ar, NewEmptySession())
|
|
require.NoError(t, err)
|
|
storedURI = response.GetRequestURI()
|
|
}
|
|
|
|
// Keep query and body parameters separate to exercise the same precedence as incoming HTTP requests
|
|
query := url.Values{}
|
|
if tt.queryURI != "" {
|
|
uri := tt.queryURI
|
|
if uri == "stored" {
|
|
uri = storedURI
|
|
}
|
|
query.Set("request_uri", uri)
|
|
}
|
|
for _, uri := range tt.bodyURI {
|
|
if uri == "stored" {
|
|
uri = storedURI
|
|
}
|
|
params.Add("request_uri", uri)
|
|
}
|
|
body := ""
|
|
contentType := "application/x-www-form-urlencoded"
|
|
if tt.method == http.MethodPost {
|
|
body = params.Encode()
|
|
} else {
|
|
for k, values := range params {
|
|
query[k] = values
|
|
}
|
|
}
|
|
if tt.multipart {
|
|
var buf bytes.Buffer
|
|
writer := multipart.NewWriter(&buf)
|
|
for k, values := range params {
|
|
for _, value := range values {
|
|
require.NoError(t, writer.WriteField(k, value))
|
|
}
|
|
}
|
|
require.NoError(t, writer.Close())
|
|
body = buf.String()
|
|
contentType = writer.FormDataContentType()
|
|
}
|
|
if tt.malformed {
|
|
body = "request_uri=%zz"
|
|
}
|
|
req := httptest.NewRequestWithContext(t.Context(), tt.method, "/authorize?"+query.Encode(), strings.NewReader(body))
|
|
req.Header.Set("Content-Type", contentType)
|
|
rec := httptest.NewRecorder()
|
|
router := gin.New()
|
|
router.Handle(tt.method, "/authorize", func(c *gin.Context) {
|
|
c.Set("userID", userID)
|
|
c.Set("authenticationTime", time.Now().UTC().Add(-time.Minute))
|
|
handler.authorize(c)
|
|
})
|
|
router.ServeHTTP(rec, req)
|
|
|
|
// Rejections must neither return nor persist a code, while valid PAR must be consumed
|
|
location, err := url.Parse(rec.Header().Get("Location"))
|
|
require.NoError(t, err)
|
|
var codeCount int64
|
|
require.NoError(t, db.Model(&OAuth2Session{}).Where("kind = ?", sessionKindAuthorizeCode).Count(&codeCount).Error)
|
|
if tt.wantCode {
|
|
require.Equal(t, http.StatusSeeOther, rec.Code)
|
|
require.NotEmpty(t, location.Query().Get("code"), rec.Body.String())
|
|
require.Empty(t, location.Query().Get("error"))
|
|
require.EqualValues(t, 1, codeCount)
|
|
} else {
|
|
require.Empty(t, location.Query().Get("code"))
|
|
require.NotEmpty(t, location.Query().Get("error"))
|
|
require.Zero(t, codeCount)
|
|
}
|
|
if tt.createPAR {
|
|
var parCount int64
|
|
require.NoError(t, db.Model(&OAuth2Session{}).Where("kind = ? AND key = ? AND active = ?", sessionKindPAR, storedURI, true).Count(&parCount).Error)
|
|
if tt.wantCode {
|
|
require.Zero(t, parCount)
|
|
} else {
|
|
require.EqualValues(t, 1, parCount)
|
|
}
|
|
}
|
|
}
|