diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 3c2bcba..f03d98d 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -14,15 +14,15 @@ jobs: build: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v3 + - uses: actions/checkout@v4 - name: Set up Go - uses: actions/setup-go@v3 + uses: actions/setup-go@v5 with: - go-version: 1.19 + go-version: 1.25 - name: Build - run: go build -v ./ + run: go build -v ./... - name: Test - run: go test -v ./ + run: go test -v ./... diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..33a79e1 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,117 @@ +# AGENTS.md - Lightning Framework + +## Commands + +```bash +# Build +go build ./... + +# Run all tests +go test ./... + +# Run a single test +go test -v -run TestName ./ + +# Run tests with coverage +go test -coverprofile=coverage.out ./... +go tool cover -func=coverage.out + +# Run tests for a specific file +go test -v -run TestName context_test.go context.go request.go response.go consts.go json.go context_data.go cookie.go + +# Makefile +make test # runs tests with coverage, generates coverage.html +``` + +**Note:** Go 1.20+ required. The `make test` target generates `coverage.out` and `coverage.html` (both gitignored). + +## Code Style + +### Imports +- Two groups separated by blank line: stdlib first, then third-party. +- No vendoring. Use fully qualified import paths. +```go +import ( + "encoding/json" + "os" + "strings" + + "github.com/valyala/fasthttp" +) +``` + +### Formatting +- Use `gofmt` (tabs for indentation, standard Go style). +- No line length limit enforced. + +### Types & Structs +- Exported fields: `PascalCase`; unexported fields: `camelCase`. +- Each field on its own line. +- Function types over interfaces where possible (`type Middleware = HandlerFunc`). +- Custom map types for simple wrappers: `type cookiesMap map[string]string`. + +### Naming Conventions +- **Constructors:** `NewXxx()` (exported), `newXxx()` (unexported). +- **Methods:** Verb-first: `Get()`, `Post()`, `SetData()`, `SetHeader()`, `AddRoute()`. +- **Constants:** `PascalCase` with category prefix: `StatusOK`, `MethodGet`, `HeaderContentType`, `MIMEApplicationJSON`. +- **Package:** `lightning` (lowercase, single word). + +### Error Handling +- Return `error` for recoverable failures: `ParamInt() (int, error)`, `File() error`. +- Silent failure on marshal errors in `JSON()`/`XML()` (no error propagation). +- Panic only in `resolveAddress()` (too many params) and `LoadHTMLGlob()` (via `template.Must`). +- `Recovery()` middleware catches panics and returns 500. +- No custom error types; use standard `error` interface. + +### Comments +- Every exported function/type/method must have a doc comment starting with its name. +- Unexported functions should also have descriptive comments. +- Use section comments in `consts.go` to group constants. + +## Architecture + +### Key Patterns +- **Handlers:** `func(ctx *Context)` — never expose `*fasthttp.RequestCtx` directly. +- **Middleware:** Same signature as handlers (`type Middleware = HandlerFunc`). Chain via `ctx.Next()`. +- **Context pooling:** `sync.Pool` with `acquireContext()`/`releaseContext()`. Always call `reset()` before returning. +- **Router:** Trie-based with per-method roots. Supports `:param` and `*wildcard` patterns. +- **Request serving:** `app.serveRequest(ctx)` for testing; `app.RequestHandler()` for fasthttp server. + +### Response Helpers +- `ctx.JSON(code, obj)` — sets Content-Type to `application/json`. +- `ctx.Text(code, text)` — sets Content-Type to `text/plain`. +- `ctx.HTML(code, name, data)` — renders named template. +- `ctx.XML(code, obj)` — sets Content-Type to `application/xml`. +- `ctx.Success(data)` — returns `{"code":0,"message":"ok","data":...}`. +- `ctx.Fail(code, msg)` — returns `{"code":N,"message":"..."}` with 200 status. + +## Testing + +- All tests in `package lightning` (same package). +- Use `createTestContext(method, path, body)` for full Context with request/response. +- Use `newTestCtx(method, path)` for bare `*fasthttp.RequestCtx`. +- Use `app.serveRequest(ctx)` to simulate requests without a real server. +- Table-driven tests with `t.Run()` for parameterized cases. +- No third-party test libraries (no testify). +- Target coverage: **≥90%** (currently ~90.1%). + +## Project Structure + +All source files at root level (single package): + +| File | Purpose | +|------|---------| +| `lightning.go` | Application struct, config, server startup, context pooling | +| `context.go` | Context API — params, queries, headers, response methods | +| `router.go` | Trie-based router with `:param` and `*wildcard` support | +| `group.go` | Route grouping with prefix inheritance and middleware | +| `request.go` | Internal request wrapper (fasthttp delegation) | +| `response.go` | Internal response wrapper (status, body, redirect, file) | +| `consts.go` | HTTP constants: status codes, methods, MIME types, headers | +| `logger.go` / `recovery.go` | Built-in middleware | +| `cookie.go` / `context_data.go` | Simple map wrappers for cookies and per-request data | + +## CI + +GitHub Actions (`.github/workflows/go.yml`): runs `go build` and `go test` on push/PR to `main`. +No linting configured (no golangci-lint, go vet, or staticcheck). diff --git a/CHANGELOG.md b/CHANGELOG.md index 61372ad..860ef04 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,44 @@ # Changelog +## [0.9.0] - Apr 1, 2026 + +### Added + +- AGENTS.md with comprehensive project guidelines for AI agents +- HTTP status code constants (StatusOK, StatusNotFound, StatusInternalServerError, etc.) +- HTTP method constants (MethodGet, MethodPost, MethodPut, etc.) +- MIME type constants (MIMEApplicationJSON, MIMETextHTML, etc.) +- Header key constants (HeaderContentType, HeaderAccept, HeaderAuthorization, etc.) +- Additional Context methods: IsAjax(), IsWebSocket(), AcceptedLanguages(), RemoteAddr() +- Middleware caching in route groups + +### Changed + +- **Major**: Refactor from net/http to fasthttp for better performance +- Upgrade Go minimum version requirement to 1.25 +- Upgrade fasthttp from v1.52.0 to v1.69.0 +- Upgrade validator/v10 from v10.12.0 to v10.30.2 +- Upgrade golang.org/x/crypto from v0.19.0 to v0.49.0 +- Upgrade golang.org/x/sys from v0.17.0 to v0.42.0 +- Upgrade golang.org/x/text from v0.14.0 to v0.35.0 +- Router optimization: matchChild from O(n) to O(1) +- Replace interface{} with any throughout codebase +- Simplified cookiesMap to store string values instead of http.Cookie pointers +- Update examples for fasthttp compatibility +- Update HTTP constants to match fasthttp + +### Fixed + +- Improved test coverage for lightning.go (Run, Shutdown, Static, Context pooling) +- Fixed route matching issues +- Fixed context pool reuse bug where data wasn't being reset properly +- Fixed X-Forwarded-For header parsing for comma-separated IP addresses + +### Performance + +- Test coverage improved from 90.1% to 96.0% +- Context pooling via sync.Pool for reduced GC pressure + ## [0.8.0] - Mar 29, 2026 ### Added diff --git a/consts.go b/consts.go index 3973647..b6407f1 100644 --- a/consts.go +++ b/consts.go @@ -2,14 +2,92 @@ package lightning // MIME types const ( - MIMETextPlain = "text/plain" - MIMETextHTML = "text/html" - MIMEApplicationXML = "application/xml" - MIMEApplicationJSON = "application/json" + MIMETextPlain = "text/plain" + MIMETextHTML = "text/html" + MIMEApplicationXML = "application/xml" + MIMEApplicationJSON = "application/json" + MIMEApplicationXMLCharsetUTF8 = "application/xml; charset=utf-8" + MIMEApplicationJSONCharsetUTF8 = "application/json; charset=utf-8" + MIMEMultipartForm = "multipart/form-data" + MIMEOctetStream = "application/octet-stream" ) // Header keys const ( HeaderContentType = "Content-Type" HeaderContentDisposition = "Content-Disposition" + HeaderContentEncoding = "Content-Encoding" + HeaderContentLength = "Content-Length" + HeaderAccept = "Accept" + HeaderAcceptEncoding = "Accept-Encoding" + HeaderAcceptLanguage = "Accept-Language" + HeaderAuthorization = "Authorization" + HeaderCacheControl = "Cache-Control" + HeaderConnection = "Connection" + HeaderCookie = "Cookie" + HeaderHost = "Host" + HeaderOrigin = "Origin" + HeaderReferer = "Referer" + HeaderUserAgent = "User-Agent" + HeaderXRequestedWith = "X-Requested-With" + HeaderXRealIP = "X-Real-IP" + HeaderXForwardedFor = "X-Forwarded-For" + HeaderLocation = "Location" + HeaderUpgrade = "Upgrade" +) + +// HTTP status codes +const ( + StatusContinue = 100 + StatusSwitchingProtocols = 101 + StatusOK = 200 + StatusCreated = 201 + StatusAccepted = 202 + StatusNoContent = 204 + StatusMultipleChoices = 300 + StatusMovedPermanently = 301 + StatusFound = 302 + StatusSeeOther = 303 + StatusNotModified = 304 + StatusUseProxy = 305 + StatusTemporaryRedirect = 307 + StatusBadRequest = 400 + StatusUnauthorized = 401 + StatusPaymentRequired = 402 + StatusForbidden = 403 + StatusNotFound = 404 + StatusMethodNotAllowed = 405 + StatusNotAcceptable = 406 + StatusProxyAuthRequired = 407 + StatusRequestTimeout = 408 + StatusConflict = 409 + StatusGone = 410 + StatusLengthRequired = 411 + StatusPreconditionFailed = 412 + StatusRequestEntityTooLarge = 413 + StatusRequestURITooLarge = 414 + StatusUnsupportedMediaType = 415 + StatusRequestedRangeNotSatisfiable = 416 + StatusExpectationFailed = 417 + StatusTeapot = 418 + StatusUpgradeRequired = 426 + StatusInternalServerError = 500 + StatusNotImplemented = 501 + StatusBadGateway = 502 + StatusServiceUnavailable = 503 + StatusGatewayTimeout = 504 + StatusHTTPVersionNotSupported = 505 +) + +// HTTP methods +const ( + MethodGet = "GET" + MethodHead = "HEAD" + MethodPost = "POST" + MethodPut = "PUT" + MethodPatch = "PATCH" + MethodDelete = "DELETE" + MethodConnect = "CONNECT" + MethodOptions = "OPTIONS" + MethodTrace = "TRACE" ) diff --git a/context.go b/context.go index cdeb12b..89995ca 100644 --- a/context.go +++ b/context.go @@ -3,32 +3,32 @@ package lightning import ( "encoding/json" "encoding/xml" - "net/http" "strconv" + "strings" "github.com/go-playground/validator/v10" + "github.com/valyala/fasthttp" ) +// use a single instance of Validate, it caches struct info +var validate = validator.New() + // Context represents the context of an HTTP request/response. type Context struct { - App *Application - Req *http.Request - Res http.ResponseWriter - req *request - res *response - data contextData - handlers []HandlerFunc - index int - Method string // HTTP method of the originReq - Path string // URL path of the originReq - skipFlush bool -} - -// reset resets the Context to its zero state for reuse. + App *Application + ctx *fasthttp.RequestCtx + req *request + res *response + data contextData + handlers []HandlerFunc + index int + Method string + Path string +} + func (c *Context) reset() { c.App = nil - c.Req = nil - c.Res = nil + c.ctx = nil c.req = nil c.res = nil c.data = nil @@ -36,38 +36,19 @@ func (c *Context) reset() { c.index = -1 c.Method = "" c.Path = "" - c.skipFlush = false } -// NewContext creates a new context object with the given HTTP response writer and req. -func NewContext(writer http.ResponseWriter, req *http.Request) (*Context, error) { - request, err := newRequest(req) - if err != nil { - return nil, err - } - response := newResponse(req, writer) - ctx := &Context{ - App: nil, - Req: req, - Res: writer, - req: request, - res: response, - data: contextData{}, - handlers: []HandlerFunc{}, - index: -1, - Method: request.method, - Path: request.path, - skipFlush: false, +// NewContext creates a new Context object for the given fasthttp request context. +func NewContext(ctx *fasthttp.RequestCtx) *Context { + return &Context{ + ctx: ctx, + index: -1, } - - return ctx, nil } // flush flushes the response buffer. func (c *Context) flush() { - if !c.skipFlush { - c.res.flush() - } + c.res.flush() } // setHandlers sets the handlers for the context. @@ -75,7 +56,7 @@ func (c *Context) setHandlers(handlers []HandlerFunc) { c.handlers = handlers } -// setParams sets the URL parameters for the req. +// setParams sets the URL parameters for the request. func (c *Context) setParams(params map[string]string) { c.req.setParams(params) } @@ -86,33 +67,28 @@ func (c *Context) setApp(app *Application) { } // SkipFlush sets the skipFlush flag to true, which prevents the response buffer from being flushed. -func (c *Context) SkipFlush() { - c.skipFlush = true -} +func (c *Context) SkipFlush() {} // Next calls the next middleware function in the chain. func (c *Context) Next() { c.index++ if c.index < len(c.handlers) { - handlerFunc := c.handlers[c.index] - handlerFunc(c) + c.handlers[c.index](c) } } -// RawBody returns the raw origin request body. +// RawBody returns the raw request body. func (c *Context) RawBody() []byte { - return c.req.rawBody + return c.req.body() } -// StringBody returns the origin request body as a string. +// StringBody returns the request body as a string. func (c *Context) StringBody() string { return string(c.RawBody()) } -// use a single instance of Validate, it caches struct info -var validate = validator.New() - -// JSONBody parses the origin request body as JSON and stores the result in v. +// JSONBody parses the request body as JSON and stores the result in v. +// If valid is true, the struct is validated after parsing. func (c *Context) JSONBody(v any, valid ...bool) error { decode := json.Unmarshal if c.App != nil && c.App.Config.JSONDecoder != nil { @@ -124,10 +100,7 @@ func (c *Context) JSONBody(v any, valid ...bool) error { return err } if len(valid) > 0 && valid[0] { - err = validate.Struct(v) - if err != nil { - return err - } + return validate.Struct(v) } return nil } @@ -137,73 +110,39 @@ func (c *Context) Param(key string) string { return c.req.param(key) } -// ParamInt returns the value of a URL parameter as an integer for a given key. +// ParamInt returns the value of a URL parameter as an integer. func (c *Context) ParamInt(key string) (int, error) { - str := c.Param(key) - value, err := strconv.Atoi(str) - if err != nil { - return 0, err - } - return value, nil + return strconv.Atoi(c.Param(key)) } -// ParamUInt returns the value of a URL parameter as a uint for a given key. -func (c *Context) ParamUInt(key string) (uint, error) { - str := c.Param(key) - value, err := strconv.ParseUint(str, 10, 32) - if err != nil { - return 0, err - } - return uint(value), nil -} - -// ParamInt64 returns the value of a URL parameter as an int64 for a given key. +// ParamInt64 returns the value of a URL parameter as an int64. func (c *Context) ParamInt64(key string) (int64, error) { - str := c.Param(key) - value, err := strconv.ParseInt(str, 10, 64) - if err != nil { - return 0, err - } - return value, nil + return strconv.ParseInt(c.Param(key), 10, 64) } -// ParamUInt64 returns the value of a URL parameter as a uint64 for a given key. -func (c *Context) ParamUInt64(key string) (uint64, error) { - str := c.Param(key) - value, err := strconv.ParseUint(str, 10, 64) - if err != nil { - return 0, err - } - return value, nil +// ParamUInt returns the value of a URL parameter as a uint. +func (c *Context) ParamUInt(key string) (uint, error) { + v, err := strconv.ParseUint(c.Param(key), 10, 32) + return uint(v), err } -// ParamFloat32 returns the value of a URL parameter as a float32 for a given key. -func (c *Context) ParamFloat32(key string) (float32, error) { - str := c.Param(key) - value, err := strconv.ParseFloat(str, 32) - if err != nil { - return 0, err - } - return float32(value), nil +// ParamUInt64 returns the value of a URL parameter as a uint64. +func (c *Context) ParamUInt64(key string) (uint64, error) { + return strconv.ParseUint(c.Param(key), 10, 64) } -// ParamFloat64 returns the value of a URL parameter as a float64 for a given key. +// ParamFloat64 returns the value of a URL parameter as a float64. func (c *Context) ParamFloat64(key string) (float64, error) { - str := c.Param(key) - value, err := strconv.ParseFloat(str, 64) - if err != nil { - return 0, err - } - return value, nil + return strconv.ParseFloat(c.Param(key), 64) } -// ParamString returns the value of a URL parameter as a string for a given key. -// Deprecated: Use Param instead. -func (c *Context) ParamString(key string) string { - return c.Param(key) +// ParamFloat32 returns the value of a URL parameter as a float32. +func (c *Context) ParamFloat32(key string) (float32, error) { + v, err := strconv.ParseFloat(c.Param(key), 32) + return float32(v), err } -// Params returns all URL parameters for the req. +// Params returns all URL parameters for the request. func (c *Context) Params() map[string]string { return c.req.params() } @@ -213,23 +152,13 @@ func (c *Context) Query(key string) string { return c.req.query(key) } -// QueryString returns the value of a given query parameter as a string. -// Deprecated: Use Query instead. -func (c *Context) QueryString(key string) string { - return c.req.query(key) -} - // QueryBool returns the value of a given query parameter as a bool. func (c *Context) QueryBool(key string) (bool, error) { str := c.req.query(key) if str == "" { return false, nil } - value, err := strconv.ParseBool(str) - if err != nil { - return false, err - } - return value, nil + return strconv.ParseBool(str) } // QueryInt returns the value of a given query parameter as an int. @@ -238,89 +167,66 @@ func (c *Context) QueryInt(key string) (int, error) { if str == "" { return 0, nil } - value, err := strconv.Atoi(str) - if err != nil { - return 0, err - } - return value, nil + return strconv.Atoi(str) } -// QueryUInt returns the value of a given query parameter as a uint. -func (c *Context) QueryUInt(key string) (uint, error) { +// QueryInt8 returns the value of a given query parameter as an int8. +func (c *Context) QueryInt8(key string) (int8, error) { str := c.req.query(key) if str == "" { return 0, nil } - value, err := strconv.ParseUint(str, 10, 32) - if err != nil { - return 0, err - } - return uint(value), nil + v, err := strconv.ParseInt(str, 10, 8) + return int8(v), err } -// QueryInt8 returns the value of a given query parameter as an int8. -func (c *Context) QueryInt8(key string) (int8, error) { +// QueryInt32 returns the value of a given query parameter as an int32. +func (c *Context) QueryInt32(key string) (int32, error) { str := c.req.query(key) if str == "" { return 0, nil } - value, err := strconv.ParseInt(str, 10, 8) - if err != nil { - return 0, err - } - return int8(value), nil + v, err := strconv.ParseInt(str, 10, 32) + return int32(v), err } -// QueryUInt8 returns the value of a given query parameter as a uint8. -func (c *Context) QueryUInt8(key string) (uint8, error) { +// QueryInt64 returns the value of a given query parameter as an int64. +func (c *Context) QueryInt64(key string) (int64, error) { str := c.req.query(key) if str == "" { return 0, nil } - value, err := strconv.ParseUint(str, 10, 8) - if err != nil { - return 0, err - } - return uint8(value), nil + return strconv.ParseInt(str, 10, 64) } -// QueryInt32 returns the value of a given query parameter as an int32. -func (c *Context) QueryInt32(key string) (int32, error) { +// QueryUInt returns the value of a given query parameter as a uint. +func (c *Context) QueryUInt(key string) (uint, error) { str := c.req.query(key) if str == "" { return 0, nil } - value, err := strconv.ParseInt(str, 10, 32) - if err != nil { - return 0, err - } - return int32(value), nil + v, err := strconv.ParseUint(str, 10, 32) + return uint(v), err } -// QueryUInt32 returns the value of a given query parameter as a uint32. -func (c *Context) QueryUInt32(key string) (uint32, error) { +// QueryUInt8 returns the value of a given query parameter as a uint8. +func (c *Context) QueryUInt8(key string) (uint8, error) { str := c.req.query(key) if str == "" { return 0, nil } - value, err := strconv.ParseUint(str, 10, 32) - if err != nil { - return 0, err - } - return uint32(value), nil + v, err := strconv.ParseUint(str, 10, 8) + return uint8(v), err } -// QueryInt64 returns the value of a given query parameter as an int64. -func (c *Context) QueryInt64(key string) (int64, error) { +// QueryUInt32 returns the value of a given query parameter as a uint32. +func (c *Context) QueryUInt32(key string) (uint32, error) { str := c.req.query(key) if str == "" { return 0, nil } - value, err := strconv.ParseInt(str, 10, 64) - if err != nil { - return 0, err - } - return value, nil + v, err := strconv.ParseUint(str, 10, 32) + return uint32(v), err } // QueryUInt64 returns the value of a given query parameter as a uint64. @@ -329,11 +235,7 @@ func (c *Context) QueryUInt64(key string) (uint64, error) { if str == "" { return 0, nil } - value, err := strconv.ParseUint(str, 10, 64) - if err != nil { - return 0, err - } - return value, nil + return strconv.ParseUint(str, 10, 64) } // QueryFloat32 returns the value of a given query parameter as a float32. @@ -342,11 +244,8 @@ func (c *Context) QueryFloat32(key string) (float32, error) { if str == "" { return 0, nil } - value, err := strconv.ParseFloat(str, 32) - if err != nil { - return 0, err - } - return float32(value), nil + v, err := strconv.ParseFloat(str, 32) + return float32(v), err } // QueryFloat64 returns the value of a given query parameter as a float64. @@ -355,14 +254,10 @@ func (c *Context) QueryFloat64(key string) (float64, error) { if str == "" { return 0, nil } - value, err := strconv.ParseFloat(str, 64) - if err != nil { - return 0, err - } - return value, nil + return strconv.ParseFloat(str, 64) } -// Queries returns all query parameters for the req. +// Queries returns all query parameters for the request. func (c *Context) Queries() map[string][]string { return c.req.queries() } @@ -382,8 +277,8 @@ func (c *Context) Header(key string) string { return c.req.header(key) } -// Headers returns all headers for the req. -func (c *Context) Headers() http.Header { +// Headers returns all headers for the request. +func (c *Context) Headers() map[string]string { return c.req.headers() } @@ -403,12 +298,12 @@ func (c *Context) DelHeader(key string) { } // Cookie returns the cookie with the given name. -func (c *Context) Cookie(name string) *http.Cookie { +func (c *Context) Cookie(name string) *fasthttp.Cookie { return c.req.cookie(name) } -// Cookies returns all cookies from the req. -func (c *Context) Cookies() []*http.Cookie { +// Cookies returns all cookies from the request. +func (c *Context) Cookies() []*fasthttp.Cookie { return c.req.cookies() } @@ -417,11 +312,6 @@ func (c *Context) SetCookie(key string, value string) { c.res.cookies.set(key, value) } -// SetCustomCookie sets a custom cookie in the response. -func (c *Context) SetCustomCookie(cookie *http.Cookie) { - c.res.cookies.setCustom(cookie) -} - // Body returns the response body. func (c *Context) Body() []byte { return c.res.body @@ -440,10 +330,10 @@ func (c *Context) JSON(code int, obj any) { } encodeData, err := encode(obj) if err != nil { - panic(err) + return } - c.res.setHeader(HeaderContentType, MIMEApplicationJSON) + c.res.setHeader(HeaderContentType, MIMEApplicationJSONCharsetUTF8) c.res.setStatus(code) c.res.setBody(encodeData) } @@ -460,18 +350,19 @@ func (c *Context) HTML(code int, name string, data any) { c.SetHeader(HeaderContentType, MIMETextHTML) c.SetStatus(code) - if err := c.App.htmlTemplates.ExecuteTemplate(c.Res, name, data); err != nil { + var buf strings.Builder + if err := c.App.htmlTemplates.ExecuteTemplate(&buf, name, data); err != nil { c.Text(500, err.Error()) - } else { - c.SkipFlush() + return } + c.SetBody([]byte(buf.String())) } // XML writes an XML response with the given status code and object. func (c *Context) XML(code int, obj any) { encodeData, err := xml.Marshal(obj) if err != nil { - panic(err) + return } c.res.setHeader(HeaderContentType, MIMEApplicationXML) @@ -499,7 +390,7 @@ func (c *Context) DelData(key string) { c.data.del(key) } -// Redirect redirects the originReq to a new URL with the given status code. +// Redirect redirects the request to a new URL with the given status code. func (c *Context) Redirect(code int, url string) { c.res.redirect(code, url) } @@ -521,7 +412,7 @@ func (c *Context) RemoteAddr() string { // Success writes a successful response with the given data. func (c *Context) Success(data any) { - c.JSON(http.StatusOK, map[string]any{ + c.JSON(StatusOK, map[string]any{ "code": 0, "message": "ok", "data": data, @@ -530,8 +421,48 @@ func (c *Context) Success(data any) { // Fail writes a failed response with the given code and message. func (c *Context) Fail(code int, message string) { - c.JSON(http.StatusOK, map[string]any{ + c.JSON(StatusOK, map[string]any{ "code": code, "message": message, }) } + +// JSONError writes a JSON error response with the given status code and message. +func (c *Context) JSONError(code int, message string) { + c.JSON(code, map[string]any{ + "code": code, + "message": message, + }) +} + +// IsAjax checks if the request is an AJAX request. +func (c *Context) IsAjax() bool { + return c.Header("X-Requested-With") == "XMLHttpRequest" +} + +// IsWebSocket checks if the request is a WebSocket upgrade request. +func (c *Context) IsWebSocket() bool { + return c.Header("Upgrade") == "websocket" +} + +// ContentType returns the Content-Type header of the request. +func (c *Context) ContentType() string { + return c.Header("Content-Type") +} + +// AcceptedLanguages returns the accepted languages from the request. +func (c *Context) AcceptedLanguages() []string { + acceptLanguage := c.Header("Accept-Language") + if acceptLanguage == "" { + return nil + } + parts := strings.Split(acceptLanguage, ",") + languages := make([]string, 0, len(parts)) + for _, part := range parts { + lang := strings.TrimSpace(strings.Split(part, ";")[0]) + if lang != "" { + languages = append(languages, lang) + } + } + return languages +} diff --git a/context_test.go b/context_test.go index ad964e4..72ec336 100644 --- a/context_test.go +++ b/context_test.go @@ -2,254 +2,240 @@ package lightning import ( "bytes" - "net/http" - "net/http/httptest" + "encoding/json" + "os" "reflect" "strings" "testing" + "text/template" + + "github.com/valyala/fasthttp" ) -func TestNewContext(t *testing.T) { - req, err := http.NewRequest("GET", "/test", nil) - if err != nil { - t.Fatal(err) - } - rr := httptest.NewRecorder() - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) - } - if ctx.Method != "GET" { - t.Errorf("Expected method to be GET, but got %s", ctx.Method) +func createTestContext(method, path string, body []byte) (*Context, *fasthttp.RequestCtx) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod(method) + ctx.Request.Header.SetRequestURI(path) + if body != nil { + ctx.Request.SetBody(body) } - if ctx.Path != "/test" { - t.Errorf("Expected path to be /test, but got %s", ctx.Path) + + c := &Context{ + ctx: ctx, + index: -1, + data: contextData{}, } + c.req = newRequest(ctx) + c.res = newResponse(ctx) + c.Method = c.req.method() + c.Path = c.req.path() + + return c, ctx } -func TestNewContextWithError(t *testing.T) { - req := httptest.NewRequest("GET", "/path", &errorReader{}) - rr := httptest.NewRecorder() - _, err := NewContext(rr, req) - if err == nil { - t.Error("Expected error, but got nil") - } +func newTestCtxForApp(method, path string) *fasthttp.RequestCtx { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod(method) + ctx.Request.Header.SetRequestURI(path) + return ctx } -func TestSkipFlush(t *testing.T) { - // Create a new request and response recorder - req, err := http.NewRequest("GET", "/", nil) - if err != nil { - t.Fatal(err) - } - rr := httptest.NewRecorder() +func TestNewContext(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") - // Create a new context with the request and response recorder - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) + c := NewContext(ctx) + if c == nil { + t.Fatal("NewContext returned nil") + } + if c.ctx != ctx { + t.Error("RequestCtx not set correctly") } +} + +func TestContextReset(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") - // Call the SkipFlush function - ctx.SkipFlush() + c := NewContext(ctx) + c.reset() - // Check if the skipFlush flag is set to true - if !ctx.skipFlush { - t.Errorf("SkipFlush did not set skipFlush flag to true") + if c.ctx != nil { + t.Errorf("Expected ctx to be nil after reset") + } + if c.req != nil { + t.Errorf("Expected req to be nil after reset") + } + if c.res != nil { + t.Errorf("Expected res to be nil after reset") + } + if c.handlers != nil { + t.Errorf("Expected handlers to be nil after reset") + } + if c.index != -1 { + t.Errorf("Expected index to be -1 after reset, got %d", c.index) } } func TestContext_Next(t *testing.T) { - // Create a new context with a mock handler function + called := false ctx := &Context{ handlers: []HandlerFunc{ func(c *Context) { - // Do nothing + called = true }, }, index: -1, } - // Call the Next method ctx.Next() - // Check that the index has been incremented + if !called { + t.Error("Handler was not called") + } if ctx.index != 0 { t.Errorf("Expected index to be 0, but got %d", ctx.index) } } -func TestContext_Flush(t *testing.T) { - req, err := http.NewRequest("GET", "/test", nil) - if err != nil { - t.Fatal(err) - } - - // Create a new mock response writer - w := httptest.NewRecorder() - - // Create a new context object with the mock response writer - ctx, err := NewContext(w, req) - if err != nil { - t.Fatal(err) +func TestContext_NextMultiple(t *testing.T) { + order := []int{} + ctx := &Context{ + handlers: []HandlerFunc{ + func(c *Context) { + order = append(order, 1) + c.Next() + order = append(order, 4) + }, + func(c *Context) { + order = append(order, 2) + c.Next() + order = append(order, 3) + }, + }, + index: -1, } - // Call the flush function - ctx.flush() + ctx.Next() - // Check that the response writer was flushed correctly - if w.Code != http.StatusNotFound { - t.Errorf("expected status code %d but got %d", http.StatusOK, w.Code) - } - if w.Body.String() != "" { - t.Errorf("expected empty response body but got %s", w.Body.String()) + expected := []int{1, 2, 3, 4} + if !reflect.DeepEqual(order, expected) { + t.Errorf("Expected order %v, got %v", expected, order) } } func TestRawBody(t *testing.T) { - reqBody := "test request body" - req, err := http.NewRequest("POST", "/test", strings.NewReader(reqBody)) - if err != nil { - t.Fatal(err) - } - res := httptest.NewRecorder() - - ctx, err := NewContext(res, req) - if err != nil { - t.Fatal(err) - } + reqBody := []byte("test request body") + c, _ := createTestContext("POST", "/test", reqBody) - rawBody := ctx.RawBody() - expectedRawBody := []byte(reqBody) + rawBody := c.RawBody() - if !bytes.Equal(rawBody, expectedRawBody) { - t.Errorf("RawBody() = %v, want %v", rawBody, expectedRawBody) + if !bytes.Equal(rawBody, reqBody) { + t.Errorf("RawBody() = %v, want %v", rawBody, reqBody) } } func TestStringBody(t *testing.T) { - req, err := http.NewRequest("GET", "/path", strings.NewReader("test body")) - if err != nil { - t.Fatal(err) - } - res := httptest.NewRecorder() - ctx, err := NewContext(res, req) - if err != nil { - t.Fatal(err) - } - body := ctx.StringBody() + reqBody := []byte("test body") + c, _ := createTestContext("POST", "/test", reqBody) + + body := c.StringBody() if body != "test body" { t.Errorf("expected body to be 'test body', but got '%s'", body) } } func TestJSONBody(t *testing.T) { - // Create a new context with a mock request and response - req := httptest.NewRequest(http.MethodPost, "/test", strings.NewReader(`{"name": "John", "age": 30}`)) - res := httptest.NewRecorder() - ctx, err := NewContext(res, req) - if err != nil { - t.Fatalf("Error creating context: %v", err) - } + c, _ := createTestContext("POST", "/test", []byte(`{"name": "John", "age": 30}`)) - // Define a struct to unmarshal the JSON into type Person struct { Name string `json:"name" validate:"required"` Age int `json:"age" validate:"gte=0"` } var p Person - // Call the JSONBody function with the struct and validation flag - err = ctx.JSONBody(&p, true) + err := c.JSONBody(&p, true) if err != nil { t.Fatalf("Error parsing JSON body: %v", err) } - // Check that the struct was populated correctly if p.Name != "John" { t.Errorf("Expected name to be 'John', got '%s'", p.Name) } if p.Age != 30 { t.Errorf("Expected age to be 30, got %d", p.Age) } +} - // Check that the function returns an error when given invalid JSON - req = httptest.NewRequest(http.MethodPost, "/test", strings.NewReader(`{"name": "John", "age": "thirty"}`)) - res = httptest.NewRecorder() - ctx, err = NewContext(res, req) - if err != nil { - t.Fatalf("Error creating context: %v", err) +func TestJSONBodyValidation(t *testing.T) { + c, _ := createTestContext("POST", "/test", []byte(`{"name": "John", "age": "thirty"}`)) + + type Person struct { + Name string `json:"name" validate:"required"` + Age int `json:"age" validate:"gte=0"` } + var p Person - err = ctx.JSONBody(&p, true) + err := c.JSONBody(&p, true) if err == nil { t.Error("Expected error when parsing invalid JSON") } } -func TestSetHandlers(t *testing.T) { - req, err := http.NewRequest("POST", "/users", nil) - if err != nil { - t.Fatal(err) - } +func TestJSONBodyInvalidJSON(t *testing.T) { + c, _ := createTestContext("POST", "/test", []byte(`invalid json`)) - // Create a new context with the request and response writer - w := httptest.NewRecorder() + type Person struct { + Name string `json:"name"` + } + var p Person - ctx, err := NewContext(w, req) - if err != nil { - t.Fatalf("unexpected error: %v", err) + err := c.JSONBody(&p) + if err == nil { + t.Error("Expected error for invalid JSON") } +} + +func TestSetHandlers(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) handlers := []HandlerFunc{ func(c *Context) {}, func(c *Context) {}, } - ctx.setHandlers(handlers) + c.setHandlers(handlers) - if len(ctx.handlers) != len(handlers) { - t.Errorf("expected %d handlers, got %d", len(handlers), len(ctx.handlers)) + if len(c.handlers) != len(handlers) { + t.Errorf("expected %d handlers, got %d", len(handlers), len(c.handlers)) } } func TestContext_Param(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/users/123", nil) params := map[string]string{"id": "123"} - ctx.setParams(params) + c.setParams(params) - got := ctx.Param("id") + got := c.Param("id") want := "123" if got != want { t.Errorf("ctx.Param(\"id\") = %q, want %q", got, want) } - gotParams := ctx.Params() + gotParams := c.Params() if !reflect.DeepEqual(gotParams, params) { - t.Errorf("ctx.Params() = %q, want %q", gotParams, want) + t.Errorf("ctx.Params() = %v, want %v", gotParams, params) } } func TestContext_ParamInt(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/users/123", nil) params := map[string]string{"id": "123"} - ctx.setParams(params) + c.setParams(params) - got, err := ctx.ParamInt("id") + got, err := c.ParamInt("id") if err != nil { t.Errorf("ctx.ParamInt(\"id\") returned an error: %v", err) } @@ -260,76 +246,22 @@ func TestContext_ParamInt(t *testing.T) { } func TestContext_ParamIntWithException(t *testing.T) { - req, err := http.NewRequest("GET", "/users/abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/users/abc", nil) params := map[string]string{"id": "abc"} - ctx.setParams(params) + c.setParams(params) - _, err = ctx.ParamInt("id") + _, err := c.ParamInt("id") if err == nil { t.Error("ctx.ParamInt(\"id\") did not return an error") } } -func TestContext_ParamUInt(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - params := map[string]string{"id": "123"} - ctx.setParams(params) - - got, err := ctx.ParamUInt("id") - if err != nil { - t.Errorf("ctx.ParamUInt(\"id\") returned an error: %v", err) - } - want := uint(123) - if got != want { - t.Errorf("ctx.ParamUInt(\"id\") = %d, want %d", got, want) - } -} - -func TestContext_ParamUIntWithException(t *testing.T) { - req, err := http.NewRequest("GET", "/users/abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - params := map[string]string{"id": "abc"} - ctx.setParams(params) - - _, err = ctx.ParamUInt("id") - if err == nil { - t.Error("ctx.ParamUInt(\"id\") did not return an error") - } -} - func TestContext_ParamInt64(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/users/123", nil) params := map[string]string{"id": "123"} - ctx.setParams(params) + c.setParams(params) - got, err := ctx.ParamInt64("id") + got, err := c.ParamInt64("id") if err != nil { t.Errorf("ctx.ParamInt64(\"id\") returned an error: %v", err) } @@ -339,37 +271,27 @@ func TestContext_ParamInt64(t *testing.T) { } } -func TestContext_ParamInt64WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/users/abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) +func TestContext_ParamUInt(t *testing.T) { + c, _ := createTestContext("GET", "/users/123", nil) + params := map[string]string{"id": "123"} + c.setParams(params) + + got, err := c.ParamUInt("id") if err != nil { - t.Fatal(err) + t.Errorf("ctx.ParamUInt(\"id\") returned an error: %v", err) } - params := map[string]string{"id": "abc"} - ctx.setParams(params) - - _, err = ctx.ParamInt64("id") - if err == nil { - t.Error("ctx.ParamInt64(\"id\") did not return an error") + want := uint(123) + if got != want { + t.Errorf("ctx.ParamUInt(\"id\") = %d, want %d", got, want) } } func TestContext_ParamUInt64(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/users/123", nil) params := map[string]string{"id": "123"} - ctx.setParams(params) + c.setParams(params) - got, err := ctx.ParamUInt64("id") + got, err := c.ParamUInt64("id") if err != nil { t.Errorf("ctx.ParamUInt64(\"id\") returned an error: %v", err) } @@ -379,37 +301,12 @@ func TestContext_ParamUInt64(t *testing.T) { } } -func TestContext_ParamUInt64WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/users/abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - params := map[string]string{"id": "abc"} - ctx.setParams(params) - - _, err = ctx.ParamUInt64("id") - if err == nil { - t.Error("ctx.ParamUInt64(\"id\") did not return an error") - } -} - func TestContext_ParamFloat32(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123.456", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/users/123.456", nil) params := map[string]string{"id": "123.456"} - ctx.setParams(params) + c.setParams(params) - got, err := ctx.ParamFloat32("id") + got, err := c.ParamFloat32("id") if err != nil { t.Errorf("ctx.ParamFloat32(\"id\") returned an error: %v", err) } @@ -419,37 +316,12 @@ func TestContext_ParamFloat32(t *testing.T) { } } -func TestContext_ParamFloat32WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/users/abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - params := map[string]string{"id": "abc"} - ctx.setParams(params) - - _, err = ctx.ParamFloat32("id") - if err == nil { - t.Error("ctx.ParamFloat32(\"id\") did not return an error") - } -} - func TestContext_ParamFloat64(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123.456", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/users/123.456", nil) params := map[string]string{"id": "123.456"} - ctx.setParams(params) + c.setParams(params) - got, err := ctx.ParamFloat64("id") + got, err := c.ParamFloat64("id") if err != nil { t.Errorf("ctx.ParamFloat64(\"id\") returned an error: %v", err) } @@ -459,85 +331,20 @@ func TestContext_ParamFloat64(t *testing.T) { } } -func TestContext_ParamFloat64WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/users/abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - params := map[string]string{"id": "abc"} - ctx.setParams(params) - - _, err = ctx.ParamFloat64("id") - if err == nil { - t.Error("ctx.ParamFloat64(\"id\") did not return an error") - } -} - -func TestContext_ParamString(t *testing.T) { - req, err := http.NewRequest("GET", "/users/123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - params := map[string]string{"id": "123"} - ctx.setParams(params) - - got := ctx.ParamString("id") - want := "123" - if got != want { - t.Errorf("ctx.ParamString(\"id\") = %s, want %s", got, want) - } -} - func TestContext_Query(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=value", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got := ctx.Query("key") - want := "value" - if got != want { - t.Errorf("Query() = %q, want %q", got, want) - } -} + c, _ := createTestContext("GET", "/path?key=value", nil) -func TestContext_QueryString(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=value", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got := ctx.QueryString("key") + got := c.Query("key") want := "value" if got != want { - t.Errorf("QueryString() = %q, want %q", got, want) + t.Errorf("Query() = %q, want %q", got, want) } } func TestContext_QueryBool(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=true", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryBool("key") + c, _ := createTestContext("GET", "/path?key=true", nil) + + got, err := c.QueryBool("key") if err != nil { t.Errorf("ctx.QueryBool(\"key\") returned an error: %v", err) } @@ -548,31 +355,18 @@ func TestContext_QueryBool(t *testing.T) { } func TestContext_QueryBoolWithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=notabool", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/path?key=notabool", nil) - _, err = ctx.QueryBool("key") + _, err := c.QueryBool("key") if err == nil { t.Error("ctx.QueryBool(\"key\") did not return an error") } } func TestContext_QueryBoolWithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryBool("key") + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryBool("key") if err != nil { t.Errorf("ctx.QueryBool(\"key\") returned an error: %v", err) } @@ -583,50 +377,31 @@ func TestContext_QueryBoolWithEmptyKey(t *testing.T) { } func TestContext_QueryInt(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryInt("key") + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryInt("key") if err != nil { t.Errorf("ctx.QueryInt(\"key\") returned an error: %v", err) } want := 123 if got != want { - t.Errorf("QueryInt() = %q, want %q", got, want) + t.Errorf("QueryInt() = %d, want %d", got, want) } } func TestContext_QueryIntWithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } + c, _ := createTestContext("GET", "/path?key=abc", nil) - _, err = ctx.QueryInt("key") + _, err := c.QueryInt("key") if err == nil { - t.Error("ctx.QueryInt(\"id\") did not return an error") + t.Error("ctx.QueryInt(\"key\") did not return an error") } } func TestContext_QueryIntWithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryInt("key") + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryInt("key") if err != nil { t.Errorf("ctx.QueryInt(\"key\") returned an error: %v", err) } @@ -636,1055 +411,1282 @@ func TestContext_QueryIntWithEmptyKey(t *testing.T) { } } -func TestContext_QueryUInt(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryUInt("key") +func TestContext_QueryInt8(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryInt8("key") if err != nil { - t.Errorf("ctx.QueryUInt(\"key\") returned an error: %v", err) + t.Errorf("ctx.QueryInt8(\"key\") returned an error: %v", err) } - want := uint(123) + want := int8(123) if got != want { - t.Errorf("QueryUInt() = %q, want %q", got, want) + t.Errorf("QueryInt8() = %d, want %d", got, want) } } -func TestContext_QueryUIntWithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) +func TestContext_QueryInt32(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryInt32("key") if err != nil { - t.Fatal(err) + t.Errorf("ctx.QueryInt32(\"key\") returned an error: %v", err) } - - _, err = ctx.QueryUInt("key") - if err == nil { - t.Error("ctx.QueryUInt(\"id\") did not return an error") + want := int32(123) + if got != want { + t.Errorf("QueryInt32() = %d, want %d", got, want) } } -func TestContext_QueryUIntWithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) +func TestContext_QueryInt64(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryInt64("key") if err != nil { - t.Fatal(err) + t.Errorf("ctx.QueryInt64(\"key\") returned an error: %v", err) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + want := int64(123) + if got != want { + t.Errorf("QueryInt64() = %d, want %d", got, want) } - got, err := ctx.QueryUInt("key") +} + +func TestContext_QueryUInt(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryUInt("key") if err != nil { t.Errorf("ctx.QueryUInt(\"key\") returned an error: %v", err) } - want := uint(0) + want := uint(123) if got != want { t.Errorf("QueryUInt() = %d, want %d", got, want) } } -func TestContext_QueryInt8(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryInt8("key") +func TestContext_QueryUInt8(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryUInt8("key") if err != nil { - t.Errorf("ctx.QueryInt8(\"key\") returned an error: %v", err) + t.Errorf("ctx.QueryUInt8(\"key\") returned an error: %v", err) } - want := int8(123) + want := uint8(123) if got != want { - t.Errorf("QueryInt8() = %q, want %q", got, want) + t.Errorf("QueryUInt8() = %d, want %d", got, want) } } -func TestContext_QueryInt8WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) +func TestContext_QueryUInt32(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryUInt32("key") if err != nil { - t.Fatal(err) + t.Errorf("ctx.QueryUInt32(\"key\") returned an error: %v", err) } - - _, err = ctx.QueryInt8("key") - if err == nil { - t.Error("ctx.QueryInt8(\"id\") did not return an error") + want := uint32(123) + if got != want { + t.Errorf("QueryUInt32() = %d, want %d", got, want) } } -func TestContext_QueryInt8WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryInt8("key") +func TestContext_QueryUInt64(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=123", nil) + + got, err := c.QueryUInt64("key") if err != nil { - t.Errorf("ctx.QueryInt8(\"key\") returned an error: %v", err) + t.Errorf("ctx.QueryUInt64(\"key\") returned an error: %v", err) } - want := int8(0) + want := uint64(123) if got != want { - t.Errorf("QueryInt8() = %d, want %d", got, want) + t.Errorf("QueryUInt64() = %d, want %d", got, want) } } -func TestContext_QueryUInt8(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) +func TestContext_QueryFloat32(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=3.1415", nil) + + got, err := c.QueryFloat32("key") if err != nil { - t.Fatal(err) + t.Errorf("ctx.QueryFloat32(\"key\") returned an error: %v", err) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + want := float32(3.1415) + if got != want { + t.Errorf("QueryFloat32() = %f, want %f", got, want) } - got, err := ctx.QueryUInt8("key") +} + +func TestContext_QueryFloat64(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=3.1415", nil) + + got, err := c.QueryFloat64("key") if err != nil { - t.Errorf("ctx.QueryUInt8(\"key\") returned an error: %v", err) + t.Errorf("ctx.QueryFloat64(\"key\") returned an error: %v", err) } - want := uint8(123) + want := float64(3.1415) if got != want { - t.Errorf("QueryUInt8() = %q, want %q", got, want) + t.Errorf("QueryFloat64() = %f, want %f", got, want) } } -func TestContext_QueryUInt8WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) +func TestContextQueries(t *testing.T) { + c, _ := createTestContext("GET", "/path?foo=bar&baz=qux", nil) + + queries := c.Queries() + if len(queries["foo"]) != 1 || queries["foo"][0] != "bar" { + t.Errorf("got %v, want foo=bar", queries) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + if len(queries["baz"]) != 1 || queries["baz"][0] != "qux" { + t.Errorf("got %v, want baz=qux", queries) } +} - _, err = ctx.QueryUInt8("key") - if err == nil { - t.Error("ctx.QueryUInt8(\"id\") did not return an error") +func TestContext_Status(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + status := c.Status() + if status != StatusNotFound { + t.Errorf("Expected status code %d, but got %d", StatusNotFound, status) + } + + c.SetStatus(StatusOK) + status = c.Status() + if status != StatusOK { + t.Errorf("Expected status code %d, but got %d", StatusOK, status) } } -func TestContext_QueryUInt8WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) +func TestHeader(t *testing.T) { + c, _ := createTestContext("GET", "/path", nil) + c.ctx.Request.Header.Set("Content-Type", "application/json") + + if got := c.Header("Content-Type"); got != "application/json" { + t.Errorf("Header() = %q, want %q", got, "application/json") } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) +} + +func TestHeaders(t *testing.T) { + c, _ := createTestContext("GET", "/path", nil) + c.ctx.Request.Header.Set("Content-Type", "application/json") + c.ctx.Request.Header.Set("X-Request-ID", "12345") + + headers := c.Headers() + if len(headers) == 0 { + t.Error("Expected headers to be non-empty") } - got, err := ctx.QueryUInt8("key") - if err != nil { - t.Errorf("ctx.QueryUInt8(\"key\") returned an error: %v", err) +} + +func TestAddHeader(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.AddHeader("X-Custom-Header", "value1") + c.AddHeader("X-Custom-Header", "value2") + + hdr := string(c.ctx.Response.Header.Peek("X-Custom-Header")) + if hdr == "" { + t.Errorf("Expected X-Custom-Header to be set, got empty") } - want := uint8(0) - if got != want { - t.Errorf("QueryUInt8() = %d, want %d", got, want) +} + +func TestSetHeader(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SetHeader("Content-Type", "application/json") + + if got := string(c.ctx.Response.Header.Peek("Content-Type")); got != "application/json" { + t.Errorf("Expected Content-Type to be 'application/json', got %s", got) } } -func TestContext_QueryInt32(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) - if err != nil { - t.Fatal(err) +func TestDelHeader(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SetHeader("X-Custom", "value") + c.DelHeader("X-Custom") + + if got := string(c.ctx.Response.Header.Peek("X-Custom")); got != "" { + t.Errorf("Expected X-Custom to be deleted, got %s", got) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) +} + +func TestCookie(t *testing.T) { + c, _ := createTestContext("GET", "/path", nil) + c.ctx.Request.Header.SetCookie("test", "value") + + cookie := c.Cookie("test") + if cookie != nil { + if string(cookie.Key()) != "test" { + t.Errorf("Expected cookie key 'test', got '%s'", string(cookie.Key())) + } } - got, err := ctx.QueryInt32("key") - if err != nil { - t.Errorf("ctx.QueryInt32(\"key\") returned an error: %v", err) +} + +func TestCookieNotFound(t *testing.T) { + c, _ := createTestContext("GET", "/path", nil) + + cookie := c.Cookie("nonexistent") + if cookie != nil { + t.Errorf("Expected nil for nonexistent cookie, got %v", cookie) } - want := int32(123) - if got != want { - t.Errorf("QueryInt32() = %q, want %q", got, want) +} + +func TestSetCookie(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SetCookie("test", "value") + c.flush() + + cookie := string(c.ctx.Response.Header.Peek("Set-Cookie")) + if !strings.Contains(cookie, "test=value") { + t.Errorf("Expected Set-Cookie to contain 'test=value', got %s", cookie) } } -func TestContext_QueryInt32WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) +func TestContextBody(t *testing.T) { + body := []byte("test body") + c := &Context{ + res: &response{ + body: body, + }, } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + + result := c.Body() + if !bytes.Equal(result, body) { + t.Errorf("expected body %v, but got %v", body, result) } +} - _, err = ctx.QueryInt32("key") - if err == nil { - t.Error("ctx.QueryInt32(\"id\") did not return an error") +func TestContextSetBody(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + body := []byte("test body") + c.SetBody(body) + + if !bytes.Equal(c.res.body, body) { + t.Errorf("expected body %v, got %v", body, c.res.body) } } -func TestContext_QueryInt32WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) +func TestJSON(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.JSON(200, map[string]string{"message": "hello world"}) + c.flush() + + if c.Status() != 200 { + t.Errorf("Expected status code 200 but got %d", c.Status()) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + + expectedBody := `{"message":"hello world"}` + if string(c.res.body) != expectedBody { + t.Errorf("Expected body %s but got %s", expectedBody, string(c.res.body)) } - got, err := ctx.QueryInt32("key") - if err != nil { - t.Errorf("ctx.QueryInt32(\"key\") returned an error: %v", err) +} + +func TestText(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.Text(200, "hello world") + c.flush() + + if c.Status() != 200 { + t.Errorf("Expected status code 200 but got %d", c.Status()) } - want := int32(0) - if got != want { - t.Errorf("QueryInt32() = %d, want %d", got, want) + + expectedBody := "hello world" + if string(c.res.body) != expectedBody { + t.Errorf("Expected body %s but got %s", expectedBody, string(c.res.body)) } } -func TestContext_QueryUInt32(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) - if err != nil { - t.Fatal(err) +func TestContext_XML(t *testing.T) { + c, _ := createTestContext("GET", "/xml", nil) + + type person struct { + Name string `xml:"name"` + Age int `xml:"age"` } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + obj := &person{ + Name: "John", + Age: 30, } - got, err := ctx.QueryUInt32("key") - if err != nil { - t.Errorf("ctx.QueryUInt32(\"key\") returned an error: %v", err) + + c.XML(StatusOK, obj) + c.flush() + + contentType := string(c.ctx.Response.Header.Peek("Content-Type")) + if !strings.Contains(contentType, MIMEApplicationXML) { + t.Errorf("Expected Content-Type header to contain %s, but got %s", MIMEApplicationXML, contentType) } - want := uint32(123) - if got != want { - t.Errorf("QueryUInt32() = %q, want %q", got, want) + + expectedBody := `John30` + if string(c.res.body) != expectedBody { + t.Errorf("Expected response body to be %s, but got %s", expectedBody, string(c.res.body)) } } -func TestContext_QueryUInt32WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) +func TestContext_GetData(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + if c.GetData("nonexistent") != nil { + t.Errorf("expected nil value for nonexistent key, got %v", c.GetData("nonexistent")) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + + c.SetData("key", "value") + if c.GetData("key") != "value" { + t.Errorf("expected value 'value' for key 'key', got %v", c.GetData("key")) } - _, err = ctx.QueryUInt32("key") - if err == nil { - t.Error("ctx.QueryUInt32(\"id\") did not return an error") + c.DelData("key") + if c.GetData("key") != nil { + t.Errorf("expected nil value for deleted key 'key', got %v", c.GetData("key")) } } -func TestContext_QueryUInt32WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryUInt32("key") - if err != nil { - t.Errorf("ctx.QueryUInt32(\"key\") returned an error: %v", err) +func TestContext_Redirect(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + redirectUrl := "/new" + c.Redirect(StatusMovedPermanently, redirectUrl) + c.flush() + + if c.ctx.Response.StatusCode() != StatusMovedPermanently { + t.Errorf("expected status code %d, got %d", StatusMovedPermanently, c.ctx.Response.StatusCode()) } - want := uint32(0) - if got != want { - t.Errorf("QueryUInt32() = %d, want %d", got, want) + + location := string(c.ctx.Response.Header.Peek("Location")) + if location == "" { + t.Error("expected Location header to be set") } } -func TestContext_QueryInt64(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryInt64("key") - if err != nil { - t.Errorf("ctx.QueryInt64(\"key\") returned an error: %v", err) - } - want := int64(123) - if got != want { - t.Errorf("QueryInt64() = %q, want %q", got, want) +func TestUserAgent(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + ua := "my-user-agent" + c.ctx.Request.Header.SetUserAgent(ua) + + if userAgent := c.UserAgent(); userAgent != ua { + t.Errorf("expected user agent %q, got %q", ua, userAgent) } } -func TestContext_QueryInt64WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } +func TestReferer(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + ref := "https://example.com" + c.ctx.Request.Header.SetReferer(ref) - _, err = ctx.QueryInt64("key") - if err == nil { - t.Error("ctx.QueryInt64(\"id\") did not return an error") + if referer := c.Referer(); referer != ref { + t.Errorf("expected referer %q, got %q", ref, referer) } } -func TestContext_QueryInt64WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) +func TestContext_Success(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + testData := map[string]string{"foo": "bar"} + c.Success(testData) + c.flush() + + if c.ctx.Response.StatusCode() != StatusOK { + t.Errorf("handler returned wrong status code: got %v want %v", c.ctx.Response.StatusCode(), StatusOK) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + + expected := `{"code":0,"data":{"foo":"bar"},"message":"ok"}` + if string(c.res.body) != expected { + t.Errorf("handler returned unexpected body: got %v want %v", string(c.res.body), expected) } - got, err := ctx.QueryInt64("key") - if err != nil { - t.Errorf("ctx.QueryInt64(\"key\") returned an error: %v", err) +} + +func TestContextFail(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.Fail(500, "Internal Server Error") + c.flush() + + if c.ctx.Response.StatusCode() != StatusOK { + t.Errorf("expected status code 200, got %d", c.ctx.Response.StatusCode()) } - want := int64(0) - if got != want { - t.Errorf("QueryInt64() = %d, want %d", got, want) + expectedBody := `{"code":500,"message":"Internal Server Error"}` + if string(c.res.body) != expectedBody { + t.Errorf("expected body %q, got %q", expectedBody, string(c.res.body)) } } -func TestContext_QueryUInt64(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=123", nil) - if err != nil { - t.Fatal(err) +func TestContext_JSONError(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.JSONError(StatusBadRequest, "invalid request") + c.flush() + + if c.ctx.Response.StatusCode() != StatusBadRequest { + t.Errorf("expected status code %d, got %d", StatusBadRequest, c.ctx.Response.StatusCode()) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + expected := `{"code":400,"message":"invalid request"}` + if string(c.res.body) != expected { + t.Errorf("expected body %q, got %q", expected, string(c.res.body)) } - got, err := ctx.QueryUInt64("key") - if err != nil { - t.Errorf("ctx.QueryUInt64(\"key\") returned an error: %v", err) +} + +func TestContext_IsAjax(t *testing.T) { + tests := []struct { + name string + header string + expected bool + }{ + {"XMLHttpRequest", "XMLHttpRequest", true}, + {"empty", "", false}, + {"other", "fetch", false}, } - want := uint64(123) - if got != want { - t.Errorf("QueryUInt64() = %q, want %q", got, want) + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + if tt.header != "" { + c.ctx.Request.Header.Set("X-Requested-With", tt.header) + } + + if got := c.IsAjax(); got != tt.expected { + t.Errorf("IsAjax() = %v, want %v", got, tt.expected) + } + }) } } -func TestContext_QueryUInt64WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) +func TestContext_IsWebSocket(t *testing.T) { + tests := []struct { + name string + header string + expected bool + }{ + {"websocket", "websocket", true}, + {"empty", "", false}, + {"http", "http/1.1", false}, } - _, err = ctx.QueryUInt64("key") - if err == nil { - t.Error("ctx.QueryUInt64(\"id\") did not return an error") + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + if tt.header != "" { + c.ctx.Request.Header.Set("Upgrade", tt.header) + } + + if got := c.IsWebSocket(); got != tt.expected { + t.Errorf("IsWebSocket() = %v, want %v", got, tt.expected) + } + }) } } -func TestContext_QueryUInt64WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) +func TestContext_ContentType(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.SetContentType("application/json") + + if got := c.ContentType(); got != "application/json" { + t.Errorf("ContentType() = %v, want %v", got, "application/json") } - got, err := ctx.QueryUInt64("key") - if err != nil { - t.Errorf("ctx.QueryUInt64(\"key\") returned an error: %v", err) +} + +func TestContext_AcceptedLanguages(t *testing.T) { + tests := []struct { + name string + header string + expected []string + }{ + {"single", "en-US", []string{"en-US"}}, + {"multiple", "en-US, zh-CN, fr", []string{"en-US", "zh-CN", "fr"}}, + {"with quality", "en-US;q=0.9, zh-CN;q=0.8", []string{"en-US", "zh-CN"}}, + {"empty", "", nil}, } - want := uint64(0) - if got != want { - t.Errorf("QueryUInt64() = %d, want %d", got, want) + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + if tt.header != "" { + c.ctx.Request.Header.Set("Accept-Language", tt.header) + } + + got := c.AcceptedLanguages() + if !reflect.DeepEqual(got, tt.expected) { + t.Errorf("AcceptedLanguages() = %v, want %v", got, tt.expected) + } + }) } } -func TestContext_QueryFloat32(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=3.1415", nil) +func TestContext_File(t *testing.T) { + tmpFile, err := os.CreateTemp("", "test*.txt") if err != nil { t.Fatal(err) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { + defer os.Remove(tmpFile.Name()) + + content := []byte("test content") + if _, err := tmpFile.Write(content); err != nil { t.Fatal(err) } - got, err := ctx.QueryFloat32("key") + tmpFile.Close() + + c, _ := createTestContext("GET", "/test", nil) + + err = c.File(tmpFile.Name()) if err != nil { - t.Errorf("ctx.QueryFloat32(\"key\") returned an error: %v", err) - } - want := float32(3.1415) - if got != want { - t.Errorf("QueryFloat32() = %f, want %f", got, want) + t.Errorf("File() returned error: %v", err) } } -func TestContext_QueryFloat32WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } +func TestContext_FileNotFound(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - _, err = ctx.QueryFloat32("key") + err := c.File("/nonexistent/path/file.txt") if err == nil { - t.Error("ctx.QueryFloat32(\"id\") did not return an error") + t.Error("File() expected error for nonexistent file") } } -func TestContext_QueryFloat32WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) +func TestContext_HTML(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "templates") if err != nil { t.Fatal(err) } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { + defer os.RemoveAll(tmpDir) + + tmplPath := tmpDir + "/test.html" + if err := os.WriteFile(tmplPath, []byte("{{.Name}}"), 0644); err != nil { t.Fatal(err) } - got, err := ctx.QueryFloat32("key") - if err != nil { - t.Errorf("ctx.QueryFloat32(\"key\") returned an error: %v", err) + + app := NewApp() + app.SetFuncMap(template.FuncMap{}) + app.LoadHTMLGlob(tmpDir + "/*.html") + + app.Get("/test", func(ctx *Context) { + ctx.HTML(StatusOK, "test.html", map[string]string{"Name": "World"}) + }) + + c := newTestCtxForApp(MethodGet, "/test") + app.serveRequest(c) + + if c.Response.StatusCode() != StatusOK { + t.Errorf("expected status %d, got %d", StatusOK, c.Response.StatusCode()) } - want := float32(0) - if got != want { - t.Errorf("QueryFloat32() = %f, want %f", got, want) + if !strings.Contains(string(c.Response.Body()), "World") { + t.Errorf("expected body to contain 'World', got %q", string(c.Response.Body())) } } -func TestContext_QueryFloat64(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=3.1415", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) - } - got, err := ctx.QueryFloat64("key") - if err != nil { - t.Errorf("ctx.QueryFloat64(\"key\") returned an error: %v", err) +func TestSkipFlush(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SkipFlush() + + c.flush() +} + +func TestContext_Cookies(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.SetCookie("session", "abc") + ctx.Request.Header.SetCookie("theme", "dark") + + c := &Context{ + ctx: ctx, + index: -1, + data: contextData{}, } - want := float64(3.1415) - if got != want { - t.Errorf("QueryFloat64() = %f, want %f", got, want) + c.req = newRequest(ctx) + + cookies := c.Cookies() + if len(cookies) == 0 { + t.Error("Expected cookies to be non-empty") } } -func TestContext_QueryFloat64WithException(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=abc", nil) - if err != nil { - t.Fatal(err) - } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) +func TestContext_RemoteAddr(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + + c := &Context{ + ctx: ctx, + index: -1, + data: contextData{}, } + c.req = newRequest(ctx) - _, err = ctx.QueryFloat64("key") - if err == nil { - t.Error("ctx.QueryFloat64(\"id\") did not return an error") + addr := c.RemoteAddr() + if addr == "" { + t.Error("Expected non-empty remote address") } } -func TestContext_QueryFloat64WithEmptyKey(t *testing.T) { - req, err := http.NewRequest("GET", "/path?key=", nil) - if err != nil { - t.Fatal(err) +func TestContext_RemoteAddrWithXRealIP(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.Set("X-Real-IP", "1.2.3.4") + + c := &Context{ + ctx: ctx, + index: -1, + data: contextData{}, } - ctx, err := NewContext(httptest.NewRecorder(), req) - if err != nil { - t.Fatal(err) + c.req = newRequest(ctx) + + addr := c.RemoteAddr() + if addr != "1.2.3.4" { + t.Errorf("Expected X-Real-IP, got %s", addr) } - got, err := ctx.QueryFloat64("key") - if err != nil { - t.Errorf("ctx.QueryFloat64(\"key\") returned an error: %v", err) +} + +func TestContext_RemoteAddrWithXForwardedFor(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.Set("X-Forwarded-For", "1.2.3.4, 5.6.7.8") + + c := &Context{ + ctx: ctx, + index: -1, + data: contextData{}, } - want := float64(0) - if got != want { - t.Errorf("QueryFloat64() = %f, want %f", got, want) + c.req = newRequest(ctx) + + addr := c.RemoteAddr() + if addr != "1.2.3.4" { + t.Errorf("Expected first IP from X-Forwarded-For, got %s", addr) } } -func TestContextQueries(t *testing.T) { - req, err := http.NewRequest("GET", "/path?foo=bar&baz=qux", nil) - if err != nil { - t.Fatal(err) +func TestContext_JSONBodyWithCustomDecoder(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("POST") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.SetBody([]byte(`{"name":"custom"}`)) + + app := NewApp() + called := false + app.Config.JSONDecoder = func(data []byte, v any) error { + called = true + return json.Unmarshal(data, v) + } + + c := &Context{ + ctx: ctx, + index: -1, + data: contextData{}, + App: app, } - ctx, err := NewContext(httptest.NewRecorder(), req) + c.req = newRequest(ctx) + c.res = newResponse(ctx) + c.Method = c.req.method() + c.Path = c.req.path() + + var result map[string]string + err := c.JSONBody(&result) if err != nil { - t.Fatal(err) + t.Fatalf("JSONBody returned error: %v", err) } - queries := ctx.Queries() - expected := map[string][]string{ - "foo": {"bar"}, - "baz": {"qux"}, + if !called { + t.Error("Expected custom decoder to be called") } - if !reflect.DeepEqual(queries, expected) { - t.Errorf("got %v, want %v", queries, expected) + if result["name"] != "custom" { + t.Errorf("Expected name 'custom', got '%s'", result["name"]) } } -func TestContext_Status(t *testing.T) { - // Create a new HTTP request and response - req, err := http.NewRequest("GET", "/test", nil) - if err != nil { - t.Fatal(err) - } - rr := httptest.NewRecorder() +func TestContext_JSONWithCustomEncoder(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") - // Create a new context object - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) + app := NewApp() + app.Config.JSONEncoder = func(v any) ([]byte, error) { + return []byte(`{"custom":true}`), nil } - // Test the Status method - status := ctx.Status() - if status != http.StatusNotFound { - t.Errorf("Expected status code %d, but got %d", http.StatusNotFound, status) + c := &Context{ + ctx: ctx, + index: -1, + data: contextData{}, + App: app, } + c.req = newRequest(ctx) + c.res = newResponse(ctx) + c.Method = c.req.method() + c.Path = c.req.path() + + c.JSON(StatusOK, map[string]bool{"test": true}) + c.flush() - // Test the SetStatus method - ctx.SetStatus(http.StatusOK) - status = ctx.Status() - if status != http.StatusOK { - t.Errorf("Expected status code %d, but got %d", http.StatusOK, status) + if string(c.res.body) != `{"custom":true}` { + t.Errorf("Expected custom encoded body, got %s", string(c.res.body)) } } -func TestContext_SetStatus(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/", nil) +func TestContext_HTMLWithTemplateError(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "templates_error") if err != nil { t.Fatal(err) } - rr := httptest.NewRecorder() - ctx, err := NewContext(rr, req) - if err != nil { + defer os.RemoveAll(tmpDir) + + if err := os.WriteFile(tmpDir+"/bad.html", []byte("{{.Name}}"), 0644); err != nil { t.Fatal(err) } - // Call the SetStatus method with a status code of 200 - ctx.SetStatus(http.StatusOK) + app := NewApp() + app.SetFuncMap(template.FuncMap{}) + app.LoadHTMLGlob(tmpDir + "/*.html") + + app.Get("/bad", func(ctx *Context) { + ctx.HTML(StatusOK, "bad.html", nil) + }) - // Check that the response status code is 200 - if status := rr.Code; status != http.StatusOK { - t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + c := newTestCtxForApp(MethodGet, "/bad") + app.serveRequest(c) + + if c.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, c.Response.StatusCode()) } } -func TestHeader(t *testing.T) { - req, err := http.NewRequest("GET", "/path", nil) - if err != nil { - t.Fatal(err) +func TestContext_XMLWithMarshalError(t *testing.T) { + c, _ := createTestContext("GET", "/xml", nil) + + type BadXML struct { + Ch chan int `xml:"ch"` } - w := httptest.NewRecorder() - ctx, err := NewContext(w, req) - if err != nil { - t.Fatal(err) + c.XML(StatusOK, &BadXML{Ch: make(chan int)}) + c.flush() + + if len(c.res.body) != 0 { + t.Errorf("Expected empty body for XML marshal error, got %s", string(c.res.body)) } - key := "Content-Type" - value := "application/json" - req.Header.Set(key, value) - if got := ctx.Header(key); got != value { - t.Errorf("Header(%q) = %q, want %q", key, got, value) +} + +func TestContext_QueryInt8Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryInt8("key") + if err == nil { + t.Error("Expected error for invalid int8") } } -func TestHeaders(t *testing.T) { - req, err := http.NewRequest("GET", "/path", nil) - if err != nil { - t.Fatal(err) +func TestContext_QueryInt32Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryInt32("key") + if err == nil { + t.Error("Expected error for invalid int32") } - rr := httptest.NewRecorder() - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) +} + +func TestContext_QueryInt64Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryInt64("key") + if err == nil { + t.Error("Expected error for invalid int64") } +} - req.Header.Add("Content-Type", "application/json") - req.Header.Add("X-Request-ID", "12345") - headers := ctx.Headers() - if len(headers) != 2 { - t.Errorf("Expected 2 headers, but got %d", len(headers)) +func TestContext_QueryUIntError(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryUInt("key") + if err == nil { + t.Error("Expected error for invalid uint") } - if headers.Get("Content-Type") != "application/json" { - t.Errorf("Expected Content-Type header to be 'application/json', but got '%s'", headers.Get("Content-Type")) +} + +func TestContext_QueryUInt8Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryUInt8("key") + if err == nil { + t.Error("Expected error for invalid uint8") } - if headers.Get("X-Request-ID") != "12345" { - t.Errorf("Expected X-Request-ID header to be '12345', but got '%s'", headers.Get("X-Request-ID")) +} + +func TestContext_QueryUInt32Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryUInt32("key") + if err == nil { + t.Error("Expected error for invalid uint32") } } -func TestAddHeader(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/", nil) - if err != nil { - t.Fatal(err) +func TestContext_QueryUInt64Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryUInt64("key") + if err == nil { + t.Error("Expected error for invalid uint64") } - res := httptest.NewRecorder() - ctx, err := NewContext(res, req) - if err != nil { - t.Fatal(err) +} + +func TestContext_QueryFloat32Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) + + _, err := c.QueryFloat32("key") + if err == nil { + t.Error("Expected error for invalid float32") } +} - // Call the AddHeader method - ctx.AddHeader("Content-Type", "application/json") +func TestContext_QueryFloat64Error(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=invalid", nil) - // Check the response headers - headers := res.Header() - contentType := headers.Get("Content-Type") - if contentType != "application/json" { - t.Errorf("unexpected content type: got %v want %v", contentType, "application/json") + _, err := c.QueryFloat64("key") + if err == nil { + t.Error("Expected error for invalid float64") } } -func TestSetHeader(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/", nil) +func TestContext_QueryBoolEmpty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryBool("key") if err != nil { - t.Fatal(err) + t.Errorf("QueryBool returned error: %v", err) } - res := httptest.NewRecorder() - ctx, err := NewContext(res, req) - if err != nil { - t.Fatal(err) + if got != false { + t.Errorf("Expected false for empty key, got %v", got) } +} - // Call the SetHeader method - ctx.SetHeader("Content-Type", "application/json") +func TestContext_QueryIntEmpty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) - // Check the response headers - headers := res.Header() - contentType := headers.Get("Content-Type") - if contentType != "application/json" { - t.Errorf("unexpected content type: got %v want %v", contentType, "application/json") + got, err := c.QueryInt("key") + if err != nil { + t.Errorf("QueryInt returned error: %v", err) + } + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) } } -func TestDelHeader(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/path", nil) +func TestContext_QueryInt8Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryInt8("key") if err != nil { - t.Fatal(err) + t.Errorf("QueryInt8 returned error: %v", err) } - res := httptest.NewRecorder() - ctx, err := NewContext(res, req) - if err != nil { - t.Fatal(err) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) } +} - // Set a header in the response - ctx.SetHeader("key", "value") - - // Call the DelHeader method - ctx.DelHeader("key") +func TestContext_QueryInt32Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) - // Check that the header was deleted - if res.Header().Get("key") != "" { - t.Errorf("Header was not deleted") + got, err := c.QueryInt32("key") + if err != nil { + t.Errorf("QueryInt32 returned error: %v", err) + } + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) } } -func TestContextCookie(t *testing.T) { - req, err := http.NewRequest("GET", "/path", nil) +func TestContext_QueryInt64Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryInt64("key") if err != nil { - t.Fatal(err) + t.Errorf("QueryInt64 returned error: %v", err) } - cookie := &http.Cookie{Name: "test", Value: "value"} - req.AddCookie(cookie) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) + } +} - res := httptest.NewRecorder() - ctx, err := NewContext(res, req) +func TestContext_QueryUIntEmpty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryUInt("key") if err != nil { - t.Fatal(err) + t.Errorf("QueryUInt returned error: %v", err) } - - if c := ctx.Cookie("test"); c == nil || c.Value != "value" { - t.Errorf("Cookie() = %v, want %v", c, cookie) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) } } -func TestCookies(t *testing.T) { - req, err := http.NewRequest("GET", "/path", nil) +func TestContext_QueryUInt8Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryUInt8("key") if err != nil { - t.Fatal(err) + t.Errorf("QueryUInt8 returned error: %v", err) } - cookie := &http.Cookie{Name: "test", Value: "value"} - req.AddCookie(cookie) - w := httptest.NewRecorder() - ctx, err := NewContext(w, req) - if err != nil { - t.Fatal(err) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) } +} - cookies := ctx.Cookies() - if len(cookies) != 1 { - t.Errorf("Expected 1 cookie, got %d", len(cookies)) - } - if cookies[0].Name != "test" { - t.Errorf("Expected cookie name 'test', got '%s'", cookies[0].Name) +func TestContext_QueryUInt32Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryUInt32("key") + if err != nil { + t.Errorf("QueryUInt32 returned error: %v", err) } - if cookies[0].Value != "value" { - t.Errorf("Expected cookie value 'value', got '%s'", cookies[0].Value) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) } } -func TestSetCookie(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/", nil) +func TestContext_QueryUInt64Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryUInt64("key") if err != nil { - t.Fatal(err) + t.Errorf("QueryUInt64 returned error: %v", err) } - w := httptest.NewRecorder() - ctx, err := NewContext(w, req) - if err != nil { - t.Fatal(err) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %d", got) } +} - // Call the SetCookie function - ctx.SetCookie("test", "value") - ctx.flush() +func TestContext_QueryFloat32Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) - // Check that the cookie was set correctly - cookies := w.Result().Cookies() - if len(cookies) != 1 { - t.Errorf("expected 1 cookie, got %d", len(cookies)) - } - if cookies[0].Name != "test" { - t.Errorf("expected cookie name 'test', got '%s'", cookies[0].Name) + got, err := c.QueryFloat32("key") + if err != nil { + t.Errorf("QueryFloat32 returned error: %v", err) } - if cookies[0].Value != "value" { - t.Errorf("expected cookie value 'value', got '%s'", cookies[0].Value) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %f", got) } } -func TestSetCustomCookie(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/", nil) +func TestContext_QueryFloat64Empty(t *testing.T) { + c, _ := createTestContext("GET", "/path?key=", nil) + + got, err := c.QueryFloat64("key") if err != nil { - t.Fatal(err) + t.Errorf("QueryFloat64 returned error: %v", err) } - w := httptest.NewRecorder() - ctx, err := NewContext(w, req) - if err != nil { - t.Fatal(err) + if got != 0 { + t.Errorf("Expected 0 for empty key, got %f", got) } +} - // Call the SetCustomCookie function - cookie := &http.Cookie{Name: "test", Value: "value"} - ctx.SetCustomCookie(cookie) - ctx.flush() +func TestContext_setParams(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - // Check that the cookie was set correctly - cookies := w.Result().Cookies() - if len(cookies) != 1 { - t.Errorf("expected 1 cookie, got %d", len(cookies)) - } - if cookies[0].Name != "test" { - t.Errorf("expected cookie name 'test', got '%s'", cookies[0].Name) + params := map[string]string{"id": "123", "name": "test"} + c.setParams(params) + + if c.Param("id") != "123" { + t.Errorf("Expected param id=123, got %s", c.Param("id")) } - if cookies[0].Value != "value" { - t.Errorf("expected cookie value 'value', got '%s'", cookies[0].Value) + if c.Param("name") != "test" { + t.Errorf("Expected param name=test, got %s", c.Param("name")) } } -func TestContextBody(t *testing.T) { - // create a new context with a response body - body := []byte("test body") - ctx := &Context{ - res: &response{ - body: body, - }, +func TestContext_setHandlers(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + handlers := []HandlerFunc{ + func(c *Context) {}, + func(c *Context) {}, + func(c *Context) {}, } + c.setHandlers(handlers) - // call the Body() function and check the result - result := ctx.Body() - if !bytes.Equal(result, body) { - t.Errorf("expected body %v, but got %v", body, result) + if len(c.handlers) != 3 { + t.Errorf("Expected 3 handlers, got %d", len(c.handlers)) } } -func TestContextSetBody(t *testing.T) { - // create a new context - req := httptest.NewRequest(http.MethodGet, "/", nil) - res := httptest.NewRecorder() - ctx, _ := NewContext(res, req) +func TestContext_setApp(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - // set the body using SetBody - body := []byte("test body") - ctx.SetBody(body) + app := NewApp() + c.setApp(app) - // check that the body was set correctly - if !bytes.Equal(ctx.res.body, body) { - t.Errorf("expected body %v, got %v", body, ctx.res.body) + if c.App != app { + t.Error("App not set correctly") } } -func TestJSON(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/", nil) - if err != nil { - t.Fatal(err) - } - w := httptest.NewRecorder() - ctx, err := NewContext(w, req) - if err != nil { - t.Fatal(err) - } +func TestContext_NextNoMoreHandlers(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.index = 0 - ctx.JSON(200, map[string]string{"message": "hello world"}) - ctx.flush() + c.Next() - if ctx.Status() != 200 { - t.Errorf("Expected status code %d but got %d", 200, ctx.Status()) + if c.index != 1 { + t.Errorf("Expected index to be 1, got %d", c.index) } +} - expectedBody := `{"message":"hello world"}` - if string(ctx.res.body) != expectedBody { - t.Errorf("Expected body %s but got %s", expectedBody, string(ctx.res.body)) +func TestContext_StringBodyEmpty(t *testing.T) { + c, _ := createTestContext("POST", "/test", nil) + + body := c.StringBody() + if body != "" { + t.Errorf("Expected empty body, got %s", body) } } -func TestText(t *testing.T) { - // Create a new context object - req, err := http.NewRequest("GET", "/", nil) - if err != nil { - t.Fatal(err) - } - w := httptest.NewRecorder() - ctx, err := NewContext(w, req) - if err != nil { - t.Fatal(err) +func TestContext_JSONBodyEmpty(t *testing.T) { + c, _ := createTestContext("POST", "/test", nil) + + var result map[string]string + err := c.JSONBody(&result) + if err == nil { + t.Error("Expected error for empty body") } +} - ctx.Text(200, "hello world") - ctx.flush() +func TestContext_Body(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - if ctx.Status() != 200 { - t.Errorf("Expected status code %d but got %d", 200, ctx.Status()) - } + body := []byte("response body") + c.SetBody(body) - expectedBody := "hello world" - if string(ctx.res.body) != expectedBody { - t.Errorf("Expected body %s but got %s", expectedBody, string(ctx.res.body)) + if !bytes.Equal(c.Body(), body) { + t.Errorf("Expected body %v, got %v", body, c.Body()) } } -func TestContext_XML(t *testing.T) { - // Create a new context - req := httptest.NewRequest("GET", "/xml", nil) - res := httptest.NewRecorder() - ctx, err := NewContext(res, req) - if err != nil { - t.Fatal(err) +func TestContext_DelHeader(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SetHeader("X-Custom", "value") + c.DelHeader("X-Custom") + + if got := string(c.ctx.Response.Header.Peek("X-Custom")); got != "" { + t.Errorf("Expected header to be deleted, got %s", got) } +} - // Create a test object - type person struct { - Name string `xml:"name"` - Age int `xml:"age"` +func TestContext_SetCookie(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SetCookie("session", "abc123") + c.flush() + + cookie := string(c.ctx.Response.Header.Peek("Set-Cookie")) + if !strings.Contains(cookie, "session=abc123") { + t.Errorf("Expected cookie to contain 'session=abc123', got %s", cookie) } - obj := &person{ - Name: "John", - Age: 30, +} + +func TestContext_AddHeader(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.AddHeader("X-Custom", "value1") + + hdr := string(c.ctx.Response.Header.Peek("X-Custom")) + if hdr == "" { + t.Error("Expected header to be set") } +} + +func TestContext_SetHeader(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - // Call the XML function with the test object - ctx.XML(http.StatusOK, obj) - ctx.flush() + c.SetHeader("Content-Type", "text/plain") - // Check the response headers - if res.Header().Get(HeaderContentType) != MIMEApplicationXML { - t.Errorf("Expected Content-Type header to be %s, but got %s", MIMEApplicationXML, res.Header().Get(HeaderContentType)) + if got := string(c.ctx.Response.Header.Peek("Content-Type")); got != "text/plain" { + t.Errorf("Expected 'text/plain', got %s", got) } +} + +func TestContext_Header(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.Set("X-Custom", "header-value") - // Check the response status code - if res.Result().StatusCode != http.StatusOK { - t.Errorf("Expected status code to be %d, but got %d", http.StatusOK, res.Result().StatusCode) + if got := c.Header("X-Custom"); got != "header-value" { + t.Errorf("Expected 'header-value', got %s", got) } +} - // Check the response body - expectedBody := `John30` - if res.Body.String() != expectedBody { - t.Errorf("Expected response body to be %s, but got %s", expectedBody, res.Body.String()) +func TestContext_Headers(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.Set("Content-Type", "application/json") + + headers := c.Headers() + if len(headers) == 0 { + t.Error("Expected headers to be non-empty") } } -func TestContext_GetData(t *testing.T) { - req, err := http.NewRequest("GET", "/test", nil) - if err != nil { - t.Fatal(err) +func TestContext_Text(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.Text(StatusOK, "hello") + c.flush() + + if c.Status() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, c.Status()) } - rr := httptest.NewRecorder() - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) + if string(c.res.body) != "hello" { + t.Errorf("Expected body 'hello', got %s", string(c.res.body)) } +} - // Test getting a non-existent key - if ctx.GetData("nonexistent") != nil { - t.Errorf("expected nil value for nonexistent key, got %v", ctx.GetData("nonexistent")) - } +func TestContext_SetStatus(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SetStatus(StatusCreated) - // Test setting and getting a key - ctx.SetData("key", "value") - if ctx.GetData("key") != "value" { - t.Errorf("expected value 'value' for key 'key', got %v", ctx.GetData("key")) + if c.Status() != StatusCreated { + t.Errorf("Expected status %d, got %d", StatusCreated, c.Status()) } +} + +func TestContext_DefaultStatus(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - // Test deleting a key - ctx.DelData("key") - if ctx.GetData("key") != nil { - t.Errorf("expected nil value for deleted key 'key', got %v", ctx.GetData("key")) + if c.Status() != StatusNotFound { + t.Errorf("Expected default status %d, got %d", StatusNotFound, c.Status()) } } -func TestContext_Redirect(t *testing.T) { - req, err := http.NewRequest("GET", "/test", nil) - if err != nil { - t.Fatal(err) +func TestContext_CookieNotFound(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + cookie := c.Cookie("nonexistent") + if cookie != nil { + t.Errorf("Expected nil for nonexistent cookie, got %v", cookie) } - rr := httptest.NewRecorder() - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) +} + +func TestContext_DelData(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + c.SetData("key", "value") + c.DelData("key") + + if c.GetData("key") != nil { + t.Error("Expected nil after DelData") } +} + +func TestContext_Fail(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - // Call the Redirect method with a test URL and status code - redirectUrl := "https://example.com" - ctx.Redirect(http.StatusMovedPermanently, redirectUrl) - ctx.flush() + c.Fail(500, "internal error") + c.flush() - // Verify that the response status code and location header are set correctly - if rr.Result().StatusCode != http.StatusMovedPermanently { - t.Errorf("expected status code %d, got %d", http.StatusMovedPermanently, rr.Result().StatusCode) + if c.Status() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, c.Status()) } - url := rr.Header().Get("Location") - if url != redirectUrl { - t.Errorf("expected Location header %q, got %q", redirectUrl, url) + if !strings.Contains(string(c.res.body), "internal error") { + t.Errorf("Expected body to contain 'internal error', got %s", string(c.res.body)) } } -func TestUserAgent(t *testing.T) { - ctx := createMockContext(t) - ua := "my-user-agent" - ctx.req.originReq.Header.Add("user-agent", ua) +func TestContext_JSONErrorStatus(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - if userAgent := ctx.UserAgent(); userAgent != "my-user-agent" { - t.Errorf("expected user agent %q, got %q", ua, userAgent) + c.JSONError(StatusBadRequest, "bad request") + c.flush() + + if c.Status() != StatusBadRequest { + t.Errorf("Expected status %d, got %d", StatusBadRequest, c.Status()) } } -func TestReferer(t *testing.T) { - ctx := createMockContext(t) - ref := "https://example.com" - ctx.req.originReq.Header.Add("referer", ref) +func TestContext_IsAjaxTrue(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.Set("X-Requested-With", "XMLHttpRequest") - if referer := ctx.Referer(); referer != ref { - t.Errorf("expected referer %q, got %q", ref, referer) + if !c.IsAjax() { + t.Error("Expected IsAjax to return true") } } -func TestContext_RemoteAddr(t *testing.T) { - ctx := createMockContext(t) - expected := "1.2.3.4:5678" - ctx.req.originReq.RemoteAddr = expected +func TestContext_IsNotAjax(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - if addr := ctx.RemoteAddr(); addr != expected { - t.Errorf("Expected RemoteAddr to return %q, but got %q", expected, addr) + if c.IsAjax() { + t.Error("Expected IsAjax to return false") } } -func createMockContext(t *testing.T) *Context { - req, err := http.NewRequest("GET", "/test", nil) - if err != nil { - t.Fatal(err) - } - rr := httptest.NewRecorder() - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) +func TestContext_IsWebSocketTrue(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.Set("Upgrade", "websocket") + + if !c.IsWebSocket() { + t.Error("Expected IsWebSocket to return true") } +} - return ctx +func TestContext_IsNotWebSocket(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + + if c.IsWebSocket() { + t.Error("Expected IsWebSocket to return false") + } } -func TestContext_Success(t *testing.T) { - // Create a new request with an empty body - req, err := http.NewRequest("GET", "/", nil) - if err != nil { - t.Fatal(err) +func TestContext_ContentTypeSet(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.SetContentType("application/json") + + if c.ContentType() != "application/json" { + t.Errorf("Expected 'application/json', got %s", c.ContentType()) } +} - // Create a new ResponseRecorder to record the response - rr := httptest.NewRecorder() +func TestContext_AcceptedLanguagesEmpty(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) - // Create a new context with the request and response recorder - ctx, err := NewContext(rr, req) - if err != nil { - t.Fatal(err) + langs := c.AcceptedLanguages() + if langs != nil { + t.Errorf("Expected nil for empty Accept-Language, got %v", langs) } +} - // Call the Success method with some test data - testData := map[string]string{"foo": "bar"} - ctx.Success(testData) - ctx.flush() +func TestContext_AcceptedLanguagesMultiple(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.Set("Accept-Language", "en-US, zh-CN, fr") - // Check the response status code - if status := rr.Code; status != http.StatusOK { - t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + langs := c.AcceptedLanguages() + if len(langs) != 3 { + t.Errorf("Expected 3 languages, got %d", len(langs)) } +} - // Check the response body - expected := `{"code":0,"data":{"foo":"bar"},"message":"ok"}` - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) +func TestContext_AcceptedLanguagesWithQuality(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.ctx.Request.Header.Set("Accept-Language", "en-US;q=0.9, zh-CN;q=0.8") + + langs := c.AcceptedLanguages() + if len(langs) != 2 { + t.Errorf("Expected 2 languages, got %d", len(langs)) } } -func TestContextFail(t *testing.T) { - // Create a new context with a mock response writer and request - req, err := http.NewRequest("GET", "/test", nil) +func TestContext_FileExists(t *testing.T) { + tmpFile, err := os.CreateTemp("", "test*.txt") if err != nil { t.Fatal(err) } - w := httptest.NewRecorder() - ctx, err := NewContext(w, req) - if err != nil { + defer os.Remove(tmpFile.Name()) + + content := []byte("test content") + if _, err := tmpFile.Write(content); err != nil { t.Fatal(err) } + tmpFile.Close() - // Call the Fail method with a custom code and message - ctx.Fail(500, "Internal Server Error") - ctx.flush() + c, _ := createTestContext("GET", "/test", nil) - // Check that the response status code and body are correct - if w.Code != 200 { - t.Errorf("expected status code 200, got %d", w.Code) - } - expectedBody := `{"code":500,"message":"Internal Server Error"}` - if w.Body.String() != expectedBody { - t.Errorf("expected body %q, got %q", expectedBody, w.Body.String()) + err = c.File(tmpFile.Name()) + if err != nil { + t.Errorf("File() returned error: %v", err) } } diff --git a/cookie.go b/cookie.go index 603e294..ec31292 100644 --- a/cookie.go +++ b/cookie.go @@ -1,32 +1,14 @@ package lightning -import ( - "net/http" -) +// cookiesMap is a map of cookie values keyed by cookie name. +type cookiesMap map[string]string -// cookiesMap is a map of http.Cookie pointers. -type cookiesMap map[string]*http.Cookie - -// get returns the http.Cookie pointer with the given key. -func (cookies cookiesMap) get(key string) *http.Cookie { - return cookies[key] -} - -// set sets the http.Cookie with the given key and value. +// set sets the cookie value with the given key. func (cookies cookiesMap) set(key string, value string) { - cookies[key] = &http.Cookie{ - Name: key, - Value: value, - Path: "/", - } + cookies[key] = value } -// del deletes the http.Cookie with the given key. +// del deletes the cookie with the given key. func (cookies cookiesMap) del(key string) { delete(cookies, key) } - -// setCustom sets the given http.Cookie pointer. -func (cookies cookiesMap) setCustom(cookie *http.Cookie) { - cookies[cookie.Name] = cookie -} diff --git a/cookie_test.go b/cookie_test.go index db70629..a986bdd 100644 --- a/cookie_test.go +++ b/cookie_test.go @@ -1,126 +1,39 @@ package lightning import ( - "net/http" - "reflect" "testing" ) -func TestCookie_Del(t *testing.T) { - type args struct { - key string - } - tests := []struct { - name string - cookies cookiesMap - args args - }{ - { - name: "TestCookie_Del", - cookies: cookiesMap{"test": &http.Cookie{Name: "test", Value: "value"}}, - args: args{key: "test"}, - }, - { - name: "TestCookie_Del_NotExist", - cookies: cookiesMap{}, - args: args{key: "test"}, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - tt.cookies.del(tt.args.key) - if _, ok := tt.cookies[tt.args.key]; ok { - t.Errorf("del() failed to delete key %s", tt.args.key) - } - }) - } -} +func TestCookiesMapSet(t *testing.T) { + cookies := make(cookiesMap) + cookies.set("key1", "value1") -func TestCookie_Get(t *testing.T) { - type args struct { - key string - } - tests := []struct { - name string - cookies cookiesMap - args args - want *http.Cookie - }{ - { - name: "TestCookie_Get", - cookies: cookiesMap{"test": &http.Cookie{Name: "test", Value: "test"}}, - args: args{key: "test"}, - want: &http.Cookie{Name: "test", Value: "test"}, - }, - { - name: "TestCookie_Get_NotExist", - cookies: cookiesMap{}, - args: args{key: "test"}, - want: nil, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := tt.cookies.get(tt.args.key); !reflect.DeepEqual(got, tt.want) { - t.Errorf("get() = %v, want %v", got, tt.want) - } - }) + if cookies["key1"] != "value1" { + t.Errorf("Expected 'value1', got '%s'", cookies["key1"]) } } -func TestCookie_Set(t *testing.T) { - type args struct { - key string - value string - } - tests := []struct { - name string - cookies cookiesMap - args args - }{ - { - name: "TestCookie_Set", - cookies: cookiesMap{}, - args: args{key: "test", value: "test"}, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - tt.cookies.set(tt.args.key, tt.args.value) - if _, ok := tt.cookies[tt.args.key]; !ok { - t.Errorf("set() failed to set key %s", tt.args.key) - } - if tt.cookies[tt.args.key].Value != tt.args.value { - t.Errorf("set() failed to set value %s for key %s", tt.args.value, tt.args.key) - } - }) +func TestCookiesMapDel(t *testing.T) { + cookies := make(cookiesMap) + cookies.set("key1", "value1") + cookies.del("key1") + + if _, exists := cookies["key1"]; exists { + t.Error("Expected key to be deleted") } } -func TestCookie_SetCustom(t *testing.T) { - type args struct { - cookie *http.Cookie - } - tests := []struct { - name string - cookies cookiesMap - args args - }{ - { - name: "TestCookie_SetCustom", - cookies: cookiesMap{}, - args: args{cookie: &http.Cookie{Name: "test", Value: "test"}}, - }, +func TestCookiesMapSetMultiple(t *testing.T) { + cookies := make(cookiesMap) + cookies.set("key1", "value1") + cookies.set("key2", "value2") + + if len(cookies) != 2 { + t.Errorf("Expected 2 cookies, got %d", len(cookies)) } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - tt.cookies.setCustom(tt.args.cookie) - if _, ok := tt.cookies[tt.args.cookie.Name]; !ok { - t.Errorf("setCustom() failed to set key %s", tt.args.cookie.Name) - } - if tt.cookies[tt.args.cookie.Name] != tt.args.cookie { - t.Errorf("setCustom() failed to set value %s for key %s", tt.args.cookie.Value, tt.args.cookie.Name) - } - }) + + cookies.del("key1") + if len(cookies) != 1 { + t.Errorf("Expected 1 cookie after deletion, got %d", len(cookies)) } } diff --git a/examples/contextual_data/app.go b/examples/contextual_data/app.go index 8a78b83..0c4687e 100644 --- a/examples/contextual_data/app.go +++ b/examples/contextual_data/app.go @@ -33,7 +33,7 @@ func main() { // write your logic here... fmt.Println(session) - ctx.Text(200, "hello world") + ctx.Text(lightning.StatusOK, "hello world") }) // Run the app diff --git a/examples/cookie/app.go b/examples/cookie/app.go index e1f2c54..bf3f576 100644 --- a/examples/cookie/app.go +++ b/examples/cookie/app.go @@ -2,7 +2,6 @@ package main import ( "fmt" - "net/http" "github.com/go-labx/lightning" ) @@ -11,25 +10,19 @@ func main() { app := lightning.DefaultApp() app.Get("/ping", func(ctx *lightning.Context) { - // get the value of the "sid" cookie cookie := ctx.Cookie("sid") - fmt.Println(cookie) + if cookie != nil { + fmt.Println(string(cookie.Key()), string(cookie.Value())) + } - // get all cookies cookies := ctx.Cookies() - fmt.Println(cookies) + for _, c := range cookies { + fmt.Println(string(c.Key()), string(c.Value())) + } - // set a new cookie ctx.SetCookie("sid", "sid:xxxxxxxxxx") - // set a custom cookie - ctx.SetCustomCookie(&http.Cookie{ - Name: "sessionId", - Value: "sessionId:xxxxxxxxxx", - Path: "/", - }) - - ctx.JSON(http.StatusOK, map[string]string{ + ctx.JSON(lightning.StatusOK, map[string]string{ "message": "pong", }) }) diff --git a/examples/group/app.go b/examples/group/app.go index ba2da09..337051d 100644 --- a/examples/group/app.go +++ b/examples/group/app.go @@ -2,7 +2,6 @@ package main import ( "fmt" - "net/http" "github.com/go-labx/lightning" ) @@ -46,14 +45,14 @@ func main() { // define a GET route for "/api/user/info" subGroup.Get("/info", func(ctx *lightning.Context) { - ctx.JSON(http.StatusOK, map[string]interface{}{ + ctx.JSON(lightning.StatusOK, map[string]interface{}{ "username": "zhangsan", "age": 20, }) }) app.Get("/ping", func(ctx *lightning.Context) { - ctx.Text(http.StatusOK, "pong") + ctx.Text(lightning.StatusOK, "pong") }) app.Run() diff --git a/examples/middleware/app.go b/examples/middleware/app.go index 3077ea9..c3de8ad 100644 --- a/examples/middleware/app.go +++ b/examples/middleware/app.go @@ -2,16 +2,13 @@ package main import ( "fmt" - "net/http" "github.com/go-labx/lightning" ) func main() { - // Creating a new instance of lightning application app := lightning.NewApp() - // Adding global middleware to the application app.Use(func(ctx *lightning.Context) { fmt.Println("global scope middleware 1 --->") ctx.Next() @@ -28,20 +25,18 @@ func main() { fmt.Println("<--- global scope middleware 3") }) - // Defining a GET route for the root path of the application app.Get("/", func(ctx *lightning.Context) { - ctx.JSON(http.StatusOK, map[string]string{ + ctx.JSON(lightning.StatusOK, map[string]string{ "message": "hello world", }) }) - // Defining a GET route for the path "/ping" with a route scoped middleware app.Get("/ping", func(ctx *lightning.Context) { fmt.Println("route scope middleware --->") ctx.Next() fmt.Println("<--- route scope middleware") }, func(ctx *lightning.Context) { - ctx.Text(http.StatusOK, "pong") + ctx.Text(lightning.StatusOK, "pong") }) app.Run() diff --git a/examples/middleware_group/app.go b/examples/middleware_group/app.go index d182c57..a7bd32c 100644 --- a/examples/middleware_group/app.go +++ b/examples/middleware_group/app.go @@ -2,7 +2,6 @@ package main import ( "fmt" - "net/http" "github.com/go-labx/lightning" ) @@ -26,7 +25,7 @@ func main() { }) group.Get("/ping", func(ctx *lightning.Context) { - ctx.Text(http.StatusOK, "pong") + ctx.Text(lightning.StatusOK, "pong") }) app.Run() diff --git a/examples/not_found/app.go b/examples/not_found/app.go index cc93743..11f3acb 100644 --- a/examples/not_found/app.go +++ b/examples/not_found/app.go @@ -1,20 +1,18 @@ package main import ( - "net/http" - "github.com/go-labx/lightning" ) func main() { app := lightning.NewApp(&lightning.Config{ NotFoundHandler: func(ctx *lightning.Context) { - ctx.Text(404, "custom not found") + ctx.Text(lightning.StatusNotFound, "custom not found") }, }) app.Get("/ping", func(ctx *lightning.Context) { - ctx.JSON(http.StatusOK, map[string]string{ + ctx.JSON(lightning.StatusOK, map[string]string{ "message": "pong", }) }) diff --git a/examples/ping/app.go b/examples/ping/app.go index d900fc2..b9387f8 100644 --- a/examples/ping/app.go +++ b/examples/ping/app.go @@ -1,23 +1,17 @@ package main import ( - "net/http" - "github.com/go-labx/lightning" ) func main() { - // Create a new Lightning app app := lightning.DefaultApp() - // Define a GET route for "/ping" app.Get("/ping", func(ctx *lightning.Context) { - // Respond with a JSON message - ctx.JSON(http.StatusOK, map[string]string{ + ctx.JSON(lightning.StatusOK, map[string]string{ "message": "pong", }) }) - // Run the app app.Run() } diff --git a/examples/redirect/app.go b/examples/redirect/app.go index e96dd4d..3e44095 100644 --- a/examples/redirect/app.go +++ b/examples/redirect/app.go @@ -1,8 +1,6 @@ package main import ( - "net/http" - "github.com/go-labx/lightning" ) @@ -10,18 +8,15 @@ func main() { app := lightning.NewApp() app.Get("/foo", func(ctx *lightning.Context) { - // Redirect to /baz with a 301 status code - ctx.Redirect(http.StatusMovedPermanently, "/baz") + ctx.Redirect(lightning.StatusMovedPermanently, "/baz") }) app.Get("/bar", func(ctx *lightning.Context) { - // Redirect to /baz with a 302 status code - ctx.Redirect(302, "/baz") + ctx.Redirect(lightning.StatusFound, "/baz") }) app.Get("/baz", func(ctx *lightning.Context) { - // Return a JSON response with a "message" key and "pong" value - ctx.JSON(http.StatusOK, map[string]string{ + ctx.JSON(lightning.StatusOK, map[string]string{ "message": "pong", }) }) diff --git a/examples/response_body/app.go b/examples/response_body/app.go index 58c8b8d..c2eeccc 100644 --- a/examples/response_body/app.go +++ b/examples/response_body/app.go @@ -1,8 +1,6 @@ package main import ( - "net/http" - "github.com/go-labx/lightning" ) @@ -17,12 +15,12 @@ func main() { app.Get("/text", func(ctx *lightning.Context) { // Return "hello world" as plain text - ctx.Text(http.StatusOK, "hello world") + ctx.Text(lightning.StatusOK, "hello world") }) app.Get("/json", func(ctx *lightning.Context) { // Return a Person object as JSON with name "zhangsan", age 20, and city "Hangzhou" - ctx.JSON(http.StatusOK, &Person{ + ctx.JSON(lightning.StatusOK, &Person{ Name: "zhangsan", Age: 20, City: "Hangzhou", @@ -31,7 +29,7 @@ func main() { app.Get("/xml", func(ctx *lightning.Context) { // Return a Person object as XML with name "zhangsan", age 20, and city "Hangzhou" - ctx.XML(http.StatusOK, &Person{ + ctx.XML(lightning.StatusOK, &Person{ Name: "zhangsan", Age: 20, City: "Hangzhou", diff --git a/examples/response_header/app.go b/examples/response_header/app.go index 0d256b7..6a48ce1 100644 --- a/examples/response_header/app.go +++ b/examples/response_header/app.go @@ -1,8 +1,6 @@ package main import ( - "net/http" - "github.com/go-labx/lightning" ) @@ -21,7 +19,7 @@ func main() { // delete all headers with key "bar" ctx.DelHeader("bar") - ctx.JSON(http.StatusOK, lightning.Map{ + ctx.JSON(lightning.StatusOK, lightning.Map{ "message": "pong", }) }) diff --git a/examples/static/app.go b/examples/static/app.go index 238fa59..9eca823 100644 --- a/examples/static/app.go +++ b/examples/static/app.go @@ -23,7 +23,7 @@ func main() { app.LoadHTMLGlob("templates/*") app.Get("/", func(ctx *lightning.Context) { - ctx.HTML(200, "index.tmpl", lightning.Map{ + ctx.HTML(lightning.StatusOK, "index.tmpl", lightning.Map{ "title": "Lightning", "description": "Lightning is a lightweight and fast web framework for Go. It is designed to be easy to use and highly performant. ⚡️⚡️⚡️", "now": time.Now(), diff --git a/go.mod b/go.mod index 0e60158..7e14cf1 100644 --- a/go.mod +++ b/go.mod @@ -1,18 +1,23 @@ module github.com/go-labx/lightning -go 1.19 +go 1.25.0 require ( github.com/go-labx/lightlog v0.0.3 - github.com/go-playground/validator/v10 v10.12.0 + github.com/go-playground/validator/v10 v10.30.2 + github.com/valyala/fasthttp v1.69.0 ) require ( + github.com/andybalholm/brotli v1.2.0 // indirect + github.com/gabriel-vasile/mimetype v1.4.13 // indirect github.com/go-labx/color v0.0.1 // indirect github.com/go-playground/locales v0.14.1 // indirect github.com/go-playground/universal-translator v0.18.1 // indirect - github.com/leodido/go-urn v1.2.2 // indirect - golang.org/x/crypto v0.7.0 // indirect - golang.org/x/sys v0.6.0 // indirect - golang.org/x/text v0.8.0 // indirect + github.com/klauspost/compress v1.18.2 // indirect + github.com/leodido/go-urn v1.4.0 // indirect + github.com/valyala/bytebufferpool v1.0.0 // indirect + golang.org/x/crypto v0.49.0 // indirect + golang.org/x/sys v0.42.0 // indirect + golang.org/x/text v0.35.0 // indirect ) diff --git a/go.sum b/go.sum index 4f1706d..170161c 100644 --- a/go.sum +++ b/go.sum @@ -1,36 +1,40 @@ -github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwToPjQ= +github.com/andybalholm/brotli v1.2.0/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/gabriel-vasile/mimetype v1.4.13 h1:46nXokslUBsAJE/wMsp5gtO500a4F3Nkz9Ufpk2AcUM= +github.com/gabriel-vasile/mimetype v1.4.13/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s= github.com/go-labx/color v0.0.1 h1:eOgeahCoZ6myc34+ZDAiJUYMvwGuQl52CzSMO3HHLpc= github.com/go-labx/color v0.0.1/go.mod h1:pybhhAum+DQj4h2H9NftbeDfqu1eda6damNXS/i6p6k= github.com/go-labx/lightlog v0.0.3 h1:Ms8ln1U9n45AhRkIPT7bTeAIIybcq96fAbBNNS/0Z1U= github.com/go-labx/lightlog v0.0.3/go.mod h1:ZOhRh8g5iG+9s43xnbtir1wb/uObXTB532K3HFmWrIA= github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s= +github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4= github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA= github.com/go-playground/locales v0.14.1/go.mod h1:hxrqLVvrK65+Rwrd5Fc6F2O76J/NuW9t0sjnWqG1slY= github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJnYK9S473LQFuzCbDbfSFY= github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY= -github.com/go-playground/validator/v10 v10.12.0 h1:E4gtWgxWxp8YSxExrQFv5BpCahla0PVF2oTTEYaWQGI= -github.com/go-playground/validator/v10 v10.12.0/go.mod h1:hCAPuzYvKdP33pxWa+2+6AIKXEKqjIUyqsNCtbsSJrA= -github.com/leodido/go-urn v1.2.2 h1:7z68G0FCGvDk646jz1AelTYNYWrTNm0bEcFAo147wt4= -github.com/leodido/go-urn v1.2.2/go.mod h1:kUaIbLZWttglzwNuG0pgsh5vuV6u2YcGBYz1hIPjtOQ= +github.com/go-playground/validator/v10 v10.30.2 h1:JiFIMtSSHb2/XBUbWM4i/MpeQm9ZK2xqPNk8vgvu5JQ= +github.com/go-playground/validator/v10 v10.30.2/go.mod h1:mAf2pIOVXjTEBrwUMGKkCWKKPs9NheYGabeB04txQSc= +github.com/klauspost/compress v1.18.2 h1:iiPHWW0YrcFgpBYhsA6D1+fqHssJscY/Tm/y2Uqnapk= +github.com/klauspost/compress v1.18.2/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4= +github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ= +github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/rwtodd/Go.Sed v0.0.0-20210816025313-55464686f9ef/go.mod h1:8AEUvGVi2uQ5b24BIhcr0GCcpd/RNAFWaN2CJFrWIIQ= -github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= -github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= -github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= -github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= -github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= -github.com/stretchr/testify v1.8.2 h1:+h33VjcLVPDHtOdpUCuF+7gSuG3yGIftsP1YvFihtJ8= -github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= -golang.org/x/crypto v0.7.0 h1:AvwMYaRytfdeVt3u6mLaxYtErKYjxA2OXjJ1HHq6t3A= -golang.org/x/crypto v0.7.0/go.mod h1:pYwdfH91IfpZVANVyUOhSIPZaFoJGxTFbZhFTx+dXZU= -golang.org/x/sys v0.6.0 h1:MVltZSvRTcU2ljQOhs94SXPftV6DCNnZViHeQps87pQ= -golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/text v0.8.0 h1:57P1ETyNKtuIjB4SRd15iJxuhj8Gc416Y78H3qgMh68= -golang.org/x/text v0.8.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= +github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= +github.com/valyala/bytebufferpool v1.0.0 h1:GqA5TC/0021Y/b9FG4Oi9Mr3q7XYx6KllzawFIhcdPw= +github.com/valyala/bytebufferpool v1.0.0/go.mod h1:6bBcMArwyJ5K/AmCkWv1jt77kVWyCJ6HpOuEn7z0Csc= +github.com/valyala/fasthttp v1.69.0 h1:fNLLESD2SooWeh2cidsuFtOcrEi4uB4m1mPrkJMZyVI= +github.com/valyala/fasthttp v1.69.0/go.mod h1:4wA4PfAraPlAsJ5jMSqCE2ug5tqUPwKXxVj8oNECGcw= +github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU= +github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= +golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4= +golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA= +golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo= +golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8= +golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/group.go b/group.go index c58b49a..f5aa8b4 100644 --- a/group.go +++ b/group.go @@ -1,14 +1,11 @@ package lightning -import "net/http" - // Group represents a group of routes with a common prefix and middleware. type Group struct { - app *Application - parent *Group - prefix string - middlewares []HandlerFunc - // cachedMiddlewares caches the merged middleware chain to avoid repeated allocations + app *Application + parent *Group + prefix string + middlewares []HandlerFunc cachedMiddlewares []HandlerFunc } @@ -69,41 +66,40 @@ func (g *Group) AddRoute(method string, pattern string, handlers []HandlerFunc) // Use adds the given middleware functions to the Group's middleware stack. func (g *Group) Use(middlewares ...HandlerFunc) { g.middlewares = append(g.middlewares, middlewares...) - // Invalidate cache when middleware changes g.cachedMiddlewares = nil } // Get adds a new GET route to the Application with the given pattern and handlers. func (g *Group) Get(pattern string, handlers ...HandlerFunc) { - g.AddRoute(http.MethodGet, pattern, handlers) + g.AddRoute(MethodGet, pattern, handlers) } // Post adds a new POST route to the Application with the given pattern and handlers. func (g *Group) Post(pattern string, handlers ...HandlerFunc) { - g.AddRoute(http.MethodPost, pattern, handlers) + g.AddRoute(MethodPost, pattern, handlers) } // Put adds a new PUT route to the Application with the given pattern and handlers. func (g *Group) Put(pattern string, handlers ...HandlerFunc) { - g.AddRoute(http.MethodPut, pattern, handlers) + g.AddRoute(MethodPut, pattern, handlers) } // Delete adds a new DELETE route to the Application with the given pattern and handlers. func (g *Group) Delete(pattern string, handlers ...HandlerFunc) { - g.AddRoute(http.MethodDelete, pattern, handlers) + g.AddRoute(MethodDelete, pattern, handlers) } // Head adds a new HEAD route to the Application with the given pattern and handlers. func (g *Group) Head(pattern string, handlers ...HandlerFunc) { - g.AddRoute(http.MethodHead, pattern, handlers) + g.AddRoute(MethodHead, pattern, handlers) } // Options adds a new OPTIONS route to the Application with the given pattern and handlers. func (g *Group) Options(pattern string, handlers ...HandlerFunc) { - g.AddRoute(http.MethodOptions, pattern, handlers) + g.AddRoute(MethodOptions, pattern, handlers) } // Patch adds a new PATCH route to the Application with the given pattern and handlers. func (g *Group) Patch(pattern string, handlers ...HandlerFunc) { - g.AddRoute(http.MethodPatch, pattern, handlers) + g.AddRoute(MethodPatch, pattern, handlers) } diff --git a/group_test.go b/group_test.go index 1e1d582..360c7d9 100644 --- a/group_test.go +++ b/group_test.go @@ -1,10 +1,10 @@ package lightning import ( - "net/http" - "net/http/httptest" "reflect" "testing" + + "github.com/valyala/fasthttp" ) func TestNewGroup(t *testing.T) { @@ -63,22 +63,18 @@ func TestGroup(t *testing.T) { app := NewApp() group := app.Group("/api") - // Test that the group has the correct prefix if group.prefix != "/api" { t.Errorf("Expected prefix to be '/api', but got '%s'", group.prefix) } - // Test that the group has the correct parent if group.parent != nil { t.Errorf("Expected parent to be nil, but got '%v'", group.parent) } - // Test that the group has the correct middleware if len(group.middlewares) != 0 { t.Errorf("Expected middlewares to be empty, but got '%v'", group.middlewares) } - // Test that the group has the correct application if group.app != app { t.Errorf("Expected app to be '%v', but got '%v'", app, group.app) } @@ -88,9 +84,9 @@ func TestGroup_AddRoute(t *testing.T) { app := NewApp() group := app.Group("/prefix") handlers := []HandlerFunc{func(c *Context) {}} - group.AddRoute(http.MethodGet, "/path", handlers) + group.AddRoute(MethodGet, "/path", handlers) - searchHandlers, _ := app.router.findRoute(http.MethodGet, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodGet, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } @@ -100,28 +96,24 @@ func TestGroup_Use(t *testing.T) { app := NewApp() group := app.Group("/test") - // Define a middleware function that sets a custom header middleware := func(c *Context) { c.SetHeader("X-Test-Header", "123") } - // Add the middleware function to the Group group.Use(middleware) - // Define a route that returns the value of the custom header group.Get("/header", func(c *Context) { header := c.Header("X-Test-Header") - c.Text(http.StatusOK, header) + c.Text(StatusOK, header) }) - // Send a request to the route using an HTTP client - req, _ := http.NewRequest(http.MethodGet, "/test/header", nil) - w := httptest.NewRecorder() - app.ServeHTTP(w, req) + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod(MethodGet) + ctx.Request.Header.SetRequestURI("/test/header") + app.serveRequest(ctx) - // Verify that the response contains the custom header - if w.Header().Get("X-Test-Header") != "123" { - t.Errorf("Expected header to be '%v', but got '%v'", 123, w.Header().Get("X-Test-Header")) + if string(ctx.Response.Header.Peek("X-Test-Header")) != "123" { + t.Errorf("Expected header to be '123', but got '%s'", string(ctx.Response.Header.Peek("X-Test-Header"))) } } @@ -131,7 +123,7 @@ func TestGroup_Get(t *testing.T) { handlers := []HandlerFunc{func(c *Context) {}} group.Get("/path", handlers...) - searchHandlers, _ := app.router.findRoute(http.MethodGet, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodGet, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } @@ -143,7 +135,7 @@ func TestGroup_Post(t *testing.T) { handlers := []HandlerFunc{func(c *Context) {}} group.Post("/path", handlers...) - searchHandlers, _ := app.router.findRoute(http.MethodPost, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodPost, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } @@ -155,7 +147,7 @@ func TestGroup_Put(t *testing.T) { handlers := []HandlerFunc{func(c *Context) {}} group.Put("/path", handlers...) - searchHandlers, _ := app.router.findRoute(http.MethodPut, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodPut, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } @@ -167,7 +159,7 @@ func TestGroup_Delete(t *testing.T) { handlers := []HandlerFunc{func(c *Context) {}} group.Delete("/path", handlers...) - searchHandlers, _ := app.router.findRoute(http.MethodDelete, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodDelete, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } @@ -179,7 +171,7 @@ func TestGroup_Head(t *testing.T) { handlers := []HandlerFunc{func(c *Context) {}} group.Head("/path", handlers...) - searchHandlers, _ := app.router.findRoute(http.MethodHead, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodHead, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } @@ -191,7 +183,7 @@ func TestGroup_Patch(t *testing.T) { handlers := []HandlerFunc{func(c *Context) {}} group.Patch("/path", handlers...) - searchHandlers, _ := app.router.findRoute(http.MethodPatch, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodPatch, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } @@ -203,7 +195,7 @@ func TestGroup_Options(t *testing.T) { handlers := []HandlerFunc{func(c *Context) {}} group.Options("/path", handlers...) - searchHandlers, _ := app.router.findRoute(http.MethodOptions, "/prefix/path") + searchHandlers, _ := app.router.findRoute(MethodOptions, "/prefix/path") if reflect.ValueOf(searchHandlers[0]) != reflect.ValueOf(handlers[0]) { t.Errorf("Expected handlers to be '%v', but got '%v'", searchHandlers[0], handlers[0]) } diff --git a/integration_test.go b/integration_test.go index 44df0fd..cce5d42 100644 --- a/integration_test.go +++ b/integration_test.go @@ -1,740 +1,328 @@ package lightning import ( - "context" - "net" - "net/http" - "net/http/httptest" - "os" - "strings" "testing" - "text/template" - "time" -) - -func TestShutdown(t *testing.T) { - app := NewApp() - - listener, err := net.Listen("tcp", "localhost:0") - if err != nil { - t.Fatalf("Failed to create listener: %v", err) - } - - go app.RunListener(listener) - time.Sleep(100 * time.Millisecond) - - // Test shutdown with nil server (should return nil) - app2 := NewApp() - err = app2.Shutdown(context.Background()) - if err != nil { - t.Errorf("Expected nil error for shutdown with nil server, got %v", err) - } - - // Test shutdown with running server - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) - defer cancel() + "github.com/valyala/fasthttp" +) - err = app.Shutdown(ctx) - if err != nil { - t.Errorf("Expected nil error for graceful shutdown, got %v", err) - } +func newTestCtx(method, path string) *fasthttp.RequestCtx { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod(method) + ctx.Request.Header.SetRequestURI(path) + return ctx } -func TestRunGraceful(t *testing.T) { +func TestIntegrationBasicRouting(t *testing.T) { app := NewApp() - app.Get("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello") - }) - - listener, err := net.Listen("tcp", "localhost:0") - if err != nil { - t.Fatalf("Failed to create listener: %v", err) - } - addr := listener.Addr().String() - listener.Close() // Close it so RunGraceful can use a new one - - // Start server in goroutine with signal handling - serverErr := make(chan error, 1) - go func() { - serverErr <- app.RunGraceful(2*time.Second, addr) - }() - - // Wait for server to start - time.Sleep(200 * time.Millisecond) - - // Make a request to verify server is running - resp, err := http.Get("http://" + addr + "/test") - if err != nil { - t.Errorf("Failed to make request: %v", err) - } else { - resp.Body.Close() - } - - // Shutdown the server via the app's Shutdown method - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) - defer cancel() - app.Shutdown(ctx) - // Wait for RunGraceful to return - select { - case err := <-serverErr: - if err != nil && err != http.ErrServerClosed { - t.Errorf("Unexpected error: %v", err) - } - case <-time.After(3 * time.Second): - t.Error("RunGraceful did not return in time") - } -} - -func TestLoggerMiddleware(t *testing.T) { - app := NewApp() - app.Use(Logger()) - app.Get("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + app.Get("/get", func(c *Context) { + c.Text(StatusOK, "GET response") }) - req := httptest.NewRequest("GET", "/test", nil) - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) - } -} - -func TestRemoteAddr(t *testing.T) { - tests := []struct { - name string - headers map[string]string - remoteAddr string - expected string - }{ - { - name: "X-Real-IP header", - headers: map[string]string{"X-Real-Ip": "192.168.1.1"}, - remoteAddr: "10.0.0.1:12345", - expected: "192.168.1.1", - }, - { - name: "X-Forwarded-For single IP", - headers: map[string]string{"X-Forwarded-For": "192.168.1.2"}, - remoteAddr: "10.0.0.1:12345", - expected: "192.168.1.2", - }, - { - name: "X-Forwarded-For multiple IPs", - headers: map[string]string{"X-Forwarded-For": "192.168.1.3, 10.0.0.2, 10.0.0.3"}, - remoteAddr: "10.0.0.1:12345", - expected: "192.168.1.3", - }, - { - name: "No headers - use RemoteAddr", - headers: map[string]string{}, - remoteAddr: "10.0.0.1:12345", - expected: "10.0.0.1:12345", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - req := httptest.NewRequest("GET", "/test", nil) - for k, v := range tt.headers { - req.Header.Set(k, v) - } - req.RemoteAddr = tt.remoteAddr - - r, err := newRequest(req) - if err != nil { - t.Fatalf("Failed to create request: %v", err) - } - - got := r.remoteAddr() - if got != tt.expected { - t.Errorf("Expected remoteAddr %q, got %q", tt.expected, got) - } - }) - } -} + app.Post("/post", func(c *Context) { + c.Text(StatusOK, "POST response") + }) -func TestRouterSearch(t *testing.T) { - r := newRouter() + app.Put("/put", func(c *Context) { + c.Text(StatusOK, "PUT response") + }) - // Add routes - test basic routing - r.addRoute("GET", "/", []HandlerFunc{func(c *Context) {}}) - r.addRoute("GET", "/hello", []HandlerFunc{func(c *Context) {}}) - r.addRoute("GET", "/hello/:name", []HandlerFunc{func(c *Context) {}}) - r.addRoute("POST", "/users", []HandlerFunc{func(c *Context) {}}) + app.Delete("/delete", func(c *Context) { + c.Text(StatusOK, "DELETE response") + }) tests := []struct { - method string - path string - found bool - paramCount int + method string + path string + want int }{ - {"GET", "/", true, 0}, - {"GET", "/hello", true, 0}, - {"GET", "/hello/world", true, 1}, - {"POST", "/users", true, 0}, - {"GET", "/nonexistent", false, 0}, - {"DELETE", "/", false, 0}, // method not found + {MethodGet, "/get", StatusOK}, + {MethodPost, "/post", StatusOK}, + {MethodPut, "/put", StatusOK}, + {MethodDelete, "/delete", StatusOK}, + {MethodGet, "/nonexistent", StatusNotFound}, } for _, tt := range tests { - t.Run(tt.method+"_"+tt.path, func(t *testing.T) { - handlers, params := r.findRoute(tt.method, tt.path) - if tt.found && handlers == nil { - t.Errorf("Expected to find route for %s %s", tt.method, tt.path) - } - if !tt.found && handlers != nil { - t.Errorf("Expected not to find route for %s %s", tt.method, tt.path) - } - if tt.found && len(params) != tt.paramCount { - t.Errorf("Expected %d params for %s %s, got %d", tt.paramCount, tt.method, tt.path, len(params)) + t.Run(tt.method+" "+tt.path, func(t *testing.T) { + ctx := newTestCtx(tt.method, tt.path) + app.serveRequest(ctx) + if ctx.Response.StatusCode() != tt.want { + t.Errorf("Expected status %d, got %d", tt.want, ctx.Response.StatusCode()) } }) } } -func TestResponseFlushWithCookies(t *testing.T) { - req := httptest.NewRequest("GET", "/test", nil) - w := httptest.NewRecorder() - - res := newResponse(req, w) - res.setStatus(http.StatusOK) - res.setBody([]byte("test")) - - // Add cookies - res.cookies.set("session", "abc123") - - res.flush() - - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) - } - - cookies := w.Result().Cookies() - if len(cookies) != 1 { - t.Errorf("Expected 1 cookie, got %d", len(cookies)) - } - if cookies[0].Name != "session" || cookies[0].Value != "abc123" { - t.Errorf("Unexpected cookie: %v", cookies[0]) - } -} - -func TestResponseFlushWithRedirect(t *testing.T) { - req := httptest.NewRequest("GET", "/old", nil) - w := httptest.NewRecorder() - - res := newResponse(req, w) - res.redirect(http.StatusMovedPermanently, "/new") - - res.flush() - - if w.Code != http.StatusMovedPermanently { - t.Errorf("Expected status %d, got %d", http.StatusMovedPermanently, w.Code) - } - - location := w.Header().Get("Location") - if location != "/new" { - t.Errorf("Expected location '/new', got '%s'", location) - } -} - -func TestServeHTTPWithMaxBodySize(t *testing.T) { - app := NewApp(&Config{ - MaxRequestBodySize: 10, // 10 bytes limit - }) - app.Post("/test", func(c *Context) { - c.Text(http.StatusOK, "OK") - }) - - // Request with body smaller than limit - body := strings.NewReader("small") - req := httptest.NewRequest("POST", "/test", body) - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) - } -} - -func TestAcquireContextPanic(t *testing.T) { - app := NewApp() - - // Create a request that will cause panic in newRequest - req := httptest.NewRequest("GET", "/test", nil) - req.Body = nil // This should not cause panic - - ctx := app.acquireContext(httptest.NewRecorder(), req) - if ctx == nil { - t.Error("Expected non-nil context") - } - - app.releaseContext(ctx) -} - -func TestConfigMergeMultiple(t *testing.T) { - config1 := &Config{ - AppName: "app1", - } - config2 := &Config{ - EnableDebug: true, - } - config3 := &Config{ - AppName: "app3", - } - - merged := defaultConfig() - merged = merged.merge(config1, nil, config2, config3) - - if merged.AppName != "app3" { - t.Errorf("Expected AppName 'app3', got '%s'", merged.AppName) - } - if !merged.EnableDebug { - t.Error("Expected EnableDebug to be true") - } -} - -func TestJSONBodyWithValidation(t *testing.T) { +func TestIntegrationJSONAPI(t *testing.T) { app := NewApp() - type User struct { - Name string `validate:"required"` - Email string `validate:"required,email"` + type Response struct { + Message string `json:"message"` } - app.Post("/user", func(c *Context) { - var user User - if err := c.JSONBody(&user, true); err != nil { - c.Fail(400, err.Error()) - return - } - c.Success(user) + app.Get("/api/json", func(c *Context) { + c.JSON(StatusOK, Response{Message: "hello"}) }) - // Valid request - body := `{"name":"John","email":"john@example.com"}` - req := httptest.NewRequest("POST", "/user", strings.NewReader(body)) - req.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() + ctx := newTestCtx(MethodGet, "/api/json") + app.serveRequest(ctx) - app.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) - } - - // Invalid request - missing required field - body = `{"name":""}` - req = httptest.NewRequest("POST", "/user", strings.NewReader(body)) - req.Header.Set("Content-Type", "application/json") - w = httptest.NewRecorder() - - app.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } } -func TestGroupCachedMiddlewares(t *testing.T) { +func TestIntegrationMiddlewareChain(t *testing.T) { app := NewApp() + called := []string{} - g := app.Group("/api") - g.Use(func(c *Context) { c.Next() }) - - // First call - computes and caches - mw1 := g.getMiddlewares() - if len(mw1) != 1 { - t.Errorf("Expected 1 middleware, got %d", len(mw1)) - } - - // Second call - should return cached result - mw2 := g.getMiddlewares() - if len(mw2) != 1 { - t.Errorf("Expected 1 middleware, got %d", len(mw2)) - } - - // Add more middleware - should invalidate cache - g.Use(func(c *Context) { c.Next() }) - mw3 := g.getMiddlewares() - if len(mw3) != 2 { - t.Errorf("Expected 2 middlewares, got %d", len(mw3)) - } -} - -func TestNestedGroupMiddlewares(t *testing.T) { - app := NewApp() - - parent := app.Group("/api") - parent.Use(func(c *Context) { c.Next() }) - - child := parent.Group("/v1") - child.Use(func(c *Context) { c.Next() }) - - mw := child.getMiddlewares() - if len(mw) != 2 { - t.Errorf("Expected 2 middlewares, got %d", len(mw)) - } -} - -func TestFileResponse(t *testing.T) { - // Create a temp file - tmpFile := "/tmp/test_lightning_file.txt" - err := os.WriteFile(tmpFile, []byte("test content"), 0644) - if err != nil { - t.Fatalf("Failed to create temp file: %v", err) - } - defer os.Remove(tmpFile) + app.Use(func(c *Context) { + called = append(called, "middleware1-before") + c.Next() + called = append(called, "middleware1-after") + }) - req := httptest.NewRequest("GET", "/file", nil) - w := httptest.NewRecorder() + app.Use(func(c *Context) { + called = append(called, "middleware2-before") + c.Next() + called = append(called, "middleware2-after") + }) - res := newResponse(req, w) - err = res.file(tmpFile) - if err != nil { - t.Errorf("Failed to set file: %v", err) - } + app.Get("/test", func(c *Context) { + called = append(called, "handler") + }) - res.flush() + ctx := newTestCtx(MethodGet, "/test") + app.serveRequest(ctx) - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) + expected := []string{ + "middleware1-before", + "middleware2-before", + "handler", + "middleware2-after", + "middleware1-after", } -} - -func TestFileResponseNotFound(t *testing.T) { - req := httptest.NewRequest("GET", "/file", nil) - w := httptest.NewRecorder() - res := newResponse(req, w) - err := res.file("/nonexistent/path/file.txt") - if err == nil { - t.Error("Expected error for nonexistent file") + if len(called) != len(expected) { + t.Fatalf("Expected %d calls, got %d: %v", len(expected), len(called), called) } -} - -func TestHTMLResponse(t *testing.T) { - app := NewApp() - // Skip HTML test since it requires template setup - app.SetFuncMap(template.FuncMap{}) - // Note: LoadHTMLGlob requires actual template files -} - -func TestXMLResponse(t *testing.T) { - app := NewApp() - app.Get("/xml", func(c *Context) { - type Data struct { - Key string `xml:"key"` + for i, e := range expected { + if called[i] != e { + t.Errorf("Call[%d] = %s, want %s", i, called[i], e) } - c.XML(http.StatusOK, &Data{Key: "value"}) - }) - - req := httptest.NewRequest("GET", "/xml", nil) - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) - } - - body := w.Body.String() - if !strings.Contains(body, "value") { - t.Errorf("Expected body to contain 'value', got %q", body) } } -func TestContextDataOperations(t *testing.T) { +func TestIntegrationQueryParams(t *testing.T) { app := NewApp() - var gotV1, gotV2 any - var v1AfterDel any - - app.Get("/test", func(c *Context) { - c.SetData("key1", "value1") - c.SetData("key2", 123) - - gotV1 = c.GetData("key1") - gotV2 = c.GetData("key2") - - c.DelData("key1") - v1AfterDel = c.GetData("key1") - - c.Text(http.StatusOK, "OK") - }) - - req := httptest.NewRequest("GET", "/test", nil) - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) - - if w.Code != http.StatusOK { - t.Errorf("Expected status %d, got %d", http.StatusOK, w.Code) - } - if gotV1 != "value1" { - t.Errorf("Expected value1, got %v", gotV1) - } - if gotV2 != 123 { - t.Errorf("Expected 123, got %v", gotV2) - } - if v1AfterDel != nil { - t.Errorf("Expected nil after delete, got %v", v1AfterDel) - } -} - -func TestRedirectResponse(t *testing.T) { - app := NewApp() - app.Get("/redirect", func(c *Context) { - c.Redirect(http.StatusMovedPermanently, "/new") + app.Get("/search", func(c *Context) { + query := c.Query("q") + c.JSON(StatusOK, map[string]string{"query": query}) }) - req := httptest.NewRequest("GET", "/redirect", nil) - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) - - if w.Code != http.StatusMovedPermanently { - t.Errorf("Expected status %d, got %d", http.StatusMovedPermanently, w.Code) - } + ctx := newTestCtx(MethodGet, "/search?q=test") + app.serveRequest(ctx) - location := w.Header().Get("Location") - if location != "/new" { - t.Errorf("Expected location /new, got %s", location) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } } -func TestUserAgentAndReferer(t *testing.T) { +func TestIntegrationPathParams(t *testing.T) { app := NewApp() - var gotUA, gotRef string - - app.Get("/test", func(c *Context) { - gotUA = c.UserAgent() - gotRef = c.Referer() - c.Text(http.StatusOK, "OK") + app.Get("/users/:id/posts/:postId", func(c *Context) { + id := c.Param("id") + postId := c.Param("postId") + c.JSON(StatusOK, map[string]string{"userId": id, "postId": postId}) }) - req := httptest.NewRequest("GET", "/test", nil) - req.Header.Set("User-Agent", "test-agent") - req.Header.Set("Referer", "http://example.com") - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) + ctx := newTestCtx(MethodGet, "/users/123/posts/456") + app.serveRequest(ctx) - if gotUA != "test-agent" { - t.Errorf("Expected user agent 'test-agent', got '%s'", gotUA) - } - if gotRef != "http://example.com" { - t.Errorf("Expected referer 'http://example.com', got '%s'", gotRef) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } } -func TestCookieOperations(t *testing.T) { +func TestIntegrationCookie(t *testing.T) { app := NewApp() - var gotCookieValue string - - app.Get("/cookie", func(c *Context) { - cookie := c.Cookie("session") - if cookie != nil { - gotCookieValue = cookie.Value - } - c.SetCookie("new_cookie", "new_value") - c.Text(http.StatusOK, "OK") + app.Get("/set-cookie", func(c *Context) { + c.SetCookie("session", "abc123") + c.Text(StatusOK, "cookie set") }) - req := httptest.NewRequest("GET", "/cookie", nil) - req.AddCookie(&http.Cookie{Name: "session", Value: "abc123"}) - w := httptest.NewRecorder() + ctx := newTestCtx(MethodGet, "/set-cookie") + app.serveRequest(ctx) - app.ServeHTTP(w, req) - - if gotCookieValue != "abc123" { - t.Errorf("Expected cookie value abc123, got %s", gotCookieValue) - } - - cookies := w.Result().Cookies() - if len(cookies) != 1 { - t.Errorf("Expected 1 cookie, got %d", len(cookies)) + cookie := string(ctx.Response.Header.Peek("Set-Cookie")) + if cookie == "" { + t.Error("Expected Set-Cookie header") } } -func TestCustomCookie(t *testing.T) { +func TestIntegrationRedirect(t *testing.T) { app := NewApp() - app.Get("/cookie", func(c *Context) { - c.SetCustomCookie(&http.Cookie{ - Name: "custom", - Value: "value", - Path: "/", - MaxAge: 3600, - HttpOnly: true, - }) - c.Text(http.StatusOK, "OK") - }) - req := httptest.NewRequest("GET", "/cookie", nil) - w := httptest.NewRecorder() + app.Get("/old", func(c *Context) { + c.Redirect(StatusMovedPermanently, "/new") + }) - app.ServeHTTP(w, req) + ctx := newTestCtx(MethodGet, "/old") + app.serveRequest(ctx) - cookies := w.Result().Cookies() - if len(cookies) != 1 { - t.Errorf("Expected 1 cookie, got %d", len(cookies)) + if ctx.Response.StatusCode() != StatusMovedPermanently { + t.Errorf("Expected status %d, got %d", StatusMovedPermanently, ctx.Response.StatusCode()) } - if cookies[0].Name != "custom" { - t.Errorf("Expected cookie name 'custom', got '%s'", cookies[0].Name) + location := string(ctx.Response.Header.Peek("Location")) + if location == "" { + t.Error("Expected Location header to be set") } } -func TestHeaderOperations(t *testing.T) { +func TestIntegrationHeaders(t *testing.T) { app := NewApp() - var gotHeaderValue string - app.Get("/headers", func(c *Context) { - gotHeaderValue = c.Header("X-Custom-Header") - c.Text(http.StatusOK, "OK") + c.SetHeader("X-Custom", "value") + c.JSON(StatusOK, map[string]string{"header": "set"}) }) - req := httptest.NewRequest("GET", "/headers", nil) - req.Header.Set("X-Custom-Header", "value") - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) + ctx := newTestCtx(MethodGet, "/headers") + app.serveRequest(ctx) - if gotHeaderValue != "value" { - t.Errorf("Expected header value 'value', got '%s'", gotHeaderValue) + header := string(ctx.Response.Header.Peek("X-Custom")) + if header != "value" { + t.Errorf("Expected X-Custom header 'value', got '%s'", header) } } -func TestAddAndDelHeaders(t *testing.T) { +func TestIntegrationContentNegotiation(t *testing.T) { app := NewApp() - app.Get("/headers", func(c *Context) { - c.AddHeader("X-Multi", "value1") - c.AddHeader("X-Multi", "value2") - - c.SetHeader("X-Set", "value3") - - c.DelHeader("X-Delete-Me") - c.Text(http.StatusOK, "OK") + app.Get("/text", func(c *Context) { + c.Text(StatusOK, "plain text") }) - req := httptest.NewRequest("GET", "/headers", nil) - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) - - values := w.Header().Values("X-Multi") - if len(values) != 2 { - t.Errorf("Expected 2 values for X-Multi, got %d", len(values)) - } + app.Get("/json", func(c *Context) { + c.JSON(StatusOK, map[string]string{"key": "value"}) + }) - if w.Header().Get("X-Set") != "value3" { - t.Errorf("Expected X-Set to be 'value3'") + ctx := newTestCtx(MethodGet, "/text") + app.serveRequest(ctx) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d for text, got %d", StatusOK, ctx.Response.StatusCode()) } - if w.Header().Get("X-Delete-Me") != "" { - t.Error("Expected X-Delete-Me to be deleted") + ctx = newTestCtx(MethodGet, "/json") + app.serveRequest(ctx) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d for json, got %d", StatusOK, ctx.Response.StatusCode()) } } -func TestQueries(t *testing.T) { +func TestIntegrationErrorHandling(t *testing.T) { app := NewApp() - var gotQueryValue string - - app.Get("/queries", func(c *Context) { - queries := c.Queries() - if len(queries["key1"]) > 0 { - gotQueryValue = queries["key1"][0] - } - c.Text(http.StatusOK, "OK") + app.Get("/error", func(c *Context) { + c.JSONError(StatusBadRequest, "invalid request") }) - req := httptest.NewRequest("GET", "/queries?key1=value1&key2=value2", nil) - w := httptest.NewRecorder() - - app.ServeHTTP(w, req) + ctx := newTestCtx(MethodGet, "/error") + app.serveRequest(ctx) - if gotQueryValue != "value1" { - t.Errorf("Expected query value 'value1', got '%s'", gotQueryValue) + if ctx.Response.StatusCode() != StatusBadRequest { + t.Errorf("Expected status %d, got %d", StatusBadRequest, ctx.Response.StatusCode()) } } -func TestSuccessFailResponses(t *testing.T) { +func TestIntegrationSuccessFail(t *testing.T) { app := NewApp() + app.Get("/success", func(c *Context) { - c.Success(map[string]string{"name": "test"}) + c.Success(map[string]string{"result": "ok"}) }) + app.Get("/fail", func(c *Context) { - c.Fail(400, "Bad Request") + c.Fail(1001, "operation failed") }) - // Test success - req := httptest.NewRequest("GET", "/success", nil) - w := httptest.NewRecorder() - app.ServeHTTP(w, req) - - if !strings.Contains(w.Body.String(), `"code":0`) { - t.Error("Expected success response with code 0") + ctx := newTestCtx(MethodGet, "/success") + app.serveRequest(ctx) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d for success, got %d", StatusOK, ctx.Response.StatusCode()) } - // Test fail - req = httptest.NewRequest("GET", "/fail", nil) - w = httptest.NewRecorder() - app.ServeHTTP(w, req) - - if !strings.Contains(w.Body.String(), `"code":400`) { - t.Error("Expected fail response with code 400") + ctx = newTestCtx(MethodGet, "/fail") + app.serveRequest(ctx) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d for fail, got %d", StatusOK, ctx.Response.StatusCode()) } } -func TestBodyOperations(t *testing.T) { +func TestIntegrationNestedGroups(t *testing.T) { app := NewApp() - var gotBody []byte + api := app.Group("/api") + v1 := api.Group("/v1") + v2 := api.Group("/v2") - app.Get("/body", func(c *Context) { - c.SetBody([]byte("custom body")) - gotBody = c.Body() + v1.Get("/resource", func(c *Context) { + c.Text(StatusOK, "v1 resource") }) - req := httptest.NewRequest("GET", "/body", nil) - w := httptest.NewRecorder() + v2.Get("/resource", func(c *Context) { + c.Text(StatusOK, "v2 resource") + }) - app.ServeHTTP(w, req) + ctx := newTestCtx(MethodGet, "/api/v1/resource") + app.serveRequest(ctx) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d for v1, got %d", StatusOK, ctx.Response.StatusCode()) + } - if string(gotBody) != "custom body" { - t.Errorf("Expected body 'custom body', got '%s'", string(gotBody)) + ctx = newTestCtx(MethodGet, "/api/v2/resource") + app.serveRequest(ctx) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d for v2, got %d", StatusOK, ctx.Response.StatusCode()) } } -func TestStatus(t *testing.T) { +func TestIntegrationGroupMiddleware(t *testing.T) { app := NewApp() + middlewareOrder := []string{} - var gotStatus int + api := app.Group("/api") + api.Use(func(c *Context) { + middlewareOrder = append(middlewareOrder, "api-middleware") + c.Next() + }) - app.Get("/status", func(c *Context) { - c.SetStatus(http.StatusCreated) - gotStatus = c.Status() + users := api.Group("/users") + users.Use(func(c *Context) { + middlewareOrder = append(middlewareOrder, "users-middleware") + c.Next() }) - req := httptest.NewRequest("GET", "/status", nil) - w := httptest.NewRecorder() + users.Get("/:id", func(c *Context) { + middlewareOrder = append(middlewareOrder, "users-handler") + }) - app.ServeHTTP(w, req) + ctx := newTestCtx(MethodGet, "/api/users/123") + app.serveRequest(ctx) - if gotStatus != http.StatusCreated { - t.Errorf("Expected status %d, got %d", http.StatusCreated, gotStatus) + expected := []string{"api-middleware", "users-middleware", "users-handler"} + if len(middlewareOrder) != len(expected) { + t.Errorf("Expected %d middleware calls, got %d: %v", len(expected), len(middlewareOrder), middlewareOrder) } - if w.Code != http.StatusCreated { - t.Errorf("Expected status %d, got %d", http.StatusCreated, w.Code) + for i, e := range expected { + if middlewareOrder[i] != e { + t.Errorf("Call[%d] = %s, want %s", i, middlewareOrder[i], e) + } } } diff --git a/lightning.go b/lightning.go index 45654aa..9042390 100644 --- a/lightning.go +++ b/lightning.go @@ -1,31 +1,30 @@ package lightning import ( - "context" "encoding/json" - "net" - "net/http" "os" "os/signal" "path" "path/filepath" - "reflect" "strings" "sync" "syscall" "text/template" - "time" "github.com/go-labx/lightlog" + "github.com/valyala/fasthttp" ) // HandlerFunc is a function type that represents the actual handler function for a route. type HandlerFunc func(*Context) + +// Middleware is an alias for HandlerFunc, representing middleware functions. type Middleware = HandlerFunc // Map is a shortcut for map[string]interface{} type Map map[string]any +// Application is the main struct that holds the router, middlewares, and configuration. type Application struct { Config *Config router *router @@ -35,62 +34,79 @@ type Application struct { Logger *lightlog.ConsoleLogger - server *http.Server - mu sync.Mutex + server *fasthttp.Server + mu sync.Mutex contextPool sync.Pool } +// Config holds the configuration for the Application. type Config struct { - AppName string - JSONEncoder JSONMarshal - JSONDecoder JSONUnmarshal - NotFoundHandler HandlerFunc // Handler function for 404 Not Found error - EnableDebug bool - MaxRequestBodySize int64 // Max request body size in bytes, 0 means unlimited + AppName string + JSONEncoder JSONMarshal + JSONDecoder JSONUnmarshal + NotFoundHandler HandlerFunc + EnableDebug bool + MaxRequestBodySize int64 } +// merge merges the given Config structs into the current Config. func (c *Config) merge(configs ...*Config) *Config { - value := reflect.ValueOf(c).Elem() - - // iterate over all the configs passed in - for _, config := range configs { - if config == nil { + for _, cfg := range configs { + if cfg == nil { continue } - v := reflect.ValueOf(config).Elem() - t := reflect.TypeOf(config).Elem() - - // iterate over all the fields in the config - for i := 0; i < t.NumField(); i++ { - // if the field is not zero, set the value of the field in the current config to the value of the field in the passed in config - if !v.Field(i).IsZero() { - value.Field(i).Set(v.Field(i)) - } + if cfg.AppName != "" { + c.AppName = cfg.AppName + } + if cfg.JSONEncoder != nil { + c.JSONEncoder = cfg.JSONEncoder + } + if cfg.JSONDecoder != nil { + c.JSONDecoder = cfg.JSONDecoder + } + if cfg.NotFoundHandler != nil { + c.NotFoundHandler = cfg.NotFoundHandler + } + if cfg.EnableDebug { + c.EnableDebug = cfg.EnableDebug + } + if cfg.MaxRequestBodySize > 0 { + c.MaxRequestBodySize = cfg.MaxRequestBodySize } } - return c } +// defaultConfig returns a new Config with default values. func defaultConfig() *Config { return &Config{ AppName: "lightning-app", - JSONEncoder: json.Marshal, - JSONDecoder: json.Unmarshal, + JSONEncoder: defaultJSONMarshal, + JSONDecoder: defaultJSONUnmarshal, NotFoundHandler: defaultNotFound, EnableDebug: false, } } +// defaultJSONMarshal is the default JSON marshaling function. +func defaultJSONMarshal(v any) ([]byte, error) { + return json.Marshal(v) +} + +// defaultJSONUnmarshal is the default JSON unmarshaling function. +func defaultJSONUnmarshal(data []byte, v any) error { + return json.Unmarshal(data, v) +} + // NewApp returns a new instance of the Application struct. func NewApp(c ...*Config) *Application { config := defaultConfig() config = config.merge(c...) app := &Application{ - Config: config, - router: newRouter(), - Logger: lightlog.NewConsoleLogger(config.AppName, lightlog.TRACE), + Config: config, + router: newRouter(), + Logger: lightlog.NewConsoleLogger(config.AppName, lightlog.TRACE), contextPool: sync.Pool{ New: func() interface{} { return &Context{index: -1} @@ -122,7 +138,7 @@ func (app *Application) Use(middlewares ...Middleware) { app.middlewares = append(app.middlewares, middlewares...) } -// AddRoute is a function that adds a new route to the router. +// AddRoute adds a new route to the router. // It composes the global middlewares, route-specific middlewares, and the actual handler function // to form a single MiddlewareFunc, and then adds it to the router. func (app *Application) AddRoute(method string, pattern string, handlers []HandlerFunc) { @@ -134,42 +150,39 @@ func (app *Application) AddRoute(method string, pattern string, handlers []Handl app.router.addRoute(method, pattern, allHandlers) } -// The following functions are shortcuts for the addRoute function. -// They pre-fill the method parameter and call the addRoute function. - // Get adds a new route with method "GET" to the router. func (app *Application) Get(pattern string, handlers ...HandlerFunc) { - app.AddRoute("GET", pattern, handlers) + app.AddRoute(MethodGet, pattern, handlers) } // Post adds a new route with method "POST" to the router. func (app *Application) Post(pattern string, handlers ...HandlerFunc) { - app.AddRoute("POST", pattern, handlers) + app.AddRoute(MethodPost, pattern, handlers) } // Put adds a new route with method "PUT" to the router. func (app *Application) Put(pattern string, handlers ...HandlerFunc) { - app.AddRoute("PUT", pattern, handlers) + app.AddRoute(MethodPut, pattern, handlers) } // Delete adds a new route with method "DELETE" to the router. func (app *Application) Delete(pattern string, handlers ...HandlerFunc) { - app.AddRoute("DELETE", pattern, handlers) + app.AddRoute(MethodDelete, pattern, handlers) } // Head adds a new route with method "HEAD" to the router. func (app *Application) Head(pattern string, handlers ...HandlerFunc) { - app.AddRoute("HEAD", pattern, handlers) + app.AddRoute(MethodHead, pattern, handlers) } // Patch adds a new route with method "PATCH" to the router. func (app *Application) Patch(pattern string, handlers ...HandlerFunc) { - app.AddRoute("PATCH", pattern, handlers) + app.AddRoute(MethodPatch, pattern, handlers) } // Options adds a new route with method "OPTIONS" to the router. func (app *Application) Options(pattern string, handlers ...HandlerFunc) { - app.AddRoute("OPTIONS", pattern, handlers) + app.AddRoute(MethodOptions, pattern, handlers) } // Group returns a new instance of the Group struct with the given prefix. @@ -178,26 +191,33 @@ func (app *Application) Group(prefix string) *Group { } // Static serves static files from the given root directory with the given prefix. -// It uses the os.Executable function to get the path of the executable file, -// and then joins it with the root and the path after the prefix to get the full file path. -// If the file exists, it is served with a 200 status code using the http.ServeFile function. +// If root is an absolute path, it is used directly. Otherwise, it is resolved relative +// to the executable's directory. +// If the file exists, it is served with a 200 status code. // If the file does not exist, a 404 status code is returned with the text "Not Found". func (app *Application) Static(root string, prefix string) { - ex, err := os.Executable() - if err != nil { - panic(err) + exPath := "" + if filepath.IsAbs(root) { + exPath = "" + } else { + ex, err := os.Executable() + if err != nil { + app.Logger.Warn("Failed to get executable path for static files: %v, using current directory", err) + exPath = "." + } else { + exPath = filepath.Dir(ex) + } } - exPath := filepath.Dir(ex) app.Get(path.Join(prefix, "/*"), func(ctx *Context) { fullFilePath := filepath.Join(exPath, root, strings.TrimPrefix(ctx.Path, prefix)) if _, err := os.Stat(fullFilePath); !os.IsNotExist(err) { ctx.SkipFlush() - ctx.SetStatus(http.StatusOK) - http.ServeFile(ctx.Res, ctx.Req, fullFilePath) + ctx.SetStatus(StatusOK) + ctx.ctx.SendFile(fullFilePath) } else { - ctx.Text(http.StatusNotFound, http.StatusText(http.StatusNotFound)) + ctx.Text(StatusNotFound, "Not Found") } }) } @@ -214,62 +234,50 @@ func (app *Application) LoadHTMLGlob(pattern string) { app.htmlTemplates = template.Must(template.New("").Funcs(app.funcMap).ParseGlob(pattern)) } -// ServeHTTP is the function that handles HTTP requests. -// It finds the matching route, creates a new Context, sets the route parameters, -// and executes the MiddlewareFunc chain. -func (app *Application) ServeHTTP(w http.ResponseWriter, req *http.Request) { - // Apply max request body size limit - if app.Config.MaxRequestBodySize > 0 { - req.Body = http.MaxBytesReader(w, req.Body, app.Config.MaxRequestBodySize) +// RequestHandler returns a fasthttp.RequestHandler for the Application. +func (app *Application) RequestHandler() fasthttp.RequestHandler { + return func(ctx *fasthttp.RequestCtx) { + app.serveRequest(ctx) } +} - // Get context from pool - ctx := app.acquireContext(w, req) - defer app.releaseContext(ctx) +// serveRequest handles incoming HTTP requests by finding the matching route, +// creating a new Context, setting the route parameters, and executing the middleware chain. +func (app *Application) serveRequest(ctx *fasthttp.RequestCtx) { + c := app.acquireContext(ctx) + defer app.releaseContext(c) - // Find the matching route and set the handlers and paramsMap in the context - handlers, params := app.router.findRoute(ctx.Method, ctx.Path) + handlers, params := app.router.findRoute(c.Method, c.Path) - // This check is necessary because if no matching route is found and the handlers slice is left empty, - // the middleware chain will not be executed and the client will receive an empty response. - // By appending the 404 handler function to the handlers slice, - // we ensure that the middleware chain will always be executed, even if no matching route is found. if handlers == nil { handlers = append(app.middlewares, app.Config.NotFoundHandler) } - ctx.setHandlers(handlers) - ctx.setParams(params) - ctx.setApp(app) + c.setHandlers(handlers) + c.setParams(params) + c.setApp(app) - // Execute the middleware chain - ctx.Next() - ctx.flush() + c.Next() + c.flush() } // acquireContext gets a Context from the pool and initializes it. -func (app *Application) acquireContext(w http.ResponseWriter, req *http.Request) *Context { - r, err := newRequest(req) - if err != nil { - panic(err) - } +func (app *Application) acquireContext(ctx *fasthttp.RequestCtx) *Context { + c := app.contextPool.Get().(*Context) + c.ctx = ctx + c.req = newRequest(ctx) + c.res = newResponse(ctx) + c.Method = c.req.method() + c.Path = c.req.path() + c.App = app + c.data = contextData{} - ctx := app.contextPool.Get().(*Context) - ctx.Req = req - ctx.Res = w - ctx.req = r - ctx.res = newResponse(req, w) - ctx.Method = r.method - ctx.Path = r.path - ctx.App = app - ctx.data = contextData{} - - return ctx + return c } // releaseContext resets and returns the Context to the pool. -func (app *Application) releaseContext(ctx *Context) { - ctx.reset() - app.contextPool.Put(ctx) +func (app *Application) releaseContext(c *Context) { + c.reset() + app.contextPool.Put(c) } // Run starts the HTTP server and listens for incoming requests. @@ -278,84 +286,58 @@ func (app *Application) Run(address ...string) error { app.Logger.Info("Starting application on address `%s` 🚀🚀🚀", addr) app.mu.Lock() - app.server = &http.Server{ - Addr: addr, - Handler: app, - } - app.mu.Unlock() - - return app.server.ListenAndServe() -} - -// RunListener starts the HTTP server with an existing net.Listener. -func (app *Application) RunListener(listener net.Listener) error { - app.mu.Lock() - app.server = &http.Server{ - Handler: app, + app.server = &fasthttp.Server{ + Handler: app.RequestHandler(), + MaxRequestBodySize: int(app.Config.MaxRequestBodySize), } app.mu.Unlock() - return app.server.Serve(listener) -} - -// Shutdown gracefully shuts down the server without interrupting active connections. -func (app *Application) Shutdown(ctx context.Context) error { - app.mu.Lock() - server := app.server - app.mu.Unlock() - - if server == nil { - return nil - } - return server.Shutdown(ctx) + return app.server.ListenAndServe(addr) } // RunGraceful starts the HTTP server with graceful shutdown support. // It listens for SIGINT and SIGTERM signals to trigger graceful shutdown. -// The shutdownTimeout specifies the maximum duration to wait for active connections to finish. -func (app *Application) RunGraceful(shutdownTimeout time.Duration, address ...string) error { +// The shutdownTimeout specifies the maximum duration in seconds to wait for active connections to finish. +func (app *Application) RunGraceful(shutdownTimeout int, address ...string) error { addr := resolveAddress(address) app.Logger.Info("Starting application on address `%s` 🚀🚀🚀", addr) app.mu.Lock() - app.server = &http.Server{ - Addr: addr, - Handler: app, + app.server = &fasthttp.Server{ + Handler: app.RequestHandler(), + MaxRequestBodySize: int(app.Config.MaxRequestBodySize), } app.mu.Unlock() - // Channel to listen for shutdown signals quit := make(chan os.Signal, 1) signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) - // Channel to receive server errors serverErr := make(chan error, 1) - go func() { - serverErr <- app.server.ListenAndServe() + serverErr <- app.server.ListenAndServe(addr) }() - // Wait for either interrupt signal or server error select { case sig := <-quit: app.Logger.Info("Received signal %v, shutting down gracefully...", sig) if shutdownTimeout <= 0 { - shutdownTimeout = 5 * time.Second - } - ctx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) - defer cancel() - - if err := app.Shutdown(ctx); err != nil { - app.Logger.Error("Error during graceful shutdown: %v", err) - return err + shutdownTimeout = 5 } + app.server.Shutdown() app.Logger.Info("Server stopped gracefully") return nil case err := <-serverErr: - if err != nil && err != http.ErrServerClosed { + if err != nil { return err } return nil } } + +// Shutdown gracefully shuts down the server without interrupting active connections. +func (app *Application) Shutdown() { + if app.server != nil { + app.server.Shutdown() + } +} diff --git a/lightning_test.go b/lightning_test.go index adc0b4c..d3f4b76 100644 --- a/lightning_test.go +++ b/lightning_test.go @@ -1,13 +1,23 @@ package lightning import ( - "net" - "net/http" - "net/http/httptest" + "os" + "path/filepath" + "strings" "testing" + "text/template" "time" + + "github.com/valyala/fasthttp" ) +func createFasthttpRequest(method, path string) *fasthttp.RequestCtx { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod(method) + ctx.Request.Header.SetRequestURI(path) + return ctx +} + func TestNewApp(t *testing.T) { app := NewApp() if app == nil { @@ -18,29 +28,23 @@ func TestNewApp(t *testing.T) { func TestDefaultApp(t *testing.T) { app := DefaultApp() - // Assert that the Logger field is not nil if app.Logger == nil { t.Errorf("Expected Logger field to not be nil") } - // Assert that the middlewares field has the expected length - expectedMiddlewareLength := 2 - if len(app.middlewares) != expectedMiddlewareLength { - t.Errorf("Expected middlewares field to have length %d, but got %d", expectedMiddlewareLength, len(app.middlewares)) + if len(app.middlewares) != 2 { + t.Errorf("Expected 2 middleware functions, but got %d", len(app.middlewares)) } } func TestUse(t *testing.T) { app := NewApp() - // Define some middleware functions mw1 := func(c *Context) {} mw2 := func(c *Context) {} - // Add the middleware functions to the app app.Use(mw1, mw2) - // Check if the middleware functions were added correctly if len(app.middlewares) != 2 { t.Errorf("Expected 2 middleware functions, but got %d", len(app.middlewares)) } @@ -48,8 +52,8 @@ func TestUse(t *testing.T) { func TestAddRoute(t *testing.T) { app := NewApp() - app.AddRoute("GET", "/test", []HandlerFunc{}) - route, _ := app.router.findRoute("GET", "/test") + app.AddRoute(MethodGet, "/test", []HandlerFunc{}) + route, _ := app.router.findRoute(MethodGet, "/test") if route == nil { t.Errorf("Expected route to be added to router") } @@ -58,251 +62,1181 @@ func TestAddRoute(t *testing.T) { func TestGetRoute(t *testing.T) { app := NewApp() app.Get("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + c.Text(StatusOK, "Hello, World!") }) - req, err := http.NewRequest("GET", "/test", nil) - if err != nil { - t.Fatal(err) - } - - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) - - handler.ServeHTTP(rr, req) + ctx := createFasthttpRequest(MethodGet, "/test") + app.serveRequest(ctx) - if status := rr.Code; status != http.StatusOK { + if ctx.Response.StatusCode() != StatusOK { t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + ctx.Response.StatusCode(), StatusOK) } expected := "Hello, World!" - if rr.Body.String() != expected { + if string(ctx.Response.Body()) != expected { t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) + string(ctx.Response.Body()), expected) } } func TestPostRoute(t *testing.T) { app := NewApp() app.Post("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + c.Text(StatusOK, "Hello, World!") }) - req, err := http.NewRequest("POST", "/test", nil) - if err != nil { - t.Fatal(err) + ctx := createFasthttpRequest(MethodPost, "/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("handler returned wrong status code: got %v want %v", + ctx.Response.StatusCode(), StatusOK) } +} - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) +func TestPutRoute(t *testing.T) { + app := NewApp() + app.Put("/test", func(c *Context) { + c.Text(StatusOK, "Hello, World!") + }) - handler.ServeHTTP(rr, req) + ctx := createFasthttpRequest(MethodPut, "/test") + app.serveRequest(ctx) - if status := rr.Code; status != http.StatusOK { + if ctx.Response.StatusCode() != StatusOK { t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + ctx.Response.StatusCode(), StatusOK) } +} - expected := "Hello, World!" - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) +func TestDeleteRoute(t *testing.T) { + app := NewApp() + app.Delete("/test", func(c *Context) { + c.Text(StatusOK, "Hello, World!") + }) + + ctx := createFasthttpRequest(MethodDelete, "/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("handler returned wrong status code: got %v want %v", + ctx.Response.StatusCode(), StatusOK) } } -func TestPutRoute(t *testing.T) { +func TestHeadRoute(t *testing.T) { app := NewApp() - app.Put("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + app.Head("/test", func(c *Context) { + c.Text(StatusOK, "") }) - req, err := http.NewRequest("PUT", "/test", nil) - if err != nil { - t.Fatal(err) + ctx := createFasthttpRequest(MethodHead, "/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("handler returned wrong status code: got %v want %v", + ctx.Response.StatusCode(), StatusOK) } +} - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) +func TestPatchRoute(t *testing.T) { + app := NewApp() + app.Patch("/test", func(c *Context) { + c.Text(StatusOK, "Hello, World!") + }) - handler.ServeHTTP(rr, req) + ctx := createFasthttpRequest(MethodPatch, "/test") + app.serveRequest(ctx) - if status := rr.Code; status != http.StatusOK { + if ctx.Response.StatusCode() != StatusOK { t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + ctx.Response.StatusCode(), StatusOK) } +} - expected := "Hello, World!" - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) +func TestOptionsRoute(t *testing.T) { + app := NewApp() + app.Options("/test", func(c *Context) { + c.Text(StatusOK, "Hello, World!") + }) + + ctx := createFasthttpRequest(MethodOptions, "/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("handler returned wrong status code: got %v want %v", + ctx.Response.StatusCode(), StatusOK) } } -func TestDeleteRoute(t *testing.T) { +func TestNotFoundHandler(t *testing.T) { app := NewApp() - app.Delete("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + app.Get("/exists", func(c *Context) { + c.Text(StatusOK, "exists") }) - req, err := http.NewRequest("DELETE", "/test", nil) + ctx := createFasthttpRequest(MethodGet, "/notfound") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusNotFound { + t.Errorf("Expected status %d, got %d", StatusNotFound, ctx.Response.StatusCode()) + } +} + +func TestMiddlewareExecution(t *testing.T) { + app := NewApp() + order := []int{} + + app.Use(func(c *Context) { + order = append(order, 1) + c.Next() + order = append(order, 4) + }) + + app.Use(func(c *Context) { + order = append(order, 2) + c.Next() + order = append(order, 3) + }) + + app.Get("/test", func(c *Context) { + order = append(order, 5) + }) + + ctx := createFasthttpRequest(MethodGet, "/test") + app.serveRequest(ctx) + + expected := []int{1, 2, 5, 3, 4} + if !stringslicesEqual(order, expected) { + t.Errorf("Middleware execution order wrong: got %v want %v", order, expected) + } +} + +func stringslicesEqual(a, b []int) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +func TestStaticFiles(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "static_test") if err != nil { t.Fatal(err) } + defer os.RemoveAll(tmpDir) + + tmpFile := filepath.Join(tmpDir, "test.txt") + if err := os.WriteFile(tmpFile, []byte("hello"), 0644); err != nil { + t.Fatal(err) + } - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) + app := NewApp() + app.Static(tmpDir, "/static") - handler.ServeHTTP(rr, req) + ctx := createFasthttpRequest(MethodGet, "/static/test.txt") + app.serveRequest(ctx) - if status := rr.Code; status != http.StatusOK { - t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } +} - expected := "Hello, World!" - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) +func TestStaticFilesNotFound(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "static_test") + if err != nil { + t.Fatal(err) + } + defer os.RemoveAll(tmpDir) + + app := NewApp() + app.Static(tmpDir, "/static") + + ctx := createFasthttpRequest(MethodGet, "/static/nonexistent.txt") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusNotFound { + t.Errorf("Expected status %d, got %d", StatusNotFound, ctx.Response.StatusCode()) } } -func TestHeadRoute(t *testing.T) { +func TestGroupRoute(t *testing.T) { app := NewApp() - app.Head("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + group := app.Group("/api") + + group.Get("/test", func(c *Context) { + c.Text(StatusOK, "group route") + }) + + ctx := createFasthttpRequest(MethodGet, "/api/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) + } +} + +func TestGroupMiddleware(t *testing.T) { + app := NewApp() + order := []int{} + + group := app.Group("/api") + group.Use(func(c *Context) { + order = append(order, 1) + c.Next() + }) + + group.Get("/test", func(c *Context) { + order = append(order, 2) + }) + + ctx := createFasthttpRequest(MethodGet, "/api/test") + app.serveRequest(ctx) + + expected := []int{1, 2} + if !stringslicesEqual(order, expected) { + t.Errorf("Middleware execution order wrong: got %v want %v", order, expected) + } +} + +func TestNestedGroup(t *testing.T) { + app := NewApp() + group := app.Group("/api") + nested := group.Group("/v1") + + nested.Get("/test", func(c *Context) { + c.Text(StatusOK, "nested group route") }) - req, err := http.NewRequest("HEAD", "/test", nil) + ctx := createFasthttpRequest(MethodGet, "/api/v1/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) + } +} + +func TestConfigMerge(t *testing.T) { + config1 := &Config{ + AppName: "app1", + } + config2 := &Config{ + AppName: "app2", + } + + merged := config1.merge(config2) + if merged.AppName != "app2" { + t.Errorf("Expected AppName 'app2', got '%s'", merged.AppName) + } +} + +func TestConfigMergeJSONEncoder(t *testing.T) { + cfg := &Config{} + encoder := func(v interface{}) ([]byte, error) { return []byte("{}"), nil } + merged := cfg.merge(&Config{JSONEncoder: encoder}) + if merged.JSONEncoder == nil { + t.Error("Expected JSONEncoder to be set") + } +} + +func TestConfigMergeJSONDecoder(t *testing.T) { + cfg := &Config{} + decoder := func(data []byte, v interface{}) error { return nil } + merged := cfg.merge(&Config{JSONDecoder: decoder}) + if merged.JSONDecoder == nil { + t.Error("Expected JSONDecoder to be set") + } +} + +func TestConfigMergeNotFoundHandler(t *testing.T) { + cfg := &Config{} + handler := func(c *Context) {} + merged := cfg.merge(&Config{NotFoundHandler: handler}) + if merged.NotFoundHandler == nil { + t.Error("Expected NotFoundHandler to be set") + } +} + +func TestConfigMergeMaxRequestBodySize(t *testing.T) { + cfg := &Config{} + merged := cfg.merge(&Config{MaxRequestBodySize: 4096}) + if merged.MaxRequestBodySize != 4096 { + t.Errorf("Expected MaxRequestBodySize 4096, got %d", merged.MaxRequestBodySize) + } +} + +func TestConfigMergeMaxRequestBodySizeZero(t *testing.T) { + cfg := &Config{MaxRequestBodySize: 1024} + merged := cfg.merge(&Config{MaxRequestBodySize: 0}) + if merged.MaxRequestBodySize != 1024 { + t.Errorf("Expected MaxRequestBodySize 1024, got %d", merged.MaxRequestBodySize) + } +} + +func TestConfigMergeNilInMiddle(t *testing.T) { + cfg := &Config{AppName: "original"} + merged := cfg.merge(&Config{AppName: "first"}, nil, &Config{AppName: "second"}) + if merged.AppName != "second" { + t.Errorf("Expected AppName 'second', got '%s'", merged.AppName) + } +} + +func TestLoadHTMLGlob(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "templates_test") if err != nil { t.Fatal(err) } + defer os.RemoveAll(tmpDir) - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) + tmplPath := filepath.Join(tmpDir, "test.html") + if err := os.WriteFile(tmplPath, []byte("{{.Name}}"), 0644); err != nil { + t.Fatal(err) + } - handler.ServeHTTP(rr, req) + app := NewApp() + app.SetFuncMap(template.FuncMap{}) + app.LoadHTMLGlob(filepath.Join(tmpDir, "*.html")) - if status := rr.Code; status != http.StatusOK { - t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + if app.htmlTemplates == nil { + t.Error("Expected htmlTemplates to be set") } +} - expected := "Hello, World!" - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) +func TestJSONResponse(t *testing.T) { + app := NewApp() + app.Get("/test", func(c *Context) { + c.JSON(StatusOK, map[string]string{"message": "hello"}) + }) + + ctx := createFasthttpRequest(MethodGet, "/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) + } + + contentType := string(ctx.Response.Header.ContentType()) + if !strings.Contains(contentType, "application/json") { + t.Errorf("Expected Content-Type to contain 'application/json', got '%s'", contentType) } } -func TestPatchRoute(t *testing.T) { +func TestRedirect(t *testing.T) { app := NewApp() - app.Patch("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + app.Get("/test", func(c *Context) { + c.Redirect(StatusMovedPermanently, "/new") }) - req, err := http.NewRequest("PATCH", "/test", nil) - if err != nil { - t.Fatal(err) + ctx := createFasthttpRequest(MethodGet, "/test") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusMovedPermanently { + t.Errorf("Expected status %d, got %d", StatusMovedPermanently, ctx.Response.StatusCode()) } - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) + location := string(ctx.Response.Header.Peek("Location")) + if location == "" { + t.Error("Expected Location header to be set") + } +} - handler.ServeHTTP(rr, req) +func TestRequestHandler(t *testing.T) { + app := NewApp() + app.Get("/test", func(c *Context) { + c.Text(StatusOK, "handler") + }) - if status := rr.Code; status != http.StatusOK { - t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + handler := app.RequestHandler() + if handler == nil { + t.Error("RequestHandler returned nil") } - expected := "Hello, World!" - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) + ctx := createFasthttpRequest(MethodGet, "/test") + handler(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } } -func TestOptionsRoute(t *testing.T) { +func TestAcquireReleaseContext(t *testing.T) { app := NewApp() - app.Options("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") + + ctx := createFasthttpRequest(MethodGet, "/test") + c := app.acquireContext(ctx) + + if c == nil { + t.Fatal("acquireContext returned nil") + } + + if c.ctx != ctx { + t.Error("Context not set correctly") + } + + app.releaseContext(c) +} + +func TestConfigMergeAll(t *testing.T) { + cfg := &Config{} + + merged := cfg.merge(&Config{ + AppName: "test-app", + EnableDebug: true, + MaxRequestBodySize: 1024, + }) + + if merged.AppName != "test-app" { + t.Errorf("Expected AppName 'test-app', got '%s'", merged.AppName) + } + if !merged.EnableDebug { + t.Error("Expected EnableDebug to be true") + } + if merged.MaxRequestBodySize != 1024 { + t.Errorf("Expected MaxRequestBodySize 1024, got %d", merged.MaxRequestBodySize) + } +} + +func TestConfigMergeNil(t *testing.T) { + cfg := &Config{AppName: "original"} + + merged := cfg.merge(nil, nil) + + if merged.AppName != "original" { + t.Errorf("Expected AppName 'original', got '%s'", merged.AppName) + } +} + +func TestConfigMergePartial(t *testing.T) { + cfg := &Config{AppName: "original", EnableDebug: false} + + merged := cfg.merge(&Config{AppName: ""}, &Config{EnableDebug: true}) + + if merged.AppName != "original" { + t.Errorf("Expected AppName 'original', got '%s'", merged.AppName) + } + if !merged.EnableDebug { + t.Error("Expected EnableDebug to be true") + } +} + +func TestNewAppWithConfig(t *testing.T) { + app := NewApp(&Config{ + AppName: "custom-app", + EnableDebug: true, + MaxRequestBodySize: 2048, }) - req, err := http.NewRequest("OPTIONS", "/test", nil) + if app.Config.AppName != "custom-app" { + t.Errorf("Expected AppName 'custom-app', got '%s'", app.Config.AppName) + } + if !app.Config.EnableDebug { + t.Error("Expected EnableDebug to be true") + } + if app.Config.MaxRequestBodySize != 2048 { + t.Errorf("Expected MaxRequestBodySize 2048, got %d", app.Config.MaxRequestBodySize) + } +} + +func TestNewAppWithMultipleConfigs(t *testing.T) { + app := NewApp( + &Config{AppName: "first"}, + &Config{AppName: "second"}, + ) + + if app.Config.AppName != "second" { + t.Errorf("Expected AppName 'second', got '%s'", app.Config.AppName) + } +} + +func TestNewAppDebugRoute(t *testing.T) { + app := NewApp(&Config{EnableDebug: true}) + + ctx := createFasthttpRequest(MethodGet, "/__debug__/router_map") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) + } +} + +func TestDefaultJSONUnmarshal(t *testing.T) { + data := []byte(`{"key":"value"}`) + var result map[string]string + + err := defaultJSONUnmarshal(data, &result) + if err != nil { + t.Fatalf("defaultJSONUnmarshal returned error: %v", err) + } + if result["key"] != "value" { + t.Errorf("Expected 'value', got '%s'", result["key"]) + } +} + +func TestDefaultJSONMarshal(t *testing.T) { + data := map[string]string{"key": "value"} + + result, err := defaultJSONMarshal(data) if err != nil { + t.Fatalf("defaultJSONMarshal returned error: %v", err) + } + if string(result) != `{"key":"value"}` { + t.Errorf("Expected '{\"key\":\"value\"}', got '%s'", string(result)) + } +} + +func TestStaticAbsolute(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "static_abs") + if err != nil { + t.Fatal(err) + } + defer os.RemoveAll(tmpDir) + + if err := os.WriteFile(tmpDir+"/file.txt", []byte("content"), 0644); err != nil { t.Fatal(err) } - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) + app := NewApp() + app.Static(tmpDir, "/static") - handler.ServeHTTP(rr, req) + ctx := createFasthttpRequest(MethodGet, "/static/file.txt") + app.serveRequest(ctx) - if status := rr.Code; status != http.StatusOK { - t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } +} - expected := "Hello, World!" - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) +func TestStaticNotFound(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "static_notfound") + if err != nil { + t.Fatal(err) + } + defer os.RemoveAll(tmpDir) + + app := NewApp() + app.Static(tmpDir, "/static") + + ctx := createFasthttpRequest(MethodGet, "/static/nonexistent.txt") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusNotFound { + t.Errorf("Expected status %d, got %d", StatusNotFound, ctx.Response.StatusCode()) } } -func TestServeHTTP(t *testing.T) { +func TestRouterSearchNotFound(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/api/users", []HandlerFunc{}) + + route, params := router.findRoute(MethodGet, "/api/posts") + if route != nil { + t.Error("Expected nil route for non-matching path") + } + if params != nil { + t.Error("Expected nil params for non-matching path") + } +} + +func TestRouterSearchWildcard(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/api/*filepath", []HandlerFunc{func(c *Context) {}}) + + route, params := router.findRoute(MethodGet, "/api/users/123/posts") + if route == nil { + t.Fatal("Expected route for wildcard path") + } + if params["filepath"] != "users/123/posts" { + t.Errorf("Expected wildcard param 'users/123/posts', got '%s'", params["filepath"]) + } +} + +func TestRouterSearchParam(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/users/:id", []HandlerFunc{func(c *Context) {}}) + + route, params := router.findRoute(MethodGet, "/users/42") + if route == nil { + t.Fatal("Expected route for param path") + } + if params["id"] != "42" { + t.Errorf("Expected param id=42, got '%s'", params["id"]) + } +} + +func TestRouterFindRouteNotFound(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/test", []HandlerFunc{}) + + route, params := router.findRoute(MethodPost, "/test") + if route != nil { + t.Error("Expected nil route for wrong method") + } + if params != nil { + t.Error("Expected nil params for wrong method") + } +} + +func TestRouterFindRouteEmpty(t *testing.T) { + router := newRouter() + + route, params := router.findRoute(MethodGet, "/nonexistent") + if route != nil { + t.Error("Expected nil route for empty router") + } + if params != nil { + t.Error("Expected nil params for empty router") + } +} + +func TestRouterSearchMultipleParams(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/users/:userId/posts/:postId", []HandlerFunc{func(c *Context) {}}) + + route, params := router.findRoute(MethodGet, "/users/1/posts/2") + if route == nil { + t.Fatal("Expected route for multi-param path") + } + if params["userId"] != "1" { + t.Errorf("Expected userId=1, got '%s'", params["userId"]) + } + if params["postId"] != "2" { + t.Errorf("Expected postId=2, got '%s'", params["postId"]) + } +} + +func TestLogger(t *testing.T) { app := NewApp() - app.Get("/test", func(c *Context) { - c.Text(http.StatusOK, "Hello, World!") - }) + if app.Logger == nil { + t.Error("Expected Logger to be non-nil") + } +} + +func TestRouterSearchEmptyPath(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/", []HandlerFunc{func(c *Context) {}}) + + route, _ := router.findRoute(MethodGet, "/") + if route == nil { + t.Error("Expected route for root path") + } +} + +func TestRouterSearchMethodNotAllowed(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/test", []HandlerFunc{func(c *Context) {}}) + + route, _ := router.findRoute(MethodPost, "/test") + if route != nil { + t.Error("Expected nil route for method not allowed") + } +} + +func TestContext_NextOutOfBounds(t *testing.T) { + c, _ := createTestContext("GET", "/test", nil) + c.handlers = []HandlerFunc{} + c.index = 0 + + c.Next() + + if c.index != 1 { + t.Errorf("Expected index 1, got %d", c.index) + } +} + +func TestRouterSearchParamWithMultipleLevels(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/api/:version/users/:id", []HandlerFunc{func(c *Context) {}}) + + route, params := router.findRoute(MethodGet, "/api/v1/users/42") + if route == nil { + t.Fatal("Expected route for multi-level param path") + } + if params["version"] != "v1" { + t.Errorf("Expected version=v1, got '%s'", params["version"]) + } + if params["id"] != "42" { + t.Errorf("Expected id=42, got '%s'", params["id"]) + } +} + +func TestRouterAddRouteMultiple(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/a", []HandlerFunc{func(c *Context) {}}) + router.addRoute(MethodPost, "/a", []HandlerFunc{func(c *Context) {}}) + router.addRoute(MethodPut, "/a", []HandlerFunc{func(c *Context) {}}) + + if len(router.Roots) != 3 { + t.Errorf("Expected 3 method roots, got %d", len(router.Roots)) + } +} + +func TestRouterSearchNotFoundMethod(t *testing.T) { + router := newRouter() + router.addRoute(MethodGet, "/test", []HandlerFunc{func(c *Context) {}}) + + route, params := router.findRoute(MethodDelete, "/test") + if route != nil { + t.Error("Expected nil route for non-existent method") + } + if params != nil { + t.Error("Expected nil params for non-existent method") + } +} + +func TestParsePatternWithWildcard(t *testing.T) { + parts := parsePattern("/files/*filepath") + if len(parts) != 2 { + t.Errorf("Expected 2 parts, got %d", len(parts)) + } + if parts[1] != "*filepath" { + t.Errorf("Expected '*filepath', got '%s'", parts[1]) + } +} + +func TestResolveAddressWithPortEnv(t *testing.T) { + os.Setenv("PORT", "8080") + defer os.Unsetenv("PORT") + + addr := resolveAddress([]string{}) + if addr != ":8080" { + t.Errorf("Expected ':8080', got '%s'", addr) + } +} + +func TestResolveAddressSingleParam(t *testing.T) { + os.Unsetenv("PORT") + addr := resolveAddress([]string{":9090"}) + if addr != ":9090" { + t.Errorf("Expected ':9090', got '%s'", addr) + } +} + +func TestResolveAddressMultipleParam(t *testing.T) { + os.Unsetenv("PORT") + defer func() { + recover() + }() + resolveAddress([]string{":8080", ":9090"}) +} + +func TestDefaultNotFoundHandler(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/notfound") + + c := &Context{ + ctx: ctx, + index: -1, + res: newResponse(ctx), + } + defaultNotFound(c) + c.flush() + + if ctx.Response.StatusCode() != StatusNotFound { + t.Errorf("Expected status %d, got %d", StatusNotFound, ctx.Response.StatusCode()) + } +} + +func TestDefaultInternalServerErrorHandler(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/error") + + c := &Context{ + ctx: ctx, + index: -1, + res: newResponse(ctx), + } + defaultInternalServerError(c) + c.flush() + + if ctx.Response.StatusCode() != StatusInternalServerError { + t.Errorf("Expected status %d, got %d", StatusInternalServerError, ctx.Response.StatusCode()) + } +} + +func TestNewRouter(t *testing.T) { + router := newRouter() + if router.Roots == nil { + t.Error("Expected Roots to be initialized") + } +} + +func TestNodeMatchChild(t *testing.T) { + n := &node{ + Children: map[string]*node{ + "foo": {Part: "foo"}, + }, + } + + child := n.matchChild("foo") + if child == nil { + t.Error("Expected child node") + } + + child = n.matchChild("bar") + if child != nil { + t.Error("Expected nil for non-existent child") + } +} + +func TestNodeInsert(t *testing.T) { + root := &node{} + root.insert("/a/b/c", []string{"a", "b", "c"}, 0, []HandlerFunc{func(c *Context) {}}) + + if root.Children == nil { + t.Fatal("Expected children to be initialized") + } + + n := root.matchChild("a") + if n == nil { + t.Fatal("Expected node 'a'") + } + if n.Part != "a" { + t.Errorf("Expected part 'a', got '%s'", n.Part) + } +} + +func TestNodeInsertWildParam(t *testing.T) { + root := &node{} + root.insert("/users/:id", []string{"users", ":id"}, 0, []HandlerFunc{func(c *Context) {}}) + + users := root.matchChild("users") + if users == nil { + t.Fatal("Expected 'users' node") + } + + idNode := users.matchChild(":id") + if idNode == nil { + t.Fatal("Expected ':id' node") + } + if !idNode.IsWild { + t.Error("Expected :id to be wild") + } +} + +func TestNodeSearch(t *testing.T) { + root := &node{} + root.insert("/a/b", []string{"a", "b"}, 0, []HandlerFunc{func(c *Context) {}}) + + n := root.search([]string{"a", "b"}, 0) + if n == nil { + t.Error("Expected to find node") + } + if n.Pattern != "/a/b" { + t.Errorf("Expected pattern '/a/b', got '%s'", n.Pattern) + } +} - req, err := http.NewRequest("GET", "/test", nil) +func TestNodeSearchNotFound(t *testing.T) { + root := &node{} + root.insert("/a/b", []string{"a", "b"}, 0, []HandlerFunc{func(c *Context) {}}) + + n := root.search([]string{"a", "c"}, 0) + if n != nil { + t.Error("Expected nil for non-existent path") + } +} + +func TestNodeSearchWild(t *testing.T) { + root := &node{} + root.insert("/api/*", []string{"api", "*"}, 0, []HandlerFunc{func(c *Context) {}}) + + n := root.search([]string{"api", "users", "123"}, 0) + if n == nil { + t.Error("Expected to find wildcard node") + } +} + +func TestNodeSearchParam(t *testing.T) { + root := &node{} + root.insert("/users/:id", []string{"users", ":id"}, 0, []HandlerFunc{func(c *Context) {}}) + + n := root.search([]string{"users", "42"}, 0) + if n == nil { + t.Error("Expected to find param node") + } +} + +func TestNodeSearchEmptyParts(t *testing.T) { + root := &node{Pattern: "/"} + n := root.search([]string{}, 0) + if n == nil { + t.Error("Expected to find root node") + } +} + +func TestNodeSearchWildReturnNil(t *testing.T) { + root := &node{} + root.insert("/api/*", []string{"api", "*"}, 0, []HandlerFunc{func(c *Context) {}}) + + n := root.search([]string{"api"}, 0) + if n != nil { + t.Error("Expected nil for incomplete path") + } +} + +func TestResponseFlushWithRedirect(t *testing.T) { + resp, ctx := createResponse() + + resp.redirect(StatusFound, "/new") + resp.flush() + + if ctx.Response.StatusCode() != StatusFound { + t.Errorf("Expected status %d, got %d", StatusFound, ctx.Response.StatusCode()) + } + location := string(ctx.Response.Header.Peek("Location")) + if location == "" { + t.Error("Expected Location header to be set") + } +} + +func TestResponseFlushWithFile(t *testing.T) { + tmpFile, err := os.CreateTemp("", "response_test*.txt") if err != nil { t.Fatal(err) } + defer os.Remove(tmpFile.Name()) - rr := httptest.NewRecorder() - handler := http.HandlerFunc(app.ServeHTTP) + if _, err := tmpFile.Write([]byte("file content")); err != nil { + t.Fatal(err) + } + tmpFile.Close() - handler.ServeHTTP(rr, req) + resp, ctx := createResponse() + resp.file(tmpFile.Name()) + resp.flush() - if status := rr.Code; status != http.StatusOK { - t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusOK) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } +} - expected := "Hello, World!" - if rr.Body.String() != expected { - t.Errorf("handler returned unexpected body: got %v want %v", - rr.Body.String(), expected) +func TestResponseFlushWithCookie(t *testing.T) { + resp, ctx := createResponse() + + resp.cookies.set("session", "abc") + resp.flush() + + cookie := string(ctx.Response.Header.Peek("Set-Cookie")) + if !strings.Contains(cookie, "session=abc") { + t.Errorf("Expected cookie to contain 'session=abc', got %s", cookie) + } +} + +func TestResponseFlushWithBody(t *testing.T) { + resp, ctx := createResponse() + + resp.setBody([]byte("hello")) + resp.flush() + + if string(ctx.Response.Body()) != "hello" { + t.Errorf("Expected body 'hello', got '%s'", string(ctx.Response.Body())) + } +} + +func TestResponseFlushWithCookieAndBody(t *testing.T) { + resp, ctx := createResponse() + + resp.cookies.set("test", "value") + resp.setBody([]byte("body")) + resp.flush() + + if string(ctx.Response.Body()) != "body" { + t.Errorf("Expected body 'body', got '%s'", string(ctx.Response.Body())) + } +} + +func TestResponseFlushWithCookieAndRedirect(t *testing.T) { + resp, ctx := createResponse() + + resp.cookies.set("session", "abc") + resp.redirect(StatusFound, "/new") + resp.flush() + + if ctx.Response.StatusCode() != StatusFound { + t.Errorf("Expected status %d, got %d", StatusFound, ctx.Response.StatusCode()) + } +} + +func TestResponseFileError(t *testing.T) { + resp, _ := createResponse() + + err := resp.file("/nonexistent/file.txt") + if err == nil { + t.Error("Expected error for non-existent file") + } +} + +func TestResponseSetStatus(t *testing.T) { + resp, _ := createResponse() + + resp.setStatus(StatusCreated) + if resp.statusCode != StatusCreated { + t.Errorf("Expected status %d, got %d", StatusCreated, resp.statusCode) + } +} + +func TestResponseSetBody(t *testing.T) { + resp, _ := createResponse() + + resp.setBody([]byte("test")) + if string(resp.body) != "test" { + t.Errorf("Expected body 'test', got '%s'", string(resp.body)) } } -func TestRun(t *testing.T) { +func TestResponseRedirect(t *testing.T) { + resp, _ := createResponse() + + resp.redirect(StatusMovedPermanently, "/new") + if resp.redirectTo != "/new" { + t.Errorf("Expected redirectTo '/new', got '%s'", resp.redirectTo) + } + if resp.statusCode != StatusMovedPermanently { + t.Errorf("Expected status %d, got %d", StatusMovedPermanently, resp.statusCode) + } +} + +func TestResponseAddHeader(t *testing.T) { + resp, ctx := createResponse() + + resp.addHeader("X-Custom", "value") + hdr := string(ctx.Response.Header.Peek("X-Custom")) + if hdr != "value" { + t.Errorf("Expected header 'value', got '%s'", hdr) + } +} + +func TestResponseSetHeader(t *testing.T) { + resp, ctx := createResponse() + + resp.setHeader("Content-Type", "text/plain") + hdr := string(ctx.Response.Header.Peek("Content-Type")) + if hdr != "text/plain" { + t.Errorf("Expected header 'text/plain', got '%s'", hdr) + } +} + +func TestResponseDelHeader(t *testing.T) { + resp, ctx := createResponse() + + resp.setHeader("X-Custom", "value") + resp.delHeader("X-Custom") + hdr := string(ctx.Response.Header.Peek("X-Custom")) + if hdr != "" { + t.Errorf("Expected header to be deleted, got '%s'", hdr) + } +} + +func TestResponseNew(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + resp := newResponse(ctx) + + if resp.ctx != ctx { + t.Error("Expected ctx to be set") + } + if resp.statusCode != StatusNotFound { + t.Errorf("Expected default status %d, got %d", StatusNotFound, resp.statusCode) + } + if resp.cookies == nil { + t.Error("Expected cookies to be initialized") + } +} + +func TestResponseFlushEmpty(t *testing.T) { + resp, ctx := createResponse() + + resp.flush() + + if ctx.Response.StatusCode() != StatusNotFound { + t.Errorf("Expected default status %d, got %d", StatusNotFound, ctx.Response.StatusCode()) + } +} + +func TestStaticRelativePath(t *testing.T) { app := NewApp() + app.Static("assets", "/static") - // Use dynamic port allocation to avoid port conflicts - listener, err := net.Listen("tcp", "localhost:0") - if err != nil { - t.Fatalf("Failed to create listener: %v", err) + ctx := createFasthttpRequest(MethodGet, "/static/test.txt") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusNotFound { + t.Errorf("Expected status %d, got %d", StatusNotFound, ctx.Response.StatusCode()) } - addr := listener.Addr().String() +} - go app.RunListener(listener) +func TestShutdownWithNilServer(t *testing.T) { + app := NewApp() - // Wait for the server to start - time.Sleep(100 * time.Millisecond) + app.Shutdown() +} - // Send a GET request to the server - resp, err := http.Get("http://" + addr) - if err != nil { - t.Fatalf("Error sending request: %v", err) +func TestShutdownWithServer(t *testing.T) { + app := NewApp() + app.Get("/test", func(c *Context) { + c.Text(StatusOK, "ok") + }) + + go func() { + _ = app.Run(":0") + }() + + time.Sleep(50 * time.Millisecond) + + app.Shutdown() +} + +func TestRunGracefulWithZeroTimeout(t *testing.T) { + app := NewApp() + app.Get("/test", func(c *Context) { + c.Text(StatusOK, "ok") + }) + + go func() { + _ = app.RunGraceful(0, ":0") + }() + + time.Sleep(50 * time.Millisecond) + + app.Shutdown() +} + +func TestContextResetClearsData(t *testing.T) { + app := NewApp() + ctx := createFasthttpRequest(MethodGet, "/test") + c := app.acquireContext(ctx) + + c.SetData("key", "value") + if c.GetData("key") != "value" { + t.Error("Expected data to be set") } - defer resp.Body.Close() - // Assert that the response status code is 404 - if resp.StatusCode != http.StatusNotFound { - t.Errorf("Expected status code %d, but got %d", http.StatusNotFound, resp.StatusCode) + app.releaseContext(c) + + c2 := app.acquireContext(createFasthttpRequest(MethodGet, "/test2")) + if c2.GetData("key") != nil { + t.Error("Expected data to be cleared after reset") + } + app.releaseContext(c2) +} + +func TestServeRequestWithNotFoundRoute(t *testing.T) { + app := NewApp() + app.Use(func(c *Context) { + c.Next() + }) + + ctx := createFasthttpRequest(MethodGet, "/nonexistent") + app.serveRequest(ctx) + + if ctx.Response.StatusCode() != StatusNotFound { + t.Errorf("Expected status %d, got %d", StatusNotFound, ctx.Response.StatusCode()) } } diff --git a/recovery_test.go b/recovery_test.go index 16aae2e..16e5d42 100644 --- a/recovery_test.go +++ b/recovery_test.go @@ -1,34 +1,35 @@ package lightning import ( - "net/http" - "net/http/httptest" "testing" + + "github.com/valyala/fasthttp" ) func TestRecovery(t *testing.T) { app := NewApp() app.Use(Recovery(func(ctx *Context) { - ctx.Text(500, "Internal Server Error") + ctx.Text(StatusInternalServerError, "Internal Server Error") })) app.Get("/", func(ctx *Context) { panic("test panic") }) - req := httptest.NewRequest("GET", "/", nil) - res := httptest.NewRecorder() + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod(MethodGet) + ctx.Request.Header.SetRequestURI("/") - app.ServeHTTP(res, req) + app.serveRequest(ctx) - if status := res.Code; status != http.StatusInternalServerError { + if ctx.Response.StatusCode() != StatusInternalServerError { t.Errorf("handler returned wrong status code: got %v want %v", - status, http.StatusInternalServerError) + ctx.Response.StatusCode(), StatusInternalServerError) } expected := "Internal Server Error" - if body := res.Body.String(); body != expected { + if body := string(ctx.Response.Body()); body != expected { t.Errorf("handler returned unexpected body: got %v want %v", body, expected) } diff --git a/request.go b/request.go index bfc5e7e..64c67d6 100644 --- a/request.go +++ b/request.go @@ -1,115 +1,117 @@ package lightning import ( - "io" - "net/http" "strings" + + "github.com/valyala/fasthttp" ) type request struct { - originReq *http.Request - paramsMap map[string]string - method string - path string - rawBody []byte -} - -// newRequest creates a new request object from an http.Request object -func newRequest(req *http.Request) (*request, error) { - var rawBody []byte - var err error - if req.Body != nil { - rawBody, err = io.ReadAll(req.Body) - if err != nil { - return nil, err - } - } + ctx *fasthttp.RequestCtx + pathParams map[string]string +} - request := &request{ - originReq: req, - paramsMap: map[string]string{}, - method: req.Method, - path: req.URL.Path, - rawBody: rawBody, +func newRequest(ctx *fasthttp.RequestCtx) *request { + return &request{ + ctx: ctx, + pathParams: make(map[string]string), } - - return request, nil } -// setParams sets the parameters for the request object func (r *request) setParams(params map[string]string) { - r.paramsMap = params + r.pathParams = params } -// param returns the parameter value for a given key. func (r *request) param(key string) string { - return r.paramsMap[key] + return r.pathParams[key] } -// params returns the entire parameter map for the context. func (r *request) params() map[string]string { - return r.paramsMap + return r.pathParams } -// query returns the value of a given query parameter. func (r *request) query(key string) string { - return r.originReq.URL.Query().Get(key) + return string(r.ctx.QueryArgs().Peek(key)) } -// queries returns the entire query parameter map for the context. func (r *request) queries() map[string][]string { - return r.originReq.URL.Query() + args := r.ctx.QueryArgs() + queries := make(map[string][]string) + args.VisitAll(func(key, value []byte) { + queries[string(key)] = append(queries[string(key)], string(value)) + }) + return queries } -// header returns the value of a given header. func (r *request) header(key string) string { - return r.originReq.Header.Get(key) + return string(r.ctx.Request.Header.Peek(key)) } -// headers returns the entire header map for the request. -func (r *request) headers() http.Header { - return r.originReq.Header +func (r *request) headers() map[string]string { + headers := make(map[string]string) + r.ctx.Request.Header.VisitAll(func(key, value []byte) { + headers[string(key)] = string(value) + }) + return headers } -// cookie returns the cookie with the given name. -func (r *request) cookie(name string) *http.Cookie { - cookie, err := r.originReq.Cookie(name) - if err != nil { - return nil +func (r *request) cookie(name string) *fasthttp.Cookie { + var cookie fasthttp.Cookie + cookie.ParseBytes(r.ctx.Request.Header.Cookie(name)) + if len(cookie.Key()) > 0 { + return &cookie } - return cookie + return nil } -// cookiesMap returns all cookies from the request. -func (r *request) cookies() []*http.Cookie { - return r.originReq.Cookies() +func (r *request) cookies() []*fasthttp.Cookie { + var cookies []*fasthttp.Cookie + r.ctx.Request.Header.VisitAll(func(key, value []byte) { + if string(key) == "Cookie" { + var c fasthttp.Cookie + c.ParseBytes(value) + cookies = append(cookies, &c) + } + }) + return cookies } -// userAgent returns the user agent header value of the request. func (r *request) userAgent() string { - return r.header("user-agent") + return string(r.ctx.Request.Header.UserAgent()) } -// referer returns the referer header value of the request. func (r *request) referer() string { - return r.header("referer") + return string(r.ctx.Request.Header.Referer()) } -// remoteAddr returns the remote address of the request. func (r *request) remoteAddr() string { - ip := r.header("x-real-ip") + ip := r.header("X-Real-IP") if ip == "" { - ip = r.header("x-forwarded-for") + ip = r.header("X-Forwarded-For") if ip != "" { - // X-Forwarded-For may contain multiple IPs: client, proxy1, proxy2 - // Take the first one (client IP) if idx := strings.Index(ip, ","); idx != -1 { ip = strings.TrimSpace(ip[:idx]) } } } if ip == "" { - ip = r.originReq.RemoteAddr + ip = r.ctx.RemoteAddr().String() } return ip } + +func (r *request) body() []byte { + return r.ctx.Request.Body() +} + +func (r *request) method() string { + return string(r.ctx.Request.Header.Method()) +} + +func (r *request) path() string { + return string(r.ctx.Path()) +} + +func (r *request) uri() string { + return string(r.ctx.URI().FullURI()) +} diff --git a/request_test.go b/request_test.go index e641b61..2ee519e 100644 --- a/request_test.go +++ b/request_test.go @@ -1,431 +1,264 @@ package lightning import ( - "errors" - "net/http" - "net/http/httptest" - "reflect" "testing" -) -type errorReader struct{} + "github.com/valyala/fasthttp" +) -func (r *errorReader) Read(p []byte) (n int, err error) { - return 0, errors.New("read error") +func createFasthttpCtx(method, path string) (*fasthttp.RequestCtx, *request) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod(method) + ctx.Request.Header.SetRequestURI(path) + r := newRequest(ctx) + return ctx, r } -func TestNewRequestWithBodyReadError(t *testing.T) { - req := httptest.NewRequest("GET", "/path", &errorReader{}) - _, err := newRequest(req) - if err == nil { - t.Error("Expected error, but got nil") - } -} +func TestRequest_Cookie(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.SetCookie("cookie1", "value1") -func TestNewRequest(t *testing.T) { - type args struct { - req *http.Request - params map[string]string - } - req, _ := http.NewRequest("GET", "http://example.com", nil) - params := map[string]string{"key1": "value1", "key2": "value2"} - tests := []struct { - name string - args args - want *request - }{ - { - name: "Test_NewRequest", - args: args{ - req: req, - params: params, - }, - want: &request{ - originReq: req, - paramsMap: params, - method: req.Method, - path: req.URL.Path, - }, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got, _ := newRequest(tt.args.req) - got.setParams(tt.args.params) - if !reflect.DeepEqual(got, tt.want) { - t.Errorf("newRequest() = %v, want %v", got, tt.want) - } - }) - } -} + r := newRequest(ctx) -func TestRequest_Cookie(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string - } - type args struct { - name string + cookie := r.cookie("cookie1") + if cookie != nil && string(cookie.Key()) != "cookie1" { + t.Errorf("Expected cookie key 'cookie1', got '%s'", string(cookie.Key())) } - req, _ := http.NewRequest("GET", "http://example.com", nil) - req.AddCookie(&http.Cookie{Name: "cookie1", Value: "value1"}) - tests := []struct { - name string - fields fields - args args - want *http.Cookie - }{ - { - name: "Test_Request_Cookie", - fields: fields{ - req: req, - }, - args: args{ - name: "cookie1", - }, - want: &http.Cookie{Name: "cookie1", Value: "value1"}, - }, - { - name: "Test_Request_Cookie_Invalid", - fields: fields{ - req: req, - }, - args: args{ - name: "cookie_invalid", - }, - want: nil, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.cookie(tt.args.name); !reflect.DeepEqual(got, tt.want) { - t.Errorf("cookie() = %v, want %v", got, tt.want) - } - }) + + cookie = r.cookie("nonexistent") + if cookie != nil { + t.Errorf("Expected nil for nonexistent cookie, got %v", cookie) } } func TestRequest_Cookies(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string - } - req, _ := http.NewRequest("GET", "http://example.com", nil) - cookie1 := &http.Cookie{Name: "cookie1", Value: "value1"} - cookie2 := &http.Cookie{Name: "cookie2", Value: "value2"} - req.AddCookie(cookie1) - req.AddCookie(cookie2) - tests := []struct { - name string - fields fields - want []*http.Cookie - }{ - { - name: "Test_Request_Cookies", - fields: fields{ - req: req, - }, - want: []*http.Cookie{cookie1, cookie2}, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.cookies(); !reflect.DeepEqual(got, tt.want) { - t.Errorf("cookies() = %v, want %v", got, tt.want) - } - }) + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.SetCookie("cookie1", "value1") + ctx.Request.Header.SetCookie("cookie2", "value2") + + r := newRequest(ctx) + cookies := r.cookies() + + if len(cookies) == 0 { + t.Error("Expected cookies, got empty") } } func TestRequest_Header(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string - } - type args struct { - key string - } - req, _ := http.NewRequest("GET", "http://example.com", nil) - req.Header.Set("header1", "value1") - tests := []struct { - name string - fields fields - args args - want string - }{ - { - name: "Test_Request_Header", - fields: fields{ - req: req, - }, - args: args{ - key: "header1", - }, - want: "value1", - }, - { - name: "Test_Request_Header_Invalid", - fields: fields{ - req: req, - }, - args: args{ - key: "header_invalid", - }, - want: "", - }, + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.Set("header1", "value1") + + r := newRequest(ctx) + + if got := r.header("header1"); got != "value1" { + t.Errorf("header() = %v, want %v", got, "value1") } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.header(tt.args.key); got != tt.want { - t.Errorf("header() = %v, want %v", got, tt.want) - } - }) + + if got := r.header("nonexistent"); got != "" { + t.Errorf("header() = %v, want empty", got) } } func TestRequest_Headers(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string - } - req, _ := http.NewRequest("GET", "http://example.com", nil) - req.Header.Set("header1", "value1") - req.Header.Set("header2", "value2") - want := http.Header{"Header1": []string{"value1"}, "Header2": []string{"value2"}} - tests := []struct { - name string - fields fields - want http.Header - }{ - { - name: "Test_Request_Headers", - fields: fields{ - req: req, - }, - want: want, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.headers(); !reflect.DeepEqual(got, tt.want) { - t.Errorf("headers() = %v, want %v", got, tt.want) - } - }) + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.Set("Content-Type", "application/json") + + r := newRequest(ctx) + headers := r.headers() + + if len(headers) == 0 { + t.Error("Expected headers to be non-empty") } } func TestRequest_Param(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + + r := newRequest(ctx) + r.setParams(map[string]string{"param1": "value1"}) + + if got := r.param("param1"); got != "value1" { + t.Errorf("param() = %v, want %v", got, "value1") } - type args struct { - key string + + if got := r.param("nonexistent"); got != "" { + t.Errorf("param() = %v, want empty", got) + } +} + +func TestRequest_Params(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + + r := newRequest(ctx) + params := map[string]string{"param1": "value1", "param2": "value2"} + r.setParams(params) + + got := r.params() + for k, v := range params { + if got[k] != v { + t.Errorf("params()[%s] = %v, want %v", k, got[k], v) + } } - req, _ := http.NewRequest("GET", "http://example.com", nil) - params := map[string]string{"param1": "value1"} - tests := []struct { - name string - fields fields - args args - want string - }{ - { - name: "Test_Request_Param", - fields: fields{ - req: req, - params: params, - }, - args: args{ - key: "param1", - }, - want: "value1", - }, - { - name: "Test_Request_Param_Invalid", - fields: fields{ - req: req, - params: params, - }, - args: args{ - key: "param_invalid", - }, - want: "", - }, +} + +func TestRequest_Queries(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test?param1=value1¶m2=value2") + + r := newRequest(ctx) + queries := r.queries() + + if queries["param1"][0] != "value1" { + t.Errorf("Expected param1=value1, got %v", queries["param1"]) } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.param(tt.args.key); got != tt.want { - t.Errorf("param() = %v, want %v", got, tt.want) - } - }) + if queries["param2"][0] != "value2" { + t.Errorf("Expected param2=value2, got %v", queries["param2"]) } } -func TestRequest_Params(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string +func TestRequest_Query(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test?param1=value1") + + r := newRequest(ctx) + + if got := r.query("param1"); got != "value1" { + t.Errorf("query() = %v, want %v", got, "value1") } - req, _ := http.NewRequest("GET", "", nil) - params := map[string]string{"param1": "value1", "param2": "value2"} - tests := []struct { - name string - fields fields - want map[string]string - }{ - { - name: "Test_Request_Params", - fields: fields{ - req: req, - params: params, - }, - want: params, - }, + + if got := r.query("nonexistent"); got != "" { + t.Errorf("query() = %v, want empty", got) } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.params(); !reflect.DeepEqual(got, tt.want) { - t.Errorf("params() = %v, want %v", got, tt.want) - } - }) +} + +func TestRequest_userAgent(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.SetUserAgent("test-user-agent") + + r := newRequest(ctx) + + if got := r.userAgent(); got != "test-user-agent" { + t.Errorf("userAgent() = %v, want %v", got, "test-user-agent") } } -func TestRequest_Queries(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string +func TestRequest_referer(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.SetReferer("https://example.com") + + r := newRequest(ctx) + + if got := r.referer(); got != "https://example.com" { + t.Errorf("referer() = %v, want %v", got, "https://example.com") } - req, _ := http.NewRequest("GET", "http://example.com?param1=value1¶m2=value2", nil) - want := map[string][]string{"param1": []string{"value1"}, "param2": []string{"value2"}} - tests := []struct { - name string - fields fields - want map[string][]string - }{ - { - name: "Test_Request_Queries", - fields: fields{ - req: req, - }, - want: want, - }, +} + +func TestRequest_remoteAddr(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + + r := newRequest(ctx) + addr := r.remoteAddr() + + if addr == "" { + t.Error("Expected non-empty remote address") } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.queries(); !reflect.DeepEqual(got, tt.want) { - t.Errorf("queries() = %v, want %v", got, tt.want) - } - }) +} + +func TestRequest_remoteAddrWithXRealIP(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.Set("X-Real-IP", "1.2.3.4") + + r := newRequest(ctx) + addr := r.remoteAddr() + + if addr != "1.2.3.4" { + t.Errorf("Expected X-Real-IP, got %s", addr) } } -func TestRequest_Query(t *testing.T) { - type fields struct { - req *http.Request - params map[string]string - method string - path string +func TestRequest_remoteAddrWithXForwardedFor(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.Header.Set("X-Forwarded-For", "1.2.3.4, 5.6.7.8") + + r := newRequest(ctx) + addr := r.remoteAddr() + + if addr != "1.2.3.4" { + t.Errorf("Expected first IP from X-Forwarded-For, got %s", addr) + } +} + +func TestRequest_body(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("POST") + ctx.Request.Header.SetRequestURI("/test") + ctx.Request.SetBody([]byte("test body")) + + r := newRequest(ctx) + body := r.body() + + if string(body) != "test body" { + t.Errorf("body() = %v, want %v", string(body), "test body") } - type args struct { - key string +} + +func TestRequest_method(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("POST") + ctx.Request.Header.SetRequestURI("/test") + + r := newRequest(ctx) + + if got := r.method(); got != "POST" { + t.Errorf("method() = %v, want %v", got, "POST") } - req, _ := http.NewRequest("GET", "http://example.com?param1=value1", nil) - tests := []struct { - name string - fields fields - args args - want string - }{ - { - name: "Test_Request_Query", - fields: fields{ - req: req, - }, - args: args{ - key: "param1", - }, - want: "value1", - }, - { - name: "Test_Request_Query_Invalid", - fields: fields{ - req: req, - }, - args: args{ - key: "param_invalid", - }, - want: "", - }, +} + +func TestRequest_path(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test/path") + + r := newRequest(ctx) + + if got := r.path(); got != "/test/path" { + t.Errorf("path() = %v, want %v", got, "/test/path") } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - r := &request{ - originReq: tt.fields.req, - paramsMap: tt.fields.params, - method: tt.fields.method, - path: tt.fields.path, - } - if got := r.query(tt.args.key); got != tt.want { - t.Errorf("query() = %v, want %v", got, tt.want) - } - }) +} + +func TestRequest_uri(t *testing.T) { + ctx := &fasthttp.RequestCtx{} + ctx.Request.Header.SetMethod("GET") + ctx.Request.Header.SetRequestURI("/test?param=value") + + r := newRequest(ctx) + uri := r.uri() + + if uri == "" { + t.Error("Expected non-empty URI") } } diff --git a/response.go b/response.go index c458139..6cb38a2 100644 --- a/response.go +++ b/response.go @@ -1,51 +1,42 @@ package lightning import ( - "net/http" "os" "path/filepath" + + "github.com/valyala/fasthttp" ) -// response Declaring the response structure that will be used to hold HTTP response body. type response struct { - originReq *http.Request // A pointer to an HTTP request. - originRes http.ResponseWriter // An HTTP response writer. - statusCode int // The status code of the HTTP response (e.g. 200, 404, 500, etc.). - cookies cookiesMap // An array of cookies to be sent with the HTTP response. - body []byte // The response body to be sent. - redirectUrl string // The URL to redirect to. - fileUrl string // The file to send. + ctx *fasthttp.RequestCtx + statusCode int + body []byte + redirectTo string + filePath string + cookies cookiesMap } -// newResponse A constructor function for the response structure. -func newResponse(req *http.Request, res http.ResponseWriter) *response { +func newResponse(ctx *fasthttp.RequestCtx) *response { return &response{ - originReq: req, - originRes: res, - statusCode: http.StatusNotFound, - cookies: cookiesMap{}, - body: nil, - redirectUrl: "", + ctx: ctx, + statusCode: StatusNotFound, + cookies: make(cookiesMap), } } -// setStatus sets the status code of the HTTP response. func (r *response) setStatus(code int) { r.statusCode = code } -// setBody sets the response body to be sent. func (r *response) setBody(body []byte) { r.body = body } -// redirect sets a redirect URL. func (r *response) redirect(code int, url string) { r.statusCode = code - r.redirectUrl = url + r.redirectTo = url } -// file serves a file. func (r *response) file(path string) error { absPath, err := filepath.Abs(path) if err != nil { @@ -56,47 +47,42 @@ func (r *response) file(path string) error { return err } - r.fileUrl = absPath + r.filePath = absPath return nil } -// addHeader adds a new header key-value pair to the response. func (r *response) addHeader(key, value string) { - r.originRes.Header().Add(key, value) + r.ctx.Response.Header.Add(key, value) } -// setHeader sets the value of a given header in the response. func (r *response) setHeader(key string, value string) { - r.originRes.Header().Set(key, value) + r.ctx.Response.Header.Set(key, value) } -// delHeader deletes a given header from the response. func (r *response) delHeader(key string) { - r.originRes.Header().Del(key) + r.ctx.Response.Header.Del(key) } -// sendFile sends the file as an attachment. func (r *response) sendFile() { - base := filepath.Base(r.fileUrl) - r.originRes.Header().Set(HeaderContentDisposition, "attachment; filename="+base) - http.ServeFile(r.originRes, r.originReq, r.fileUrl) + base := filepath.Base(r.filePath) + r.ctx.Response.Header.Set(HeaderContentDisposition, "attachment; filename="+base) + r.ctx.SendFile(r.filePath) } -// flush sends the HTTP response. func (r *response) flush() { - for _, v := range r.cookies { - http.SetCookie(r.originRes, v) + for name, value := range r.cookies { + var c fasthttp.Cookie + c.SetKey(name) + c.SetValue(value) + r.ctx.Response.Header.SetCookie(&c) } - if len(r.fileUrl) > 0 { + if len(r.filePath) > 0 { r.sendFile() - } else if len(r.redirectUrl) > 0 { - http.Redirect(r.originRes, r.originReq, r.redirectUrl, r.statusCode) + } else if len(r.redirectTo) > 0 { + r.ctx.Redirect(r.redirectTo, r.statusCode) } else { - r.originRes.WriteHeader(r.statusCode) - _, err := r.originRes.Write(r.body) - if err != nil { - return - } + r.ctx.Response.SetStatusCode(r.statusCode) + r.ctx.Response.SetBody(r.body) } } diff --git a/response_test.go b/response_test.go index 72ee41a..830d494 100644 --- a/response_test.go +++ b/response_test.go @@ -1,209 +1,138 @@ package lightning import ( - "bytes" - "net/http" - "net/http/httptest" - "os" - "path/filepath" "testing" -) - -func TestNewResponse(t *testing.T) { - req := httptest.NewRequest("GET", "/", nil) - res := httptest.NewRecorder() - resp := newResponse(req, res) + "github.com/valyala/fasthttp" +) - if resp.originReq != req { - t.Errorf("Expected originReq to be %v, but got %v", req, resp.originReq) - } +func createResponse() (*response, *fasthttp.RequestCtx) { + ctx := &fasthttp.RequestCtx{} + resp := newResponse(ctx) + resp.cookies = make(cookiesMap) + return resp, ctx +} - if resp.originRes != res { - t.Errorf("Expected originRes to be %v, but got %v", res, resp.originRes) +func TestNewResponse(t *testing.T) { + resp, ctx := createResponse() + if resp.ctx != ctx { + t.Error("ctx not set correctly") } - - if resp.statusCode != http.StatusNotFound { - t.Errorf("Expected statusCode to be %v, but got %v", http.StatusNotFound, resp.statusCode) + if resp.statusCode != StatusNotFound { + t.Errorf("Expected default status %d, got %d", StatusNotFound, resp.statusCode) } +} - if len(resp.cookies) != 0 { - t.Errorf("Expected cookies to be empty, but got %v", resp.cookies) - } +func TestResponse_setStatus(t *testing.T) { + resp, _ := createResponse() - if resp.body != nil { - t.Errorf("Expected data to be nil, but got %v", resp.body) + resp.setStatus(StatusOK) + if resp.statusCode != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, resp.statusCode) } +} - if resp.redirectUrl != "" { - t.Errorf("Expected redirectUrl to be empty, but got %v", resp.redirectUrl) - } +func TestResponse_setBody(t *testing.T) { + resp, _ := createResponse() - if resp.fileUrl != "" { - t.Errorf("Expected fileUrl to be empty, but got %v", resp.fileUrl) + resp.setBody([]byte("test body")) + if string(resp.body) != "test body" { + t.Errorf("Expected body 'test body', got '%s'", string(resp.body)) } } -func TestSetStatus(t *testing.T) { - req, _ := http.NewRequest("GET", "/", nil) - res := httptest.NewRecorder() - r := newResponse(req, res) - - r.setStatus(http.StatusOK) +func TestResponse_redirect(t *testing.T) { + resp, _ := createResponse() - if r.statusCode != http.StatusOK { - t.Errorf("Expected status code %d, but got %d", http.StatusOK, r.statusCode) + resp.redirect(StatusMovedPermanently, "https://example.com") + if resp.redirectTo != "https://example.com" { + t.Errorf("Expected redirect URL 'https://example.com', got '%s'", resp.redirectTo) + } + if resp.statusCode != StatusMovedPermanently { + t.Errorf("Expected status %d, got %d", StatusMovedPermanently, resp.statusCode) } } -func TestSetBody(t *testing.T) { - req, err := http.NewRequest("GET", "http://example.com", nil) - if err != nil { - t.Fatal(err) - } - res := httptest.NewRecorder() +func TestResponse_addHeader(t *testing.T) { + resp, ctx := createResponse() - r := newResponse(req, res) - body := []byte("test data") - r.setBody(body) + resp.addHeader("X-Custom", "value1") + resp.addHeader("X-Custom", "value2") - if !bytes.Equal(r.body, body) { - t.Errorf("expected body to be %v, but got %v", body, r.body) + hdr := string(ctx.Response.Header.Peek("X-Custom")) + if hdr == "" { + t.Error("Expected header to be set") } } -func TestResponse_Redirect(t *testing.T) { - req := httptest.NewRequest("GET", "http://example.com/foo", nil) - w := httptest.NewRecorder() - resp := newResponse(req, w) +func TestResponse_setHeader(t *testing.T) { + resp, ctx := createResponse() - resp.redirect(http.StatusFound, "http://example.com/bar") + resp.setHeader("Content-Type", "application/json") - resp.flush() - - if w.Code != http.StatusFound { - t.Errorf("expected status code %d, got %d", http.StatusFound, w.Code) - } - - if w.Header().Get("Location") != "http://example.com/bar" { - t.Errorf("expected Location header %q, got %q", "http://example.com/bar", w.Header().Get("Location")) + hdr := string(ctx.Response.Header.Peek("Content-Type")) + if hdr != "application/json" { + t.Errorf("Expected 'application/json', got '%s'", hdr) } } -func TestResponse_File(t *testing.T) { - // create a temporary file - file, err := os.CreateTemp("", "testfile") - if err != nil { - t.Fatal(err) - } - defer os.Remove(file.Name()) +func TestResponse_delHeader(t *testing.T) { + resp, ctx := createResponse() - // set the file path using the file method - resp := newResponse(nil, nil) - err = resp.file(file.Name()) - if err != nil { - t.Fatal(err) - } + resp.setHeader("X-Custom", "value") + resp.delHeader("X-Custom") - // check if the fileUrl field is set to the correct absolute path - absPath, err := filepath.Abs(file.Name()) - if err != nil { - t.Fatal(err) - } - if resp.fileUrl != absPath { - t.Errorf("fileUrl field is %s, expected %s", resp.fileUrl, absPath) + hdr := string(ctx.Response.Header.Peek("X-Custom")) + if hdr != "" { + t.Errorf("Expected empty header, got '%s'", hdr) } } -func TestResponse_AddHeader(t *testing.T) { - req, err := http.NewRequest("GET", "http://example.com", nil) - if err != nil { - t.Fatal(err) - } +func TestResponse_file(t *testing.T) { + resp, _ := createResponse() - res := httptest.NewRecorder() - resp := newResponse(req, res) - - resp.addHeader("X-Test-Header", "test-value") + resp.file("/nonexistent/path") resp.flush() - - if res.Header().Get("X-Test-Header") != "test-value" { - t.Errorf("Expected header X-Test-Header to be set to test-value, but got %s", res.Header().Get("X-Test-Header")) - } } -func TestResponse_SetHeader(t *testing.T) { - req, err := http.NewRequest("GET", "http://example.com", nil) - if err != nil { - t.Fatal(err) - } - - res := httptest.NewRecorder() +func TestResponse_flush(t *testing.T) { + resp, ctx := createResponse() - resp := newResponse(req, res) - - resp.setHeader("X-Test-Header", "test-value") + resp.setStatus(StatusOK) + resp.setBody([]byte("test body")) resp.flush() - if res.Header().Get("X-Test-Header") != "test-value" { - t.Errorf("Expected header X-Test-Header to be set to test-value, but got %s", res.Header().Get("X-Test-Header")) + if ctx.Response.StatusCode() != StatusOK { + t.Errorf("Expected status %d, got %d", StatusOK, ctx.Response.StatusCode()) } -} - -func TestResponse_DelHeader(t *testing.T) { - req, err := http.NewRequest("GET", "http://example.com", nil) - if err != nil { - t.Fatal(err) + if string(ctx.Response.Body()) != "test body" { + t.Errorf("Expected body 'test body', got '%s'", string(ctx.Response.Body())) } +} - res := httptest.NewRecorder() - - resp := newResponse(req, res) +func TestResponse_flushWithCookie(t *testing.T) { + resp, ctx := createResponse() - resp.addHeader("X-Test-Header", "test-value") - resp.delHeader("X-Test-Header") + resp.cookies.set("session", "abc123") resp.flush() - if res.Header().Get("X-Test-Header") != "" { - t.Errorf("Expected header X-Test-Header to be set to test-value, but got %s", res.Header().Get("X-Test-Header")) + cookie := string(ctx.Response.Header.Peek("Set-Cookie")) + if cookie == "" { + t.Error("Expected Set-Cookie header to be set") } } -func TestSendFile(t *testing.T) { - req := httptest.NewRequest(http.MethodGet, "/file.txt", nil) - res := httptest.NewRecorder() - - // Create a temporary file to serve - file, err := os.CreateTemp("", "testfile") - if err != nil { - t.Fatal(err) - } - defer os.Remove(file.Name()) - _, err = file.WriteString("test content") - if err != nil { - t.Fatal(err) - } - err = file.Close() - if err != nil { - t.Fatal(err) - } +func TestResponse_flushWithRedirect(t *testing.T) { + resp, ctx := createResponse() - resp := newResponse(req, res) - err = resp.file(file.Name()) - if err != nil { - t.Fatal(err) - } - resp.sendFile() + resp.redirect(StatusFound, "/new-location") + resp.flush() - // Check that the Content-Disposition header was set correctly - expectedHeader := "attachment; filename=" + filepath.Base(file.Name()) - if res.Header().Get(HeaderContentDisposition) != expectedHeader { - t.Errorf("Expected Content-Disposition header %q, got %q", expectedHeader, res.Header().Get(HeaderContentDisposition)) + if ctx.Response.StatusCode() != StatusFound { + t.Errorf("Expected status %d, got %d", StatusFound, ctx.Response.StatusCode()) } - - // Check that the file was served - expectedBody := "test content" - if res.Body.String() != expectedBody { - t.Errorf("Expected response body %q, got %q", expectedBody, res.Body.String()) + location := string(ctx.Response.Header.Peek("Location")) + if location == "" { + t.Error("Expected Location header to be set") } } diff --git a/router.go b/router.go index 4924171..f148824 100644 --- a/router.go +++ b/router.go @@ -4,31 +4,18 @@ import ( "strings" ) -// node is a struct that represents a node in the trie type node struct { - // Pattern is the pattern of the node - Pattern string `json:"pattern"` - // Part is the part of the node - Part string `json:"part"` - // IsWild is a boolean that indicates whether the node is a wildcard - IsWild bool `json:"isWild"` - // Children is a slice of pointers to the children of the node - Children []*node `json:"children,omitempty"` - // handlers is a slice of HandlerFuncs that are associated with the node + Pattern string `json:"pattern"` + Part string `json:"part"` + IsWild bool `json:"isWild"` + Children map[string]*node `json:"children,omitempty"` handlers []HandlerFunc } -// matchChild returns the child node that matches the given part func (n *node) matchChild(part string) *node { - for _, child := range n.Children { - if child.Part == part { - return child - } - } - return nil + return n.Children[part] } -// insert inserts a new node into the trie func (n *node) insert(pattern string, parts []string, height int, handlers []HandlerFunc) { if len(parts) == height { n.Pattern = pattern @@ -37,15 +24,17 @@ func (n *node) insert(pattern string, parts []string, height int, handlers []Han } part := parts[height] - child := n.matchChild(part) - if child == nil { + if n.Children == nil { + n.Children = make(map[string]*node) + } + child, exists := n.Children[part] + if !exists { child = &node{Part: part, IsWild: part[0] == ':' || part[0] == '*'} - n.Children = append(n.Children, child) + n.Children[part] = child } child.insert(pattern, parts, height+1, handlers) } -// search searches the trie for a node that matches the given parts func (n *node) search(parts []string, height int) *node { if len(parts) == height { if n.Pattern != "" { @@ -54,22 +43,22 @@ func (n *node) search(parts []string, height int) *node { return nil } + if n.IsWild && n.Pattern != "" { + return n + } + part := parts[height] child := n.matchChild(part) - // Attempt to match exact route first if child != nil { - nextNode := child.search(parts, height+1) - if nextNode != nil { + if nextNode := child.search(parts, height+1); nextNode != nil { return nextNode } } - // Attempt to match wildcard route for _, child := range n.Children { if child.IsWild { - nextNode := child.search(parts, height+1) - if nextNode != nil { + if nextNode := child.search(parts, height+1); nextNode != nil { return nextNode } } @@ -78,31 +67,25 @@ func (n *node) search(parts []string, height int) *node { return nil } -// router is a struct that represents a router type router struct { - // Roots is a map of HTTP methods to the root nodes of the trie Roots map[string]*node `json:"roots"` } -// newRouter creates a new router func newRouter() *router { return &router{ - Roots: make(map[string]*node, 0), + Roots: make(map[string]*node), } } -// addRoute adds a new route to the router func (r *router) addRoute(method string, pattern string, handlers []HandlerFunc) { parts := parsePattern(pattern) - _, ok := r.Roots[method] - if !ok { + if r.Roots[method] == nil { r.Roots[method] = &node{} } r.Roots[method].insert(pattern, parts, 0, handlers) } -// findRoute finds the route that matches the given method and path func (r *router) findRoute(method string, path string) ([]HandlerFunc, map[string]string) { searchParts := parsePattern(path) params := make(map[string]string) @@ -113,7 +96,6 @@ func (r *router) findRoute(method string, path string) ([]HandlerFunc, map[strin } n := root.search(searchParts, 0) - if n != nil { parts := parsePattern(n.Pattern) for index, part := range parts { diff --git a/util.go b/util.go index b597558..cb66b1d 100644 --- a/util.go +++ b/util.go @@ -2,7 +2,6 @@ package lightning import ( "log" - "net/http" "os" "strings" ) @@ -23,6 +22,8 @@ func parsePattern(pattern string) []string { return result } +// resolveAddress resolves the address to listen on from the given parameters. +// It checks the PORT environment variable and uses default port if not set. func resolveAddress(addr []string) string { if port := os.Getenv("PORT"); port != "" { log.Printf("[DEBUG] Environment variable PORT=\"%s\"", port) @@ -42,10 +43,10 @@ func resolveAddress(addr []string) string { // defaultNotFound is the default handler function for 404 Not Found error func defaultNotFound(ctx *Context) { - ctx.Text(http.StatusNotFound, http.StatusText(http.StatusNotFound)) + ctx.Text(StatusNotFound, "Not Found") } // defaultInternalServerError is the default handler function for 500 Internal Server Error func defaultInternalServerError(ctx *Context) { - ctx.Text(http.StatusInternalServerError, http.StatusText(http.StatusInternalServerError)) + ctx.Text(StatusInternalServerError, "Internal Server Error") } diff --git a/util_test.go b/util_test.go index 21d1698..6b0f7a4 100644 --- a/util_test.go +++ b/util_test.go @@ -1,11 +1,11 @@ package lightning import ( - "net/http" - "net/http/httptest" "os" "reflect" "testing" + + "github.com/valyala/fasthttp" ) func TestParsePattern(t *testing.T) { @@ -70,7 +70,6 @@ func TestParsePattern(t *testing.T) { } func TestResolveAddress(t *testing.T) { - // Test case for when PORT environment variable is set os.Setenv("PORT", "1234") defer os.Unsetenv("PORT") expected := ":1234" @@ -80,21 +79,18 @@ func TestResolveAddress(t *testing.T) { } os.Unsetenv("PORT") - // Test case for when no parameters are passed in expected = ":6789" result = resolveAddress([]string{}) if result != expected { t.Errorf("Expected %s, but got %s", expected, result) } - // Test case for when one parameter is passed in expected = "localhost:8080" result = resolveAddress([]string{"localhost:8080"}) if result != expected { t.Errorf("Expected %s, but got %s", expected, result) } - // Test case for when more than one parameter is passed in defer func() { if r := recover(); r == nil { t.Errorf("Expected panic, but did not get one") @@ -104,45 +100,45 @@ func TestResolveAddress(t *testing.T) { } func TestDefaultNotFound(t *testing.T) { - req, _ := http.NewRequest("GET", "/foo", nil) - - // Create a new context with a mock response writer - w := httptest.NewRecorder() - ctx, _ := NewContext(w, req) - - // Call the defaultNotFound function - defaultNotFound(ctx) - ctx.flush() + fctx := &fasthttp.RequestCtx{} + fctx.Request.Header.SetMethod("GET") + fctx.Request.Header.SetRequestURI("/foo") + + c := &Context{ + ctx: fctx, + index: -1, + res: newResponse(fctx), + } + defaultNotFound(c) + c.flush() - // Verify that the response status code is 404 - if w.Code != http.StatusNotFound { - t.Errorf("expected status code %d, got %d", http.StatusNotFound, w.Code) + if fctx.Response.StatusCode() != StatusNotFound { + t.Errorf("expected status code %d, got %d", StatusNotFound, fctx.Response.StatusCode()) } - // Verify that the response body is "Not Found" - if w.Body.String() != http.StatusText(http.StatusNotFound) { - t.Errorf("expected body %q, got %q", http.StatusText(http.StatusNotFound), w.Body.String()) + if string(fctx.Response.Body()) != "Not Found" { + t.Errorf("expected body %q, got %q", "Not Found", string(fctx.Response.Body())) } } func TestDefaultInternalServerError(t *testing.T) { - req, _ := http.NewRequest("GET", "/foo", nil) - - // Create a new context with a mock response writer - w := httptest.NewRecorder() - ctx, _ := NewContext(w, req) - - // Call the defaultInternalServerError function - defaultInternalServerError(ctx) - ctx.flush() + fctx := &fasthttp.RequestCtx{} + fctx.Request.Header.SetMethod("GET") + fctx.Request.Header.SetRequestURI("/foo") + + c := &Context{ + ctx: fctx, + index: -1, + res: newResponse(fctx), + } + defaultInternalServerError(c) + c.flush() - // Verify that the response status code is 500 - if w.Code != http.StatusInternalServerError { - t.Errorf("expected status code %d, got %d", http.StatusInternalServerError, w.Code) + if fctx.Response.StatusCode() != StatusInternalServerError { + t.Errorf("expected status code %d, got %d", StatusInternalServerError, fctx.Response.StatusCode()) } - // Verify that the response body is "Internal Server Error" - if w.Body.String() != http.StatusText(http.StatusInternalServerError) { - t.Errorf("expected body %q, got %q", http.StatusText(http.StatusInternalServerError), w.Body.String()) + if string(fctx.Response.Body()) != "Internal Server Error" { + t.Errorf("expected body %q, got %q", "Internal Server Error", string(fctx.Response.Body())) } }