diff --git a/server/cmd/api/server.go b/server/cmd/api/server.go index ecc99f1..e6c1613 100644 --- a/server/cmd/api/server.go +++ b/server/cmd/api/server.go @@ -8,6 +8,7 @@ import ( "os/signal" "time" + "github.com/casbin/casbin/v2" "github.com/gin-gonic/gin" "github.com/go-admin-team/go-admin-core/config/source/file" log "github.com/go-admin-team/go-admin-core/logger" @@ -95,6 +96,9 @@ func run() error { return fmt.Errorf("reconcile GoAuto purchaser permissions for database %q: %w", name, err) } } + if err := reloadCasbinPolicies(sdk.Runtime.GetCasbin()); err != nil { + return err + } initRouter() for _, f := range AppRouters { @@ -189,6 +193,30 @@ func run() error { return nil } +type policyLoader interface { + LoadPolicy() error +} + +func reloadCasbinPolicies(enforcers map[string]*casbin.SyncedEnforcer) error { + loaders := make(map[string]policyLoader, len(enforcers)) + for key, enforcer := range enforcers { + loaders[key] = enforcer + } + return reloadPolicies(loaders) +} + +func reloadPolicies(loaders map[string]policyLoader) error { + for key, loader := range loaders { + if loader == nil { + return fmt.Errorf("reload GoAuto purchaser permissions for casbin %q: enforcer is nil", key) + } + if err := loader.LoadPolicy(); err != nil { + return fmt.Errorf("reload GoAuto purchaser permissions for casbin %q: %w", key, err) + } + } + return nil +} + //var Router runtime.Router func tip() { diff --git a/server/cmd/api/server_policy_test.go b/server/cmd/api/server_policy_test.go new file mode 100644 index 0000000..7bfc9ef --- /dev/null +++ b/server/cmd/api/server_policy_test.go @@ -0,0 +1,38 @@ +package api + +import ( + "errors" + "strings" + "testing" +) + +type stubPolicyLoader struct { + calls int + err error +} + +func (loader *stubPolicyLoader) LoadPolicy() error { + loader.calls++ + return loader.err +} + +func TestReloadPoliciesRefreshesEveryRuntimeEnforcer(t *testing.T) { + first, second := &stubPolicyLoader{}, &stubPolicyLoader{} + if err := reloadPolicies(map[string]policyLoader{"*": first, "tenant": second}); err != nil { + t.Fatal(err) + } + if first.calls != 1 || second.calls != 1 { + t.Fatalf("reload calls first=%d second=%d", first.calls, second.calls) + } +} + +func TestReloadPoliciesFailsClosed(t *testing.T) { + failure := &stubPolicyLoader{err: errors.New("adapter unavailable")} + err := reloadPolicies(map[string]policyLoader{"*": failure}) + if err == nil || !strings.Contains(err.Error(), "adapter unavailable") || failure.calls != 1 { + t.Fatalf("unexpected reload failure: %v calls=%d", err, failure.calls) + } + if err := reloadPolicies(map[string]policyLoader{"*": nil}); err == nil { + t.Fatal("nil enforcer was accepted") + } +}