Harden CIMD client resolution and caching

Tighten redirect URI validation for metadata documents, honor
Cache-Control no-store when caching fetched documents, and resolve
clients on the same transaction as authorization. Load
external_client_id from the database and parse unbounded max-stale
directives in cachecontrol.

Signed-off-by: Bryan Frimin <bryan@probo.com>
This commit is contained in:
Bryan Frimin
2026-06-19 16:51:30 +02:00
parent 5b0d3e5052
commit d7e23fd890
7 changed files with 257 additions and 54 deletions

View File

@@ -17,6 +17,7 @@ package oauth2
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
@@ -196,12 +197,8 @@ func validateClientMetadataDocument(clientIDURL string, doc *ClientMetadataDocum
}
for _, redirectURI := range doc.RedirectURIs {
parsed, err := url.Parse(redirectURI)
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
return NewError(
ErrInvalidClient,
WithDescription("client metadata document contains invalid redirect_uri"),
)
if err := validateCIMDRedirectURI(redirectURI); err != nil {
return err
}
}
@@ -223,6 +220,41 @@ func validateClientMetadataDocument(clientIDURL string, doc *ClientMetadataDocum
return nil
}
func validateCIMDRedirectURI(redirectURI string) error {
parsed, err := url.Parse(redirectURI)
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
return NewError(
ErrInvalidClient,
WithDescription("client metadata document contains invalid redirect_uri"),
)
}
if parsed.User != nil || parsed.Fragment != "" {
return NewError(
ErrInvalidClient,
WithDescription("client metadata document contains invalid redirect_uri"),
)
}
switch parsed.Scheme {
case "https":
case "http":
if !net.IsLoopback(parsed.Hostname()) {
return NewError(
ErrInvalidClient,
WithDescription("client metadata document contains invalid redirect_uri"),
)
}
default:
return NewError(
ErrInvalidClient,
WithDescription("client metadata document contains invalid redirect_uri"),
)
}
return nil
}
func cimdRedirectURIAllowed(doc *ClientMetadataDocument, redirectURI string) bool {
for _, allowed := range doc.RedirectURIs {
if redirectURI == allowed {
@@ -280,8 +312,14 @@ func (f *cimdFetcher) loadCache(clientIDURL string) (*ClientMetadataDocument, bo
}
func (f *cimdFetcher) storeCache(clientIDURL string, doc *ClientMetadataDocument, cacheControl string) {
dir, err := cachecontrol.ParseResponse(cacheControl)
if err == nil && dir.NoStore() {
return
}
ttl := cimdDefaultCacheTTL
if dir, err := cachecontrol.ParseResponse(cacheControl); err == nil {
if err == nil {
if maxAge, ok := dir.MaxAgeDuration(); ok {
ttl = min(ttl, maxAge)
}
@@ -300,8 +338,30 @@ func (s *Service) ResolveClient(
ctx context.Context,
clientIDRaw string,
redirectURI string,
) (*coredata.OAuth2Client, error) {
return s.resolveClient(ctx, nil, clientIDRaw, redirectURI)
}
func (s *Service) resolveClient(
ctx context.Context,
tx pg.Tx,
clientIDRaw string,
redirectURI string,
) (*coredata.OAuth2Client, error) {
if clientID, err := gid.ParseGID(clientIDRaw); err == nil {
if tx != nil {
client := coredata.OAuth2Client{}
if err := client.LoadByID(ctx, tx, coredata.NewNoScope(), clientID); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return nil, NewError(ErrInvalidClient, WithDescription("client not found"))
}
return nil, fmt.Errorf("cannot load oauth2 client: %w", err)
}
return &client, nil
}
return s.GetClientByID(ctx, clientID)
}
@@ -325,7 +385,7 @@ func (s *Service) ResolveClient(
return nil, ErrInvalidRedirectURI
}
client, err := s.upsertCIMDClient(ctx, clientIDRaw, doc)
client, err := s.upsertCIMDClient(ctx, tx, clientIDRaw, doc)
if err != nil {
return nil, err
}
@@ -335,6 +395,7 @@ func (s *Service) ResolveClient(
func (s *Service) upsertCIMDClient(
ctx context.Context,
tx pg.Tx,
externalClientID string,
doc *ClientMetadataDocument,
) (*coredata.OAuth2Client, error) {
@@ -358,6 +419,7 @@ func (s *Service) upsertCIMDClient(
)
now := time.Now()
candidate, err := coredata.NewCIMDClient(
externalClientID,
doc.ClientName,
@@ -373,16 +435,28 @@ func (s *Service) upsertCIMDClient(
var client coredata.OAuth2Client
upsert := func(ctx context.Context, conn pg.Tx) error {
client = *candidate
if err := client.UpsertCIMD(ctx, conn); err != nil {
return fmt.Errorf("cannot upsert cimd oauth2 client: %w", err)
}
return nil
}
if tx != nil {
if err := upsert(ctx, tx); err != nil {
return nil, err
}
return &client, nil
}
err = s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
client = *candidate
if err := client.UpsertCIMD(ctx, tx); err != nil {
return fmt.Errorf("cannot upsert cimd oauth2 client: %w", err)
}
return nil
func(ctx context.Context, conn pg.Tx) error {
return upsert(ctx, conn)
},
)
if err != nil {