211 lines
5.5 KiB
Go
211 lines
5.5 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"net/http"
|
|
"strings"
|
|
|
|
"google.golang.org/grpc"
|
|
"google.golang.org/grpc/codes"
|
|
"google.golang.org/grpc/metadata"
|
|
"google.golang.org/grpc/status"
|
|
)
|
|
|
|
type contextKey string
|
|
|
|
const (
|
|
ContextUserKey contextKey = "auth_user"
|
|
ContextRoleKey contextKey = "auth_role"
|
|
ContextClaimsKey contextKey = "auth_claims"
|
|
)
|
|
|
|
type AuthedUser struct {
|
|
Username string
|
|
Role string
|
|
Claims *Claims
|
|
Token string
|
|
}
|
|
|
|
func FromContext(ctx context.Context) (*AuthedUser, bool) {
|
|
u, ok := ctx.Value(ContextUserKey).(*AuthedUser)
|
|
return u, ok
|
|
}
|
|
|
|
func WithContext(ctx context.Context, u *AuthedUser) context.Context {
|
|
ctx = context.WithValue(ctx, ContextUserKey, u)
|
|
ctx = context.WithValue(ctx, ContextRoleKey, u.Role)
|
|
return ctx
|
|
}
|
|
|
|
func (m *Manager) HTTPMiddleware(next http.Handler, requireAuth bool) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
token := extractToken(r)
|
|
var authed *AuthedUser
|
|
if token != "" {
|
|
if claims, user, err := m.VerifyToken(token); err == nil {
|
|
authed = &AuthedUser{Username: claims.Username, Role: user.Role, Claims: claims, Token: token}
|
|
_ = m.Enforcer.AddUserRole(authed.Username, authed.Role)
|
|
}
|
|
}
|
|
if authed == nil {
|
|
authed = &AuthedUser{Username: "", Role: RoleGuest}
|
|
if requireAuth {
|
|
http.Error(w, "unauthorized: missing or invalid token", http.StatusUnauthorized)
|
|
return
|
|
}
|
|
}
|
|
ctx := WithContext(r.Context(), authed)
|
|
r = r.WithContext(ctx)
|
|
r.Header.Set("X-Auth-User", authed.Username)
|
|
r.Header.Set("X-Auth-Role", authed.Role)
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
func (m *Manager) RequireAuth(next http.Handler) http.Handler {
|
|
return m.HTTPMiddleware(next, true)
|
|
}
|
|
|
|
func extractToken(r *http.Request) string {
|
|
if h := r.Header.Get("Authorization"); h != "" {
|
|
if strings.HasPrefix(strings.ToLower(h), "bearer ") {
|
|
return strings.TrimSpace(h[7:])
|
|
}
|
|
}
|
|
if c, err := r.Cookie("token"); err == nil && c.Value != "" {
|
|
return c.Value
|
|
}
|
|
if q := r.URL.Query().Get("token"); q != "" {
|
|
return q
|
|
}
|
|
if h := r.Header.Get("X-Token"); h != "" {
|
|
return h
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func (m *Manager) Authorize(r *http.Request, obj, act string) (bool, error) {
|
|
u, _ := FromContext(r.Context())
|
|
username := ""
|
|
role := RoleGuest
|
|
if u != nil {
|
|
username = u.Username
|
|
role = u.Role
|
|
}
|
|
return m.Check(username, role, obj, act)
|
|
}
|
|
|
|
func (m *Manager) AuthorizeMiddleware(obj, act string, needAuth bool) func(http.Handler) http.Handler {
|
|
return func(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
token := extractToken(r)
|
|
var authed *AuthedUser
|
|
if token != "" {
|
|
if claims, user, err := m.VerifyToken(token); err == nil {
|
|
authed = &AuthedUser{Username: claims.Username, Role: user.Role, Claims: claims, Token: token}
|
|
}
|
|
}
|
|
if authed == nil {
|
|
authed = &AuthedUser{Username: "", Role: RoleGuest}
|
|
if needAuth {
|
|
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
|
return
|
|
}
|
|
}
|
|
ctx := WithContext(r.Context(), authed)
|
|
r = r.WithContext(ctx)
|
|
|
|
ok, err := m.Check(authed.Username, authed.Role, obj, act)
|
|
if err != nil {
|
|
http.Error(w, "auth error: "+err.Error(), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
if !ok {
|
|
http.Error(w, "forbidden: role "+authed.Role+" cannot "+act+" "+obj, http.StatusForbidden)
|
|
return
|
|
}
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
}
|
|
|
|
func (m *Manager) UnaryAuthInterceptor() grpc.UnaryServerInterceptor {
|
|
return func(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (interface{}, error) {
|
|
ctx = m.contextWithGRPCAuth(ctx)
|
|
return handler(ctx, req)
|
|
}
|
|
}
|
|
|
|
func (m *Manager) StreamAuthInterceptor() grpc.StreamServerInterceptor {
|
|
return func(srv interface{}, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
|
ctx := m.contextWithGRPCAuth(ss.Context())
|
|
wrapped := &grpcWrappedStream{ServerStream: ss, ctx: ctx}
|
|
return handler(srv, wrapped)
|
|
}
|
|
}
|
|
|
|
func (m *Manager) contextWithGRPCAuth(ctx context.Context) context.Context {
|
|
token := extractGRPC_TOKEN(ctx)
|
|
var authed *AuthedUser
|
|
if token != "" {
|
|
if claims, user, err := m.VerifyToken(token); err == nil {
|
|
authed = &AuthedUser{Username: claims.Username, Role: user.Role, Claims: claims, Token: token}
|
|
}
|
|
}
|
|
if authed == nil {
|
|
authed = &AuthedUser{Username: "", Role: RoleGuest}
|
|
}
|
|
return WithContext(ctx, authed)
|
|
}
|
|
|
|
func extractGRPC_TOKEN(ctx context.Context) string {
|
|
md, ok := metadata.FromIncomingContext(ctx)
|
|
if !ok {
|
|
return ""
|
|
}
|
|
for _, key := range []string{"authorization", "token", "x-token"} {
|
|
if vals := md.Get(key); len(vals) > 0 {
|
|
v := vals[0]
|
|
if strings.HasPrefix(strings.ToLower(v), "bearer ") {
|
|
return strings.TrimSpace(v[7:])
|
|
}
|
|
return strings.TrimSpace(v)
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func RequireGRPCAuth(ctx context.Context) error {
|
|
u, ok := FromContext(ctx)
|
|
if !ok || u.Username == "" || u.Role == RoleGuest {
|
|
return status.Error(codes.Unauthenticated, "unauthorized: missing token")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (m *Manager) CheckGRPC(ctx context.Context, obj, act string) error {
|
|
u, _ := FromContext(ctx)
|
|
username := ""
|
|
role := RoleGuest
|
|
if u != nil {
|
|
username = u.Username
|
|
role = u.Role
|
|
}
|
|
ok, err := m.Check(username, role, obj, act)
|
|
if err != nil {
|
|
return status.Errorf(codes.Internal, "auth error: %v", err)
|
|
}
|
|
if !ok {
|
|
return status.Errorf(codes.PermissionDenied, "forbidden: role %s cannot %s %s", role, act, obj)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
type grpcWrappedStream struct {
|
|
grpc.ServerStream
|
|
ctx context.Context
|
|
}
|
|
|
|
func (w *grpcWrappedStream) Context() context.Context { return w.ctx }
|
|
|