diff --git a/go.mod b/go.mod index aaf6118f59..8270de9072 100644 --- a/go.mod +++ b/go.mod @@ -297,7 +297,7 @@ require ( github.com/gabriel-vasile/mimetype v1.4.13 // indirect github.com/go-chi/chi/v5 v5.2.5 // indirect github.com/go-jose/go-jose/v3 v3.0.5 // indirect - github.com/go-jose/go-jose/v4 v4.1.4 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 github.com/go-logr/logr v1.4.3 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/go-ole/go-ole v1.3.0 // indirect diff --git a/internal/pkg/service/appsproxy/config/config.go b/internal/pkg/service/appsproxy/config/config.go index f1135e165a..89efdeabdb 100644 --- a/internal/pkg/service/appsproxy/config/config.go +++ b/internal/pkg/service/appsproxy/config/config.go @@ -29,6 +29,7 @@ type Config struct { CsrfTokenSalt string `configKey:"csrfTokenSalt" configUsage:"Salt used for generating CSRF tokens" validate:"required" sensitive:"true"` StorageAPIURL *url.URL `configKey:"storageApiUrl" configUsage:"Base URL of the Keboola Storage API for this stack, used for Storage token verification (kai-preview flow). Must match the stack the proxy fronts — e.g. https://connection.eu-central-1.keboola.com for an EU stack. No default; required." validate:"required"` KaiPreview KaiPreview `configKey:"kaiPreview" configUsage:"kai-preview iframe-auth configuration."` + Preview Preview `configKey:"preview" configUsage:"Dev-mode app preview links minted by sandboxes-service."` K8s K8s `configKey:"k8s" configUsage:"Kubernetes configuration."` E2bWebhook E2BWebhook `configKey:"e2bWebhook"` Sessions Sessions `configKey:"sessions" configUsage:"End-user session tracking for data apps."` @@ -87,6 +88,53 @@ type KaiPreview struct { AllowedOrigins []string `configKey:"allowedOrigins" configUsage:"Origins allowed to embed apps via kai-preview and mint handshake tokens (e.g. https://connection.keboola.com). Drives both the CORS allowlist and the bootstrap CSP frame-ancestors directive." validate:"required,min=1,dive,http_url"` } +type Preview struct { + JWKSURL string `configKey:"jwksURL" configUsage:"In-cluster URL of the sandboxes-service JWKS. Empty disables preview links." validate:"omitempty,http_url"` + Issuer string `configKey:"issuer" configUsage:"Expected iss claim of preview links, e.g. https://apps.. Required when jwksURL is set."` + SessionSigningKey string `configKey:"sessionSigningKey" configUsage:"HMAC key for the preview session cookie, at least 32 characters. Required when jwksURL is set. Generate with 'openssl rand -hex 32'." sensitive:"true"` + AllowedFrameAncestors []string `configKey:"allowedFrameAncestors" configUsage:"Origins allowed to frame the preview landing page (CSP frame-ancestors), e.g. https://connection.keboola.com."` +} + +func (c Preview) Enabled() bool { + return c.JWKSURL != "" +} + +func (c *Preview) Normalize() { + for i, o := range c.AllowedFrameAncestors { + c.AllowedFrameAncestors[i] = strings.TrimRight(strings.TrimSpace(o), "/") + } +} + +func (c *Preview) Validate() error { + if !c.Enabled() { + return nil + } + errs := errors.NewMultiError() + if c.Issuer == "" { + errs.Append(errors.New("preview.issuer is required when preview.jwksURL is set")) + } + if len(c.SessionSigningKey) < 32 { + errs.Append(errors.New("preview.sessionSigningKey must have at least 32 characters when preview.jwksURL is set")) + } + for _, origin := range c.AllowedFrameAncestors { + if !isPlainOrigin(origin) { + errs.Append(errors.Errorf(`preview.allowedFrameAncestors: "%s" must be a scheme and host only`, origin)) + } + } + return errs.ErrorOrNil() +} + +func isPlainOrigin(origin string) bool { + if strings.ContainsAny(origin, " ;,'\"") { + return false + } + u, err := url.Parse(origin) + if err != nil { + return false + } + return (u.Scheme == "https" || u.Scheme == "http") && u.Host != "" && u.Path == "" && u.RawQuery == "" && u.Fragment == "" && u.User == nil +} + type API struct { Listen string `configKey:"listen" configUsage:"Listen address of the configuration HTTP API." validate:"required,hostname_port"` PublicURL *url.URL `configKey:"publicUrl" configUsage:"Public URL of the configuration HTTP API for link generation." validate:"required"` diff --git a/internal/pkg/service/appsproxy/config/preview_test.go b/internal/pkg/service/appsproxy/config/preview_test.go new file mode 100644 index 0000000000..484f68a4e0 --- /dev/null +++ b/internal/pkg/service/appsproxy/config/preview_test.go @@ -0,0 +1,95 @@ +package config_test + +import ( + "net/url" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/keboola/keboola-as-code/internal/pkg/env" + "github.com/keboola/keboola-as-code/internal/pkg/service/appsproxy/config" + "github.com/keboola/keboola-as-code/internal/pkg/service/common/configmap" +) + +func requiredConfig(t *testing.T) config.Config { + t.Helper() + cfg := config.New() + cfg.CookieSecretSalt = "x" + cfg.CsrfTokenSalt = "x" + cfg.SandboxesAPI.URL = "https://example" + cfg.K8s = config.K8s{AppsNamespace: "ns"} + storageURL, err := url.Parse("https://connection.keboola.com") + require.NoError(t, err) + cfg.StorageAPIURL = storageURL + cfg.KaiPreview = config.KaiPreview{ + HandshakeSigningKey: "k1", + SessionSigningKey: "k2", + SessionTTL: 4 * time.Hour, + AllowedOrigins: []string{"https://connection.keboola.com"}, + } + return cfg +} + +func TestPreviewConfig_Defaults(t *testing.T) { + t.Parallel() + cfg := config.New() + assert.Empty(t, cfg.Preview.JWKSURL) + assert.False(t, cfg.Preview.Enabled()) +} + +func TestPreviewConfig_DisabledNeedsNothing(t *testing.T) { + t.Parallel() + cfg := requiredConfig(t) + require.NoError(t, configmap.ValidateAndNormalize(&cfg)) +} + +func TestPreviewConfig_EnvNames(t *testing.T) { + t.Parallel() + cfg := requiredConfig(t) + + envs := env.Empty() + envs.Set("APPS_PROXY_PREVIEW_JWKS_URL", "http://sandboxes-service-api.default.svc.cluster.local/.well-known/jwks.json") + envs.Set("APPS_PROXY_PREVIEW_ISSUER", "https://apps.keboola.com") + envs.Set("APPS_PROXY_PREVIEW_SESSION_SIGNING_KEY", strings.Repeat("k", 64)) + envs.Set("APPS_PROXY_PREVIEW_ALLOWED_FRAME_ANCESTORS", "https://connection.keboola.com,https://connection.north-europe.azure.keboola.com/") + require.NoError(t, configmap.GenerateAndBind(configmap.GenerateAndBindConfig{ + EnvNaming: env.NewNamingConvention("APPS_PROXY_"), + Envs: envs, + }, &cfg)) + + assert.True(t, cfg.Preview.Enabled()) + assert.Equal(t, "http://sandboxes-service-api.default.svc.cluster.local/.well-known/jwks.json", cfg.Preview.JWKSURL) + assert.Equal(t, "https://apps.keboola.com", cfg.Preview.Issuer) + assert.Equal(t, strings.Repeat("k", 64), cfg.Preview.SessionSigningKey) + assert.Equal(t, []string{"https://connection.keboola.com", "https://connection.north-europe.azure.keboola.com"}, cfg.Preview.AllowedFrameAncestors) +} + +func TestPreviewConfig_EnabledValidation(t *testing.T) { + t.Parallel() + cases := []struct { + name string + mutate func(p *config.Preview) + wantErr string + }{ + {name: "missing-issuer", mutate: func(p *config.Preview) { p.Issuer = "" }, wantErr: "preview.issuer"}, + {name: "short-key", mutate: func(p *config.Preview) { p.SessionSigningKey = "short" }, wantErr: "preview.sessionSigningKey"}, + {name: "ancestor-with-path", mutate: func(p *config.Preview) { p.AllowedFrameAncestors = []string{"https://connection.keboola.com/admin"} }, wantErr: "preview.allowedFrameAncestors"}, + {name: "ancestor-with-semicolon", mutate: func(p *config.Preview) { p.AllowedFrameAncestors = []string{"https://a.com;script-src *"} }, wantErr: "preview.allowedFrameAncestors"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + cfg := requiredConfig(t) + cfg.Preview.JWKSURL = "http://sandboxes-service-api.default.svc.cluster.local/.well-known/jwks.json" + cfg.Preview.Issuer = "https://apps.keboola.com" + cfg.Preview.SessionSigningKey = strings.Repeat("k", 64) + tc.mutate(&cfg.Preview) + err := configmap.ValidateAndNormalize(&cfg) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErr) + }) + } +} diff --git a/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/jwks.go b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/jwks.go new file mode 100644 index 0000000000..0e99d13bc8 --- /dev/null +++ b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/jwks.go @@ -0,0 +1,218 @@ +package preview + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "encoding/json" + "io" + "net/http" + "sort" + "sync/atomic" + "time" + + "github.com/go-jose/go-jose/v4" + "github.com/jonboulle/clockwork" + "golang.org/x/sync/singleflight" + + "github.com/keboola/keboola-as-code/internal/pkg/log" + "github.com/keboola/keboola-as-code/internal/pkg/utils/errors" +) + +const ( + unknownKidRefetchInterval = time.Minute + maxJWKSBodySize = 64 << 10 + jwksFetchTimeout = 10 * time.Second +) + +type KeySetConfig struct { + URL string + RefreshInterval time.Duration + MaxStaleness time.Duration +} + +type KeySet struct { + cfg KeySetConfig + client *http.Client + clock clockwork.Clock + logger log.Logger + + snapshot atomic.Pointer[keySnapshot] + lastAttempt atomic.Pointer[time.Time] + refetch singleflight.Group +} + +type keySnapshot struct { + keys map[string]*ecdsa.PublicKey + fetchedAt time.Time +} + +type jwksDocument struct { + Keys *[]json.RawMessage `json:"keys"` +} + +func NewKeySet(cfg KeySetConfig, clock clockwork.Clock, logger log.Logger) *KeySet { + return &KeySet{ + cfg: cfg, + clock: clock, + logger: logger, + client: &http.Client{ + Timeout: jwksFetchTimeout, + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + } +} + +func (s *KeySet) Run(ctx context.Context) { + s.refreshAndLog(ctx) + ticker := s.clock.NewTicker(s.cfg.RefreshInterval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.Chan(): + s.refreshAndLog(ctx) + } + } +} + +func (s *KeySet) Refresh(ctx context.Context) error { + s.markAttempt() + keys, skipped, err := s.fetch(ctx) + if err != nil { + return err + } + for _, kid := range skipped { + s.logger.Warnf(ctx, `preview: JWKS key "%s" skipped: only kty=EC, crv=P-256, use=sig, alg ES256 or absent, and 32-byte x/y are accepted`, log.Sanitize(kid)) + } + s.store(keys) + s.logger.Debugf(ctx, "preview: JWKS loaded, kids=%v", sortedKids(keys)) + return nil +} + +func (s *KeySet) Key(ctx context.Context, kid string) (*ecdsa.PublicKey, error) { + if key, ok := s.lookup(kid); ok { + return key, nil + } + <-s.refetch.DoChan("jwks", func() (any, error) { + if s.claimRefetch() { + s.refreshAndLog(context.WithoutCancel(ctx)) + } + return nil, nil + }) + if key, ok := s.lookup(kid); ok { + return key, nil + } + return nil, errors.Errorf(`preview: no usable key for kid "%s"`, log.Sanitize(kid)) +} + +func (s *KeySet) refreshAndLog(ctx context.Context) { + err := s.Refresh(ctx) + if err == nil || ctx.Err() != nil { + return + } + s.logger.Warnf(ctx, "preview: JWKS refresh failed: %s", err) +} + +func (s *KeySet) markAttempt() { + now := s.clock.Now() + s.lastAttempt.Store(&now) +} + +func (s *KeySet) store(keys map[string]*ecdsa.PublicKey) { + s.snapshot.Store(&keySnapshot{keys: keys, fetchedAt: s.clock.Now()}) +} + +func (s *KeySet) lookup(kid string) (*ecdsa.PublicKey, bool) { + snap := s.snapshot.Load() + if snap == nil || s.clock.Since(snap.fetchedAt) > s.cfg.MaxStaleness { + return nil, false + } + key, ok := snap.keys[kid] + return key, ok +} + +func (s *KeySet) claimRefetch() bool { + now := s.clock.Now() + last := s.lastAttempt.Load() + if last != nil && now.Sub(*last) < unknownKidRefetchInterval { + return false + } + return s.lastAttempt.CompareAndSwap(last, &now) +} + +func (s *KeySet) fetch(ctx context.Context) (map[string]*ecdsa.PublicKey, []string, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, s.cfg.URL, nil) + if err != nil { + return nil, nil, errors.Errorf("preview: JWKS request: %w", err) + } + resp, err := s.client.Do(req) + if err != nil { + return nil, nil, errors.Errorf("preview: JWKS fetch: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return nil, nil, errors.Errorf("preview: JWKS answered HTTP %d", resp.StatusCode) + } + body, err := io.ReadAll(io.LimitReader(resp.Body, maxJWKSBodySize+1)) + if err != nil { + return nil, nil, errors.Errorf("preview: JWKS read: %w", err) + } + if len(body) > maxJWKSBodySize { + return nil, nil, errors.New("preview: JWKS response is too large") + } + return parseJWKS(body) +} + +func parseJWKS(body []byte) (map[string]*ecdsa.PublicKey, []string, error) { + var doc jwksDocument + if err := json.Unmarshal(body, &doc); err != nil { + return nil, nil, errors.Errorf("preview: invalid JWKS: %w", err) + } + if doc.Keys == nil { + return nil, nil, errors.New("preview: JWKS has no keys field") + } + keys := make(map[string]*ecdsa.PublicKey, len(*doc.Keys)) + var skipped []string + for _, raw := range *doc.Keys { + kid, pub, ok := parseJWK(raw) + if !ok { + skipped = append(skipped, kid) + continue + } + keys[kid] = pub + } + return keys, skipped, nil +} + +func parseJWK(raw json.RawMessage) (string, *ecdsa.PublicKey, bool) { + var ident struct { + Kid string `json:"kid"` + } + _ = json.Unmarshal(raw, &ident) + + var k jose.JSONWebKey + if err := json.Unmarshal(raw, &k); err != nil { + return ident.Kid, nil, false + } + if k.KeyID == "" || k.Use != "sig" || (k.Algorithm != "" && k.Algorithm != "ES256") { + return ident.Kid, nil, false + } + pub, ok := k.Key.(*ecdsa.PublicKey) + if !ok || pub.Curve != elliptic.P256() { + return ident.Kid, nil, false + } + return k.KeyID, pub, true +} + +func sortedKids(keys map[string]*ecdsa.PublicKey) []string { + kids := make([]string, 0, len(keys)) + for kid := range keys { + kids = append(kids, log.Sanitize(kid)) + } + sort.Strings(kids) + return kids +} diff --git a/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/jwks_test.go b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/jwks_test.go new file mode 100644 index 0000000000..5f48e1b164 --- /dev/null +++ b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/jwks_test.go @@ -0,0 +1,366 @@ +package preview_test + +import ( + "context" + "encoding/base64" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/jonboulle/clockwork" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/keboola/keboola-as-code/internal/pkg/log" + "github.com/keboola/keboola-as-code/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview" + "github.com/keboola/keboola-as-code/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/previewtest" +) + +func newKeySet(t *testing.T, url string, clock clockwork.Clock) *preview.KeySet { + t.Helper() + return preview.NewKeySet(preview.KeySetConfig{ + URL: url, + RefreshInterval: 10 * time.Minute, + MaxStaleness: time.Hour, + }, clock, log.NewNopLogger()) +} + +func withField(jwk map[string]any, key string, value any) map[string]any { + out := make(map[string]any, len(jwk)) + for k, v := range jwk { + out[k] = v + } + if value == nil { + delete(out, key) + return out + } + out[key] = value + return out +} + +func TestKeySet_KeyFilter(t *testing.T) { + t.Parallel() + ctx := t.Context() + good := previewtest.NewSigner(t, "good") + noAlg := previewtest.NewSigner(t, "no-alg") + base := previewtest.NewSigner(t, "x").JWK() + zero := base64.RawURLEncoding.EncodeToString(make([]byte, 32)) + private := previewtest.NewSigner(t, "private") + scalar, err := private.Key.Bytes() + require.NoError(t, err) + privateJWK := withField(private.JWK(), "d", base64.RawURLEncoding.EncodeToString(scalar)) + + server := previewtest.NewJWKSServer(t, + good.JWK(), + privateJWK, + withField(noAlg.JWK(), "alg", nil), + withField(withField(base, "kid", "rsa"), "kty", "RSA"), + withField(withField(base, "kid", "p384"), "crv", "P-384"), + withField(withField(base, "kid", "enc"), "use", "enc"), + withField(withField(base, "kid", "rs256"), "alg", "RS256"), + withField(base, "kid", ""), + withField(withField(withField(base, "kid", "off-curve"), "x", zero), "y", zero), + withField(withField(base, "kid", "short"), "x", base64.RawURLEncoding.EncodeToString(make([]byte, 31))), + ) + keys := newKeySet(t, server.JWKSURL(), clockwork.NewFakeClock()) + require.NoError(t, keys.Refresh(ctx)) + + for _, kid := range []string{"good", "no-alg"} { + _, err := keys.Key(ctx, kid) + require.NoError(t, err, kid) + } + for _, kid := range []string{"rsa", "p384", "enc", "rs256", "", "off-curve", "short", "private"} { + _, err := keys.Key(ctx, kid) + assert.Error(t, err, kid) + } +} + +func TestKeySet_FailedFetchKeepsLastGoodSet(t *testing.T) { + t.Parallel() + ctx := t.Context() + signer := previewtest.NewSigner(t, "k1") + server := previewtest.NewJWKSServer(t, signer.JWK()) + keys := newKeySet(t, server.JWKSURL(), clockwork.NewFakeClock()) + require.NoError(t, keys.Refresh(ctx)) + + server.SetStatus(http.StatusInternalServerError) + require.Error(t, keys.Refresh(ctx)) + server.SetStatus(http.StatusOK) + server.SetBody([]byte("not json")) + require.Error(t, keys.Refresh(ctx)) + server.SetBody([]byte(`{"keys":[` + strings.Repeat(`{"kty":"EC"},`, 8000) + `{}]}`)) + err := keys.Refresh(ctx) + require.Error(t, err) + assert.Contains(t, err.Error(), "too large") + + _, err = keys.Key(ctx, "k1") + assert.NoError(t, err) +} + +func TestKeySet_BodyWithoutKeysFieldKeepsLastGoodSet(t *testing.T) { + t.Parallel() + for _, body := range []string{`{}`, `null`, `{"error":"x"}`, `{"keys":null}`} { + t.Run(body, func(t *testing.T) { + t.Parallel() + ctx := t.Context() + signer := previewtest.NewSigner(t, "k1") + server := previewtest.NewJWKSServer(t, signer.JWK()) + keys := newKeySet(t, server.JWKSURL(), clockwork.NewFakeClock()) + require.NoError(t, keys.Refresh(ctx)) + + server.SetBody([]byte(body)) + require.Error(t, keys.Refresh(ctx)) + + _, err := keys.Key(ctx, "k1") + assert.NoError(t, err) + }) + } +} + +func TestKeySet_EmptySetRemovesKeys(t *testing.T) { + t.Parallel() + ctx := t.Context() + signer := previewtest.NewSigner(t, "k1") + server := previewtest.NewJWKSServer(t, signer.JWK()) + keys := newKeySet(t, server.JWKSURL(), clockwork.NewFakeClock()) + require.NoError(t, keys.Refresh(ctx)) + server.SetKeys() + require.NoError(t, keys.Refresh(ctx)) + _, err := keys.Key(ctx, "k1") + assert.Error(t, err) +} + +func TestKeySet_DoesNotFollowRedirects(t *testing.T) { + t.Parallel() + ctx := t.Context() + signer := previewtest.NewSigner(t, "k1") + other := previewtest.NewJWKSServer(t, signer.JWK()) + redirect := httptest.NewServer(http.RedirectHandler(other.JWKSURL(), http.StatusFound)) + t.Cleanup(redirect.Close) + + keys := newKeySet(t, redirect.URL, clockwork.NewFakeClock()) + require.Error(t, keys.Refresh(ctx)) + assert.Equal(t, int64(0), other.Hits()) +} + +func TestKeySet_StaleSetIsRejected(t *testing.T) { + t.Parallel() + ctx := t.Context() + clock := clockwork.NewFakeClock() + signer := previewtest.NewSigner(t, "k1") + server := previewtest.NewJWKSServer(t, signer.JWK()) + keys := newKeySet(t, server.JWKSURL(), clock) + require.NoError(t, keys.Refresh(ctx)) + server.SetStatus(http.StatusServiceUnavailable) + + clock.Advance(59 * time.Minute) + _, err := keys.Key(ctx, "k1") + require.NoError(t, err, "within maxStaleness the last good set is trusted") + + clock.Advance(2 * time.Minute) + _, err = keys.Key(ctx, "k1") + require.Error(t, err, "past maxStaleness the set is no longer trusted") + + server.SetStatus(http.StatusOK) + clock.Advance(time.Minute) + _, err = keys.Key(ctx, "k1") + require.NoError(t, err, "a successful refetch makes the set fresh again") +} + +func TestKeySet_UnknownKidRefetchAtMostOncePerMinute(t *testing.T) { + t.Parallel() + ctx := t.Context() + clock := clockwork.NewFakeClock() + k1 := previewtest.NewSigner(t, "k1") + k2 := previewtest.NewSigner(t, "k2") + server := previewtest.NewJWKSServer(t, k1.JWK()) + keys := newKeySet(t, server.JWKSURL(), clock) + require.NoError(t, keys.Refresh(ctx)) + require.Equal(t, int64(1), server.Hits()) + + server.SetKeys(k1.JWK(), k2.JWK()) + _, err := keys.Key(ctx, "k2") + require.Error(t, err) + assert.Equal(t, int64(1), server.Hits(), "no refetch within a minute of the last fetch") + + clock.Advance(61 * time.Second) + _, err = keys.Key(ctx, "k2") + require.NoError(t, err) + assert.Equal(t, int64(2), server.Hits()) + + _, err = keys.Key(ctx, "k3") + require.Error(t, err) + _, err = keys.Key(ctx, "k3") + require.Error(t, err) + assert.Equal(t, int64(2), server.Hits()) +} + +func TestKeySet_ConcurrentUnknownKidCallersShareOneRefetch(t *testing.T) { + t.Parallel() + ctx := t.Context() + clock := clockwork.NewFakeClock() + k1 := previewtest.NewSigner(t, "k1") + k2 := previewtest.NewSigner(t, "k2") + + var served atomic.Value + served.Store([]map[string]any{k1.JWK()}) + var hits atomic.Int64 + entered := make(chan struct{}, 1) + release := make(chan struct{}) + var block atomic.Bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + hits.Add(1) + if block.Load() { + entered <- struct{}{} + <-release + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{"keys": served.Load()}) + })) + t.Cleanup(server.Close) + + keys := newKeySet(t, server.URL, clock) + require.NoError(t, keys.Refresh(ctx)) + require.Equal(t, int64(1), hits.Load()) + + served.Store([]map[string]any{k1.JWK(), k2.JWK()}) + block.Store(true) + clock.Advance(61 * time.Second) + + const callers = 8 + errs := make(chan error, callers+1) + var wg sync.WaitGroup + call := func() { + defer wg.Done() + _, err := keys.Key(ctx, "k2") + errs <- err + } + wg.Add(1) + go call() + <-entered + + for range callers { + wg.Add(1) + go call() + } + time.Sleep(100 * time.Millisecond) + close(release) + wg.Wait() + close(errs) + + for err := range errs { + require.NoError(t, err) + } + assert.Equal(t, int64(2), hits.Load()) +} + +func TestKeySet_UnknownKidRefetchIgnoresRequestCancellation(t *testing.T) { + t.Parallel() + clock := clockwork.NewFakeClock() + k1 := previewtest.NewSigner(t, "k1") + k2 := previewtest.NewSigner(t, "k2") + server := previewtest.NewJWKSServer(t, k1.JWK()) + keys := newKeySet(t, server.JWKSURL(), clock) + require.NoError(t, keys.Refresh(t.Context())) + + server.SetKeys(k1.JWK(), k2.JWK()) + clock.Advance(61 * time.Second) + cancelled, cancel := context.WithCancelCause(t.Context()) + cancel(nil) + + _, err := keys.Key(cancelled, "k2") + require.NoError(t, err) + assert.Equal(t, int64(2), server.Hits()) +} + +func TestKeySet_RunDoesNotWarnWhenCancelledDuringFetch(t *testing.T) { + t.Parallel() + started := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + close(started) + <-r.Context().Done() + })) + t.Cleanup(server.Close) + logger := log.NewDebugLogger() + keys := preview.NewKeySet(preview.KeySetConfig{ + URL: server.URL, + RefreshInterval: 10 * time.Minute, + MaxStaleness: time.Hour, + }, clockwork.NewFakeClock(), logger) + + ctx, cancel := context.WithCancelCause(t.Context()) + defer cancel(nil) + done := make(chan struct{}) + go func() { + keys.Run(ctx) + close(done) + }() + <-started + cancel(nil) + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("Run did not stop on context cancellation") + } + assert.Empty(t, logger.WarnAndErrorMessages()) +} + +func TestKeySet_RunWarnsWhenRefreshFails(t *testing.T) { + t.Parallel() + signer := previewtest.NewSigner(t, "k1") + server := previewtest.NewJWKSServer(t, signer.JWK()) + server.SetStatus(http.StatusServiceUnavailable) + logger := log.NewDebugLogger() + keys := preview.NewKeySet(preview.KeySetConfig{ + URL: server.JWKSURL(), + RefreshInterval: 10 * time.Minute, + MaxStaleness: time.Hour, + }, clockwork.NewFakeClock(), logger) + + ctx, cancel := context.WithCancelCause(t.Context()) + defer cancel(nil) + go keys.Run(ctx) + require.Eventually(t, func() bool { + return strings.Contains(logger.WarnMessages(), "JWKS refresh failed") + }, 5*time.Second, 10*time.Millisecond) +} + +func TestKeySet_RunSurvivesDownJWKS(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancelCause(t.Context()) + defer cancel(nil) + clock := clockwork.NewFakeClock() + signer := previewtest.NewSigner(t, "k1") + server := previewtest.NewJWKSServer(t, signer.JWK()) + server.SetStatus(http.StatusServiceUnavailable) + keys := newKeySet(t, server.JWKSURL(), clock) + + done := make(chan struct{}) + go func() { + keys.Run(ctx) + close(done) + }() + require.Eventually(t, func() bool { return server.Hits() >= 1 }, 5*time.Second, 10*time.Millisecond) + _, err := keys.Key(ctx, "k1") + require.Error(t, err) + + server.SetStatus(http.StatusOK) + require.NoError(t, clock.BlockUntilContext(ctx, 1)) + clock.Advance(10 * time.Minute) + require.Eventually(t, func() bool { + _, err := keys.Key(ctx, "k1") + return err == nil + }, 5*time.Second, 10*time.Millisecond) + + cancel(nil) + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("Run did not stop on context cancellation") + } +} diff --git a/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/origin.go b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/origin.go new file mode 100644 index 0000000000..30af5dddc9 --- /dev/null +++ b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/origin.go @@ -0,0 +1,37 @@ +// Package preview implements dev-mode app preview links: a link token minted by +// sandboxes-service is redeemed on the app host for a host-only session cookie. +package preview + +import ( + "net/url" + "strings" + + "github.com/keboola/keboola-as-code/internal/pkg/utils/errors" +) + +func NormalizeOrigin(raw string) (string, error) { + u, err := url.Parse(raw) + if err != nil { + return "", errors.Errorf("invalid origin: %w", err) + } + scheme := strings.ToLower(u.Scheme) + if scheme != "https" && scheme != "http" { + return "", errors.Errorf(`invalid origin scheme "%s"`, u.Scheme) + } + if u.Hostname() == "" || u.User != nil || u.Opaque != "" || u.RawQuery != "" || u.Fragment != "" || (u.Path != "" && u.Path != "/") { + return "", errors.New("origin must consist of a scheme and a host only") + } + port := u.Port() + if (scheme == "https" && port == "443") || (scheme == "http" && port == "80") { + port = "" + } + var b strings.Builder + b.WriteString(scheme) + b.WriteString("://") + b.WriteString(strings.ToLower(u.Hostname())) + if port != "" { + b.WriteByte(':') + b.WriteString(port) + } + return b.String(), nil +} diff --git a/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/origin_test.go b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/origin_test.go new file mode 100644 index 0000000000..9419890c2e --- /dev/null +++ b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/origin_test.go @@ -0,0 +1,48 @@ +package preview_test + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/keboola/keboola-as-code/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview" +) + +func TestNormalizeOrigin(t *testing.T) { + t.Parallel() + valid := []struct{ in, want string }{ + {"https://my-app-123.hub.keboola.com", "https://my-app-123.hub.keboola.com"}, + {"https://my-app-123.hub.keboola.com/", "https://my-app-123.hub.keboola.com"}, + {"HTTPS://MY-APP-123.HUB.KEBOOLA.COM", "https://my-app-123.hub.keboola.com"}, + {"https://my-app-123.hub.keboola.com:443", "https://my-app-123.hub.keboola.com"}, + {"http://localhost:80", "http://localhost"}, + {"https://basic-auth.hub.keboola.local:18443", "https://basic-auth.hub.keboola.local:18443"}, + } + for _, tc := range valid { + got, err := preview.NormalizeOrigin(tc.in) + require.NoError(t, err, tc.in) + assert.Equal(t, tc.want, got, tc.in) + } + + invalid := []string{ + "", + "my-app-123.hub.keboola.com", + "ftp://my-app-123.hub.keboola.com", + "https://my-app-123.hub.keboola.com/x", + "https://my-app-123.hub.keboola.com?a=b", + "https://my-app-123.hub.keboola.com#t=x", + "https://user@my-app-123.hub.keboola.com", + "https://", + } + for _, in := range invalid { + _, err := preview.NormalizeOrigin(in) + require.Error(t, err, in) + } + + a, _ := preview.NormalizeOrigin("https://a.hub.keboola.com:8443") + b, _ := preview.NormalizeOrigin("https://a.hub.keboola.com") + assert.NotEqual(t, a, b) + c, _ := preview.NormalizeOrigin("http://a.hub.keboola.com") + assert.NotEqual(t, b, c) +} diff --git a/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/previewtest/previewtest.go b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/previewtest/previewtest.go new file mode 100644 index 0000000000..2222d4352d --- /dev/null +++ b/internal/pkg/service/appsproxy/proxy/apphandler/authproxy/preview/previewtest/previewtest.go @@ -0,0 +1,161 @@ +// Package previewtest provides an ES256 link signer and a fake JWKS endpoint for preview tests. +package previewtest + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "encoding/base64" + "encoding/json" + "net/http" + "net/http/httptest" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" + "github.com/stretchr/testify/require" +) + +const Issuer = "https://apps.keboola.local" + +type Signer struct { + Kid string + Key *ecdsa.PrivateKey +} + +func NewSigner(tb testing.TB, kid string) *Signer { + tb.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(tb, err) + return &Signer{Kid: kid, Key: key} +} + +func (s *Signer) JWK() map[string]any { + raw, err := s.Key.PublicKey.Bytes() + if err != nil { + panic(err) + } + return map[string]any{ + "kty": "EC", + "crv": "P-256", + "use": "sig", + "alg": "ES256", + "kid": s.Kid, + "x": base64.RawURLEncoding.EncodeToString(raw[1:33]), + "y": base64.RawURLEncoding.EncodeToString(raw[33:65]), + } +} + +type Claims struct { + Issuer string + Audience any + Subject string + Purpose string + Ver int + ID string + IssuedAt time.Time + ExpiresAt time.Time + OmitIssuedAt bool + OmitExpiresAt bool + ExtraHeaders map[string]any +} + +func ValidClaims(now time.Time, sub string) Claims { + return Claims{ + Issuer: Issuer, + Audience: "apps-proxy", + Subject: sub, + Purpose: "app-preview-link", + Ver: 1, + ID: "0123456789abcdef0123456789abcdef", + IssuedAt: now, + ExpiresAt: now.Add(60 * time.Second), + } +} + +func (c Claims) Map() jwt.MapClaims { + m := jwt.MapClaims{ + "iss": c.Issuer, + "aud": c.Audience, + "sub": c.Subject, + "purpose": c.Purpose, + "ver": c.Ver, + "jti": c.ID, + } + if !c.OmitIssuedAt { + m["iat"] = c.IssuedAt.Unix() + } + if !c.OmitExpiresAt { + m["exp"] = c.ExpiresAt.Unix() + } + return m +} + +func (s *Signer) Mint(tb testing.TB, c Claims) string { + tb.Helper() + token := jwt.NewWithClaims(jwt.SigningMethodES256, c.Map()) + token.Header["kid"] = s.Kid + for k, v := range c.ExtraHeaders { + token.Header[k] = v + } + signed, err := token.SignedString(s.Key) + require.NoError(tb, err) + return signed +} + +type JWKSServer struct { + *httptest.Server + lock sync.Mutex + body []byte + status int + hits atomic.Int64 +} + +func NewJWKSServer(tb testing.TB, keys ...map[string]any) *JWKSServer { + tb.Helper() + s := &JWKSServer{status: http.StatusOK} + s.SetKeys(keys...) + s.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + s.hits.Add(1) + s.lock.Lock() + defer s.lock.Unlock() + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(s.status) + _, _ = w.Write(s.body) + })) + tb.Cleanup(s.Close) + return s +} + +func (s *JWKSServer) SetKeys(keys ...map[string]any) { + if keys == nil { + keys = []map[string]any{} + } + body, err := json.Marshal(map[string]any{"keys": keys}) + if err != nil { + panic(err) + } + s.SetBody(body) +} + +func (s *JWKSServer) SetBody(body []byte) { + s.lock.Lock() + defer s.lock.Unlock() + s.body = body +} + +func (s *JWKSServer) SetStatus(code int) { + s.lock.Lock() + defer s.lock.Unlock() + s.status = code +} + +func (s *JWKSServer) Hits() int64 { + return s.hits.Load() +} + +func (s *JWKSServer) JWKSURL() string { + return s.URL + "/.well-known/jwks.json" +}