mirror of
https://github.com/dcarrillo/whatismyip.git
synced 2026-07-23 22:45:46 +00:00
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:
+49
-34
@@ -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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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) {
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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 != "" {
|
||||||
|
|||||||
@@ -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"))
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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())
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user