Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions internal/entrypoint/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,9 @@ type Config struct {
}

func (cfg *Config) Validate() error {
if err := cfg.Rules.NotFound.Validate(); err != nil {
return err
}
if cfg.ProxyProtocol == nil {
return nil
}
Expand Down
139 changes: 139 additions & 0 deletions internal/entrypoint/not_found_middleware_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,139 @@
package entrypoint

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/stretchr/testify/require"
"github.com/yusing/godoxy/internal/logging/accesslog"
"github.com/yusing/godoxy/internal/route/rules"
)

type remoteAddrAccessLogger struct {
accesslog.AccessLogger
remoteAddr string
}

func (logger *remoteAddrAccessLogger) LogRequest(request *http.Request, _ *http.Response) {
logger.remoteAddr = request.RemoteAddr
}

func TestNotFoundRuleRequestMiddlewareUpdatesAccessLogAddress(t *testing.T) {
ep := NewTestEntrypoint(t, nil)
server := newTestHTTPServer(t, ep)
logger := &remoteAddrAccessLogger{}
ep.accessLogger = logger

var notFoundRules rules.Rules
require.NoError(t, notFoundRules.Parse(`
default {
middleware CloudflareRealIP {
}
}
`))
ep.SetNotFoundRules(notFoundRules)

request := httptest.NewRequest(http.MethodGet, "http://unknown.example/garbage", nil)
request.RemoteAddr = "127.0.0.1:1234"
request.Header.Set("CF-Connecting-IP", "198.51.100.10")
response := httptest.NewRecorder()

server.ServeHTTP(response, request)

require.Equal(t, http.StatusNotFound, response.Code)
require.Equal(t, "198.51.100.10:1234", request.RemoteAddr)
require.Equal(t, "198.51.100.10:1234", logger.remoteAddr)
require.Equal(t, "198.51.100.10", request.Header.Get("X-Real-IP"))
}

func TestNotFoundRuleRequestMiddlewareIsOptIn(t *testing.T) {
ep := NewTestEntrypoint(t, nil)
server := newTestHTTPServer(t, ep)
logger := &remoteAddrAccessLogger{}
ep.accessLogger = logger

request := httptest.NewRequest(http.MethodGet, "http://unknown.example/garbage", nil)
request.RemoteAddr = "127.0.0.1:1234"
request.Header.Set("CF-Connecting-IP", "198.51.100.10")
response := httptest.NewRecorder()

server.ServeHTTP(response, request)

require.Equal(t, http.StatusNotFound, response.Code)
require.Equal(t, "127.0.0.1:1234", request.RemoteAddr)
require.Equal(t, "127.0.0.1:1234", logger.remoteAddr)
}

func TestNotFoundRuleRequestMiddlewareDoesNotRunForMatchedRoute(t *testing.T) {
ep := NewTestEntrypoint(t, nil)
server := newTestHTTPServer(t, ep)
logger := &remoteAddrAccessLogger{}
ep.accessLogger = logger

matchedRoute := newFakeHTTPRoute(t, "matched.example", "")
matchedRoute.handler = func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
}
server.AddRoute(matchedRoute)

var notFoundRules rules.Rules
require.NoError(t, notFoundRules.Parse(`
default {
middleware CloudflareRealIP {
}
}
`))
ep.SetNotFoundRules(notFoundRules)

request := httptest.NewRequest(http.MethodGet, "http://matched.example/garbage", nil)
request.RemoteAddr = "127.0.0.1:1234"
request.Header.Set("CF-Connecting-IP", "198.51.100.10")
response := httptest.NewRecorder()

server.ServeHTTP(response, request)

require.Equal(t, http.StatusNoContent, response.Code)
require.Equal(t, "127.0.0.1:1234", request.RemoteAddr)
require.Equal(t, "127.0.0.1:1234", logger.remoteAddr)
}

func TestConfigValidateRejectsNotFoundMiddlewareInResponsePhase(t *testing.T) {
tests := []struct {
name string
rules string
wantError bool
}{
{
name: "request phase",
rules: `default {
middleware CloudflareRealIP {
}
}`,
},
{
name: "response phase",
rules: `status 404 {
middleware CloudflareRealIP {
}
}`,
wantError: true,
},
}

for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
var notFoundRules rules.Rules
require.NoError(t, notFoundRules.Parse(test.rules))

cfg := Config{}
cfg.Rules.NotFound = notFoundRules
err := cfg.Validate()
if test.wantError {
require.ErrorContains(t, err, "request middleware cannot be used in a response-phase rule or action block")
} else {
require.NoError(t, err)
}
})
}
}
40 changes: 40 additions & 0 deletions internal/net/gphttp/middleware/middlewares.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,13 @@ package middleware

import (
"errors"
"fmt"
"io/fs"
"path"

"github.com/rs/zerolog/log"
"github.com/yusing/godoxy/internal/common"
"github.com/yusing/godoxy/internal/route/rules"
gperr "github.com/yusing/goutils/errs"
fsutils "github.com/yusing/goutils/fs"
strutils "github.com/yusing/goutils/strings"
Expand Down Expand Up @@ -48,6 +50,44 @@ var (
ErrMiddlewareAlreadyExists = errors.New("middleware with the same name already exists")
)

func init() {
rules.InitRequestMiddlewareResolver(func(name string, options map[string]any) (rules.RequestMiddleware, error) {
definition, err := Get(name)
if err != nil {
return nil, err
}

middleware, err := definition.New(OptionsRaw(options))
if err != nil {
return nil, err
}

if !hasRequestPhase(middleware.impl) {
return nil, fmt.Errorf("middleware %q has no request phase", name)
}

return NewMiddlewareChain(name, []*Middleware{middleware}).TryModifyRequest, nil
})
}

func hasRequestPhase(impl any) bool {
switch impl := impl.(type) {
case *checkBypass:
return impl.modReq != nil && hasRequestPhase(impl.modReq)
case *middlewareChain:
for _, before := range impl.befores {
if hasRequestPhase(before) {
return true
}
}
return false
case RequestModifier:
return true
default:
return false
}
}

func Get(name string) (*Middleware, error) {
middleware, ok := allMiddlewares[strutils.ToLowerNoSnake(name)]
if !ok {
Expand Down
156 changes: 156 additions & 0 deletions internal/net/gphttp/middleware/rule_middleware_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,156 @@
package middleware

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/stretchr/testify/require"
"github.com/yusing/godoxy/internal/net/gphttp"
"github.com/yusing/godoxy/internal/route/rules"
)

func TestRuleMiddlewareResolverSupportsRequestPhase(t *testing.T) {
var configured rules.Rules
require.NoError(t, configured.Parse(`
default {
middleware CloudflareRealIP
}
`))
}

func TestRuleMiddlewareResolverSupportsRequestCompose(t *testing.T) {
component, err := CloudflareRealIP.New(nil)
require.NoError(t, err)

const name = "testrulerequestcompose"
previous, existed := allMiddlewares[name]
allMiddlewares[name] = NewMiddlewareChain(name, []*Middleware{component})
t.Cleanup(func() {
if existed {
allMiddlewares[name] = previous
} else {
delete(allMiddlewares, name)
}
})

var configured rules.Rules
require.NoError(t, configured.Parse(`
default {
middleware testrulerequestcompose
}
`))
}

func TestRuleMiddlewareResolverRejectsUnsupportedMiddleware(t *testing.T) {
responseWithBypass, err := ModifyResponse.New(OptionsRaw{
"bypass": []string{"path /health"},
})
require.NoError(t, err)

const wrappedResponseName = "testrulewrappedresponse"
const responseComposeName = "testruleresponsecompose"
allMiddlewares[wrappedResponseName] = responseWithBypass
allMiddlewares[responseComposeName] = NewMiddlewareChain(responseComposeName, []*Middleware{responseWithBypass})
t.Cleanup(func() {
delete(allMiddlewares, wrappedResponseName)
delete(allMiddlewares, responseComposeName)
})

const emptyComposeName = "testruleemptycompose"
previous, existed := allMiddlewares[emptyComposeName]
allMiddlewares[emptyComposeName] = NewMiddlewareChain(emptyComposeName, nil)
t.Cleanup(func() {
if existed {
allMiddlewares[emptyComposeName] = previous
} else {
delete(allMiddlewares, emptyComposeName)
}
})

tests := []struct {
name string
middleware string
errorText string
}{
{name: "response only", middleware: "ModifyResponse", errorText: "has no request phase"},
{name: "empty compose", middleware: emptyComposeName, errorText: "has no request phase"},
{name: "wrapped response only", middleware: wrappedResponseName, errorText: "has no request phase"},
{name: "response-only compose", middleware: responseComposeName, errorText: "has no request phase"},
{name: "unknown", middleware: "does-not-exist", errorText: "unknown middleware"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
var configured rules.Rules
err := configured.Parse("default {\n middleware " + test.middleware + "\n}")
require.ErrorContains(t, err, test.errorText)
})
}
}

func TestRuleMiddlewareResolverAppliesBlockProperties(t *testing.T) {
var configured rules.Rules
require.NoError(t, configured.Parse(`
default {
middleware RealIP {
header: X-Forwarded-For
from:
- 127.0.0.1/32
}
}
`))

request := httptest.NewRequest(http.MethodGet, "http://unknown.example/", nil)
request.RemoteAddr = "127.0.0.1:1234"
request.Header.Set("X-Forwarded-For", "198.51.100.10")
response := httptest.NewRecorder()
configured.BuildHandler(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
}).ServeHTTP(response, request)

require.Equal(t, http.StatusNoContent, response.Code)
require.Equal(t, "198.51.100.10:1234", request.RemoteAddr)
require.Equal(t, "198.51.100.10", request.Header.Get("X-Real-IP"))
}

func TestRuleMiddlewareResolverPreservesNonUserAuthBypass(t *testing.T) {
var configured rules.Rules
require.NoError(t, configured.Parse(`
default {
middleware CIDRWhitelist {
}
}
`))

t.Run("non-user request bypasses access control", func(t *testing.T) {
fallbackCalled := false
handler := configured.BuildHandler(func(w http.ResponseWriter, _ *http.Request) {
fallbackCalled = true
w.WriteHeader(http.StatusNoContent)
})
request := httptest.NewRequest(http.MethodGet, "http://unknown.example/", nil)
request.RemoteAddr = "203.0.113.10:1234"
request = request.WithContext(gphttp.WithNonUserRequest(request.Context()))
response := httptest.NewRecorder()

handler.ServeHTTP(response, request)

require.True(t, fallbackCalled)
require.Equal(t, http.StatusNoContent, response.Code)
})

t.Run("user request still applies access control", func(t *testing.T) {
fallbackCalled := false
handler := configured.BuildHandler(func(http.ResponseWriter, *http.Request) {
fallbackCalled = true
})
request := httptest.NewRequest(http.MethodGet, "http://unknown.example/", nil)
request.RemoteAddr = "203.0.113.10:1234"
response := httptest.NewRecorder()

handler.ServeHTTP(response, request)

require.False(t, fallbackCalled)
require.Equal(t, http.StatusForbidden, response.Code)
})
}
Loading
Loading