File size: 1,737 Bytes
1f10f31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
package common

import (
	"fmt"
	"net/http"

	"github.com/go-chi/chi/v5/middleware"
	"github.com/google/wire"

	"github.com/openmeterio/openmeter/app/config"
	"github.com/openmeterio/openmeter/openmeter/server"
	pkgserver "github.com/openmeterio/openmeter/pkg/server"
)

var Server = wire.NewSet(
	NewTelemetryRouterHook,
	NewFFXConfigContextMiddleware,
	NewRouterHooks,
	NewPostAuthMiddlewares,
	NewClientIPMiddleware,
)

func NewRouterHooks(
	telemetry TelemetryMiddlewareHook,
) *server.RouterHooks {
	return &server.RouterHooks{
		Middlewares: []server.MiddlewareHook{
			server.MiddlewareHook(telemetry),
		},
	}
}

func NewPostAuthMiddlewares(
	ffx FFXConfigContextMiddleware,
) server.PostAuthMiddlewares {
	return server.PostAuthMiddlewares{
		func(h http.Handler) http.Handler {
			return ffx(h)
		},
	}
}

// ClientIPMiddleware is a defined type (not an alias) so the wire graph does not
// provide the ubiquitous pkgserver.MiddlewareFunc type directly.
type ClientIPMiddleware pkgserver.MiddlewareFunc

func NewClientIPMiddleware(cfg config.ClientIPMiddlewareConfig) (ClientIPMiddleware, error) {
	if err := cfg.Validate(); err != nil {
		return nil, fmt.Errorf("invalid client ip middleware config: %w", err)
	}

	switch cfg.Source {
	case config.ClientIPSourceRemoteAddr:
		return middleware.ClientIPFromRemoteAddr, nil
	case config.ClientIPSourceHeader:
		return middleware.ClientIPFromHeader(cfg.Header), nil
	case config.ClientIPSourceXFF:
		if len(cfg.TrustedIPPrefixes) > 0 {
			return middleware.ClientIPFromXFF(cfg.TrustedIPPrefixes...), nil
		}

		return middleware.ClientIPFromXFFTrustedProxies(cfg.TrustedProxies), nil
	default:
		return nil, fmt.Errorf("invalid client ip middleware source: %s", cfg.Source)
	}
}