198 lines
5.2 KiB
Go
198 lines
5.2 KiB
Go
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()
|
|
}
|
|
}
|