package usertoken import ( "context" "testing" "time" "google.golang.org/grpc" "google.golang.org/grpc/codes" "google.golang.org/grpc/metadata" "google.golang.org/grpc/status" ) func TestUnaryInterceptorForwardsBearer(t *testing.T) { key := testKey(t) signer := testSigner(t, key) auth := testAuth(t, key) groups := []string{"family", "mgmnt-admin"} raw, err := signer.Mint(time.Now(), MintInput{Subject: "konrad", Groups: &groups}) if err != nil { t.Fatalf("Mint: %v", err) } var gotSubject string var gotGroups []string var gotRaw string interceptor := auth.UnaryServerInterceptor() ctx := metadata.NewIncomingContext(t.Context(), metadata.Pairs("authorization", "Bearer "+raw)) _, err = interceptor(ctx, nil, &grpc.UnaryServerInfo{FullMethod: "/mgmnt.lists.v1.ListService/GetMe"}, func(ctx context.Context, _ any) (any, error) { c, ok := CallerFromContext(ctx) if !ok { t.Fatal("missing caller") } gotSubject = c.Subject gotGroups = c.Groups var okRaw bool gotRaw, okRaw = BearerFromContext(ctx) if !okRaw { t.Fatal("missing bearer") } return nil, nil }) if err != nil { t.Fatalf("interceptor: %v", err) } if gotSubject != "konrad" || gotRaw != raw { t.Fatalf("subject %q raw match %v", gotSubject, gotRaw == raw) } if len(gotGroups) != 2 || gotGroups[0] != "family" || gotGroups[1] != "mgmnt-admin" { t.Fatalf("groups = %#v", gotGroups) } } func TestUnaryInterceptorRejects(t *testing.T) { key := testKey(t) signer := testSigner(t, key) auth := testAuth(t, key) raw, err := signer.Mint(time.Now(), MintInput{Subject: "konrad"}) if err != nil { t.Fatalf("Mint: %v", err) } interceptor := auth.UnaryServerInterceptor() handler := func(context.Context, any) (any, error) { t.Fatal("handler ran") return nil, nil } info := &grpc.UnaryServerInfo{FullMethod: "/mgmnt.lists.v1.ListService/GetMe"} cases := []struct { name string ctx context.Context }{ {name: "missing metadata", ctx: t.Context()}, {name: "missing bearer", ctx: metadata.NewIncomingContext(t.Context(), metadata.Pairs("authorization", ""))}, {name: "bad signature", ctx: metadata.NewIncomingContext(t.Context(), metadata.Pairs("authorization", "Bearer "+raw+"x"))}, } for _, tt := range cases { t.Run(tt.name, func(t *testing.T) { _, err := interceptor(tt.ctx, nil, info, handler) if status.Code(err) != codes.Unauthenticated { t.Fatalf("code = %v, err = %v", status.Code(err), err) } }) } wrongIss, err := New(t.Context(), Config{PublicKey: &key.PublicKey, Issuer: "other", Audience: "mgmnt"}) if err != nil { t.Fatalf("New: %v", err) } ctx := metadata.NewIncomingContext(t.Context(), metadata.Pairs("authorization", "Bearer "+raw)) _, err = wrongIss.UnaryServerInterceptor()(ctx, nil, info, handler) if status.Code(err) != codes.Unauthenticated { t.Fatalf("wrong issuer code = %v, err = %v", status.Code(err), err) } wrongAud, err := New(t.Context(), Config{PublicKey: &key.PublicKey, Issuer: "manager", Audience: "other"}) if err != nil { t.Fatalf("New: %v", err) } _, err = wrongAud.UnaryServerInterceptor()(ctx, nil, info, handler) if status.Code(err) != codes.Unauthenticated { t.Fatalf("wrong audience code = %v, err = %v", status.Code(err), err) } }