diff --git a/access_gate.go b/access_gate.go index 29cf92b..586ccab 100644 --- a/access_gate.go +++ b/access_gate.go @@ -81,8 +81,9 @@ func (g *AccessGate) Check(r *http.Request) CheckResult { } func (g *AccessGate) memberOfRequired(userID string) (bool, error) { + key := g.cacheKey(userID) if g.CacheTTL > 0 { - if g.cachedAllowed(userID) { + if g.cachedAllowed(key) { return true, nil } } @@ -94,7 +95,7 @@ func (g *AccessGate) memberOfRequired(userID string) (bool, error) { for _, have := range groups { if have == need { if g.CacheTTL > 0 { - g.storeAllowed(userID) + g.storeAllowed(key) } return true, nil } @@ -103,6 +104,10 @@ func (g *AccessGate) memberOfRequired(userID string) (bool, error) { return false, nil } +func (g *AccessGate) cacheKey(userID string) string { + return userID + "\x00" + strings.Join(g.Groups, "\x00") +} + func (g *AccessGate) fetchUserGroups(userID string) ([]string, error) { ocs := g.ocsClient(userID) path := "cloud/users/" + url.PathEscape(userID) + "/groups" @@ -139,24 +144,24 @@ func decodeUserGroups(raw []byte) ([]string, error) { return parsed.OCS.Data.Groups, nil } -func (g *AccessGate) cachedAllowed(userID string) bool { +func (g *AccessGate) cachedAllowed(key string) bool { g.mu.Lock() defer g.mu.Unlock() - ent, ok := g.cache[userID] + ent, ok := g.cache[key] if !ok { return false } if g.Now().After(ent.until) { - delete(g.cache, userID) + delete(g.cache, key) return false } return true } -func (g *AccessGate) storeAllowed(userID string) { +func (g *AccessGate) storeAllowed(key string) { g.mu.Lock() defer g.mu.Unlock() - g.cache[userID] = cacheEntry{until: g.Now().Add(g.CacheTTL)} + g.cache[key] = cacheEntry{until: g.Now().Add(g.CacheTTL)} } func (g AccessGate) normalized() *AccessGate { @@ -185,27 +190,28 @@ func acceptsHTML(accept string) bool { } func (g *AccessGate) shouldSkip(path string) bool { - path = strings.TrimSuffix(path, "/") - if path == "" { - path = "/" - } + path = normalizeGatePath(path) for _, p := range defaultSkipPaths { if path == p { return true } } for _, p := range g.ExtraSkipPaths { - p = strings.TrimSuffix(p, "/") - if p == "" { - p = "/" - } - if path == p { + if path == normalizeGatePath(p) { return true } } return false } +func normalizeGatePath(path string) string { + path = strings.TrimSuffix(path, "/") + if path == "" { + return "/" + } + return path +} + var defaultSkipPaths = []string{"/heartbeat", "/enabled", "/init"} var deniedHTML = []byte(` diff --git a/access_gate_test.go b/access_gate_test.go index e12aea3..f098151 100644 --- a/access_gate_test.go +++ b/access_gate_test.go @@ -213,6 +213,35 @@ func TestAccessGateCachesPositiveMembership(t *testing.T) { } } +func TestAccessGateCacheKeyedByGroupSet(t *testing.T) { + var calls atomic.Int32 + srv := groupsOCSServer(t, func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + _ = json.NewEncoder(w).Encode(map[string]any{ + "ocs": map[string]any{"data": map[string]any{"groups": []string{"dns-ops"}}}, + }) + }) + cred := gonexapp.Credentials{ + BaseURL: srv.URL, AppID: "app", AppVersion: "0.1.0", AAVersion: "1.0.0", AppSecret: "s", + } + gate := &gonexapp.AccessGate{ + Cred: cred, Groups: []string{"dns-ops"}, CacheTTL: time.Minute, + Client: srv.Client(), OCS: gonexapp.OCSClient{Cred: cred, Client: srv.Client()}, + } + req := httptest.NewRequest(http.MethodGet, "/api/zones", nil) + req.Header.Set("AUTHORIZATION-APP-API", authHeader("alice")) + if gate.Check(req) != gonexapp.CheckAllowed { + t.Fatal("first allow") + } + gate.Groups = []string{"other-group"} + if gate.Check(req) != gonexapp.CheckDenied { + t.Fatalf("after group-set change want denied, calls=%d", calls.Load()) + } + if calls.Load() != 2 { + t.Fatalf("ocs calls=%d want 2 (cache must not reuse prior group-set)", calls.Load()) + } +} + func TestAccessGateDoesNotCacheDenial(t *testing.T) { var calls atomic.Int32 srv := groupsOCSServer(t, func(w http.ResponseWriter, r *http.Request) {