diff --git a/main_test.go b/main_test.go index 238dcb1..24014d4 100644 --- a/main_test.go +++ b/main_test.go @@ -81,6 +81,35 @@ func TestCollectTargetsSkipsCoreWhenIncludeCoreTargetsFalse(t *testing.T) { } } +func TestResolveScanRootsDetectsConsolidatedBackendDir(t *testing.T) { + repoRoot := t.TempDir() + backendDir := filepath.Join(repoRoot, "backend") + protoDir := filepath.Join(repoRoot, "proto") + + for _, dir := range []string{backendDir, protoDir} { + if err := os.MkdirAll(dir, 0755); err != nil { + t.Fatalf("mkdir %s: %v", dir, err) + } + } + if err := os.WriteFile(filepath.Join(backendDir, "go.mod"), []byte("module example.com/cms/backend\n"), 0644); err != nil { + t.Fatalf("write go.mod: %v", err) + } + if err := os.WriteFile(filepath.Join(repoRoot, "buf.yaml"), []byte("version: v2\n"), 0644); err != nil { + t.Fatalf("write buf.yaml: %v", err) + } + + gotBackend, gotRepo, err := resolveScanRoots(backendDir) + if err != nil { + t.Fatalf("resolveScanRoots() error = %v", err) + } + if gotBackend != backendDir { + t.Fatalf("backendDir = %q, want %q", gotBackend, backendDir) + } + if gotRepo != repoRoot { + t.Fatalf("repoRoot = %q, want %q", gotRepo, repoRoot) + } +} + func TestDiscoverOrchestratorTestTargetsRequiresCoreScan(t *testing.T) { repoRoot := blockNinjaRepoRoot() backendDir := filepath.Join(repoRoot, "backend") @@ -189,3 +218,39 @@ var MethodRoles = map[string]Role{ t.Fatalf("admin role = %q, want admin", methods["/example.v1.Service/Admin"]) } } + +func TestExtractRBACMethodRolesParsesSeedTable(t *testing.T) { + // The CMS renamed the static RBAC table to methodRolesSeed, kept behind an + // atomic-pointer copy-on-write live table for thread safety. The extractor must + // recognise methodRolesSeed as well as MethodRoles, or every CMS RPC reads as + // "missing RBAC entry". + dir := t.TempDir() + file := filepath.Join(dir, "interceptor.go") + if err := os.WriteFile(file, []byte(`package rbac + +type Role string + +const ( + RoleAdmin Role = "admin" +) + +var methodRolesSeed = map[string]Role{ + "/example.v1.Service/Public": "", + "/example.v1.Service/Admin": RoleAdmin, +} +`), 0644); err != nil { + t.Fatalf("write interceptor.go: %v", err) + } + + methods, err := extractRBACMethodRoles(file) + if err != nil { + t.Fatalf("extractRBACMethodRoles() error = %v", err) + } + + if methods["/example.v1.Service/Public"] != "public" { + t.Fatalf("public role = %q, want public", methods["/example.v1.Service/Public"]) + } + if methods["/example.v1.Service/Admin"] != "admin" { + t.Fatalf("admin role = %q, want admin", methods["/example.v1.Service/Admin"]) + } +} diff --git a/proto_rbac.go b/proto_rbac.go index e86aef9..5480b18 100644 --- a/proto_rbac.go +++ b/proto_rbac.go @@ -630,7 +630,9 @@ func extractConnectProceduresRecursive(root string, allowedPackagePrefixes ...st return procedures, nil } -// extractRBACMethodRoles parses interceptor.go and extracts all keys and role values from the MethodRoles map. +// extractRBACMethodRoles parses interceptor.go and extracts all keys and role values from the +// MethodRoles map (orchestrator/helpdesk) or the methodRolesSeed map (CMS — the static seed +// table behind the atomic-pointer copy-on-write live table; see interceptor.go). func extractRBACMethodRoles(file string) (map[string]string, error) { fset := token.NewFileSet() f, err := parser.ParseFile(fset, file, nil, 0) @@ -641,9 +643,13 @@ func extractRBACMethodRoles(file string) (map[string]string, error) { methods := make(map[string]string) ast.Inspect(f, func(n ast.Node) bool { - // Find: var MethodRoles = map[string]Role{ ... } + // Find: var MethodRoles = map[string]Role{ ... } (orchestrator/helpdesk) + // or: var methodRolesSeed = map[string]Role{ ... } (CMS copy-on-write seed) vs, ok := n.(*ast.ValueSpec) - if !ok || len(vs.Names) == 0 || vs.Names[0].Name != "MethodRoles" { + if !ok || len(vs.Names) == 0 { + return true + } + if name := vs.Names[0].Name; name != "MethodRoles" && name != "methodRolesSeed" { return true } if len(vs.Values) == 0 { diff --git a/targets.go b/targets.go index f055c76..4bbb2a2 100644 --- a/targets.go +++ b/targets.go @@ -62,6 +62,13 @@ func resolveScanRoots(target string) (backendDir string, repoRoot string, err er return filepath.Join(absTarget, "backend"), absTarget, nil } + if filepath.Base(absTarget) == "backend" && fileExists(filepath.Join(absTarget, "go.mod")) { + parent := filepath.Clean(filepath.Join(absTarget, "..")) + if fileExists(filepath.Join(parent, "buf.yaml")) || dirExists(filepath.Join(parent, "proto")) { + return absTarget, parent, nil + } + } + if dirExists(filepath.Join(absTarget, "cmd", "check-safety")) { parent := filepath.Clean(filepath.Join(absTarget, "..")) if fileExists(filepath.Join(parent, "buf.yaml")) {