refactor: dependency injection instead of package globals (#57)

* refactor: export Settings type, Setup returns value

* refactor: resolver.Setup takes explicit Settings struct

* refactor: GetHeadersWithoutTrustedHeaders takes explicit header params

* refactor: server constructors take narrow config

* refactor: Router struct with handler methods, remove geoSvc global

* refactor: wire DI through main, remove setting.App references

* refactor: remove App global, use returned Settings

* refactor: update router tests for DI

* chore: fix lint issues — rename ServerTimeouts to Timeouts, fix shadowed variables
This commit is contained in:
2026-07-21 18:37:21 +02:00
committed by GitHub
parent 3874f6af3f
commit 2e32a20f60
19 changed files with 298 additions and 253 deletions
+49 -34
View File
@@ -15,17 +15,17 @@ import (
"github.com/dcarrillo/whatismyip/internal/metrics" "github.com/dcarrillo/whatismyip/internal/metrics"
"github.com/dcarrillo/whatismyip/internal/setting" "github.com/dcarrillo/whatismyip/internal/setting"
"github.com/dcarrillo/whatismyip/resolver" "github.com/dcarrillo/whatismyip/resolver"
"github.com/dcarrillo/whatismyip/router"
"github.com/dcarrillo/whatismyip/server" "github.com/dcarrillo/whatismyip/server"
"github.com/dcarrillo/whatismyip/service" "github.com/dcarrillo/whatismyip/service"
"github.com/gin-contrib/secure" "github.com/gin-contrib/secure"
"github.com/patrickmn/go-cache" "github.com/patrickmn/go-cache"
"github.com/dcarrillo/whatismyip/router"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func main() { func main() {
o, err := setting.Setup(os.Args[1:]) cfg, o, err := setting.Setup(os.Args[1:])
if err != nil { if err != nil {
if errors.Is(err, flag.ErrHelp) || errors.Is(err, setting.ErrVersion) { if errors.Is(err, flag.ErrHelp) || errors.Is(err, setting.ErrVersion) {
fmt.Print(o) fmt.Print(o)
@@ -35,33 +35,44 @@ func main() {
os.Exit(1) os.Exit(1)
} }
servers := []server.Server{}
engine := setupEngine()
if setting.App.Resolver.Domain != "" {
store := cache.New(1*time.Minute, 10*time.Minute)
var dnsEngine *resolver.Resolver
if dnsEngine, err = resolver.Setup(store); err != nil {
log.Fatalf("Invalid resolver configuration: %s", err)
}
nameServer := server.NewDNSServer(context.Background(), dnsEngine.Handler())
servers = append(servers, nameServer)
engine.Use(router.GetDNSDiscoveryHandler(store, setting.App.Resolver.Domain, setting.App.Resolver.RedirectPort))
}
var geoSvc *service.Geo var geoSvc *service.Geo
if setting.App.GeodbPath.City != "" || setting.App.GeodbPath.ASN != "" { if cfg.GeodbPath.City != "" || cfg.GeodbPath.ASN != "" {
if geoSvc, err = service.NewGeo(context.Background(), setting.App.GeodbPath.City, setting.App.GeodbPath.ASN); err != nil { if geoSvc, err = service.NewGeo(context.Background(), cfg.GeodbPath.City, cfg.GeodbPath.ASN); err != nil {
log.Fatalf("Failed to load geo databases: %s", err) log.Fatalf("Failed to load geo databases: %s", err)
} }
} }
router.SetupTemplate(engine) servers := []server.Server{}
router.Setup(engine, geoSvc) engine := setupEngine(cfg)
servers = slices.Concat(servers, setupHTTPServers(context.Background(), engine.Handler()))
if setting.App.PrometheusAddress != "" { if cfg.Resolver.Domain != "" {
prometheusServer := server.NewPrometheusServer(context.Background()) store := cache.New(1*time.Minute, 10*time.Minute)
var dnsEngine *resolver.Resolver
if dnsEngine, err = resolver.Setup(store, resolver.Settings{
Domain: cfg.Resolver.Domain,
ResourceRecords: cfg.Resolver.ResourceRecords,
RedirectPort: cfg.Resolver.RedirectPort,
IPv4: cfg.Resolver.Ipv4,
IPv6: cfg.Resolver.Ipv6,
}); err != nil {
log.Fatalf("Invalid resolver configuration: %s", err)
}
nameServer := server.NewDNSServer(context.Background(), dnsEngine.Handler())
servers = append(servers, nameServer)
engine.Use(router.GetDNSDiscoveryHandler(store, geoSvc, cfg.Resolver.Domain, cfg.Resolver.RedirectPort))
}
rt := router.NewRouter(geoSvc, cfg.TrustedHeader, cfg.TrustedPortHeader, cfg.TemplatePath, cfg.DisableTCPScan)
router.SetupTemplate(engine, cfg.TemplatePath)
router.Setup(engine, rt)
servers = slices.Concat(servers, setupHTTPServers(context.Background(), engine.Handler(), cfg))
if cfg.PrometheusAddress != "" {
prometheusServer := server.NewPrometheusServer(context.Background(), cfg.PrometheusAddress,
server.Timeouts{
ReadTimeout: cfg.Server.ReadTimeout,
WriteTimeout: cfg.Server.WriteTimeout,
})
servers = append(servers, prometheusServer) servers = append(servers, prometheusServer)
} }
@@ -69,18 +80,18 @@ func main() {
whatismyip.Run() whatismyip.Run()
} }
func setupEngine() *gin.Engine { func setupEngine(cfg setting.Settings) *gin.Engine {
gin.DisableConsoleColor() gin.DisableConsoleColor()
if os.Getenv(gin.EnvGinMode) == "" { if os.Getenv(gin.EnvGinMode) == "" {
gin.SetMode(gin.ReleaseMode) gin.SetMode(gin.ReleaseMode)
} }
engine := gin.New() engine := gin.New()
engine.Use(gin.LoggerWithFormatter(httputils.GetLogFormatter), gin.Recovery()) engine.Use(gin.LoggerWithFormatter(httputils.GetLogFormatter), gin.Recovery())
if setting.App.PrometheusAddress != "" { if cfg.PrometheusAddress != "" {
metrics.Enable() metrics.Enable()
engine.Use(metrics.GinMiddleware()) engine.Use(metrics.GinMiddleware())
} }
if setting.App.EnableSecureHeaders { if cfg.EnableSecureHeaders {
engine.Use(secure.New(secure.Config{ engine.Use(secure.New(secure.Config{
BrowserXssFilter: true, BrowserXssFilter: true,
ContentTypeNosniff: true, ContentTypeNosniff: true,
@@ -88,24 +99,28 @@ func setupEngine() *gin.Engine {
})) }))
} }
_ = engine.SetTrustedProxies(nil) _ = engine.SetTrustedProxies(nil)
engine.TrustedPlatform = setting.App.TrustedHeader engine.TrustedPlatform = cfg.TrustedHeader
return engine return engine
} }
func setupHTTPServers(ctx context.Context, handler http.Handler) []server.Server { func setupHTTPServers(ctx context.Context, handler http.Handler, cfg setting.Settings) []server.Server {
var servers []server.Server var servers []server.Server
timeouts := server.Timeouts{
ReadTimeout: cfg.Server.ReadTimeout,
WriteTimeout: cfg.Server.WriteTimeout,
}
if setting.App.BindAddress != "" { if cfg.BindAddress != "" {
tcpServer := server.NewTCPServer(ctx, handler) tcpServer := server.NewTCPServer(ctx, handler, cfg.BindAddress, timeouts)
servers = append(servers, tcpServer) servers = append(servers, tcpServer)
} }
if setting.App.TLSAddress != "" { if cfg.TLSAddress != "" {
tlsServer := server.NewTLSServer(ctx, handler) tlsServer := server.NewTLSServer(ctx, handler, cfg.TLSAddress, cfg.TLSCrtPath, cfg.TLSKeyPath, timeouts)
servers = append(servers, tlsServer) servers = append(servers, tlsServer)
if setting.App.EnableHTTP3 { if cfg.EnableHTTP3 {
quicServer := server.NewQuicServer(ctx, tlsServer) quicServer := server.NewQuicServer(ctx, tlsServer, cfg.TLSAddress, cfg.TLSCrtPath, cfg.TLSKeyPath)
servers = append(servers, quicServer) servers = append(servers, quicServer)
} }
} }
+2 -3
View File
@@ -7,7 +7,6 @@ import (
"sort" "sort"
"strings" "strings"
"github.com/dcarrillo/whatismyip/internal/setting"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -30,10 +29,10 @@ func HeadersToSortedString(headers http.Header) string {
} }
// GetHeadersWithoutTrustedHeaders returns a copy of the request headers with the trusted headers removed // GetHeadersWithoutTrustedHeaders returns a copy of the request headers with the trusted headers removed
func GetHeadersWithoutTrustedHeaders(ctx *gin.Context) http.Header { func GetHeadersWithoutTrustedHeaders(ctx *gin.Context, trustedHeader, trustedPortHeader string) http.Header {
h := ctx.Request.Header.Clone() h := ctx.Request.Header.Clone()
for _, k := range []string{setting.App.TrustedHeader, setting.App.TrustedPortHeader} { for _, k := range []string{trustedHeader, trustedPortHeader} {
delete(h, textproto.CanonicalMIMEHeaderKey(k)) delete(h, textproto.CanonicalMIMEHeaderKey(k))
} }
+39 -42
View File
@@ -29,7 +29,7 @@ type resolver struct {
Ipv6 []string `yaml:"ipv6,omitempty"` Ipv6 []string `yaml:"ipv6,omitempty"`
} }
type settings struct { type Settings struct {
GeodbPath geodbConf GeodbPath geodbConf
TemplatePath string TemplatePath string
BindAddress string BindAddress string
@@ -51,75 +51,72 @@ const defaultAddress = ":8080"
var ErrVersion = errors.New("setting: version requested") var ErrVersion = errors.New("setting: version requested")
var App = settings{ func Setup(args []string) (cfg Settings, output string, err error) {
// hard-coded for the time being
Server: serverSettings{
ReadTimeout: 10 * time.Second,
WriteTimeout: 10 * time.Second,
},
}
func Setup(args []string) (output string, err error) {
flags := flag.NewFlagSet("whatismyip", flag.ContinueOnError) flags := flag.NewFlagSet("whatismyip", flag.ContinueOnError)
var buf bytes.Buffer var buf bytes.Buffer
var resolverConf string var resolverConf string
flags.SetOutput(&buf) flags.SetOutput(&buf)
flags.StringVar(&App.GeodbPath.City, "geoip2-city", "", "Path to GeoIP2 city database. Enables geo information (--geoip2-asn becomes mandatory)") cfg.Server = serverSettings{
flags.StringVar(&App.GeodbPath.ASN, "geoip2-asn", "", "Path to GeoIP2 ASN database. Enables ASN information. (--geoip2-city becomes mandatory)") ReadTimeout: 10 * time.Second,
flags.StringVar(&App.TemplatePath, "template", "", "Path to the template file") WriteTimeout: 10 * time.Second,
}
flags.StringVar(&cfg.GeodbPath.City, "geoip2-city", "", "Path to GeoIP2 city database. Enables geo information (--geoip2-asn becomes mandatory)")
flags.StringVar(&cfg.GeodbPath.ASN, "geoip2-asn", "", "Path to GeoIP2 ASN database. Enables ASN information. (--geoip2-city becomes mandatory)")
flags.StringVar(&cfg.TemplatePath, "template", "", "Path to the template file")
flags.StringVar( flags.StringVar(
&resolverConf, &resolverConf,
"resolver", "resolver",
"", "",
"Path to the resolver configuration. It actually enables the resolver for DNS client discovery.") "Path to the resolver configuration. It actually enables the resolver for DNS client discovery.")
flags.StringVar( flags.StringVar(
&App.BindAddress, &cfg.BindAddress,
"bind", "bind",
defaultAddress, defaultAddress,
"Listening address (see https://pkg.go.dev/net?#Listen)", "Listening address (see https://pkg.go.dev/net?#Listen)",
) )
flags.StringVar( flags.StringVar(
&App.TLSAddress, &cfg.TLSAddress,
"tls-bind", "tls-bind",
"", "",
"Listening address for TLS (see https://pkg.go.dev/net?#Listen)", "Listening address for TLS (see https://pkg.go.dev/net?#Listen)",
) )
flags.StringVar(&App.TLSCrtPath, "tls-crt", "", "When using TLS, path to certificate file") flags.StringVar(&cfg.TLSCrtPath, "tls-crt", "", "When using TLS, path to certificate file")
flags.StringVar(&App.TLSKeyPath, "tls-key", "", "When using TLS, path to private key file") flags.StringVar(&cfg.TLSKeyPath, "tls-key", "", "When using TLS, path to private key file")
flags.StringVar( flags.StringVar(
&App.PrometheusAddress, &cfg.PrometheusAddress,
"metrics-bind", "metrics-bind",
"", "",
"Listening address for Prometheus metrics endpoint (see https://pkg.go.dev/net?#Listen). It enables the metrics available at the given address/port via the /metrics endpoint.", "Listening address for Prometheus metrics endpoint (see https://pkg.go.dev/net?#Listen). It enables the metrics available at the given address/port via the /metrics endpoint.",
) )
flags.StringVar( flags.StringVar(
&App.TrustedHeader, &cfg.TrustedHeader,
"trusted-header", "trusted-header",
"", "",
"Trusted request header for remote IP (e.g. X-Real-IP). When using this feature if -trusted-port-header is not set the client port is shown as 'unknown'", "Trusted request header for remote IP (e.g. X-Real-IP). When using this feature if -trusted-port-header is not set the client port is shown as 'unknown'",
) )
flags.StringVar( flags.StringVar(
&App.TrustedPortHeader, &cfg.TrustedPortHeader,
"trusted-port-header", "trusted-port-header",
"", "",
"Trusted request header for remote client port (e.g. X-Real-Port). When this parameter is set -trusted-header becomes mandatory", "Trusted request header for remote client port (e.g. X-Real-Port). When this parameter is set -trusted-header becomes mandatory",
) )
flags.BoolVar(&App.version, "version", false, "Output version information and exit") flags.BoolVar(&cfg.version, "version", false, "Output version information and exit")
flags.BoolVar( flags.BoolVar(
&App.EnableSecureHeaders, &cfg.EnableSecureHeaders,
"enable-secure-headers", "enable-secure-headers",
false, false,
"Add sane security-related headers to every response", "Add sane security-related headers to every response",
) )
flags.BoolVar( flags.BoolVar(
&App.EnableHTTP3, &cfg.EnableHTTP3,
"enable-http3", "enable-http3",
false, false,
"Enable HTTP/3 protocol. HTTP/3 requires --tls-bind set, as HTTP/3 starts as a TLS connection that then gets upgraded to UDP. The UDP port is the same as the one used for the TLS server.", "Enable HTTP/3 protocol. HTTP/3 requires --tls-bind set, as HTTP/3 starts as a TLS connection that then gets upgraded to UDP. The UDP port is the same as the one used for the TLS server.",
) )
flags.BoolVar( flags.BoolVar(
&App.DisableTCPScan, &cfg.DisableTCPScan,
"disable-scan", "disable-scan",
false, false,
"Disable TCP port scanning functionality", "Disable TCP port scanning functionality",
@@ -127,48 +124,48 @@ func Setup(args []string) (output string, err error) {
err = flags.Parse(args) err = flags.Parse(args)
if err != nil { if err != nil {
return buf.String(), err return cfg, buf.String(), err
} }
if App.version { if cfg.version {
return fmt.Sprintf("whatismyip version %s", core.Version), ErrVersion return cfg, fmt.Sprintf("whatismyip version %s", core.Version), ErrVersion
} }
if (App.GeodbPath.City != "" && App.GeodbPath.ASN == "") || (App.GeodbPath.City == "" && App.GeodbPath.ASN != "") { if (cfg.GeodbPath.City != "" && cfg.GeodbPath.ASN == "") || (cfg.GeodbPath.City == "" && cfg.GeodbPath.ASN != "") {
return "", errors.New("both --geoip2-city and --geoip2-asn are mandatory to enable geo information") return cfg, "", errors.New("both --geoip2-city and --geoip2-asn are mandatory to enable geo information")
} }
if App.TrustedPortHeader != "" && App.TrustedHeader == "" { if cfg.TrustedPortHeader != "" && cfg.TrustedHeader == "" {
return "", errors.New("trusted-header is mandatory when trusted-port-header is set") return cfg, "", errors.New("trusted-header is mandatory when trusted-port-header is set")
} }
if (App.TLSAddress != "") && (App.TLSCrtPath == "" || App.TLSKeyPath == "") { if (cfg.TLSAddress != "") && (cfg.TLSCrtPath == "" || cfg.TLSKeyPath == "") {
return "", errors.New("in order to use TLS, the -tls-crt and -tls-key flags are mandatory") return cfg, "", errors.New("in order to use TLS, the -tls-crt and -tls-key flags are mandatory")
} }
if App.EnableHTTP3 && App.TLSAddress == "" { if cfg.EnableHTTP3 && cfg.TLSAddress == "" {
return "", errors.New("in order to use HTTP3, the -tls-bind is mandatory") return cfg, "", errors.New("in order to use HTTP3, the -tls-bind is mandatory")
} }
if App.TemplatePath != "" { if cfg.TemplatePath != "" {
info, err := os.Stat(App.TemplatePath) info, err := os.Stat(cfg.TemplatePath)
if err != nil { if err != nil {
return "", fmt.Errorf("template path: %w", err) return cfg, "", fmt.Errorf("template path: %w", err)
} }
if info.IsDir() { if info.IsDir() {
return "", fmt.Errorf("%s must be a file", App.TemplatePath) return cfg, "", fmt.Errorf("%s must be a file", cfg.TemplatePath)
} }
} }
if resolverConf != "" { if resolverConf != "" {
var err error var err error
App.Resolver, err = readYAML(resolverConf) cfg.Resolver, err = readYAML(resolverConf)
if err != nil { if err != nil {
return "", fmt.Errorf("reading resolver configuration: %w", err) return cfg, "", fmt.Errorf("reading resolver configuration: %w", err)
} }
} }
return buf.String(), nil return cfg, buf.String(), nil
} }
func readYAML(path string) (resolver resolver, err error) { func readYAML(path string) (resolver resolver, err error) {
+13 -13
View File
@@ -55,7 +55,7 @@ func TestParseMandatoryFlags(t *testing.T) {
for _, tt := range mandatoryFlags { for _, tt := range mandatoryFlags {
t.Run(strings.Join(tt.args, " "), func(t *testing.T) { t.Run(strings.Join(tt.args, " "), func(t *testing.T) {
_, err := Setup(tt.args) _, _, err := Setup(tt.args)
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "mandatory") assert.Contains(t, err.Error(), "mandatory")
}) })
@@ -65,11 +65,11 @@ func TestParseMandatoryFlags(t *testing.T) {
func TestParseFlags(t *testing.T) { func TestParseFlags(t *testing.T) {
flags := []struct { flags := []struct {
args []string args []string
conf settings conf Settings
}{ }{
{ {
[]string{}, []string{},
settings{ Settings{
BindAddress: ":8080", BindAddress: ":8080",
Server: serverSettings{ Server: serverSettings{
ReadTimeout: 10 * time.Second, ReadTimeout: 10 * time.Second,
@@ -79,7 +79,7 @@ func TestParseFlags(t *testing.T) {
}, },
{ {
[]string{"-disable-scan"}, []string{"-disable-scan"},
settings{ Settings{
BindAddress: ":8080", BindAddress: ":8080",
Server: serverSettings{ Server: serverSettings{
ReadTimeout: 10 * time.Second, ReadTimeout: 10 * time.Second,
@@ -90,7 +90,7 @@ func TestParseFlags(t *testing.T) {
}, },
{ {
[]string{"-bind", ":8001", "-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path"}, []string{"-bind", ":8001", "-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path"},
settings{ Settings{
GeodbPath: geodbConf{ GeodbPath: geodbConf{
City: "/city-path", City: "/city-path",
ASN: "/asn-path", ASN: "/asn-path",
@@ -107,7 +107,7 @@ func TestParseFlags(t *testing.T) {
"-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path", "-tls-bind", ":9000", "-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path", "-tls-bind", ":9000",
"-tls-crt", "/crt-path", "-tls-key", "/key-path", "-tls-crt", "/crt-path", "-tls-key", "/key-path",
}, },
settings{ Settings{
GeodbPath: geodbConf{ GeodbPath: geodbConf{
City: "/city-path", City: "/city-path",
ASN: "/asn-path", ASN: "/asn-path",
@@ -127,7 +127,7 @@ func TestParseFlags(t *testing.T) {
"-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path", "-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path",
"-trusted-header", "header", "-trusted-port-header", "port-header", "-trusted-header", "header", "-trusted-port-header", "port-header",
}, },
settings{ Settings{
GeodbPath: geodbConf{ GeodbPath: geodbConf{
City: "/city-path", City: "/city-path",
ASN: "/asn-path", ASN: "/asn-path",
@@ -146,7 +146,7 @@ func TestParseFlags(t *testing.T) {
"-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path", "-geoip2-city", "/city-path", "-geoip2-asn", "/asn-path",
"-trusted-header", "header", "-enable-secure-headers", "-trusted-header", "header", "-enable-secure-headers",
}, },
settings{ Settings{
GeodbPath: geodbConf{ GeodbPath: geodbConf{
City: "/city-path", City: "/city-path",
ASN: "/asn-path", ASN: "/asn-path",
@@ -164,9 +164,9 @@ func TestParseFlags(t *testing.T) {
for _, tt := range flags { for _, tt := range flags {
t.Run(strings.Join(tt.args, " "), func(t *testing.T) { t.Run(strings.Join(tt.args, " "), func(t *testing.T) {
_, err := Setup(tt.args) cfg, _, err := Setup(tt.args)
require.Nil(t, err) require.Nil(t, err)
assert.True(t, reflect.DeepEqual(App, tt.conf)) assert.True(t, reflect.DeepEqual(cfg, tt.conf))
}) })
} }
} }
@@ -176,7 +176,7 @@ func TestParseFlagsUsage(t *testing.T) {
for _, arg := range usageArgs { for _, arg := range usageArgs {
t.Run(arg, func(t *testing.T) { t.Run(arg, func(t *testing.T) {
output, err := Setup([]string{arg}) _, output, err := Setup([]string{arg})
assert.ErrorIs(t, err, flag.ErrHelp) assert.ErrorIs(t, err, flag.ErrHelp)
assert.Contains(t, output, "Usage of") assert.Contains(t, output, "Usage of")
}) })
@@ -184,7 +184,7 @@ func TestParseFlagsUsage(t *testing.T) {
} }
func TestParseFlagVersion(t *testing.T) { func TestParseFlagVersion(t *testing.T) {
output, err := Setup([]string{"-version"}) _, output, err := Setup([]string{"-version"})
assert.ErrorIs(t, err, ErrVersion) assert.ErrorIs(t, err, ErrVersion)
assert.Contains(t, output, "whatismyip version") assert.Contains(t, output, "whatismyip version")
} }
@@ -209,7 +209,7 @@ func TestParseFlagTemplate(t *testing.T) {
for _, tc := range testCases { for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
_, err := Setup(tc.flags) _, _, err := Setup(tc.flags)
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), tc.errMsg) assert.Contains(t, err.Error(), tc.errMsg)
}) })
+14 -7
View File
@@ -7,12 +7,19 @@ import (
"strings" "strings"
"github.com/dcarrillo/whatismyip/internal/metrics" "github.com/dcarrillo/whatismyip/internal/metrics"
"github.com/dcarrillo/whatismyip/internal/setting"
"github.com/dcarrillo/whatismyip/internal/validator/uuid" "github.com/dcarrillo/whatismyip/internal/validator/uuid"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/patrickmn/go-cache" "github.com/patrickmn/go-cache"
) )
type Settings struct {
Domain string
ResourceRecords []string
RedirectPort string
IPv4 []string
IPv6 []string
}
type Resolver struct { type Resolver struct {
handler *dns.ServeMux handler *dns.ServeMux
store *cache.Cache store *cache.Cache
@@ -29,11 +36,11 @@ func ensureDotSuffix(s string) string {
return s return s
} }
func Setup(store *cache.Cache) (*Resolver, error) { func Setup(store *cache.Cache, cfg Settings) (*Resolver, error) {
domain := ensureDotSuffix(setting.App.Resolver.Domain) domain := ensureDotSuffix(cfg.Domain)
rr := make([]dns.RR, 0, len(setting.App.Resolver.ResourceRecords)) rr := make([]dns.RR, 0, len(cfg.ResourceRecords))
for _, res := range setting.App.Resolver.ResourceRecords { for _, res := range cfg.ResourceRecords {
record, err := dns.NewRR(domain + " " + res) record, err := dns.NewRR(domain + " " + res)
if err != nil { if err != nil {
return nil, fmt.Errorf("parsing resource record %q: %w", res, err) return nil, fmt.Errorf("parsing resource record %q: %w", res, err)
@@ -41,11 +48,11 @@ func Setup(store *cache.Cache) (*Resolver, error) {
rr = append(rr, record) rr = append(rr, record)
} }
ipv4, err := parseIPs(setting.App.Resolver.Ipv4) ipv4, err := parseIPs(cfg.IPv4)
if err != nil { if err != nil {
return nil, fmt.Errorf("parsing ipv4 addresses: %w", err) return nil, fmt.Errorf("parsing ipv4 addresses: %w", err)
} }
ipv6, err := parseIPs(setting.App.Resolver.Ipv6) ipv6, err := parseIPs(cfg.IPv6)
if err != nil { if err != nil {
return nil, fmt.Errorf("parsing ipv6 addresses: %w", err) return nil, fmt.Errorf("parsing ipv6 addresses: %w", err)
} }
+4 -3
View File
@@ -7,6 +7,7 @@ import (
"strings" "strings"
validator "github.com/dcarrillo/whatismyip/internal/validator/uuid" validator "github.com/dcarrillo/whatismyip/internal/validator/uuid"
"github.com/dcarrillo/whatismyip/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/google/uuid" "github.com/google/uuid"
"github.com/patrickmn/go-cache" "github.com/patrickmn/go-cache"
@@ -27,7 +28,7 @@ type dnsData struct {
// TODO // TODO
// Implement a proper vhost manager instead of using a middleware // Implement a proper vhost manager instead of using a middleware
func GetDNSDiscoveryHandler(store *cache.Cache, domain string, redirectPort string) gin.HandlerFunc { func GetDNSDiscoveryHandler(store *cache.Cache, geoSvc *service.Geo, domain string, redirectPort string) gin.HandlerFunc {
return func(ctx *gin.Context) { return func(ctx *gin.Context) {
host := hostWithoutPort(ctx.Request.Host) host := hostWithoutPort(ctx.Request.Host)
if host != domain && !strings.HasSuffix(host, "."+domain) { if host != domain && !strings.HasSuffix(host, "."+domain) {
@@ -41,7 +42,7 @@ func GetDNSDiscoveryHandler(store *cache.Cache, domain string, redirectPort stri
return return
} }
handleDNS(ctx, store) handleDNS(ctx, store, geoSvc)
ctx.Abort() ctx.Abort()
} }
} }
@@ -53,7 +54,7 @@ func hostWithoutPort(host string) string {
return host return host
} }
func handleDNS(ctx *gin.Context, store *cache.Cache) { func handleDNS(ctx *gin.Context, store *cache.Cache, geoSvc *service.Geo) {
d := strings.Split(ctx.Request.Host, ".")[0] d := strings.Split(ctx.Request.Host, ".")[0]
if !validator.IsValid(d) { if !validator.IsValid(d) {
ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound)) ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound))
+3 -3
View File
@@ -16,7 +16,7 @@ import (
func TestGetDNSDiscoveryHandler(t *testing.T) { func TestGetDNSDiscoveryHandler(t *testing.T) {
store := cache.New(cache.NoExpiration, cache.NoExpiration) store := cache.New(cache.NoExpiration, cache.NoExpiration)
handler := GetDNSDiscoveryHandler(store, domain, "") handler := GetDNSDiscoveryHandler(store, rt.geo, domain, "")
t.Run("calls next if host does not have domain suffix", func(t *testing.T) { t.Run("calls next if host does not have domain suffix", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/", nil) req, _ := http.NewRequest("GET", "/", nil)
@@ -106,7 +106,7 @@ func TestHandleDNS(t *testing.T) {
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, _ := gin.CreateTestContext(w)
c.Request = req c.Request = req
handleDNS(c, store) handleDNS(c, store, rt.geo)
assert.Equal(t, http.StatusNotFound, w.Code) assert.Equal(t, http.StatusNotFound, w.Code)
}) })
} }
@@ -144,7 +144,7 @@ func TestAcceptDNSRequest(t *testing.T) {
c.Request = req c.Request = req
store.Set(u, testIP.ipv4, cache.DefaultExpiration) store.Set(u, testIP.ipv4, cache.DefaultExpiration)
handleDNS(c, store) handleDNS(c, store, rt.geo)
assert.Equal(t, http.StatusOK, w.Code) assert.Equal(t, http.StatusOK, w.Code)
assert.Equal(t, tt.want, w.Body.String()) assert.Equal(t, tt.want, w.Body.String())
+25 -26
View File
@@ -6,7 +6,6 @@ import (
"path/filepath" "path/filepath"
"github.com/dcarrillo/whatismyip/internal/httputils" "github.com/dcarrillo/whatismyip/internal/httputils"
"github.com/dcarrillo/whatismyip/internal/setting"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -31,31 +30,31 @@ type JSONResponse struct {
GeoResponse GeoResponse
} }
func getRoot(ctx *gin.Context) { func (rt *Router) getRoot(ctx *gin.Context) {
switch ctx.NegotiateFormat(gin.MIMEPlain, gin.MIMEHTML, gin.MIMEJSON) { switch ctx.NegotiateFormat(gin.MIMEPlain, gin.MIMEHTML, gin.MIMEJSON) {
case gin.MIMEHTML: case gin.MIMEHTML:
name := "home" name := "home"
if setting.App.TemplatePath != "" { if rt.templatePath != "" {
name = filepath.Base(setting.App.TemplatePath) name = filepath.Base(rt.templatePath)
} }
ctx.HTML(http.StatusOK, name, jsonOutput(ctx)) ctx.HTML(http.StatusOK, name, rt.jsonOutput(ctx))
case gin.MIMEJSON: case gin.MIMEJSON:
getJSON(ctx) rt.getJSON(ctx)
default: default:
ctx.String(http.StatusOK, ctx.ClientIP()+"\n") ctx.String(http.StatusOK, ctx.ClientIP()+"\n")
} }
} }
func getClientPort(ctx *gin.Context) string { func (rt *Router) getClientPort(ctx *gin.Context) string {
var port string var port string
if setting.App.TrustedPortHeader == "" { if rt.trustedPortHeader == "" {
if setting.App.TrustedHeader != "" { if rt.trustedHeader != "" {
port = "unknown" port = "unknown"
} else { } else {
_, port, _ = net.SplitHostPort(ctx.Request.RemoteAddr) _, port, _ = net.SplitHostPort(ctx.Request.RemoteAddr)
} }
} else { } else {
port = ctx.GetHeader(setting.App.TrustedPortHeader) port = ctx.GetHeader(rt.trustedPortHeader)
if port == "" { if port == "" {
port = "unknown" port = "unknown"
} }
@@ -64,37 +63,37 @@ func getClientPort(ctx *gin.Context) string {
return port return port
} }
func getClientPortAsString(ctx *gin.Context) { func (rt *Router) getClientPortAsString(ctx *gin.Context) {
ctx.String(http.StatusOK, getClientPort(ctx)+"\n") ctx.String(http.StatusOK, rt.getClientPort(ctx)+"\n")
} }
func getAllAsString(ctx *gin.Context) { func (rt *Router) getAllAsString(ctx *gin.Context) {
ip := net.ParseIP(ctx.ClientIP()) ip := net.ParseIP(ctx.ClientIP())
output := "IP: " + ip.String() + "\n" output := "IP: " + ip.String() + "\n"
output += "Client Port: " + getClientPort(ctx) + "\n" output += "Client Port: " + rt.getClientPort(ctx) + "\n"
if geoSvc != nil { if rt.geo != nil {
if cityRecord := geoSvc.LookUpCity(ip); cityRecord != nil { if cityRecord := rt.geo.LookUpCity(ip); cityRecord != nil {
output += geoCityRecordToString(cityRecord) + "\n" output += geoCityRecordToString(cityRecord) + "\n"
} }
if asnRecord := geoSvc.LookUpASN(ip); asnRecord != nil { if asnRecord := rt.geo.LookUpASN(ip); asnRecord != nil {
output += geoASNRecordToString(asnRecord) + "\n" output += geoASNRecordToString(asnRecord) + "\n"
} }
} }
h := httputils.GetHeadersWithoutTrustedHeaders(ctx) h := httputils.GetHeadersWithoutTrustedHeaders(ctx, rt.trustedHeader, rt.trustedPortHeader)
h.Set("Host", ctx.Request.Host) h.Set("Host", ctx.Request.Host)
output += httputils.HeadersToSortedString(h) output += httputils.HeadersToSortedString(h)
ctx.String(http.StatusOK, output) ctx.String(http.StatusOK, output)
} }
func getJSON(ctx *gin.Context) { func (rt *Router) getJSON(ctx *gin.Context) {
ctx.JSON(http.StatusOK, jsonOutput(ctx)) ctx.JSON(http.StatusOK, rt.jsonOutput(ctx))
} }
func jsonOutput(ctx *gin.Context) JSONResponse { func (rt *Router) jsonOutput(ctx *gin.Context) JSONResponse {
ip := net.ParseIP(ctx.ClientIP()) ip := net.ParseIP(ctx.ClientIP())
var version byte = 4 var version byte = 4
@@ -103,8 +102,8 @@ func jsonOutput(ctx *gin.Context) JSONResponse {
} }
geoResp := GeoResponse{} geoResp := GeoResponse{}
if geoSvc != nil { if rt.geo != nil {
if cityRecord := geoSvc.LookUpCity(ip); cityRecord != nil { if cityRecord := rt.geo.LookUpCity(ip); cityRecord != nil {
geoResp.Country = cityRecord.Country.Names["en"] geoResp.Country = cityRecord.Country.Names["en"]
geoResp.CountryCode = cityRecord.Country.ISOCode geoResp.CountryCode = cityRecord.Country.ISOCode
geoResp.City = cityRecord.City.Names["en"] geoResp.City = cityRecord.City.Names["en"]
@@ -113,7 +112,7 @@ func jsonOutput(ctx *gin.Context) JSONResponse {
geoResp.PostalCode = cityRecord.Postal.Code geoResp.PostalCode = cityRecord.Postal.Code
geoResp.TimeZone = cityRecord.Location.TimeZone geoResp.TimeZone = cityRecord.Location.TimeZone
} }
if asnRecord := geoSvc.LookUpASN(ip); asnRecord != nil { if asnRecord := rt.geo.LookUpASN(ip); asnRecord != nil {
geoResp.ASN = asnRecord.AutonomousSystemNumber geoResp.ASN = asnRecord.AutonomousSystemNumber
geoResp.ASNOrganization = asnRecord.AutonomousSystemOrganization geoResp.ASNOrganization = asnRecord.AutonomousSystemOrganization
} }
@@ -122,9 +121,9 @@ func jsonOutput(ctx *gin.Context) JSONResponse {
return JSONResponse{ return JSONResponse{
IP: ip.String(), IP: ip.String(),
IPVersion: version, IPVersion: version,
ClientPort: getClientPort(ctx), ClientPort: rt.getClientPort(ctx),
Host: ctx.Request.Host, Host: ctx.Request.Host,
Headers: httputils.GetHeadersWithoutTrustedHeaders(ctx), Headers: httputils.GetHeadersWithoutTrustedHeaders(ctx, rt.trustedHeader, rt.trustedPortHeader),
GeoResponse: geoResp, GeoResponse: geoResp,
} }
} }
+28 -40
View File
@@ -1,12 +1,14 @@
package router package router
import ( import (
"context"
"net" "net"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"testing" "testing"
"github.com/dcarrillo/whatismyip/internal/setting" "github.com/dcarrillo/whatismyip/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
@@ -99,7 +101,8 @@ func TestHost(t *testing.T) {
func TestClientPort(t *testing.T) { func TestClientPort(t *testing.T) {
type args struct { type args struct {
params []string trustedHeader string
trustedPortHeader string
headers map[string][]string headers map[string][]string
} }
tests := []struct { tests := []struct {
@@ -114,35 +117,23 @@ func TestClientPort(t *testing.T) {
{ {
name: "Trusted header only set", name: "Trusted header only set",
args: args{ args: args{
params: []string{ trustedHeader: trustedHeader,
"-geoip2-city", "city",
"-geoip2-asn", "asn",
"-trusted-header", trustedHeader,
},
}, },
expected: "unknown\n", expected: "unknown\n",
}, },
{ {
name: "Trusted and port header set but not included in headers", name: "Trusted and port header set but not included in headers",
args: args{ args: args{
params: []string{ trustedHeader: trustedHeader,
"-geoip2-city", "city", trustedPortHeader: trustedPortHeader,
"-geoip2-asn", "asn",
"-trusted-header", trustedHeader,
"-trusted-port-header", trustedPortHeader,
},
}, },
expected: "unknown\n", expected: "unknown\n",
}, },
{ {
name: "Trusted and port header set and included in headers", name: "Trusted and port header set and included in headers",
args: args{ args: args{
params: []string{ trustedHeader: trustedHeader,
"-geoip2-city", "city", trustedPortHeader: trustedPortHeader,
"-geoip2-asn", "asn",
"-trusted-header", trustedHeader,
"-trusted-port-header", trustedPortHeader,
},
headers: map[string][]string{ headers: map[string][]string{
trustedHeader: {testIP.ipv4}, trustedHeader: {testIP.ipv4},
trustedPortHeader: {"1001"}, trustedPortHeader: {"1001"},
@@ -153,19 +144,22 @@ func TestClientPort(t *testing.T) {
} }
for _, tt := range tests { for _, tt := range tests {
_, _ = setting.Setup(tt.args.params)
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
engine := gin.Default()
engine.TrustedPlatform = tt.args.trustedHeader
r := NewRouter(nil, tt.args.trustedHeader, tt.args.trustedPortHeader, "", false)
Setup(engine, r)
req, _ := http.NewRequest("GET", "/client-port", nil) req, _ := http.NewRequest("GET", "/client-port", nil)
req.RemoteAddr = net.JoinHostPort(testIP.ipv4, "1000") req.RemoteAddr = net.JoinHostPort(testIP.ipv4, "1000")
req.Header = tt.args.headers req.Header = tt.args.headers
w := httptest.NewRecorder() w := httptest.NewRecorder()
app.ServeHTTP(w, req) engine.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code) assert.Equal(t, 200, w.Code)
assert.Equal(t, contentType.text, w.Header().Get("Content-Type")) assert.Equal(t, contentType.text, w.Header().Get("Content-Type"))
assert.Equal(t, tt.expected, w.Body.String()) assert.Equal(t, tt.expected, w.Body.String())
t.Log(w.Header())
}) })
} }
} }
@@ -181,14 +175,11 @@ func TestNotFound(t *testing.T) {
} }
func TestJSON(t *testing.T) { func TestJSON(t *testing.T) {
_, _ = setting.Setup( svc, _ := service.NewGeo(context.Background(), "../test/GeoIP2-City-Test.mmdb", "../test/GeoLite2-ASN-Test.mmdb")
[]string{ engine := gin.Default()
"-geoip2-city", "city", engine.TrustedPlatform = trustedHeader
"-geoip2-asn", "asn", r := NewRouter(svc, trustedHeader, trustedPortHeader, "", false)
"-trusted-header", trustedHeader, Setup(engine, r)
"-trusted-port-header", trustedPortHeader,
},
)
type args struct { type args struct {
ip string ip string
@@ -222,7 +213,7 @@ func TestJSON(t *testing.T) {
req.Header.Set(trustedPortHeader, "1001") req.Header.Set(trustedPortHeader, "1001")
w := httptest.NewRecorder() w := httptest.NewRecorder()
app.ServeHTTP(w, req) engine.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code) assert.Equal(t, 200, w.Code)
assert.Equal(t, contentType.json, w.Header().Get("Content-Type")) assert.Equal(t, contentType.json, w.Header().Get("Content-Type"))
@@ -248,14 +239,11 @@ ASN Organization:
Header1: one Header1: one
Host: test Host: test
` `
_, _ = setting.Setup( svc, _ := service.NewGeo(context.Background(), "../test/GeoIP2-City-Test.mmdb", "../test/GeoLite2-ASN-Test.mmdb")
[]string{ engine := gin.Default()
"-geoip2-city", "city", engine.TrustedPlatform = trustedHeader
"-geoip2-asn", "asn", r := NewRouter(svc, trustedHeader, trustedPortHeader, "", false)
"-trusted-header", trustedHeader, Setup(engine, r)
"-trusted-port-header", trustedPortHeader,
},
)
req, _ := http.NewRequest("GET", "/all", nil) req, _ := http.NewRequest("GET", "/all", nil)
req.RemoteAddr = net.JoinHostPort(testIP.ipv4, "1000") req.RemoteAddr = net.JoinHostPort(testIP.ipv4, "1000")
@@ -265,7 +253,7 @@ Host: test
req.Header.Set("Header1", "one") req.Header.Set("Header1", "one")
w := httptest.NewRecorder() w := httptest.NewRecorder()
app.ServeHTTP(w, req) engine.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code) assert.Equal(t, 200, w.Code)
assert.Equal(t, contentType.text, w.Header().Get("Content-Type")) assert.Equal(t, contentType.text, w.Header().Get("Content-Type"))
+6 -6
View File
@@ -81,13 +81,13 @@ var asnOutput = map[string]asnDataFormatter{
}, },
} }
func getGeoAsString(ctx *gin.Context) { func (rt *Router) getGeoAsString(ctx *gin.Context) {
if geoSvc == nil { if rt.geo == nil {
ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound)) ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound))
return return
} }
record := geoSvc.LookUpCity(net.ParseIP(ctx.ClientIP())) record := rt.geo.LookUpCity(net.ParseIP(ctx.ClientIP()))
if record == nil { if record == nil {
ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound)) ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound))
return return
@@ -107,13 +107,13 @@ func getGeoAsString(ctx *gin.Context) {
ctx.String(http.StatusOK, g.format(record)) ctx.String(http.StatusOK, g.format(record))
} }
func getASNAsString(ctx *gin.Context) { func (rt *Router) getASNAsString(ctx *gin.Context) {
if geoSvc == nil { if rt.geo == nil {
ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound)) ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound))
return return
} }
record := geoSvc.LookUpASN(net.ParseIP(ctx.ClientIP())) record := rt.geo.LookUpASN(net.ParseIP(ctx.ClientIP()))
if record == nil { if record == nil {
ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound)) ctx.String(http.StatusNotFound, http.StatusText(http.StatusNotFound))
return return
+4 -5
View File
@@ -9,15 +9,14 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
func getHeadersAsSortedString(ctx *gin.Context) { func (rt *Router) getHeadersAsSortedString(ctx *gin.Context) {
h := httputils.GetHeadersWithoutTrustedHeaders(ctx) h := httputils.GetHeadersWithoutTrustedHeaders(ctx, rt.trustedHeader, rt.trustedPortHeader)
h.Set("Host", ctx.Request.Host) h.Set("Host", ctx.Request.Host)
ctx.String(http.StatusOK, httputils.HeadersToSortedString(h)) ctx.String(http.StatusOK, httputils.HeadersToSortedString(h))
} }
func getHeaderAsString(ctx *gin.Context) { func (rt *Router) getHeaderAsString(ctx *gin.Context) {
headers := httputils.GetHeadersWithoutTrustedHeaders(ctx) headers := httputils.GetHeadersWithoutTrustedHeaders(ctx, rt.trustedHeader, rt.trustedPortHeader)
h := ctx.Params.ByName("header") h := ctx.Params.ByName("header")
if v := headers.Get(ctx.Params.ByName("header")); v != "" { if v := headers.Get(ctx.Params.ByName("header")); v != "" {
+9 -8
View File
@@ -1,11 +1,13 @@
package router package router
import ( import (
"context"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"testing" "testing"
"github.com/dcarrillo/whatismyip/internal/setting" "github.com/dcarrillo/whatismyip/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
) )
@@ -27,12 +29,11 @@ Header2: value22
Header3: value3 Header3: value3
Host: Host:
` `
_, _ = setting.Setup([]string{ svc, _ := service.NewGeo(context.Background(), "../test/GeoIP2-City-Test.mmdb", "../test/GeoLite2-ASN-Test.mmdb")
"-geoip2-city", "city", engine := gin.Default()
"-geoip2-asn", "asn", engine.TrustedPlatform = trustedHeader
"-trusted-header", trustedHeader, r := NewRouter(svc, trustedHeader, trustedPortHeader, "", false)
"-trusted-port-header", trustedPortHeader, Setup(engine, r)
})
req, _ := http.NewRequest("GET", "/headers", nil) req, _ := http.NewRequest("GET", "/headers", nil)
req.Header = map[string][]string{ req.Header = map[string][]string{
"Header1": {"value1"}, "Header1": {"value1"},
@@ -43,7 +44,7 @@ Host:
req.Header.Set(trustedPortHeader, "1025") req.Header.Set(trustedPortHeader, "1025")
w := httptest.NewRecorder() w := httptest.NewRecorder()
app.ServeHTTP(w, req) engine.ServeHTTP(w, req)
assert.Equal(t, 200, w.Code) assert.Equal(t, 200, w.Code)
assert.Equal(t, contentType.text, w.Header().Get("Content-Type")) assert.Equal(t, contentType.text, w.Header().Get("Content-Type"))
+1 -1
View File
@@ -17,7 +17,7 @@ type JSONScanResponse struct {
Reason string `json:"reason"` Reason string `json:"reason"`
} }
func scanTCPPort(ctx *gin.Context) { func (rt *Router) scanTCPPort(ctx *gin.Context) {
port, err := strconv.Atoi(ctx.Params.ByName("port")) port, err := strconv.Atoi(ctx.Params.ByName("port"))
if err == nil && (port < 1 || port > 65535) { if err == nil && (port < 1 || port > 65535) {
err = fmt.Errorf("%d is not a valid port number", port) err = fmt.Errorf("%d is not a valid port number", port)
+34 -20
View File
@@ -4,35 +4,49 @@ import (
"html/template" "html/template"
"log" "log"
"github.com/dcarrillo/whatismyip/internal/setting"
"github.com/dcarrillo/whatismyip/service" "github.com/dcarrillo/whatismyip/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
var geoSvc *service.Geo type Router struct {
geo *service.Geo
trustedHeader string
trustedPortHeader string
templatePath string
disableScan bool
}
func SetupTemplate(r *gin.Engine) { func NewRouter(geo *service.Geo, trustedHeader, trustedPortHeader, templatePath string, disableScan bool) *Router {
if setting.App.TemplatePath == "" { return &Router{
geo: geo,
trustedHeader: trustedHeader,
trustedPortHeader: trustedPortHeader,
templatePath: templatePath,
disableScan: disableScan,
}
}
func SetupTemplate(r *gin.Engine, templatePath string) {
if templatePath == "" {
r.SetHTMLTemplate(template.Must(template.New("home").Parse(home))) r.SetHTMLTemplate(template.Must(template.New("home").Parse(home)))
} else { } else {
log.Printf("Template %s has been loaded", setting.App.TemplatePath) log.Printf("Template %s has been loaded", templatePath)
r.LoadHTMLFiles(setting.App.TemplatePath) r.LoadHTMLFiles(templatePath)
} }
} }
func Setup(r *gin.Engine, geo *service.Geo) { func Setup(r *gin.Engine, rt *Router) {
geoSvc = geo r.GET("/", rt.getRoot)
r.GET("/", getRoot) if !rt.disableScan {
if !setting.App.DisableTCPScan { r.GET("/scan/tcp/:port", rt.scanTCPPort)
r.GET("/scan/tcp/:port", scanTCPPort)
} }
r.GET("/client-port", getClientPortAsString) r.GET("/client-port", rt.getClientPortAsString)
r.GET("/geo", getGeoAsString) r.GET("/geo", rt.getGeoAsString)
r.GET("/geo/:field", getGeoAsString) r.GET("/geo/:field", rt.getGeoAsString)
r.GET("/asn", getASNAsString) r.GET("/asn", rt.getASNAsString)
r.GET("/asn/:field", getASNAsString) r.GET("/asn/:field", rt.getASNAsString)
r.GET("/headers", getHeadersAsSortedString) r.GET("/headers", rt.getHeadersAsSortedString)
r.GET("/all", getAllAsString) r.GET("/all", rt.getAllAsString)
r.GET("/json", getJSON) r.GET("/json", rt.getJSON)
r.GET("/:header", getHeaderAsString) r.GET("/:header", rt.getHeaderAsString)
} }
+4 -1
View File
@@ -47,11 +47,14 @@ const (
domain = "dns.example.com" domain = "dns.example.com"
) )
var rt *Router
func TestMain(m *testing.M) { func TestMain(m *testing.M) {
app = gin.Default() app = gin.Default()
app.TrustedPlatform = trustedHeader app.TrustedPlatform = trustedHeader
svc, _ := service.NewGeo(context.Background(), "../test/GeoIP2-City-Test.mmdb", "../test/GeoLite2-ASN-Test.mmdb") svc, _ := service.NewGeo(context.Background(), "../test/GeoIP2-City-Test.mmdb", "../test/GeoLite2-ASN-Test.mmdb")
Setup(app, svc) rt = NewRouter(svc, trustedHeader, "", "", false)
Setup(app, rt)
os.Exit(m.Run()) os.Exit(m.Run())
} }
+9 -6
View File
@@ -6,18 +6,21 @@ import (
"log" "log"
"net/http" "net/http"
"github.com/dcarrillo/whatismyip/internal/setting"
"github.com/prometheus/client_golang/prometheus/promhttp" "github.com/prometheus/client_golang/prometheus/promhttp"
) )
type Prometheus struct { type Prometheus struct {
server *http.Server server *http.Server
ctx context.Context ctx context.Context
addr string
timeouts Timeouts
} }
func NewPrometheusServer(ctx context.Context) *Prometheus { func NewPrometheusServer(ctx context.Context, addr string, timeouts Timeouts) *Prometheus {
return &Prometheus{ return &Prometheus{
ctx: ctx, ctx: ctx,
addr: addr,
timeouts: timeouts,
} }
} }
@@ -26,13 +29,13 @@ func (p *Prometheus) Start() {
mux.Handle("/metrics", promhttp.Handler()) mux.Handle("/metrics", promhttp.Handler())
p.server = &http.Server{ p.server = &http.Server{
Addr: setting.App.PrometheusAddress, Addr: p.addr,
Handler: mux, Handler: mux,
ReadTimeout: setting.App.Server.ReadTimeout, ReadTimeout: p.timeouts.ReadTimeout,
WriteTimeout: setting.App.Server.WriteTimeout, WriteTimeout: p.timeouts.WriteTimeout,
} }
log.Printf("Starting Prometheus server listening on %s", setting.App.PrometheusAddress) log.Printf("Starting Prometheus server listening on %s", p.addr)
go func() { go func() {
if err := p.server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { if err := p.server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Fatal(err) log.Fatal(err)
+10 -5
View File
@@ -6,7 +6,6 @@ import (
"log" "log"
"net/http" "net/http"
"github.com/dcarrillo/whatismyip/internal/setting"
"github.com/quic-go/quic-go/http3" "github.com/quic-go/quic-go/http3"
) )
@@ -14,18 +13,24 @@ type Quic struct {
server *http3.Server server *http3.Server
tlsServer *TLS tlsServer *TLS
ctx context.Context ctx context.Context
addr string
crtPath string
keyPath string
} }
func NewQuicServer(ctx context.Context, tlsServer *TLS) *Quic { func NewQuicServer(ctx context.Context, tlsServer *TLS, addr, crt, key string) *Quic {
return &Quic{ return &Quic{
tlsServer: tlsServer, tlsServer: tlsServer,
ctx: ctx, ctx: ctx,
addr: addr,
crtPath: crt,
keyPath: key,
} }
} }
func (q *Quic) Start() { func (q *Quic) Start() {
q.server = &http3.Server{ q.server = &http3.Server{
Addr: setting.App.TLSAddress, Addr: q.addr,
Handler: q.tlsServer.server.Handler, Handler: q.tlsServer.server.Handler,
} }
@@ -38,9 +43,9 @@ func (q *Quic) Start() {
parentHandler.ServeHTTP(rw, req) parentHandler.ServeHTTP(rw, req)
}) })
log.Printf("Starting QUIC server listening on %s (udp)", setting.App.TLSAddress) log.Printf("Starting QUIC server listening on %s (udp)", q.addr)
go func() { go func() {
if err := q.server.ListenAndServeTLS(setting.App.TLSCrtPath, setting.App.TLSKeyPath); err != nil && if err := q.server.ListenAndServeTLS(q.crtPath, q.keyPath); err != nil &&
!errors.Is(err, http.ErrServerClosed) { !errors.Is(err, http.ErrServerClosed) {
log.Fatal(err) log.Fatal(err)
} }
+15 -7
View File
@@ -5,32 +5,40 @@ import (
"errors" "errors"
"log" "log"
"net/http" "net/http"
"time"
"github.com/dcarrillo/whatismyip/internal/setting"
) )
type Timeouts struct {
ReadTimeout time.Duration
WriteTimeout time.Duration
}
type TCP struct { type TCP struct {
server *http.Server server *http.Server
handler http.Handler handler http.Handler
ctx context.Context ctx context.Context
addr string
timeouts Timeouts
} }
func NewTCPServer(ctx context.Context, handler http.Handler) *TCP { func NewTCPServer(ctx context.Context, handler http.Handler, addr string, timeouts Timeouts) *TCP {
return &TCP{ return &TCP{
handler: handler, handler: handler,
ctx: ctx, ctx: ctx,
addr: addr,
timeouts: timeouts,
} }
} }
func (t *TCP) Start() { func (t *TCP) Start() {
t.server = &http.Server{ t.server = &http.Server{
Addr: setting.App.BindAddress, Addr: t.addr,
Handler: t.handler, Handler: t.handler,
ReadTimeout: setting.App.Server.ReadTimeout, ReadTimeout: t.timeouts.ReadTimeout,
WriteTimeout: setting.App.Server.WriteTimeout, WriteTimeout: t.timeouts.WriteTimeout,
} }
log.Printf("Starting TCP server listening on %s", setting.App.BindAddress) log.Printf("Starting TCP server listening on %s", t.addr)
go func() { go func() {
if err := t.server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { if err := t.server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Fatal(err) log.Fatal(err)
+14 -8
View File
@@ -6,37 +6,43 @@ import (
"errors" "errors"
"log" "log"
"net/http" "net/http"
"github.com/dcarrillo/whatismyip/internal/setting"
) )
type TLS struct { type TLS struct {
server *http.Server server *http.Server
handler http.Handler handler http.Handler
ctx context.Context ctx context.Context
addr string
crtPath string
keyPath string
timeouts Timeouts
} }
func NewTLSServer(ctx context.Context, handler http.Handler) *TLS { func NewTLSServer(ctx context.Context, handler http.Handler, addr, crt, key string, timeouts Timeouts) *TLS {
return &TLS{ return &TLS{
handler: handler, handler: handler,
ctx: ctx, ctx: ctx,
addr: addr,
crtPath: crt,
keyPath: key,
timeouts: timeouts,
} }
} }
func (t *TLS) Start() { func (t *TLS) Start() {
t.server = &http.Server{ t.server = &http.Server{
Addr: setting.App.TLSAddress, Addr: t.addr,
Handler: t.handler, Handler: t.handler,
ReadTimeout: setting.App.Server.ReadTimeout, ReadTimeout: t.timeouts.ReadTimeout,
WriteTimeout: setting.App.Server.WriteTimeout, WriteTimeout: t.timeouts.WriteTimeout,
TLSConfig: &tls.Config{ TLSConfig: &tls.Config{
MinVersion: tls.VersionTLS12, MinVersion: tls.VersionTLS12,
}, },
} }
log.Printf("Starting TLS server listening on %s", setting.App.TLSAddress) log.Printf("Starting TLS server listening on %s", t.addr)
go func() { go func() {
if err := t.server.ListenAndServeTLS(setting.App.TLSCrtPath, setting.App.TLSKeyPath); err != nil && if err := t.server.ListenAndServeTLS(t.crtPath, t.keyPath); err != nil &&
!errors.Is(err, http.ErrServerClosed) { !errors.Is(err, http.ErrServerClosed) {
log.Fatal(err) log.Fatal(err)
} }