From 848b9fc51a296abdeca9573b78bb98ba605631d2 Mon Sep 17 00:00:00 2001 From: unintendedfraud Date: Fri, 8 Mar 2024 03:28:24 +0100 Subject: [PATCH] experiment --- context/main.go | 19 +++++++++++++------ context/server.go | 40 +++++++++++++++++++++++++++++++++++++--- 2 files changed, 50 insertions(+), 9 deletions(-) diff --git a/context/main.go b/context/main.go index 615e675..d3d5dd1 100644 --- a/context/main.go +++ b/context/main.go @@ -1,19 +1,21 @@ package main import ( + "context" "fmt" "math/rand" "net/http" + "time" ) const PORT = 8080 func main() { - server := Server{} + server := NewServer(3 * time.Second) server.Handle( "/get-random-number", - printHello, + addValueToContext, handleGetRandomNumber, ) @@ -25,14 +27,19 @@ func main() { func handleGetRandomNumber(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - fmt.Println("Random number: ", rand.Int()) + fmt.Println("Random number: ", r.Context().Value("number")) + w.Write([]byte(`{"error": "process timeout"}`)) }) } -func printHello(next http.Handler) http.Handler { +func addValueToContext(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - fmt.Println("hello") + n := rand.Intn(50) + ctx := context.WithValue(r.Context(), "number", n) + fmt.Println("add value to context: ", n) - next.ServeHTTP(w, r) + time.Sleep(10 * time.Second) + + next.ServeHTTP(w, r.WithContext(ctx)) }) } diff --git a/context/server.go b/context/server.go index a8aa1f0..13430bc 100644 --- a/context/server.go +++ b/context/server.go @@ -1,20 +1,32 @@ package main import ( + "context" "fmt" + "log" "net/http" + "time" ) type ServerHandler func(http.Handler) http.Handler type Server struct { + timeout time.Duration } -func (s Server) Handle(addr string, handlers ...ServerHandler) { - http.Handle(addr, handleMiddlewares(handlers)) +func NewServer(timeout time.Duration) *Server { + return &Server{ + timeout, + } } -func (s Server) Listen(port int) error { +func (s *Server) Handle(addr string, handlers ...ServerHandler) { + middlewares := append([]ServerHandler{s.ctxMiddleware}, handlers...) + + http.Handle(addr, handleMiddlewares(middlewares)) +} + +func (s *Server) Listen(port int) error { if err := http.ListenAndServe(fmt.Sprintf(":%d", port), nil); err != nil { return fmt.Errorf("server broke: %s", err.Error()) } @@ -22,6 +34,28 @@ func (s Server) Listen(port int) error { return nil } +func (s *Server) ctxMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ctx, cancel := context.WithTimeout(r.Context(), s.timeout) + defer cancel() + + processDone := make(chan bool) + + go func() { + next.ServeHTTP(w, r.WithContext(ctx)) + processDone <- true + }() + + select { + case <-ctx.Done(): + w.Write([]byte(`{"error": "context expired"}`)) + log.Panicln("context expired") + case <-processDone: + fmt.Println("process done") + } + }) +} + func handleMiddlewares(handlers []ServerHandler) http.Handler { var handler http.Handler