package httpserver import ( "context" "errors" "fmt" "net" "net/http" "sync" "time" ) var ( ErrInvalidEndpoint = errors.New("invalid HTTP endpoint") ErrInvalidOptions = errors.New("invalid HTTP server options") ) type Options struct { ReadHeaderTimeout time.Duration ReadTimeout time.Duration WriteTimeout time.Duration IdleTimeout time.Duration ShutdownTimeout time.Duration MaxHeaderBytes int } type Endpoint struct { Name string Listener net.Listener Handler http.Handler } type Binding struct { Name string Address string Handler http.Handler } func DefaultOptions() Options { return Options{ ReadHeaderTimeout: 5 * time.Second, ReadTimeout: 15 * time.Second, WriteTimeout: 30 * time.Second, IdleTimeout: 60 * time.Second, ShutdownTimeout: 15 * time.Second, MaxHeaderBytes: 16 << 10, } } func Serve(ctx context.Context, options Options, endpoints ...Endpoint) error { if ctx == nil || len(endpoints) == 0 { return ErrInvalidEndpoint } resolved, err := resolveOptions(options) if err != nil { return err } seen := make(map[string]struct{}, len(endpoints)) servers := make([]*http.Server, 0, len(endpoints)) for _, endpoint := range endpoints { if endpoint.Name == "" || endpoint.Listener == nil || endpoint.Handler == nil { return ErrInvalidEndpoint } if _, exists := seen[endpoint.Name]; exists { return fmt.Errorf("%w: duplicate name %q", ErrInvalidEndpoint, endpoint.Name) } seen[endpoint.Name] = struct{}{} servers = append(servers, &http.Server{ Handler: endpoint.Handler, ReadHeaderTimeout: resolved.ReadHeaderTimeout, ReadTimeout: resolved.ReadTimeout, WriteTimeout: resolved.WriteTimeout, IdleTimeout: resolved.IdleTimeout, MaxHeaderBytes: resolved.MaxHeaderBytes, }) } type serveResult struct { name string err error } results := make(chan serveResult, len(endpoints)) for index, endpoint := range endpoints { server := servers[index] go func() { results <- serveResult{name: endpoint.Name, err: server.Serve(endpoint.Listener)} }() } var firstErr error received := 0 select { case <-ctx.Done(): case result := <-results: received++ if !errors.Is(result.err, http.ErrServerClosed) { firstErr = fmt.Errorf("serve HTTP endpoint %q: %w", result.name, result.err) } } shutdownCtx, cancel := context.WithTimeout(context.Background(), resolved.ShutdownTimeout) defer cancel() shutdownErrors := make(chan error, len(servers)) var wait sync.WaitGroup for _, server := range servers { wait.Add(1) go func() { defer wait.Done() if shutdownErr := server.Shutdown(shutdownCtx); shutdownErr != nil { _ = server.Close() shutdownErrors <- shutdownErr } }() } wait.Wait() close(shutdownErrors) if firstErr == nil { for shutdownErr := range shutdownErrors { if firstErr == nil { firstErr = fmt.Errorf("shutdown HTTP endpoints: %w", shutdownErr) } } } for received < len(endpoints) { result := <-results received++ if firstErr == nil && !errors.Is(result.err, http.ErrServerClosed) { firstErr = fmt.Errorf("serve HTTP endpoint %q: %w", result.name, result.err) } } return firstErr } func ListenAndServe(ctx context.Context, options Options, bindings ...Binding) error { if ctx == nil || len(bindings) == 0 { return ErrInvalidEndpoint } if _, err := resolveOptions(options); err != nil { return err } seen := make(map[string]struct{}, len(bindings)) for _, binding := range bindings { if binding.Name == "" || binding.Address == "" || binding.Handler == nil { return ErrInvalidEndpoint } if _, exists := seen[binding.Name]; exists { return fmt.Errorf("%w: duplicate name %q", ErrInvalidEndpoint, binding.Name) } seen[binding.Name] = struct{}{} } endpoints := make([]Endpoint, 0, len(bindings)) defer func() { closeEndpoints(endpoints) }() for _, binding := range bindings { listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", binding.Address) if err != nil { return fmt.Errorf("listen HTTP endpoint %q: %w", binding.Name, err) } endpoints = append(endpoints, Endpoint{ Name: binding.Name, Listener: listener, Handler: binding.Handler, }) } return Serve(ctx, options, endpoints...) } func resolveOptions(options Options) (Options, error) { defaults := DefaultOptions() if options.ReadHeaderTimeout == 0 { options.ReadHeaderTimeout = defaults.ReadHeaderTimeout } if options.ReadTimeout == 0 { options.ReadTimeout = defaults.ReadTimeout } if options.WriteTimeout == 0 { options.WriteTimeout = defaults.WriteTimeout } if options.IdleTimeout == 0 { options.IdleTimeout = defaults.IdleTimeout } if options.ShutdownTimeout == 0 { options.ShutdownTimeout = defaults.ShutdownTimeout } if options.MaxHeaderBytes == 0 { options.MaxHeaderBytes = defaults.MaxHeaderBytes } if options.ReadHeaderTimeout < 0 || options.ReadTimeout < 0 || options.WriteTimeout < 0 || options.IdleTimeout < 0 || options.ShutdownTimeout <= 0 || options.MaxHeaderBytes < 0 { return Options{}, ErrInvalidOptions } return options, nil } func closeEndpoints(endpoints []Endpoint) { for _, endpoint := range endpoints { _ = endpoint.Listener.Close() } }