Refactoring (#325)
* Refactor core * Re-added tests * Small fixes * Add tests for acmetxt cidrslice and util funcs * Remove the last dangling reference to old logging package * Refactoring (#327) * chore: enable more linters and fix linter issues * ci: enable linter checks on all branches and disable recurring checks recurring linter checks don't make that much sense. The code & linter checks should not change on their own over night ;) * chore: update packages * Revert "chore: update packages" This reverts commit 30250bf28c4b39e9e5b3af012a4e28ab036bf9af. * chore: manually upgrade some packages * Updated dependencies, wrote changelog entry and fixed namespace for release * Refactoring - improving coverage (#371) * Increase code coverage in acmedns * More testing of ReadConfig() and its fallback mechanism * Found that if someone put a '"' double quote into the filename that we configure zap to log to, it would cause the the JSON created to be invalid. I have replaced the JSON string with proper config * Better handling of config options for api.TLS - we now error on an invalid value instead of silently failing. added a basic test for api.setupTLS() (to increase test coverage) * testing nameserver isOwnChallenge and isAuthoritative methods * add a unit test for nameserver answerOwnChallenge * fix linting errors * bump go and golangci-lint versions in github actions * Update golangci-lint.yml Bumping github-actions workflow versions to accommodate some changes in upstream golanci-lint * Bump Golang version to 1.23 (currently the oldest supported version) Bump golanglint-ci to 2.0.2 and migrate the config file. This should resolve the math/rand/v2 issue * bump golanglint-ci action version * Fixing up new golanglint-ci warnings and errors --------- Co-authored-by: Joona Hoikkala <5235109+joohoi@users.noreply.github.com> * Minor refactoring, error returns and e2e testing suite * Add a few tests * Fix linter and umask setting * Update github actions * Refine concurrency configuration for GitHub actions * HTTP timeouts to API, and self-validation mutex to nameserver ops --------- Co-authored-by: Florian Ritterhoff <32478819+fritterhoff@users.noreply.github.com> Co-authored-by: Jason Playne <jason@jasonplayne.com>
This commit is contained in:
co-authored by
Florian Ritterhoff
Jason Playne
parent
b7a0a8a7bc
commit
5a7bc230b8
@@ -0,0 +1,47 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"net"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// AllowedFrom Check if IP belongs to an allowed net
|
||||
func (a ACMETxt) AllowedFrom(ip string) bool {
|
||||
remoteIP := net.ParseIP(ip)
|
||||
// Range not limited
|
||||
if len(a.AllowFrom.ValidEntries()) == 0 {
|
||||
return true
|
||||
}
|
||||
for _, v := range a.AllowFrom.ValidEntries() {
|
||||
_, vnet, _ := net.ParseCIDR(v)
|
||||
if vnet.Contains(remoteIP) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// AllowedFromList Go through list (most likely from headers) to check for the IP.
|
||||
// Reason for this is that some setups use reverse proxy in front of acme-dns
|
||||
func (a ACMETxt) AllowedFromList(ips []string) bool {
|
||||
if len(ips) == 0 {
|
||||
// If no IP provided, check if no whitelist present (everyone has access)
|
||||
return a.AllowedFrom("")
|
||||
}
|
||||
for _, v := range ips {
|
||||
if a.AllowedFrom(v) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func NewACMETxt() ACMETxt {
|
||||
var a = ACMETxt{}
|
||||
password := generatePassword(40)
|
||||
a.Username = uuid.New()
|
||||
a.Password = password
|
||||
a.Subdomain = uuid.New().String()
|
||||
return a
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package acmedns
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestAllowedFrom(t *testing.T) {
|
||||
testslice := NewACMETxt()
|
||||
testslice.AllowFrom = []string{"192.168.1.0/24", "2001:db8::/32"}
|
||||
for _, test := range []struct {
|
||||
input string
|
||||
expected bool
|
||||
}{
|
||||
{"192.168.1.42", true},
|
||||
{"192.168.2.42", false},
|
||||
{"2001:db8:aaaa::", true},
|
||||
{"2001:db9:aaaa::", false},
|
||||
} {
|
||||
if testslice.AllowedFrom(test.input) != test.expected {
|
||||
t.Errorf("Was expecting AllowedFrom to return %t for %s but got %t instead.", test.expected, test.input, !test.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllowedFromList(t *testing.T) {
|
||||
testslice := ACMETxt{AllowFrom: []string{"192.168.1.0/24", "2001:db8::/32"}}
|
||||
if testslice.AllowedFromList([]string{"192.168.2.2", "1.1.1.1"}) != false {
|
||||
t.Errorf("Was expecting AllowedFromList to return false")
|
||||
}
|
||||
if testslice.AllowedFromList([]string{"192.168.1.2", "1.1.1.1"}) != true {
|
||||
t.Errorf("Was expecting AllowedFromList to return true")
|
||||
}
|
||||
allowfromall := ACMETxt{AllowFrom: []string{}}
|
||||
if allowfromall.AllowedFromList([]string{"192.168.1.2", "1.1.1.1"}) != true {
|
||||
t.Errorf("Expected non-restricted AlloFrom to be allowed")
|
||||
}
|
||||
if allowfromall.AllowedFromList([]string{}) != true {
|
||||
t.Errorf("Expected non-restricted AlloFrom to be allowed for empty list")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net"
|
||||
)
|
||||
|
||||
// cidrslice is a list of allowed cidr ranges
|
||||
type Cidrslice []string
|
||||
|
||||
func (c *Cidrslice) JSON() string {
|
||||
ret, _ := json.Marshal(c.ValidEntries())
|
||||
return string(ret)
|
||||
}
|
||||
|
||||
func (c *Cidrslice) IsValid() error {
|
||||
for _, v := range *c {
|
||||
_, _, err := net.ParseCIDR(sanitizeIPv6addr(v))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Cidrslice) ValidEntries() []string {
|
||||
valid := []string{}
|
||||
for _, v := range *c {
|
||||
_, _, err := net.ParseCIDR(sanitizeIPv6addr(v))
|
||||
if err == nil {
|
||||
valid = append(valid, sanitizeIPv6addr(v))
|
||||
}
|
||||
}
|
||||
return valid
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCidrSlice(t *testing.T) {
|
||||
for i, test := range []struct {
|
||||
input Cidrslice
|
||||
expectedErr bool
|
||||
expectedLen int
|
||||
}{
|
||||
{[]string{"192.168.1.0/24"}, false, 1},
|
||||
{[]string{"shoulderror"}, true, 0},
|
||||
{[]string{"2001:db8:aaaaa::"}, true, 0},
|
||||
{[]string{"192.168.1.0/24", "2001:db8::/32"}, false, 2},
|
||||
} {
|
||||
err := test.input.IsValid()
|
||||
if test.expectedErr && err == nil {
|
||||
t.Errorf("Expected test %d to generate IsValid() error but it didn't", i)
|
||||
}
|
||||
if !test.expectedErr && err != nil {
|
||||
t.Errorf("Expected test %d to pass IsValid() but it generated an error %s", i, err)
|
||||
}
|
||||
outSlice := []string{}
|
||||
err = json.Unmarshal([]byte(test.input.JSON()), &outSlice)
|
||||
if err != nil {
|
||||
t.Errorf("Unexpected error when unmarshaling Cidrslice JSON: %s", err)
|
||||
}
|
||||
if len(outSlice) != test.expectedLen {
|
||||
t.Errorf("Expected cidrslice JSON to be of length %d, but got %d instead for test %d", test.expectedLen, len(outSlice), i)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/BurntSushi/toml"
|
||||
)
|
||||
|
||||
const (
|
||||
ApiTlsProviderNone = "none"
|
||||
ApiTlsProviderLetsEncrypt = "letsencrypt"
|
||||
ApiTlsProviderLetsEncryptStaging = "letsencryptstaging"
|
||||
ApiTlsProviderCert = "cert"
|
||||
)
|
||||
|
||||
func FileIsAccessible(fname string) bool {
|
||||
_, err := os.Stat(fname)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
f, err := os.Open(fname)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
f.Close()
|
||||
return true
|
||||
}
|
||||
|
||||
func readTomlConfig(fname string) (AcmeDnsConfig, error) {
|
||||
var conf AcmeDnsConfig
|
||||
_, err := toml.DecodeFile(fname, &conf)
|
||||
if err != nil {
|
||||
// Return with config file parsing errors from toml package
|
||||
return conf, err
|
||||
}
|
||||
return prepareConfig(conf)
|
||||
}
|
||||
|
||||
// prepareConfig checks that mandatory values exist, and can be used to set default values in the future
|
||||
func prepareConfig(conf AcmeDnsConfig) (AcmeDnsConfig, error) {
|
||||
if conf.Database.Engine == "" {
|
||||
return conf, errors.New("missing database configuration option \"engine\"")
|
||||
}
|
||||
if conf.Database.Connection == "" {
|
||||
return conf, errors.New("missing database configuration option \"connection\"")
|
||||
}
|
||||
|
||||
// Default values for options added to config to keep backwards compatibility with old config
|
||||
if conf.API.ACMECacheDir == "" {
|
||||
conf.API.ACMECacheDir = "api-certs"
|
||||
}
|
||||
|
||||
switch conf.API.TLS {
|
||||
case ApiTlsProviderCert, ApiTlsProviderLetsEncrypt, ApiTlsProviderLetsEncryptStaging, ApiTlsProviderNone:
|
||||
// we have a good value
|
||||
default:
|
||||
return conf, fmt.Errorf("invalid value for api.tls, expected one of [%s, %s, %s, %s]", ApiTlsProviderCert, ApiTlsProviderLetsEncrypt, ApiTlsProviderLetsEncryptStaging, ApiTlsProviderNone)
|
||||
}
|
||||
|
||||
return conf, nil
|
||||
}
|
||||
|
||||
func ReadConfig(configFile, fallback string) (AcmeDnsConfig, string, error) {
|
||||
var usedConfigFile string
|
||||
var config AcmeDnsConfig
|
||||
var err error
|
||||
if FileIsAccessible(configFile) {
|
||||
usedConfigFile = configFile
|
||||
config, err = readTomlConfig(configFile)
|
||||
} else if FileIsAccessible(fallback) {
|
||||
usedConfigFile = fallback
|
||||
config, err = readTomlConfig(fallback)
|
||||
} else {
|
||||
err = fmt.Errorf("configuration file not found")
|
||||
}
|
||||
if err != nil {
|
||||
err = fmt.Errorf("encountered an error while trying to read configuration file: %w", err)
|
||||
}
|
||||
return config, usedConfigFile, err
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
type AcmednsDB interface {
|
||||
Register(cidrslice Cidrslice) (ACMETxt, error)
|
||||
GetByUsername(uuid.UUID) (ACMETxt, error)
|
||||
GetTXTForDomain(string) ([]string, error)
|
||||
Update(ACMETxtPost) error
|
||||
GetBackend() *sql.DB
|
||||
SetBackend(*sql.DB)
|
||||
Close()
|
||||
}
|
||||
|
||||
type AcmednsNS interface {
|
||||
Start(errorChannel chan error)
|
||||
SetOwnAuthKey(key string)
|
||||
SetNotifyStartedFunc(func())
|
||||
ParseRecords()
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"go.uber.org/zap/zapcore"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func SetupLogging(config AcmeDnsConfig) (*zap.Logger, error) {
|
||||
var (
|
||||
logger *zap.Logger
|
||||
zapCfg zap.Config
|
||||
err error
|
||||
)
|
||||
|
||||
logformat := "console"
|
||||
if config.Logconfig.Format == "json" {
|
||||
logformat = "json"
|
||||
}
|
||||
outputPath := "stdout"
|
||||
if config.Logconfig.Logtype == "file" {
|
||||
outputPath = config.Logconfig.File
|
||||
}
|
||||
errorPath := "stderr"
|
||||
if config.Logconfig.Logtype == "file" {
|
||||
errorPath = config.Logconfig.File
|
||||
}
|
||||
|
||||
zapCfg.Level, err = zap.ParseAtomicLevel(config.Logconfig.Level)
|
||||
if err != nil {
|
||||
return logger, err
|
||||
}
|
||||
zapCfg.Encoding = logformat
|
||||
zapCfg.OutputPaths = []string{outputPath}
|
||||
zapCfg.ErrorOutputPaths = []string{errorPath}
|
||||
zapCfg.EncoderConfig = zapcore.EncoderConfig{
|
||||
TimeKey: "time",
|
||||
MessageKey: "msg",
|
||||
LevelKey: "level",
|
||||
EncodeLevel: zapcore.LowercaseLevelEncoder,
|
||||
EncodeTime: zapcore.ISO8601TimeEncoder,
|
||||
}
|
||||
|
||||
logger, err = zapCfg.Build()
|
||||
return logger, err
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
[general]
|
||||
listen = "127.0.0.1:53"
|
||||
protocol = "both"
|
||||
domain = "test.example.org"
|
||||
nsname = "test.example.org"
|
||||
nsadmin = "test.example.org"
|
||||
records = [
|
||||
"test.example.org. A 127.0.0.1",
|
||||
"test.example.org. NS test.example.org.",
|
||||
]
|
||||
debug = true
|
||||
|
||||
[database]
|
||||
engine = "dinosaur"
|
||||
connection = "roar"
|
||||
|
||||
[api]
|
||||
ip = "0.0.0.0"
|
||||
disable_registration = false
|
||||
port = "443"
|
||||
tls = "none"
|
||||
tls_cert_privkey = "/etc/tls/example.org/privkey.pem"
|
||||
tls_cert_fullchain = "/etc/tls/example.org/fullchain.pem"
|
||||
acme_cache_dir = "api-certs"
|
||||
notification_email = ""
|
||||
corsorigins = [
|
||||
"*"
|
||||
]
|
||||
use_header = true
|
||||
header_name = "X-is-gonna-give-it-to-ya"
|
||||
|
||||
[logconfig]
|
||||
loglevel = "info"
|
||||
logtype = "stdout"
|
||||
logfile = "./acme-dns.log"
|
||||
logformat = "json"
|
||||
@@ -0,0 +1,72 @@
|
||||
package acmedns
|
||||
|
||||
import "github.com/google/uuid"
|
||||
|
||||
type Account struct {
|
||||
Username string
|
||||
Password string
|
||||
Subdomain string
|
||||
}
|
||||
|
||||
// AcmeDnsConfig holds the config structure
|
||||
type AcmeDnsConfig struct {
|
||||
General general
|
||||
Database dbsettings
|
||||
API httpapi
|
||||
Logconfig logconfig
|
||||
}
|
||||
|
||||
// Config file general section
|
||||
type general struct {
|
||||
Listen string
|
||||
Proto string `toml:"protocol"`
|
||||
Domain string
|
||||
Nsname string
|
||||
Nsadmin string
|
||||
Debug bool
|
||||
StaticRecords []string `toml:"records"`
|
||||
}
|
||||
|
||||
type dbsettings struct {
|
||||
Engine string
|
||||
Connection string
|
||||
}
|
||||
|
||||
// API config
|
||||
type httpapi struct {
|
||||
Domain string `toml:"api_domain"`
|
||||
IP string
|
||||
DisableRegistration bool `toml:"disable_registration"`
|
||||
AutocertPort string `toml:"autocert_port"`
|
||||
Port string `toml:"port"`
|
||||
TLS string
|
||||
TLSCertPrivkey string `toml:"tls_cert_privkey"`
|
||||
TLSCertFullchain string `toml:"tls_cert_fullchain"`
|
||||
ACMECacheDir string `toml:"acme_cache_dir"`
|
||||
NotificationEmail string `toml:"notification_email"`
|
||||
CorsOrigins []string
|
||||
UseHeader bool `toml:"use_header"`
|
||||
HeaderName string `toml:"header_name"`
|
||||
}
|
||||
|
||||
// Logging config
|
||||
type logconfig struct {
|
||||
Level string `toml:"loglevel"`
|
||||
Logtype string `toml:"logtype"`
|
||||
File string `toml:"logfile"`
|
||||
Format string `toml:"logformat"`
|
||||
}
|
||||
|
||||
// ACMETxt is the default structure for the user controlled record
|
||||
type ACMETxt struct {
|
||||
Username uuid.UUID
|
||||
Password string
|
||||
ACMETxtPost
|
||||
AllowFrom Cidrslice
|
||||
}
|
||||
|
||||
// ACMETxtPost holds the DNS part of the ACMETxt struct
|
||||
type ACMETxtPost struct {
|
||||
Subdomain string `json:"subdomain"`
|
||||
Value string `json:"txt"`
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"math/big"
|
||||
"regexp"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func sanitizeIPv6addr(s string) string {
|
||||
// Remove brackets from IPv6 addresses, net.ParseCIDR needs this
|
||||
re, _ := regexp.Compile(`[\[\]]+`)
|
||||
return re.ReplaceAllString(s, "")
|
||||
}
|
||||
|
||||
func SanitizeString(s string) string {
|
||||
// URL safe base64 alphabet without padding as defined in ACME
|
||||
re, _ := regexp.Compile(`[^A-Za-z\-\_0-9]+`)
|
||||
return re.ReplaceAllString(s, "")
|
||||
}
|
||||
|
||||
func generatePassword(length int) string {
|
||||
ret := make([]byte, length)
|
||||
const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz1234567890-_"
|
||||
alphalen := big.NewInt(int64(len(alphabet)))
|
||||
for i := 0; i < length; i++ {
|
||||
c, _ := rand.Int(rand.Reader, alphalen)
|
||||
r := int(c.Int64())
|
||||
ret[i] = alphabet[r]
|
||||
}
|
||||
return string(ret)
|
||||
}
|
||||
|
||||
func CorrectPassword(pw string, hash string) bool {
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(pw)); err == nil {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,326 @@
|
||||
package acmedns
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand/v2"
|
||||
"os"
|
||||
"reflect"
|
||||
"syscall"
|
||||
"testing"
|
||||
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func fakeConfig() AcmeDnsConfig {
|
||||
conf := AcmeDnsConfig{}
|
||||
conf.Logconfig.Logtype = "stdout"
|
||||
return conf
|
||||
}
|
||||
|
||||
func TestSetupLogging(t *testing.T) {
|
||||
conf := fakeConfig()
|
||||
for i, test := range []struct {
|
||||
format string
|
||||
level string
|
||||
expected zapcore.Level
|
||||
}{
|
||||
{"text", "warn", zap.WarnLevel},
|
||||
{"json", "debug", zap.DebugLevel},
|
||||
{"text", "info", zap.InfoLevel},
|
||||
{"json", "error", zap.ErrorLevel},
|
||||
} {
|
||||
conf.Logconfig.Format = test.format
|
||||
conf.Logconfig.Level = test.level
|
||||
logger, err := SetupLogging(conf)
|
||||
if err != nil {
|
||||
t.Errorf("Got unexpected error: %s", err)
|
||||
} else {
|
||||
if logger.Sugar().Level() != test.expected {
|
||||
t.Errorf("Test %d: Expected loglevel %s but got %s", i, test.expected, logger.Sugar().Level())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetupLoggingError(t *testing.T) {
|
||||
conf := fakeConfig()
|
||||
for _, test := range []struct {
|
||||
format string
|
||||
level string
|
||||
file string
|
||||
errexpected bool
|
||||
}{
|
||||
{"text", "warn", "", false},
|
||||
{"json", "debug", "", false},
|
||||
{"text", "info", "", false},
|
||||
{"json", "error", "", false},
|
||||
{"text", "something", "", true},
|
||||
{"text", "info", "a path with\" in its name.txt", false},
|
||||
} {
|
||||
conf.Logconfig.Format = test.format
|
||||
conf.Logconfig.Level = test.level
|
||||
if test.file != "" {
|
||||
conf.Logconfig.File = test.file
|
||||
conf.Logconfig.Logtype = "file"
|
||||
|
||||
}
|
||||
_, err := SetupLogging(conf)
|
||||
if test.errexpected && err == nil {
|
||||
t.Errorf("Expected error but did not get one for loglevel: %s", err)
|
||||
} else if !test.errexpected && err != nil {
|
||||
t.Errorf("Unexpected error: %s", err)
|
||||
}
|
||||
|
||||
// clean up the file zap creates
|
||||
if test.file != "" {
|
||||
_ = os.Remove(test.file)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadConfig(t *testing.T) {
|
||||
for i, test := range []struct {
|
||||
inFile []byte
|
||||
output AcmeDnsConfig
|
||||
}{
|
||||
{
|
||||
[]byte("[general]\nlisten = \":53\"\ndebug = true\n[api]\napi_domain = \"something.strange\""),
|
||||
AcmeDnsConfig{
|
||||
General: general{
|
||||
Listen: ":53",
|
||||
Debug: true,
|
||||
},
|
||||
API: httpapi{
|
||||
Domain: "something.strange",
|
||||
},
|
||||
},
|
||||
},
|
||||
|
||||
{
|
||||
[]byte("[\x00[[[[[[[[[de\nlisten =]"),
|
||||
AcmeDnsConfig{},
|
||||
},
|
||||
} {
|
||||
tmpfile, err := os.CreateTemp("", "acmedns")
|
||||
if err != nil {
|
||||
t.Fatalf("Could not create temporary file: %s", err)
|
||||
}
|
||||
defer os.Remove(tmpfile.Name())
|
||||
|
||||
if _, err := tmpfile.Write(test.inFile); err != nil {
|
||||
t.Error("Could not write to temporary file")
|
||||
}
|
||||
|
||||
if err := tmpfile.Close(); err != nil {
|
||||
t.Error("Could not close temporary file")
|
||||
}
|
||||
ret, _, _ := ReadConfig(tmpfile.Name(), "")
|
||||
if ret.General.Listen != test.output.General.Listen {
|
||||
t.Errorf("Test %d: Expected listen value %s, but got %s", i, test.output.General.Listen, ret.General.Listen)
|
||||
}
|
||||
if ret.API.Domain != test.output.API.Domain {
|
||||
t.Errorf("Test %d: Expected HTTP API domain %s, but got %s", i, test.output.API.Domain, ret.API.Domain)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadConfigFallback(t *testing.T) {
|
||||
var (
|
||||
path string
|
||||
err error
|
||||
)
|
||||
|
||||
testPath := "testdata/test_read_fallback_config.toml"
|
||||
|
||||
path, err = getNonExistentPath()
|
||||
if err != nil {
|
||||
t.Errorf("failed getting non existant path: %s", err)
|
||||
}
|
||||
|
||||
cfg, used, err := ReadConfig(path, testPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read a config file when we should have: %s", err)
|
||||
}
|
||||
|
||||
if used != testPath {
|
||||
t.Fatalf("we read from the wrong file. got: %s, want: %s", used, testPath)
|
||||
}
|
||||
|
||||
expected := AcmeDnsConfig{
|
||||
General: general{
|
||||
Listen: "127.0.0.1:53",
|
||||
Proto: "both",
|
||||
Domain: "test.example.org",
|
||||
Nsname: "test.example.org",
|
||||
Nsadmin: "test.example.org",
|
||||
Debug: true,
|
||||
StaticRecords: []string{
|
||||
"test.example.org. A 127.0.0.1",
|
||||
"test.example.org. NS test.example.org.",
|
||||
},
|
||||
},
|
||||
Database: dbsettings{
|
||||
Engine: "dinosaur",
|
||||
Connection: "roar",
|
||||
},
|
||||
API: httpapi{
|
||||
Domain: "",
|
||||
IP: "0.0.0.0",
|
||||
DisableRegistration: false,
|
||||
AutocertPort: "",
|
||||
Port: "443",
|
||||
TLS: "none",
|
||||
TLSCertPrivkey: "/etc/tls/example.org/privkey.pem",
|
||||
TLSCertFullchain: "/etc/tls/example.org/fullchain.pem",
|
||||
ACMECacheDir: "api-certs",
|
||||
NotificationEmail: "",
|
||||
CorsOrigins: []string{"*"},
|
||||
UseHeader: true,
|
||||
HeaderName: "X-is-gonna-give-it-to-ya",
|
||||
},
|
||||
Logconfig: logconfig{
|
||||
Level: "info",
|
||||
Logtype: "stdout",
|
||||
File: "./acme-dns.log",
|
||||
Format: "json",
|
||||
},
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(cfg, expected) {
|
||||
t.Errorf("Did not read the config correctly: got %+v, want: %+v", cfg, expected)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func getNonExistentPath() (string, error) {
|
||||
path := fmt.Sprintf("/some/path/that/should/not/exist/on/any/filesystem/%10d.cfg", rand.Int())
|
||||
|
||||
if _, err := os.Stat(path); os.IsNotExist(err) {
|
||||
return path, nil
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("attempted non existant file exists!?: %s", path)
|
||||
}
|
||||
|
||||
// TestReadConfigFallbackError makes sure we error when we do not have a fallback config file
|
||||
func TestReadConfigFallbackError(t *testing.T) {
|
||||
var (
|
||||
badPaths []string
|
||||
i int
|
||||
)
|
||||
for len(badPaths) < 2 && i < 10 {
|
||||
i++
|
||||
|
||||
if path, err := getNonExistentPath(); err == nil {
|
||||
badPaths = append(badPaths, path)
|
||||
}
|
||||
}
|
||||
|
||||
if len(badPaths) != 2 {
|
||||
t.Fatalf("did not create exactly 2 bad paths")
|
||||
}
|
||||
|
||||
_, _, err := ReadConfig(badPaths[0], badPaths[1])
|
||||
if err == nil {
|
||||
t.Errorf("Should have failed reading non existant file: %s", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileCheckPermissionDenied(t *testing.T) {
|
||||
tmpfile, err := os.CreateTemp("", "acmedns")
|
||||
if err != nil {
|
||||
t.Fatalf("Could not create temporary file: %s", err)
|
||||
}
|
||||
defer os.Remove(tmpfile.Name())
|
||||
_ = syscall.Chmod(tmpfile.Name(), 0000)
|
||||
if FileIsAccessible(tmpfile.Name()) {
|
||||
t.Errorf("File should not be accessible")
|
||||
}
|
||||
_ = syscall.Chmod(tmpfile.Name(), 0644)
|
||||
}
|
||||
|
||||
func TestFileCheckNotExists(t *testing.T) {
|
||||
if FileIsAccessible("/path/that/does/not/exist") {
|
||||
t.Errorf("File should not be accessible")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileCheckOK(t *testing.T) {
|
||||
tmpfile, err := os.CreateTemp("", "acmedns")
|
||||
if err != nil {
|
||||
t.Fatalf("Could not create temporary file: %s", err)
|
||||
}
|
||||
defer os.Remove(tmpfile.Name())
|
||||
if !FileIsAccessible(tmpfile.Name()) {
|
||||
t.Errorf("File should be accessible")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareConfig(t *testing.T) {
|
||||
for i, test := range []struct {
|
||||
input AcmeDnsConfig
|
||||
shoulderror bool
|
||||
}{
|
||||
{AcmeDnsConfig{
|
||||
Database: dbsettings{Engine: "whatever", Connection: "whatever_too"},
|
||||
API: httpapi{TLS: ApiTlsProviderNone},
|
||||
}, false},
|
||||
{AcmeDnsConfig{Database: dbsettings{Engine: "", Connection: "whatever_too"},
|
||||
API: httpapi{TLS: ApiTlsProviderNone},
|
||||
}, true},
|
||||
{AcmeDnsConfig{Database: dbsettings{Engine: "whatever", Connection: ""},
|
||||
API: httpapi{TLS: ApiTlsProviderNone},
|
||||
}, true},
|
||||
{AcmeDnsConfig{
|
||||
Database: dbsettings{Engine: "whatever", Connection: "whatever_too"},
|
||||
API: httpapi{TLS: "whatever"},
|
||||
}, true},
|
||||
} {
|
||||
_, err := prepareConfig(test.input)
|
||||
if test.shoulderror {
|
||||
if err == nil {
|
||||
t.Errorf("Test %d: Expected error with prepareConfig input data [%v]", i, test.input)
|
||||
}
|
||||
} else {
|
||||
if err != nil {
|
||||
t.Errorf("Test %d: Expected no error with prepareConfig input data [%v]", i, test.input)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeString(t *testing.T) {
|
||||
for i, test := range []struct {
|
||||
input string
|
||||
expected string
|
||||
}{
|
||||
{"abcd!abcd", "abcdabcd"},
|
||||
{"ABCDEFGHIJKLMNOPQRSTUVXYZabcdefghijklmnopqrstuvwxyz0123456789", "ABCDEFGHIJKLMNOPQRSTUVXYZabcdefghijklmnopqrstuvwxyz0123456789"},
|
||||
{"ABCDEFGHIJKLMNOPQRSTUVXYZabcdefghijklmnopq=@rstuvwxyz0123456789", "ABCDEFGHIJKLMNOPQRSTUVXYZabcdefghijklmnopqrstuvwxyz0123456789"},
|
||||
} {
|
||||
if SanitizeString(test.input) != test.expected {
|
||||
t.Errorf("Expected SanitizeString to return %s for test %d, but got %s instead", test.expected, i, SanitizeString(test.input))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCorrectPassword(t *testing.T) {
|
||||
testPass, _ := bcrypt.GenerateFromPassword([]byte("nevergonnagiveyouup"), 10)
|
||||
for i, test := range []struct {
|
||||
input string
|
||||
expected bool
|
||||
}{
|
||||
{"abcd", false},
|
||||
{"nevergonnagiveyouup", true},
|
||||
{"@rstuvwxyz0123456789", false},
|
||||
} {
|
||||
if test.expected && !CorrectPassword(test.input, string(testPass)) {
|
||||
t.Errorf("Expected CorrectPassword to return %t for test %d", test.expected, i)
|
||||
}
|
||||
if !test.expected && CorrectPassword(test.input, string(testPass)) {
|
||||
t.Errorf("Expected CorrectPassword to return %t for test %d", test.expected, i)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user