diff --git a/context/main.go b/context/main.go index d3d5dd1..b965b5a 100644 --- a/context/main.go +++ b/context/main.go @@ -3,7 +3,6 @@ package main import ( "context" "fmt" - "math/rand" "net/http" "time" ) @@ -11,12 +10,12 @@ import ( const PORT = 8080 func main() { - server := NewServer(3 * time.Second) + server := NewServer(1 * time.Second) server.Handle( - "/get-random-number", + "/get-value", addValueToContext, - handleGetRandomNumber, + handleGetValue, ) if err := server.Listen(PORT); err != nil { @@ -25,21 +24,76 @@ func main() { fmt.Println("listening on :", PORT) } -func handleGetRandomNumber(next http.Handler) http.Handler { +func handleGetValue(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - fmt.Println("Random number: ", r.Context().Value("number")) - w.Write([]byte(`{"error": "process timeout"}`)) + fmt.Println("starting handleGetValue") + defer fmt.Println("finished executing handleGetNumber") + + message, err := executeHandleGetValue(r.Context(), r) + if err != nil { + w.Write([]byte(fmt.Sprintf("\nError happened: %s", err.Error()))) + fmt.Println("Error happened: ", err) + return + } + + w.Write([]byte(message)) + }) } +func executeHandleGetValue(ctx context.Context, r *http.Request) (string, error) { + chanErr := make(chan error, 1) + + go func() { + fmt.Println("executeHandleGetValue started") + + time.Sleep(2 * time.Second) + + // simulate errors + shouldError := true + + if shouldError { + chanErr <- fmt.Errorf("an error happened during goroutine") + } else { + chanErr <- nil + } + + fmt.Println("executeHandleGetValue ended") + }() + + select { + case <-ctx.Done(): + <-chanErr + return "", ctx.Err() + case err := <-chanErr: + return "finished successfully", err + } +} + func addValueToContext(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - n := rand.Intn(50) - ctx := context.WithValue(r.Context(), "number", n) - fmt.Println("add value to context: ", n) + // Add simple values + ctx1 := context.WithValue(r.Context(), "number", 420) + ctx2 := context.WithValue(ctx1, "sad_message", "RIP Toriyama :(") - time.Sleep(10 * time.Second) + // Add a more complex one + cs := ComplexStruct{ + question: "What is your favourite Dragon Ball character?", + possibleAnswers: []string{ + "Goku", + "Gohan", + "Vegeta", + "You get the idea zz", + }, + } - next.ServeHTTP(w, r.WithContext(ctx)) + ctx3 := context.WithValue(ctx2, "complex_struct", cs) + + next.ServeHTTP(w, r.WithContext(ctx3)) }) } + +type ComplexStruct struct { + question string + possibleAnswers []string +} diff --git a/context/server.go b/context/server.go index 13430bc..c21abf1 100644 --- a/context/server.go +++ b/context/server.go @@ -3,7 +3,6 @@ package main import ( "context" "fmt" - "log" "net/http" "time" ) @@ -21,7 +20,13 @@ func NewServer(timeout time.Duration) *Server { } func (s *Server) Handle(addr string, handlers ...ServerHandler) { - middlewares := append([]ServerHandler{s.ctxMiddleware}, handlers...) + var middlewares []ServerHandler + + if s.timeout == 0 { + middlewares = handlers + } else { + middlewares = append([]ServerHandler{s.timeoutMiddleware}, handlers...) + } http.Handle(addr, handleMiddlewares(middlewares)) } @@ -34,25 +39,12 @@ func (s *Server) Listen(port int) error { return nil } -func (s *Server) ctxMiddleware(next http.Handler) http.Handler { +func (s *Server) timeoutMiddleware(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") - } + next.ServeHTTP(w, r.WithContext(ctx)) }) }