live-sfu-demo/internal/auth/middleware.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 }