diff --git a/redis/middleware.go b/redis/middleware.go new file mode 100644 index 0000000000000000000000000000000000000000..af718d138ab8c542708b50c6a42e23936cb53121 --- /dev/null +++ b/redis/middleware.go @@ -0,0 +1,34 @@ +package redis + +import ( + "context" + "errors" + "net/http" + + goRedis "github.com/go-redis/redis/v8" +) + +var redisCtxKey = &contextKey{"redis"} + +type contextKey struct { + name string +} + +func Middleware(client *goRedis.Client) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ctx := context.WithValue(r.Context(), redisCtxKey, client) + + r = r.WithContext(ctx) + next.ServeHTTP(w, r) + }) + } +} + +func ForContext(ctx context.Context) *goRedis.Client { + raw, ok := ctx.Value(redisCtxKey).(*goRedis.Client) + if !ok { + panic(errors.New("Invalid redis context")) + } + return raw +}