feat(backend): establish gin sqlite service skeleton

This commit is contained in:
QiuSW
2026-07-25 22:29:56 +08:00
parent 84f2da3cf7
commit ab21219a07
32 changed files with 1854 additions and 42 deletions
+127
View File
@@ -0,0 +1,127 @@
package config
import (
"errors"
"net"
"path/filepath"
"strconv"
"strings"
"time"
)
const (
HTTPAddressEnvironment = "CMROUBAO_HTTP_ADDR"
DatabasePathEnvironment = "CMROUBAO_DATABASE_PATH"
defaultHTTPAddress = "127.0.0.1:8080"
defaultDatabasePath = "var/cmroubao.db"
)
type LookupEnvironment func(string) (string, bool)
type Config struct {
HTTPAddress string
DatabasePath string
ReadHeaderTimeout time.Duration
ReadTimeout time.Duration
WriteTimeout time.Duration
IdleTimeout time.Duration
ShutdownTimeout time.Duration
MaxHeaderBytes int
}
func Load(lookup LookupEnvironment) (Config, error) {
httpAddress, err := environmentValue(
lookup,
HTTPAddressEnvironment,
defaultHTTPAddress,
)
if err != nil {
return Config{}, err
}
if err := validateHTTPAddress(httpAddress); err != nil {
return Config{}, err
}
databasePath, err := environmentValue(
lookup,
DatabasePathEnvironment,
defaultDatabasePath,
)
databasePath, err = validatedDatabasePath(databasePath, err)
if err != nil {
return Config{}, err
}
return Config{
HTTPAddress: httpAddress,
DatabasePath: filepath.Clean(databasePath),
ReadHeaderTimeout: 5 * time.Second,
ReadTimeout: 15 * time.Second,
WriteTimeout: 30 * time.Second,
IdleTimeout: 60 * time.Second,
ShutdownTimeout: 10 * time.Second,
MaxHeaderBytes: 1 << 20,
}, nil
}
func LoadDatabasePath(lookup LookupEnvironment) (string, error) {
databasePath, err := environmentValue(
lookup,
DatabasePathEnvironment,
defaultDatabasePath,
)
return validatedDatabasePath(databasePath, err)
}
func environmentValue(
lookup LookupEnvironment,
name string,
defaultValue string,
) (string, error) {
value, exists := lookup(name)
if !exists {
return defaultValue, nil
}
value = strings.TrimSpace(value)
if value == "" {
return "", errors.New(name + " must not be blank")
}
return value, nil
}
func validateHTTPAddress(address string) error {
host, portValue, err := net.SplitHostPort(address)
if err != nil || strings.TrimSpace(host) == "" {
return errors.New(HTTPAddressEnvironment + " must include a host and port")
}
port, err := strconv.Atoi(portValue)
if err != nil || port < 1 || port > 65535 {
return errors.New(HTTPAddressEnvironment + " port must be between 1 and 65535")
}
return nil
}
func validatedDatabasePath(path string, previousError error) (string, error) {
if previousError != nil {
return "", previousError
}
if strings.ContainsRune(path, '\x00') {
return "", errors.New(
DatabasePathEnvironment + " contains an invalid character",
)
}
lowerPath := strings.ToLower(path)
cleanPath := filepath.Clean(path)
extension := strings.ToLower(filepath.Ext(cleanPath))
if cleanPath == "." ||
cleanPath == string(filepath.Separator) ||
lowerPath == ":memory:" ||
strings.HasPrefix(lowerPath, "file:") ||
(extension != ".db" && extension != ".sqlite" && extension != ".sqlite3") {
return "", errors.New(
DatabasePathEnvironment + " must be a SQLite file path",
)
}
return cleanPath, nil
}
+133
View File
@@ -0,0 +1,133 @@
package config
import (
"path/filepath"
"testing"
"time"
)
func TestLoadUsesSafeDefaults(t *testing.T) {
cfg, err := Load(emptyEnvironment)
if err != nil {
t.Fatalf("Load() error = %v", err)
}
if cfg.HTTPAddress != "127.0.0.1:8080" {
t.Fatalf("HTTPAddress = %q", cfg.HTTPAddress)
}
if cfg.DatabasePath != filepath.FromSlash("var/cmroubao.db") {
t.Fatalf("DatabasePath = %q", cfg.DatabasePath)
}
if cfg.ReadHeaderTimeout <= 0 ||
cfg.ReadTimeout <= 0 ||
cfg.WriteTimeout <= 0 ||
cfg.IdleTimeout <= 0 ||
cfg.ShutdownTimeout <= 0 ||
cfg.MaxHeaderBytes <= 0 {
t.Fatal("server safety limits must all be positive")
}
if cfg.ShutdownTimeout > 30*time.Second {
t.Fatalf("ShutdownTimeout = %s", cfg.ShutdownTimeout)
}
}
func TestLoadAcceptsExplicitConfiguration(t *testing.T) {
values := map[string]string{
HTTPAddressEnvironment: "192.0.2.10:9090",
DatabasePathEnvironment: "tmp/test.db",
}
cfg, err := Load(mapEnvironment(values))
if err != nil {
t.Fatalf("Load() error = %v", err)
}
if cfg.HTTPAddress != values[HTTPAddressEnvironment] {
t.Fatalf("HTTPAddress = %q", cfg.HTTPAddress)
}
if cfg.DatabasePath != filepath.Clean(values[DatabasePathEnvironment]) {
t.Fatalf("DatabasePath = %q", cfg.DatabasePath)
}
}
func TestLoadRejectsUnsafeOrInvalidValues(t *testing.T) {
tests := []struct {
name string
values map[string]string
}{
{
name: "blank explicit address",
values: map[string]string{
HTTPAddressEnvironment: " ",
},
},
{
name: "address without host",
values: map[string]string{
HTTPAddressEnvironment: ":8080",
},
},
{
name: "invalid port",
values: map[string]string{
HTTPAddressEnvironment: "127.0.0.1:70000",
},
},
{
name: "memory database",
values: map[string]string{
DatabasePathEnvironment: ":memory:",
},
},
{
name: "database DSN",
values: map[string]string{
DatabasePathEnvironment: "file:test.db?mode=memory",
},
},
{
name: "database directory",
values: map[string]string{
DatabasePathEnvironment: "./",
},
},
{
name: "non SQLite extension",
values: map[string]string{
DatabasePathEnvironment: "var/database.txt",
},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if _, err := Load(mapEnvironment(test.values)); err == nil {
t.Fatal("Load() error = nil")
}
})
}
}
func TestLoadDatabasePathIgnoresHTTPConfiguration(t *testing.T) {
path, err := LoadDatabasePath(mapEnvironment(map[string]string{
HTTPAddressEnvironment: "invalid",
DatabasePathEnvironment: "tmp/migration.sqlite",
}))
if err != nil {
t.Fatalf("LoadDatabasePath() error = %v", err)
}
if path != filepath.Clean("tmp/migration.sqlite") {
t.Fatalf("path = %q", path)
}
}
func emptyEnvironment(string) (string, bool) {
return "", false
}
func mapEnvironment(values map[string]string) LookupEnvironment {
return func(name string) (string, bool) {
value, exists := values[name]
return value, exists
}
}
@@ -0,0 +1,87 @@
package database
import (
"context"
"database/sql"
"errors"
"net/url"
"os"
"path/filepath"
"strings"
_ "github.com/mattn/go-sqlite3"
)
const busyTimeoutMilliseconds = 5000
type safeError struct {
message string
cause error
}
func (e safeError) Error() string {
return e.message
}
func (e safeError) Unwrap() error {
return e.cause
}
func Open(ctx context.Context, path string) (*sql.DB, error) {
absolutePath, err := filepath.Abs(path)
if err != nil {
return nil, errors.New("resolve SQLite path")
}
if err := os.MkdirAll(filepath.Dir(absolutePath), 0o700); err != nil {
return nil, errors.New("create SQLite directory")
}
if info, err := os.Stat(absolutePath); err == nil && info.IsDir() {
return nil, errors.New("SQLite path is a directory")
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, errors.New("inspect SQLite path")
}
db, err := sql.Open("sqlite3", dataSourceName(absolutePath))
if err != nil {
return nil, errors.New("initialize SQLite driver")
}
db.SetMaxOpenConns(1)
db.SetMaxIdleConns(1)
if err := db.PingContext(ctx); err != nil {
_ = db.Close()
return nil, safeError{
message: "connect to SQLite",
cause: err,
}
}
return db, nil
}
func dataSourceName(absolutePath string) string {
slashPath := filepath.ToSlash(absolutePath)
if filepath.VolumeName(absolutePath) != "" &&
!strings.HasPrefix(slashPath, "/") {
slashPath = "/" + slashPath
}
dsnURL := &url.URL{
Scheme: "file",
Path: slashPath,
}
query := dsnURL.Query()
query.Set("_busy_timeout", "5000")
query.Set("_foreign_keys", "on")
query.Set("_journal_mode", "WAL")
query.Set("_txlock", "immediate")
dsnURL.RawQuery = query.Encode()
return dsnURL.String()
}
func isSafeDataSourceName(value string) bool {
lowerValue := strings.ToLower(value)
return strings.HasPrefix(lowerValue, "file:") &&
strings.Contains(lowerValue, "_busy_timeout=") &&
strings.Contains(lowerValue, "_foreign_keys=") &&
strings.Contains(lowerValue, "_journal_mode=wal") &&
strings.Contains(lowerValue, "_txlock=immediate")
}
@@ -0,0 +1,86 @@
package database
import (
"context"
"database/sql"
"errors"
"path/filepath"
"testing"
)
func TestOpenConfiguresSQLiteAndClosesCleanly(t *testing.T) {
path := filepath.Join(t.TempDir(), "nested", "test.db")
db, err := Open(context.Background(), path)
if err != nil {
t.Fatalf("Open() error = %v; cause = %v", err, errors.Unwrap(err))
}
assertPragmaInt(t, db, "foreign_keys", 1)
assertPragmaInt(t, db, "busy_timeout", busyTimeoutMilliseconds)
assertPragmaString(t, db, "journal_mode", "wal")
if db.Stats().MaxOpenConnections != 1 {
t.Fatalf("MaxOpenConnections = %d", db.Stats().MaxOpenConnections)
}
if err := db.Close(); err != nil {
t.Fatalf("Close() error = %v", err)
}
if err := db.PingContext(context.Background()); err == nil {
t.Fatal("PingContext() after Close() error = nil")
}
}
func TestOpenEnforcesForeignKeys(t *testing.T) {
db, err := Open(
context.Background(),
filepath.Join(t.TempDir(), "foreign-keys.db"),
)
if err != nil {
t.Fatalf("Open() error = %v; cause = %v", err, errors.Unwrap(err))
}
t.Cleanup(func() { _ = db.Close() })
if _, err := db.Exec(`
CREATE TABLE parent (id INTEGER PRIMARY KEY);
CREATE TABLE child (
id INTEGER PRIMARY KEY,
parent_id INTEGER NOT NULL REFERENCES parent(id)
);
`); err != nil {
t.Fatalf("create tables: %v", err)
}
if _, err := db.Exec(
"INSERT INTO child (id, parent_id) VALUES (1, 999)",
); err == nil {
t.Fatal("foreign key violation error = nil")
}
}
func TestDataSourceNameIncludesRequiredOptions(t *testing.T) {
dsn := dataSourceName(filepath.Join(t.TempDir(), "test.db"))
if !isSafeDataSourceName(dsn) {
t.Fatalf("unsafe DSN options")
}
}
func assertPragmaInt(t *testing.T, db *sql.DB, name string, want int) {
t.Helper()
var got int
if err := db.QueryRow("PRAGMA " + name).Scan(&got); err != nil {
t.Fatalf("PRAGMA %s: %v", name, err)
}
if got != want {
t.Fatalf("PRAGMA %s = %d, want %d", name, got, want)
}
}
func assertPragmaString(t *testing.T, db *sql.DB, name, want string) {
t.Helper()
var got string
if err := db.QueryRow("PRAGMA " + name).Scan(&got); err != nil {
t.Fatalf("PRAGMA %s: %v", name, err)
}
if got != want {
t.Fatalf("PRAGMA %s = %q, want %q", name, got, want)
}
}
@@ -0,0 +1,57 @@
package migration
import (
"context"
"database/sql"
"cmroubao/backend-api/migrations"
"github.com/pressly/goose/v3"
)
type Status struct {
Version int64
Applied bool
}
type Runner struct {
provider *goose.Provider
}
func New(db *sql.DB) (*Runner, error) {
provider, err := goose.NewProvider(
goose.DialectSQLite3,
db,
migrations.Files,
goose.WithDisableGlobalRegistry(true),
)
if err != nil {
return nil, err
}
return &Runner{provider: provider}, nil
}
func (r *Runner) Up(ctx context.Context) (int, error) {
results, err := r.provider.Up(ctx)
return len(results), err
}
func (r *Runner) Down(ctx context.Context) error {
_, err := r.provider.Down(ctx)
return err
}
func (r *Runner) Status(ctx context.Context) ([]Status, error) {
results, err := r.provider.Status(ctx)
if err != nil {
return nil, err
}
statuses := make([]Status, 0, len(results))
for _, result := range results {
statuses = append(statuses, Status{
Version: result.Source.Version,
Applied: result.State == goose.StateApplied,
})
}
return statuses, nil
}
@@ -0,0 +1,69 @@
package migration
import (
"context"
"path/filepath"
"testing"
"cmroubao/backend-api/internal/platform/database"
)
func TestRunnerSupportsUpStatusDownAndIdempotentUp(t *testing.T) {
db, err := database.Open(
context.Background(),
filepath.Join(t.TempDir(), "migration.db"),
)
if err != nil {
t.Fatalf("database.Open() error = %v", err)
}
t.Cleanup(func() { _ = db.Close() })
runner, err := New(db)
if err != nil {
t.Fatalf("New() error = %v", err)
}
applied, err := runner.Up(context.Background())
if err != nil {
t.Fatalf("Up() error = %v", err)
}
if applied != 1 {
t.Fatalf("Up() applied = %d, want 1", applied)
}
assertStatus(t, runner, true)
applied, err = runner.Up(context.Background())
if err != nil {
t.Fatalf("second Up() error = %v", err)
}
if applied != 0 {
t.Fatalf("second Up() applied = %d, want 0", applied)
}
if err := runner.Down(context.Background()); err != nil {
t.Fatalf("Down() error = %v", err)
}
assertStatus(t, runner, false)
applied, err = runner.Up(context.Background())
if err != nil {
t.Fatalf("final Up() error = %v", err)
}
if applied != 1 {
t.Fatalf("final Up() applied = %d, want 1", applied)
}
}
func assertStatus(t *testing.T, runner *Runner, applied bool) {
t.Helper()
statuses, err := runner.Status(context.Background())
if err != nil {
t.Fatalf("Status() error = %v", err)
}
if len(statuses) != 1 {
t.Fatalf("Status() count = %d, want 1", len(statuses))
}
if statuses[0].Version != 1 || statuses[0].Applied != applied {
t.Fatalf("Status() = %+v", statuses[0])
}
}
@@ -0,0 +1,140 @@
package httpapi
import (
"context"
"crypto/rand"
"encoding/binary"
"errors"
"fmt"
"net/http"
"sync/atomic"
"time"
"github.com/gin-gonic/gin"
)
type DatabasePinger interface {
PingContext(context.Context) error
}
type EventLogger func(string)
func NewRouter(
database DatabasePinger,
logEvent EventLogger,
) (http.Handler, error) {
if database == nil {
return nil, errors.New("database pinger is required")
}
if logEvent == nil {
return nil, errors.New("event logger is required")
}
gin.SetMode(gin.ReleaseMode)
router := gin.New()
router.Use(requestIDMiddleware())
router.Use(safeRecovery(logEvent))
router.HandleMethodNotAllowed = true
if err := router.SetTrustedProxies(nil); err != nil {
return nil, err
}
router.GET("/healthz", healthHandler(database))
router.NoRoute(func(ctx *gin.Context) {
ctx.JSON(http.StatusNotFound, errorResponse(
ctx,
"NOT_FOUND",
"resource not found",
))
})
router.NoMethod(func(ctx *gin.Context) {
ctx.JSON(http.StatusMethodNotAllowed, errorResponse(
ctx,
"METHOD_NOT_ALLOWED",
"method not allowed",
))
})
return router, nil
}
func safeRecovery(logEvent EventLogger) gin.HandlerFunc {
return func(ctx *gin.Context) {
defer func() {
if recover() != nil {
logEvent("HTTP handler panic recovered")
ctx.AbortWithStatusJSON(
http.StatusInternalServerError,
errorResponse(
ctx,
"INTERNAL_ERROR",
"internal server error",
),
)
}
}()
ctx.Next()
}
}
func requestIDMiddleware() gin.HandlerFunc {
return func(ctx *gin.Context) {
requestID := newRequestID()
ctx.Set(requestIDContextKey, requestID)
ctx.Header(requestIDHeader, requestID)
ctx.Next()
}
}
func healthHandler(database DatabasePinger) gin.HandlerFunc {
return func(ctx *gin.Context) {
pingContext, cancel := context.WithTimeout(ctx.Request.Context(), time.Second)
defer cancel()
ctx.Header("Cache-Control", "no-store")
if err := database.PingContext(pingContext); err != nil {
ctx.JSON(http.StatusServiceUnavailable, gin.H{
"status": "unavailable",
})
return
}
ctx.JSON(http.StatusOK, gin.H{"status": "ok"})
}
}
func errorResponse(ctx *gin.Context, code, message string) gin.H {
requestID, _ := ctx.Get(requestIDContextKey)
return gin.H{
"error": gin.H{
"code": code,
"message": message,
"retryable": false,
"details": gin.H{},
},
"request_id": requestID,
}
}
func newRequestID() string {
var value [16]byte
if _, err := rand.Read(value[:]); err != nil {
now := uint64(time.Now().UnixNano())
binary.BigEndian.PutUint64(value[:8], now)
binary.BigEndian.PutUint64(value[8:], fallbackRequestID.Add(1))
}
value[6] = (value[6] & 0x0f) | 0x40
value[8] = (value[8] & 0x3f) | 0x80
return fmt.Sprintf(
"%08x-%04x-%04x-%04x-%012x",
value[0:4],
value[4:6],
value[6:8],
value[8:10],
value[10:16],
)
}
const (
requestIDContextKey = "request_id"
requestIDHeader = "X-Request-ID"
)
var fallbackRequestID atomic.Uint64
@@ -0,0 +1,199 @@
package httpapi
import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"regexp"
"strings"
"testing"
)
type fakePinger struct {
err error
}
func (p fakePinger) PingContext(context.Context) error {
return p.err
}
func TestHealthzReturnsStableHealthyResponse(t *testing.T) {
router, err := NewRouter(fakePinger{}, discardEvent)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
response := performRequest(t, router, http.MethodGet, "/healthz")
if response.Code != http.StatusOK {
t.Fatalf("status = %d", response.Code)
}
assertJSON(t, response, map[string]any{"status": "ok"})
if response.Header().Get("Cache-Control") != "no-store" {
t.Fatalf("Cache-Control = %q", response.Header().Get("Cache-Control"))
}
if response.Header().Get("Access-Control-Allow-Origin") != "" {
t.Fatal("default CORS must remain disabled")
}
}
func TestHealthzReturns503WithoutLeakingDatabaseError(t *testing.T) {
router, err := NewRouter(fakePinger{
err: errors.New("private database path and driver details"),
}, discardEvent)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
response := performRequest(t, router, http.MethodGet, "/healthz")
if response.Code != http.StatusServiceUnavailable {
t.Fatalf("status = %d", response.Code)
}
if strings.Contains(response.Body.String(), "private database") {
t.Fatal("health response leaked the database error")
}
assertJSON(t, response, map[string]any{"status": "unavailable"})
}
func TestUnknownRouteAndMethodUseStableErrors(t *testing.T) {
router, err := NewRouter(fakePinger{}, discardEvent)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
notFound := performRequest(t, router, http.MethodGet, "/missing")
if notFound.Code != http.StatusNotFound {
t.Fatalf("not found status = %d", notFound.Code)
}
assertErrorCode(t, notFound, "NOT_FOUND")
notAllowed := performRequest(t, router, http.MethodPost, "/healthz")
if notAllowed.Code != http.StatusMethodNotAllowed {
t.Fatalf("not allowed status = %d", notAllowed.Code)
}
assertErrorCode(t, notAllowed, "METHOD_NOT_ALLOWED")
}
func TestSafeRecoveryReturnsStableErrorWithoutLoggingRequestHeaders(t *testing.T) {
var events []string
router, err := NewRouter(
panicPinger{},
func(event string) {
events = append(events, event)
},
)
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
request := httptest.NewRequest(http.MethodGet, "/healthz", nil)
request.Header.Set("Authorization", "Bearer private-token")
request.Header.Set("Cookie", "session=private-cookie")
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusInternalServerError {
t.Fatalf("status = %d", response.Code)
}
assertErrorCode(t, response, "INTERNAL_ERROR")
if len(events) != 1 || events[0] != "HTTP handler panic recovered" {
t.Fatalf("events = %#v", events)
}
eventText := strings.Join(events, " ")
if strings.Contains(eventText, "private") ||
strings.Contains(eventText, "Bearer") ||
strings.Contains(eventText, "session") {
t.Fatalf("event log leaked request or panic content")
}
}
func TestRouterRequiresDependencies(t *testing.T) {
if _, err := NewRouter(nil, discardEvent); err == nil {
t.Fatal("NewRouter(nil database) error = nil")
}
if _, err := NewRouter(fakePinger{}, nil); err == nil {
t.Fatal("NewRouter(nil logger) error = nil")
}
}
func performRequest(
t *testing.T,
handler http.Handler,
method string,
path string,
) *httptest.ResponseRecorder {
t.Helper()
request := httptest.NewRequest(method, path, nil)
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
return response
}
func assertJSON(
t *testing.T,
response *httptest.ResponseRecorder,
want map[string]any,
) {
t.Helper()
var got map[string]any
if err := json.Unmarshal(response.Body.Bytes(), &got); err != nil {
t.Fatalf("decode JSON: %v", err)
}
if len(got) != len(want) {
t.Fatalf("JSON = %#v, want %#v", got, want)
}
for key, wantValue := range want {
if got[key] != wantValue {
t.Fatalf("JSON[%q] = %#v, want %#v", key, got[key], wantValue)
}
}
}
func assertErrorCode(
t *testing.T,
response *httptest.ResponseRecorder,
want string,
) {
t.Helper()
var body struct {
Error struct {
Code string `json:"code"`
Retryable bool `json:"retryable"`
Details map[string]any `json:"details"`
} `json:"error"`
RequestID string `json:"request_id"`
}
if err := json.Unmarshal(response.Body.Bytes(), &body); err != nil {
t.Fatalf("decode error JSON: %v", err)
}
if body.Error.Code != want {
t.Fatalf("error code = %q, want %q", body.Error.Code, want)
}
if body.Error.Retryable {
t.Fatal("retryable = true")
}
if body.Error.Details == nil || len(body.Error.Details) != 0 {
t.Fatalf("details = %#v", body.Error.Details)
}
if !requestIDPattern.MatchString(body.RequestID) {
t.Fatalf("request_id = %q", body.RequestID)
}
if response.Header().Get(requestIDHeader) != body.RequestID {
t.Fatalf("request ID header does not match body")
}
}
type panicPinger struct{}
func (panicPinger) PingContext(context.Context) error {
panic("private failure detail")
}
func discardEvent(string) {}
var requestIDPattern = regexp.MustCompile(
`^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$`,
)
@@ -0,0 +1,19 @@
package httpapi
import (
"net/http"
"cmroubao/backend-api/internal/config"
)
func NewServer(cfg config.Config, handler http.Handler) *http.Server {
return &http.Server{
Addr: cfg.HTTPAddress,
Handler: handler,
ReadHeaderTimeout: cfg.ReadHeaderTimeout,
ReadTimeout: cfg.ReadTimeout,
WriteTimeout: cfg.WriteTimeout,
IdleTimeout: cfg.IdleTimeout,
MaxHeaderBytes: cfg.MaxHeaderBytes,
}
}
@@ -0,0 +1,88 @@
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()")
}
}