89 lines
2.1 KiB
Go
89 lines
2.1 KiB
Go
package httpapi
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"errors"
|
||
|
|
"io"
|
||
|
|
"net"
|
||
|
|
"net/http"
|
||
|
|
"testing"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"cmroubao/backend-api/internal/config"
|
||
|
|
)
|
||
|
|
|
||
|
|
func TestNewServerAppliesAllSafetyLimits(t *testing.T) {
|
||
|
|
cfg, err := config.Load(func(string) (string, bool) {
|
||
|
|
return "", false
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("config.Load() error = %v", err)
|
||
|
|
}
|
||
|
|
handler := http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})
|
||
|
|
|
||
|
|
server := NewServer(cfg, handler)
|
||
|
|
|
||
|
|
if server.Addr != cfg.HTTPAddress || server.Handler == nil {
|
||
|
|
t.Fatalf("server = %+v", server)
|
||
|
|
}
|
||
|
|
if server.ReadHeaderTimeout != cfg.ReadHeaderTimeout ||
|
||
|
|
server.ReadTimeout != cfg.ReadTimeout ||
|
||
|
|
server.WriteTimeout != cfg.WriteTimeout ||
|
||
|
|
server.IdleTimeout != cfg.IdleTimeout ||
|
||
|
|
server.MaxHeaderBytes != cfg.MaxHeaderBytes {
|
||
|
|
t.Fatal("server safety limits do not match config")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestServerSupportsBoundedGracefulShutdown(t *testing.T) {
|
||
|
|
cfg, err := config.Load(func(string) (string, bool) {
|
||
|
|
return "", false
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("config.Load() error = %v", err)
|
||
|
|
}
|
||
|
|
server := NewServer(
|
||
|
|
cfg,
|
||
|
|
http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) {
|
||
|
|
response.WriteHeader(http.StatusNoContent)
|
||
|
|
}),
|
||
|
|
)
|
||
|
|
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("net.Listen() error = %v", err)
|
||
|
|
}
|
||
|
|
serverErrors := make(chan error, 1)
|
||
|
|
go func() {
|
||
|
|
serverErrors <- server.Serve(listener)
|
||
|
|
}()
|
||
|
|
|
||
|
|
client := http.Client{Timeout: time.Second}
|
||
|
|
response, err := client.Get("http://" + listener.Addr().String())
|
||
|
|
if err != nil {
|
||
|
|
_ = listener.Close()
|
||
|
|
t.Fatalf("GET server: %v", err)
|
||
|
|
}
|
||
|
|
_, _ = io.Copy(io.Discard, response.Body)
|
||
|
|
_ = response.Body.Close()
|
||
|
|
if response.StatusCode != http.StatusNoContent {
|
||
|
|
t.Fatalf("status = %d", response.StatusCode)
|
||
|
|
}
|
||
|
|
|
||
|
|
shutdownContext, cancel := context.WithTimeout(
|
||
|
|
context.Background(),
|
||
|
|
time.Second,
|
||
|
|
)
|
||
|
|
defer cancel()
|
||
|
|
if err := server.Shutdown(shutdownContext); err != nil {
|
||
|
|
t.Fatalf("Shutdown() error = %v", err)
|
||
|
|
}
|
||
|
|
select {
|
||
|
|
case err := <-serverErrors:
|
||
|
|
if !errors.Is(err, http.ErrServerClosed) {
|
||
|
|
t.Fatalf("Serve() error = %v", err)
|
||
|
|
}
|
||
|
|
case <-time.After(time.Second):
|
||
|
|
t.Fatal("Serve() did not stop after Shutdown()")
|
||
|
|
}
|
||
|
|
}
|