diff --git a/Makefile b/Makefile index 649d6d4..a17adf1 100644 --- a/Makefile +++ b/Makefile @@ -1,30 +1,41 @@ -all:: clean build +all:: clean build test build:: get compile -clean: +motd: + @echo + @echo ' ___ __' + @echo ' ____/ (_)_____________ ____/ /___ _____' + @echo ' / __ / / ___/ ___/ __ \/ __ / __ \/ ___/' + @echo '/ /_/ / (__ ) /__/ /_/ / /_/ / / / (__ )' + @echo '\__,_/_/____/\___/\____/\__,_/_/ /_/____/' + @echo + @echo '© Copyright DueDil 2015. Licensed under MIT.' + @echo + +clean: motd @echo "\033[34m●\033[39m Cleaning out the build folder ./build" rm -rf build/* @echo "\033[32m✔\033[39m Cleaned ./build" -get: +get: motd @echo "\033[34m●\033[39m Downloading go packages" go get github.com/tools/godep go get -d godep restore @echo "\033[32m✔\033[39m Finished downloading packages" -compile: +compile: motd get @echo "\033[34m●\033[39m Building into ./build" mkdir -p build/bin go build -o build/bin/discodns *.go @echo "\033[32m✔\033[39m Successfully built into ./build" -test: +test: motd @echo "\033[34m●\033[39m Running tests" - go test -race + go test -race ./ @echo "\033[32m✔\033[39m Tests passed" -install: +install: motd compile @echo "\033[34m●\033[39m Installing into /usr/local/bin" cp build/bin/discodns /usr/local/bin/ @echo "\033[32m✔\033[39m Successfully installed into /usr/local/bin/discodns" diff --git a/README.md b/README.md index 4f36a88..e8c983f 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,6 @@ discodns -====== +======== [![Build Status](https://travis-ci.org/duedil-ltd/discodns.png?branch=master)](https://travis-ci.org/duedil-ltd/discodns) @@ -18,7 +18,7 @@ An authoritative DNS nameserver that queries an [etcd](http://github.com/coreos/ - Support for TTLs - Global default on all records - Individual TTL values for individual records -- Runtime and application metrics are captured regularly for monitoring (stdout or grahite) +- Runtime and application metrics are captured regularly for monitoring (stdout or graphite) - Incoming query filters #### Production Readyness diff --git a/error.go b/error.go index ab3d8be..5f9882a 100644 --- a/error.go +++ b/error.go @@ -14,7 +14,7 @@ type NodeConversionError struct { func (e *NodeConversionError) Error() string { return fmt.Sprintf( - "Unable to convert etc Node into a RR of type %d ('%s'): %s. Node details: %+v", + "Unable to convert etcd Node into a RR of type %d ('%s'): %s. Node details: %+v", e.AttemptedType, dns.TypeToString[e.AttemptedType], e.Message, diff --git a/filter.go b/filter.go index 6c89a59..d872b3c 100644 --- a/filter.go +++ b/filter.go @@ -68,3 +68,32 @@ func (f *QueryFilterer) ShouldAcceptQuery(req *dns.Msg) bool { return accepted } + +// parseFilters will convert a string into a Query Filter structure. The accepted +// format for input is [domain]:[type,type,...]. For example... +// +// - "domain:A,AAAA" # Match all A and AAAA queries within `domain` +// - ":TXT" # Matches only TXT queries for any domain +// - "domain:" # Matches any query within `domain` +func parseFilters(filters []string) []QueryFilter { + parsedFilters := make([]QueryFilter, 0) + for _, filter := range filters { + components := strings.Split(filter, ":") + if len(components) != 2 { + logger.Printf("Expected only one colon ([domain]:[type,type...])") + continue + } + + domain := dns.Fqdn(components[0]) + types := strings.Split(components[1], ",") + + if len(types) == 1 && len(types[0]) == 0 { + types = make([]string, 0) + } + + debugMsg("Adding filter with domain '" + domain + "' and types '" + strings.Join(types, ",") + "'") + parsedFilters = append(parsedFilters, QueryFilter{domain, types}) + } + + return parsedFilters +} diff --git a/lock.go b/lock.go new file mode 100644 index 0000000..53ecf8a --- /dev/null +++ b/lock.go @@ -0,0 +1,151 @@ +package main + +import ( + "code.google.com/p/go-uuid/uuid" + "errors" + "github.com/coreos/go-etcd/etcd" + "time" +) + +const ( + ETCD_LOCK_TTL = 10 + ETCD_LOCK_HEARTBEAT = 5 +) + +// EtcdKeyLock represents a lock on a single key. Its semantics: once asked to +// Acquire, it will try to grab hold of the key in etcd, if it doesnt exist. +// If it does exist, then the lock waits for the other party to release it. +// Once Acquired, it will hold on to the lock indefinitely until Abandoned +type EtcdKeyLock struct { + uuid string + key string + etcdClient *etcd.Client + killChan chan bool + killed bool +} + +func NewEtcdKeyLock(etcdClient *etcd.Client, key string) *EtcdKeyLock { + uuid := uuid.New() + return &EtcdKeyLock{uuid: uuid, key: key, etcdClient: etcdClient, killChan: make(chan bool)} +} + +// Start the process of trying to acquire a key lock. Returns a channel that +// will be sent true when the lock is aquired then closed. Callers can use this +// as their signal to proceed. The lock will kept indefinitely until abandoned. +func (l *EtcdKeyLock) Acquire() chan bool { + inner_acq := make(chan bool) + acquired := make(chan bool) + go func() { + _, ok := <-inner_acq + if ok { + acquired <- true + go heartbeat(l, nil) + go removeWhenCancelled(l) + } + close(acquired) + }() + go tryCreate(l, inner_acq) + return acquired +} + +// Abandons the lock. This just means closing the internal cancellation +// channel, causing all the child goros to do whatever they need to do. +func (l *EtcdKeyLock) Abandon() { + if !l.killed { + l.killed = true + close(l.killChan) + } +} + +// Blocking version of Acquire, hiding the channels from callers who just want +// to synchronously wait +func (l *EtcdKeyLock) WaitForAcquire(timeout int) (bool, error) { + timeoutKiller := time.AfterFunc(time.Duration(timeout) * time.Second, func(){ + l.Abandon() + }) + acq := l.Acquire() + ok, open := <- acq + if ok && open { + stopped := timeoutKiller.Stop() + // Stopped == false means the timer already fired: this shouldn't be + // possible (we shouldn't have been able to get an OK message in that + // case). Erroring mostly out of paranoia: I'm positive this race can't + // happen (famous last words though) + if !stopped { + return false, errors.New("Acquired a lock that was also killed by a timeout: this should not be possible!") + } + return true, nil + } else { + return false, errors.New("Couldn't aqcuire lock in time") + } +} + +// The internals of trying to get a lock: Try to PUT to the lock key iff it +// doesn't exist. If that suceeds, the lock is owned; signal the chan and +// return. If it fails, watch the etcd key until it changes. When it does +// change, try again. Repeat indefinitely until cancelled. +func tryCreate(l *EtcdKeyLock, acq chan bool) { + defer close(acq) + for { + select { + case _, chOpen := <-l.killChan: + if !chOpen { + return + } + default: + _, err := l.etcdClient.Create(l.key, l.uuid, ETCD_LOCK_TTL) + if err == nil { + acq <- true + return + } else { + err, cast := err.(*etcd.EtcdError) + if cast && err.ErrorCode == 105 { + // Watch until it changes. (The current index is given to + // make sure we don't miss any changes in between) + _, err := l.etcdClient.Watch(l.key, err.Index+1, false, nil, l.killChan) + if err == nil { + // Skip the sleep and attempt a retry asap + continue + } + } + // if not created and not watching, pause briefly + time.Sleep(1 * time.Second) + } + } + } +} + +func heartbeat(l *EtcdKeyLock, ping chan bool) { + if ping == nil { + ping = make(chan bool) + } + defer close(ping) + for { + time.Sleep(ETCD_LOCK_HEARTBEAT * time.Second) + select { + case _, chOpen := <-l.killChan: + if !chOpen { + return + } + default: + _, err := l.etcdClient.Set(l.key, l.uuid, ETCD_LOCK_TTL) + // non-blocking write on the ping channel + select { + case ping <- (err == nil): + default: + } + } + } +} + +func removeWhenCancelled(l *EtcdKeyLock) { + defer func() { + l.etcdClient.CompareAndDelete(l.key, l.uuid, 0) + }() + for { + _, chOpen := <-l.killChan + if !chOpen { + return + } + } +} diff --git a/lock_test.go b/lock_test.go new file mode 100644 index 0000000..cd909e5 --- /dev/null +++ b/lock_test.go @@ -0,0 +1,59 @@ +package main + +import ( + "testing" + "time" +) + +func TestSimpleLockUnlock(t *testing.T) { + testKey := "TestSimpleLockUnlock/.lock" + client.Delete(testKey, true) + + lock := NewEtcdKeyLock(client, testKey) + locked, err := lock.WaitForAcquire(1) + + if !locked || err != nil { + t.Error("Expected to acquire lock, failed") + t.Fatal() + } + _, err = client.Get(testKey, false, true) + if err != nil { + t.Error("Lock claimed to succeed but etcd record missing/broken") + t.Fatal() + } + + lock.Abandon() + time.Sleep(500 * time.Millisecond) + _, err = client.Get(testKey, false, true) + if err == nil { + t.Error("Lock abandoned, but key exists") + t.Fatal() + } +} + +func TestConflictingLock(t *testing.T) { + testKey := "TestConflictingLock/.lock" + client.Delete(testKey, true) + + lock_a := NewEtcdKeyLock(client, testKey) + lock_a.WaitForAcquire(30) + + lock_b := NewEtcdKeyLock(client, testKey) + b_locked, b_err := lock_b.WaitForAcquire(1) + if b_locked || b_err == nil { + t.Error("Expected second lock to timeout") + t.Fatal() + } + + lock_c := NewEtcdKeyLock(client, testKey) + go func(){ + time.Sleep(500 * time.Millisecond) + lock_a.Abandon() + }() + + c_locked, c_err := lock_c.WaitForAcquire(5) + if !c_locked || c_err != nil { + t.Error("Expected third lock to succeed in time") + t.Fatal() + } +} diff --git a/main.go b/main.go index 66b7686..b4121a4 100644 --- a/main.go +++ b/main.go @@ -30,6 +30,8 @@ var ( DefaultTtl uint32 `short:"t" long:"default-ttl" description:"Default TTL to return on records without an explicit TTL" default:"300"` Accept []string `long:"accept" description:"Limit DNS queries to a set of domain:[type,...] pairs"` Reject []string `long:"reject" description:"Limit DNS queries to a set of domain:[type,...] pairs"` + TsigSecret []string `short:"s" long:"tsig" description:"Transaction signature secret in the format zone:secret"` + TsigFreeZones []string `long:"unauth" description:"Zone names that can be updated without TSIG authentication"` } ) @@ -79,6 +81,26 @@ func main() { logger.Printf("Metric logging disabled") } + // Parse the tsig arguments, these are formatted as "zone:secret" + tsigSecret := map[string]string{} + for _, arg := range Options.TsigSecret { + components := strings.SplitN(arg, ":", 2) + if len(components) != 2 { + logger.Printf("Failed to parse TSIG argument") + continue + } + tsigSecret[dns.Fqdn(components[0])] = components[1] + } + + // create a unique list of zone names that allow unauthenticated access + tsigFreeZones := make(map[string]struct{}, len(Options.TsigFreeZones)) + for _, zone := range Options.TsigFreeZones { + if zone[len(zone)-1] != '.' { + zone = zone + "." + } + tsigFreeZones[zone] = struct{}{} + } + // Start up the DNS resolver server server := &Server{ addr: Options.ListenAddress, @@ -87,6 +109,8 @@ func main() { rTimeout: time.Duration(5) * time.Second, wTimeout: time.Duration(5) * time.Second, defaultTtl: Options.DefaultTtl, + tsigSecret: tsigSecret, + tsigFreeZones: tsigFreeZones, queryFilterer: &QueryFilterer{acceptFilters: parseFilters(Options.Accept), rejectFilters: parseFilters(Options.Reject)}} @@ -116,35 +140,6 @@ func debugMsg(v ...interface{}) { } } -// parseFilters will convert a string into a Query Filter structure. The accepted -// format for input is [domain]:[type,type,...]. For example... -// -// - "domain:A,AAAA" # Match all A and AAAA queries within `domain` -// - ":TXT" # Matches only TXT queries for any domain -// - "domain:" # Matches any query within `domain` -func parseFilters(filters []string) []QueryFilter { - parsedFilters := make([]QueryFilter, 0) - for _, filter := range filters { - components := strings.Split(filter, ":") - if len(components) != 2 { - logger.Printf("Expected only one colon ([domain]:[type,type...])") - continue - } - - domain := dns.Fqdn(components[0]) - types := strings.Split(components[1], ",") - - if len(types) == 1 && len(types[0]) == 0 { - types = make([]string, 0) - } - - debugMsg("Adding filter with domain '" + domain + "' and types '" + strings.Join(types, ",") + "'") - parsedFilters = append(parsedFilters, QueryFilter{domain, types}) - } - - return parsedFilters -} - func init() { runtime.GOMAXPROCS(runtime.NumCPU()) } diff --git a/record.go b/record.go new file mode 100644 index 0000000..912ae35 --- /dev/null +++ b/record.go @@ -0,0 +1,277 @@ +package main + +import ( + "github.com/coreos/go-etcd/etcd" + "github.com/miekg/dns" + "bytes" + "strings" + "fmt" + "strconv" + "net" +) + +type EtcdRecord struct { + node *etcd.Node + ttl uint32 +} + +// convertNodeToRR will convert an etcd node with a raw value into a dns.RR +// record, returning an error if the conversion fails +func convertNodeToRR(node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + rr, err = convertersToRR[header.Rrtype](node, header) + return +} + +// convertRRToNode will convert a DNS RR and it's type specific values to an +// etcd node with a raw value and key path, returning an error if the conversion +// fails +func convertRRToNode(rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + node, err = convertersFromRR[header.Rrtype](rr, header) + return +} + +// nameToKey returns a string representing the etcd version of a domain, replacing dots with slashes +// and reversing it (foo.net. -> /net/foo) +func nameToKey(name string, suffix string) string { + segments := strings.Split(name, ".") + + var keyBuffer bytes.Buffer + var writtenSegment bool + for i := len(segments) - 1; i >= 0; i-- { + if len(segments[i]) > 0 { + // We never want to write a leading slash + if writtenSegment { + keyBuffer.WriteString("/") + } + + keyBuffer.WriteString(segments[i]) + writtenSegment = true + } + } + + keyBuffer.WriteString(suffix) + return keyBuffer.String() +} + +// Map of conversion functions that turn individual etcd nodes into dns.RR answers +var convertersToRR = map[uint16]func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + + dns.TypeA: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + + ip := net.ParseIP(node.Value) + if ip == nil { + err = &NodeConversionError{ + Node: node, + Message: fmt.Sprintf("Failed to parse %s as IP Address", node.Value), + AttemptedType: dns.TypeA, + } + } else if ip.To4() == nil { + err = &NodeConversionError{ + Node: node, + Message: fmt.Sprintf("Value %s isn't an IPv4 address", node.Value), + AttemptedType: dns.TypeA, + } + } else { + rr = &dns.A{header, ip} + } + + return + }, + + dns.TypeAAAA: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + + ip := net.ParseIP(node.Value) + if ip == nil { + err = &NodeConversionError{ + Node: node, + Message: fmt.Sprintf("Failed to parse IP Address %s", node.Value), + AttemptedType: dns.TypeAAAA} + } else if ip.To16() == nil { + err = &NodeConversionError{ + Node: node, + Message: fmt.Sprintf("Value %s isn't an IPv6 address", node.Value), + AttemptedType: dns.TypeA} + } else { + rr = &dns.AAAA{header, ip} + } + return + }, + + dns.TypeTXT: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + rr = &dns.TXT{header, []string{node.Value}} + return + }, + + dns.TypeCNAME: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + rr = &dns.CNAME{header, dns.Fqdn(node.Value)} + return + }, + + dns.TypeNS: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + rr = &dns.NS{header, dns.Fqdn(node.Value)} + return + }, + + dns.TypePTR: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + labels, ok := dns.IsDomainName(node.Value) + + if (ok && labels > 0) { + rr = &dns.PTR{header, dns.Fqdn(node.Value)} + } else { + err = &NodeConversionError{ + Node: node, + Message: fmt.Sprintf("Value '%s' isn't a valid domain name", node.Value), + AttemptedType: dns.TypePTR} + } + return + }, + + dns.TypeSRV: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + parts := strings.SplitN(node.Value, "\t", 4) + + if len(parts) != 4 { + err = &NodeConversionError{ + Node: node, + Message: fmt.Sprintf("Value %s isn't valid for SRV", node.Value), + AttemptedType: dns.TypeSRV} + } else { + + priority, err := strconv.ParseUint(parts[0], 10, 16) + if err != nil { + return nil, err + } + + weight, err := strconv.ParseUint(parts[1], 10, 16) + if err != nil { + return nil, err + } + + port, err := strconv.ParseUint(parts[2], 10, 16) + if err != nil { + return nil, err + } + + target := dns.Fqdn(parts[3]) + + rr = &dns.SRV{ + header, + uint16(priority), + uint16(weight), + uint16(port), + target} + } + return + }, + + dns.TypeSOA: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { + parts := strings.SplitN(node.Value, "\t", 6) + + if len(parts) < 6 { + err = &NodeConversionError{ + Node: node, + Message: fmt.Sprintf("Value %s isn't valid for SOA", node.Value), + AttemptedType: dns.TypeSOA} + } else { + refresh, err := strconv.ParseUint(parts[2], 10, 32) + if err != nil { + return nil, err + } + + retry, err := strconv.ParseUint(parts[3], 10, 32) + if err != nil { + return nil, err + } + + expire, err := strconv.ParseUint(parts[4], 10, 32) + if err != nil { + return nil, err + } + + minttl, err := strconv.ParseUint(parts[5], 10, 32) + if err != nil { + return nil, err + } + + rr = &dns.SOA{ + Hdr: header, + Ns: dns.Fqdn(parts[0]), + Mbox: dns.Fqdn(parts[1]), + Refresh: uint32(refresh), + Retry: uint32(retry), + Expire: uint32(expire), + Minttl: uint32(minttl)} + } + + return + }, +} + +// Map of conversion functions that turn dns.RR answers into individual etcd nodes +var convertersFromRR = map[uint16]func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + + dns.TypeANY: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + node = &etcd.Node{Key: nameToKey(header.Name, "")} + + return + }, + + dns.TypeA: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + if record, ok := rr.(*dns.A); ok { + node = &etcd.Node{ + Key: nameToKey(header.Name, "/.A"), + Value: record.A.String()} + } + + return + }, + + // dns.TypeAAAA: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + // panic("Not implemented") + // }, + + dns.TypeTXT: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + if record, ok := rr.(*dns.TXT); ok { + node = &etcd.Node{ + Key: nameToKey(header.Name, "/.TXT"), + Value: strings.Join(record.Txt, "\n")} + } + return + }, + + dns.TypeCNAME: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + if record, ok := rr.(*dns.CNAME); ok { + node = &etcd.Node{ + Key: nameToKey(header.Name, "/.CNAME"), + Value: record.Target} + } + return + }, + + // dns.TypeNS: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + // panic("Not implemented") + // }, + + dns.TypePTR: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + if record, ok := rr.(*dns.PTR); ok { + node = &etcd.Node{ + Key: nameToKey(header.Name, "/.PTR"), + Value: record.Ptr} + } + + return + }, + + dns.TypeSRV: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + if record, ok := rr.(*dns.SRV); ok { + node = &etcd.Node{ + Key: nameToKey(header.Name, "/.SRV"), + Value: fmt.Sprintf("%d\t%d\t%d\t%s", record.Priority, record.Weight, record.Port, record.Target)} + } + + return + }, + + // dns.TypeSOA: func (rr dns.RR, header dns.RR_Header) (node *etcd.Node, err error) { + // panic("Not implemented") + // }, +} diff --git a/record_test.go b/record_test.go new file mode 100644 index 0000000..b979fd2 --- /dev/null +++ b/record_test.go @@ -0,0 +1,46 @@ +package main + +import ( + "testing" +) + +func TestRecord(t *testing.T) { + // Enable debug logging + log_debug = true +} + +func TestNameToKey(t *testing.T) { + + key := nameToKey("foo.disco.net", "") + if key != "net/disco/foo" { + t.Error("Expected key to be /net/disco/foo but got " + key) + t.Fatal() + } +} + +func TestNameToKeyFQDN(t *testing.T) { + + key := nameToKey("foo.disco.net.", "") + if key != "net/disco/foo" { + t.Error("Expected key to be /net/disco/foo but got " + key) + t.Fatal() + } +} + +func TestNameToKeyWithSuffix(t *testing.T) { + + key := nameToKey("foo.disco.net", "/.A") + if key != "net/disco/foo/.A" { + t.Error("Expected key to be /net/disco/foo/.A but got " + key) + t.Fatal() + } +} + +func TestNameToKeyFQDNWithSuffix(t *testing.T) { + + key := nameToKey("foo.disco.net.", "/.A") + if key != "net/disco/foo/.A" { + t.Error("Expected key to be /net/disco/foo/.A but got " + key) + t.Fatal() + } +} diff --git a/resolver.go b/resolver.go index f291459..035de97 100644 --- a/resolver.go +++ b/resolver.go @@ -1,12 +1,9 @@ package main import ( - "bytes" - "fmt" "github.com/coreos/go-etcd/etcd" "github.com/miekg/dns" "github.com/rcrowley/go-metrics" - "net" "strconv" "strings" "sync" @@ -19,11 +16,6 @@ type Resolver struct { defaultTtl uint32 } -type EtcdRecord struct { - node *etcd.Node - ttl uint32 -} - // GetFromStorage looks up a key in etcd and returns a slice of nodes. It supports two storage structures; // - File: /foo/bar/.A -> "value" // - Directory: /foo/bar/.A/0 -> "value-0" @@ -34,7 +26,7 @@ func (r *Resolver) GetFromStorage(key string) (nodes []*EtcdRecord, err error) { error_counter := metrics.GetOrRegisterCounter("resolver.etcd.query_error_count", metrics.DefaultRegistry) counter.Inc(1) - debugMsg("Querying etcd for " + key) + debugMsg("Querying etcd for /" + r.etcdPrefix + key) response, err := r.etcd.Get(r.etcdPrefix + key, true, true) if err != nil { @@ -79,7 +71,7 @@ func (r *Resolver) GetFromStorage(key string) (nodes []*EtcdRecord, err error) { return } - // If we don't have a TLL try and find one + // If we don't have a TTL try and find one if tryTtl { ttlKey := node.Key + ".ttl" @@ -151,7 +143,7 @@ func (r *Resolver) Lookup(req *dns.Msg) (msg *dns.Msg) { var eChan chan error if q.Qclass == dns.ClassINET { - aChan, eChan = r.AnswerQuestion(q) + aChan, eChan = r.AnswerQuestion(q, true) answers, errors = gatherFromChannels(aChan, eChan) } @@ -168,7 +160,7 @@ func (r *Resolver) Lookup(req *dns.Msg) (msg *dns.Msg) { Qtype: q.Qtype, Qclass: q.Qclass} - aChan, eChan = r.AnswerQuestion(question) + aChan, eChan = r.AnswerQuestion(question, true) answers, errors = gatherFromChannels(aChan, eChan) errored = errored || len(errors) > 0 @@ -239,7 +231,7 @@ func gatherFromChannels(rrsIn chan dns.RR, errsIn chan error) (rrs []dns.RR, err // the way. The function will return immediately, and spawn off a bunch of goroutines // to do the work, when using this function one should use a WaitGroup to know when all work // has been completed. -func (r *Resolver) AnswerQuestion(q dns.Question) (answers chan dns.RR, errors chan error) { +func (r *Resolver) AnswerQuestion(q dns.Question, resolveAliases bool) (answers chan dns.RR, errors chan error) { answers = make(chan dns.RR) errors = make(chan error) @@ -251,15 +243,15 @@ func (r *Resolver) AnswerQuestion(q dns.Question) (answers chan dns.RR, errors c if q.Qtype == dns.TypeANY { wg := sync.WaitGroup{} - wg.Add(len(converters)) + wg.Add(len(convertersToRR)) go func(){ wg.Wait() close(answers) close(errors) }() - for rrType, _ := range converters { + for rrType, _ := range convertersToRR { go func(rrType uint16) { - defer func() { recover() }() + defer recover() defer wg.Done() results, err := r.LookupAnswersForType(q.Name, rrType) @@ -272,7 +264,7 @@ func (r *Resolver) AnswerQuestion(q dns.Question) (answers chan dns.RR, errors c } }(rrType) } - } else if _, ok := converters[q.Qtype]; ok { + } else if _, ok := convertersToRR[q.Qtype]; ok { go func() { defer func(){ close(answers) @@ -286,7 +278,7 @@ func (r *Resolver) AnswerQuestion(q dns.Question) (answers chan dns.RR, errors c for _, rr := range records { answers <- rr } - } else { + } else if resolveAliases { cnames, err := r.LookupAnswersForType(q.Name, dns.TypeCNAME) if err != nil { errors <- err @@ -331,7 +323,7 @@ func (r *Resolver) LookupAnswersForType(name string, rrType uint16) (answers []d for i, node := range nodes { header := dns.RR_Header{Name: name, Class: dns.ClassINET, Rrtype: rrType, Ttl: node.ttl} - answer, err := converters[rrType](node.node, header) + answer, err := convertersToRR[rrType](node.node, header) if err != nil { debugMsg("Error converting type: ", err) @@ -344,172 +336,55 @@ func (r *Resolver) LookupAnswersForType(name string, rrType uint16) (answers []d return } -// nameToKey returns a string representing the etcd version of a domain, replacing dots with slashes -// and reversing it (foo.net. -> /net/foo) -func nameToKey(name string, suffix string) string { - segments := strings.Split(name, ".") +// NameExists will return true if the given domain name exists and has any +// resource records in the database. If an error occurs while querying for +// data the function will return false and an error. +func (r *Resolver) NameExists(name string) (exists bool, err error) { - var keyBuffer bytes.Buffer - for i := len(segments) - 1; i >= 0; i-- { - if len(segments[i]) > 0 { - keyBuffer.WriteString("/") - keyBuffer.WriteString(segments[i]) - } - } + question := dns.Question{dns.Fqdn(name), dns.TypeANY, dns.ClassINET} + aChan, eChan := r.AnswerQuestion(question, true) + answers, errors := gatherFromChannels(aChan, eChan) - keyBuffer.WriteString(suffix) - return keyBuffer.String() + if len(errors) > 0 { + return false, errors[0] + } + return len(answers) > 0, nil } -// Map of conversion functions that turn individual etcd nodes into dns.RR answers -var converters = map[uint16]func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - - dns.TypeA: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - - ip := net.ParseIP(node.Value) - if ip == nil { - err = &NodeConversionError{ - Node: node, - Message: fmt.Sprintf("Failed to parse %s as IP Address", node.Value), - AttemptedType: dns.TypeA, - } - } else if ip.To4() == nil { - err = &NodeConversionError{ - Node: node, - Message: fmt.Sprintf("Value %s isn't an IPv4 address", node.Value), - AttemptedType: dns.TypeA, - } - } else { - rr = &dns.A{header, ip} - } - - return - }, - - dns.TypeAAAA: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - - ip := net.ParseIP(node.Value) - if ip == nil { - err = &NodeConversionError{ - Node: node, - Message: fmt.Sprintf("Failed to parse IP Address %s", node.Value), - AttemptedType: dns.TypeAAAA} - } else if ip.To16() == nil { - err = &NodeConversionError{ - Node: node, - Message: fmt.Sprintf("Value %s isn't an IPv6 address", node.Value), - AttemptedType: dns.TypeA} - } else { - rr = &dns.AAAA{header, ip} - } - return - }, - - dns.TypeTXT: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - rr = &dns.TXT{header, []string{node.Value}} - return - }, - - dns.TypeCNAME: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - rr = &dns.CNAME{header, dns.Fqdn(node.Value)} - return - }, - - dns.TypeNS: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - rr = &dns.NS{header, dns.Fqdn(node.Value)} - return - }, - - dns.TypePTR: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - labels, ok := dns.IsDomainName(node.Value) - - if (ok && labels > 0) { - rr = &dns.PTR{header, dns.Fqdn(node.Value)} - } else { - err = &NodeConversionError{ - Node: node, - Message: fmt.Sprintf("Value '%s' isn't a valid domain name", node.Value), - AttemptedType: dns.TypePTR} - } - return - }, - - dns.TypeSRV: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - parts := strings.SplitN(node.Value, "\t", 4) - - if len(parts) != 4 { - err = &NodeConversionError{ - Node: node, - Message: fmt.Sprintf("Value %s isn't valid for SRV", node.Value), - AttemptedType: dns.TypeSRV} - } else { - - priority, err := strconv.ParseUint(parts[0], 10, 16) - if err != nil { - return nil, err - } - - weight, err := strconv.ParseUint(parts[1], 10, 16) - if err != nil { - return nil, err - } - - port, err := strconv.ParseUint(parts[2], 10, 16) - if err != nil { - return nil, err - } - - target := dns.Fqdn(parts[3]) - - rr = &dns.SRV{ - header, - uint16(priority), - uint16(weight), - uint16(port), - target} - } - return - }, - - dns.TypeSOA: func (node *etcd.Node, header dns.RR_Header) (rr dns.RR, err error) { - parts := strings.SplitN(node.Value, "\t", 6) - - if len(parts) < 6 { - err = &NodeConversionError{ - Node: node, - Message: fmt.Sprintf("Value %s isn't valid for SOA", node.Value), - AttemptedType: dns.TypeSOA} - } else { - refresh, err := strconv.ParseUint(parts[2], 10, 32) - if err != nil { - return nil, err - } - - retry, err := strconv.ParseUint(parts[3], 10, 32) - if err != nil { - return nil, err - } +// RRSetExists returns true if RRs exist for the given name and type (value independent) +func (r *Resolver) RRSetExists(name string, rrType uint16) (exists bool, err error) { + answers, err := r.LookupAnswersForType(dns.Fqdn(name), rrType) + if err != nil { + return false, err + } - expire, err := strconv.ParseUint(parts[4], 10, 32) - if err != nil { - return nil, err - } + return len(answers) > 0, nil +} - minttl, err := strconv.ParseUint(parts[5], 10, 32) - if err != nil { - return nil, err +// RRSetMatches checks that the set of records in the DNS for the given name +// and type *exactly* match the given RRs: their data must match, and there must +// be no more or less RRs +func (r *Resolver) RRSetMatches(name string, rrType uint16, rrs []dns.RR) (matches bool, err error) { + answers, err := r.LookupAnswersForType(dns.Fqdn(name), rrType) + if err != nil { + return false, err + } + if len(answers) != len(rrs) { + return false, nil + } + matched := 0 + // I'm sure theres a neater/faster way than comparing all to all, but meh + for _, rr := range rrs { + for _, answer := range answers { + // TTLS are explicitly excluded from comparison + cmp := dns.Copy(answer) + cmp.Header().Ttl = 0 + // TODO(orls): is string() enough? is there any relevant info not in the string reprs? + if cmp.String() == rr.String() { + matched++ + break } - - rr = &dns.SOA{ - Hdr: header, - Ns: dns.Fqdn(parts[0]), - Mbox: dns.Fqdn(parts[1]), - Refresh: uint32(refresh), - Retry: uint32(retry), - Expire: uint32(expire), - Minttl: uint32(minttl)} } - - return - }, + } + return matched == len(rrs), nil } diff --git a/resolver_test.go b/resolver_test.go index ab12a80..a4a4356 100644 --- a/resolver_test.go +++ b/resolver_test.go @@ -8,7 +8,7 @@ import ( ) var ( - client = etcd.NewClient([]string{"127.0.0.1:4001"}) + client = etcd.NewClient([]string{"http://127.0.0.1:4001"}) resolver = &Resolver{etcd: client} ) @@ -80,25 +80,6 @@ func TestGetFromStorageNestedKeys(t *testing.T) { } } -func TestNameToKeyConverter(t *testing.T) { - var key string - - key = nameToKey("foo.net.", "") - if key != "/net/foo" { - t.Error("Expected key /net/foo") - } - - key = nameToKey("foo.net", "") - if key != "/net/foo" { - t.Error("Expected key /net/foo") - } - - key = nameToKey("foo.net.", "/.A") - if key != "/net/foo/.A" { - t.Error("Expected key /net/foo/.A") - } -} - /** * Test that the right authority is being returned for different types of DNS * queries. @@ -924,3 +905,68 @@ func TestLookupAnswerForSRVInvalidValues(t *testing.T) { } } } + +func TestNameExistsDoesExist(t *testing.T) { + + resolver.etcdPrefix = "TestNameExistsDoesExist/" + client.Set("TestNameExistsDoesExist/net/disco/bar/.A", "127.0.0.1", 0) + + exists, err := resolver.NameExists("bar.disco.net") + if exists != true { + t.Error("Expected domain to exist (true), got (false)") + t.Fatal() + } + + if err != nil { + t.Error("Expected error to be nil") + t.Fatal() + } +} + +func TestNameExistsDoesNotExist(t *testing.T) { + + resolver.etcdPrefix = "TestNameExistsDoesNotExist/" + exists, err := resolver.NameExists("bar.disco.net") + if exists != false { + t.Error("Expected domain to not exist (false), got (true)") + t.Fatal() + } + + if err != nil { + t.Error("Expected error to be nil") + t.Fatal() + } +} + +func TestRRSetExistsDoesExist(t *testing.T) { + + resolver.etcdPrefix = "TestRRSetExistsDoesExist/" + client.Set("TestRRSetExistsDoesExist/net/disco/bar/.A", "127.0.0.1", 0) + + exists, err := resolver.RRSetExists("bar.disco.net", dns.TypeA) + if exists != true { + t.Error("Expected RRset to exist (true), got (false)") + t.Fatal() + } + + if err != nil { + t.Error("Expected error to be nil") + t.Fatal() + } +} + +func TestRRSetExistsDoesNotExist(t *testing.T) { + + resolver.etcdPrefix = "TestRRSetExistsDoesNotExist/" + exists, err := resolver.RRSetExists("bar.disco.net", dns.TypeA) + if exists != false { + t.Error("Expected RRset to not exist (false), got (true)") + t.Fatal() + } + + if err != nil { + t.Error("Expected error to be nil") + t.Fatal() + } +} + diff --git a/server.go b/server.go index 3538f8a..a50db62 100644 --- a/server.go +++ b/server.go @@ -15,51 +15,95 @@ type Server struct { rTimeout time.Duration wTimeout time.Duration defaultTtl uint32 + tsigSecret map[string]string + tsigFreeZones map[string]struct{} queryFilterer *QueryFilterer } type Handler struct { resolver *Resolver queryFilterer *QueryFilterer + updateManager *DynamicUpdateManager + tsigFreeZones map[string]struct{} // Metrics requestCounter metrics.Counter acceptCounter metrics.Counter rejectCounter metrics.Counter + // authFailCounter metrics.Counter + // authSuccessCounter metrics.Counter responseTimer metrics.Timer } func (h *Handler) Handle(response dns.ResponseWriter, req *dns.Msg) { h.requestCounter.Inc(1) h.responseTimer.Time(func() { - debugMsg("Handling incoming query for domain " + req.Question[0].Name) - - // Lookup the dns record for the request - // This method will add any answers to the message - var msg *dns.Msg - if h.queryFilterer.ShouldAcceptQuery(req) != true { - debugMsg("Query not accepted") - - h.rejectCounter.Inc(1) - - msg = new(dns.Msg) - msg.SetReply(req) - msg.SetRcode(req, dns.RcodeNameError) - msg.Authoritative = true - msg.RecursionAvailable = false - - // Add a useful TXT record - header := dns.RR_Header{Name: req.Question[0].Name, - Class: dns.ClassINET, - Rrtype: dns.TypeTXT} - msg.Ns = []dns.RR{&dns.TXT{header, []string{"Rejected query based on matched filters"}}} + debugMsg("Incoming message with opcode " + dns.OpcodeToString[req.MsgHdr.Opcode]) + + var res *dns.Msg + if req.MsgHdr.Opcode == dns.OpcodeQuery { + // TODO(tarnfeld): Support for multiple questions? + debugMsg("Handling incoming query for domain " + req.Question[0].Name) + + // Lookup the dns record for the request + // This method will add any answers to the message + if h.queryFilterer.ShouldAcceptQuery(req) != true { + debugMsg("Query not accepted") + + h.rejectCounter.Inc(1) + + res = new(dns.Msg) + res.SetReply(req) + res.SetRcode(req, dns.RcodeNameError) + res.Authoritative = true + res.RecursionAvailable = false + + // Add a useful TXT record + header := dns.RR_Header{Name: req.Question[0].Name, + Class: dns.ClassINET, + Rrtype: dns.TypeTXT} + res.Ns = []dns.RR{&dns.TXT{header, []string{"Rejected query based on matched filters"}}} + } else { + h.acceptCounter.Inc(1) + res = h.resolver.Lookup(req) + } + } else if req.MsgHdr.Opcode == dns.OpcodeUpdate { + zone := req.Question[0].Name + debugMsg("Handling incoming update for zone " + zone) + + res = new(dns.Msg) + res.SetReply(req) + + // Authenticate the request + tsig := req.IsTsig() + if tsig != nil && response.TsigStatus() == nil { + sig := req.IsTsig() + debugMsg("Authenticated update request") + + // Verify the tsig is for the correct zone + if sig.Hdr.Name != zone { + res.SetRcode(req, dns.RcodeBadSig) + } else { + res = h.updateManager.Update(zone, req) + res.SetTsig(tsig.Header().Name, dns.HmacMD5, 300, time.Now().Unix()) + } + } else { + if _, ok := h.tsigFreeZones[zone]; ok { + debugMsg("allowing unauthenticated update") + res = h.updateManager.Update(zone, req) + } else { + debugMsg("Update authentication failed") + res.SetRcode(req, dns.RcodeNotAuth) + } + } } else { - h.acceptCounter.Inc(1) - msg = h.resolver.Lookup(req) + res = new(dns.Msg) + res.SetReply(req) + res.SetRcode(req, dns.RcodeNotImplemented) } - if msg != nil { - err := response.WriteMsg(msg) + if res != nil { + err := response.WriteMsg(res) if err != nil { debugMsg("Error writing message: ", err) } @@ -94,20 +138,25 @@ func (s *Server) Run() { metrics.Register("request.handler.udp.filter_rejects", udpRejectCounter) resolver := Resolver{etcd: s.etcd, defaultTtl: s.defaultTtl} + updateManager := DynamicUpdateManager{etcd: s.etcd, resolver: &resolver} tcpDNShandler := &Handler{ resolver: &resolver, requestCounter: tcpRequestCounter, acceptCounter: tcpAcceptCounter, rejectCounter: tcpRejectCounter, responseTimer: tcpResponseTimer, - queryFilterer: s.queryFilterer} + queryFilterer: s.queryFilterer, + updateManager: &updateManager, + tsigFreeZones: s.tsigFreeZones} udpDNShandler := &Handler{ resolver: &resolver, requestCounter: udpRequestCounter, acceptCounter: udpAcceptCounter, rejectCounter: udpRejectCounter, responseTimer: udpResponseTimer, - queryFilterer: s.queryFilterer} + queryFilterer: s.queryFilterer, + updateManager: &updateManager, + tsigFreeZones: s.tsigFreeZones} udpHandler := dns.NewServeMux() tcpHandler := dns.NewServeMux() @@ -119,14 +168,16 @@ func (s *Server) Run() { Net: "tcp", Handler: tcpHandler, ReadTimeout: s.rTimeout, - WriteTimeout: s.wTimeout} + WriteTimeout: s.wTimeout, + TsigSecret: s.tsigSecret} udpServer := &dns.Server{Addr: s.Addr(), Net: "udp", Handler: udpHandler, UDPSize: 65535, ReadTimeout: s.rTimeout, - WriteTimeout: s.wTimeout} + WriteTimeout: s.wTimeout, + TsigSecret: s.tsigSecret} go s.start(udpServer) go s.start(tcpServer) diff --git a/update.go b/update.go new file mode 100644 index 0000000..26f7ddb --- /dev/null +++ b/update.go @@ -0,0 +1,444 @@ +package main + +import ( + "crypto/md5" + "encoding/hex" + "github.com/coreos/go-etcd/etcd" + "github.com/miekg/dns" + "fmt" + "strconv" + "strings" +) + +type DynamicUpdateManager struct { + etcd *etcd.Client + etcdPrefix string + resolver *Resolver +} + +// Update will perform the necessary logic to update the DNS database with +// the changes described in the RFC-2136 formatted DNS message given. +// The return value will be the response message to send back to the client. +// It is assumed at this level the client has already authenticated and proven +// their right to update records in the given zone. +func (u *DynamicUpdateManager) Update(zone string, req *dns.Msg) (msg *dns.Msg) { + + rrsets := [][]dns.RR{req.Answer, req.Ns} + msg = new(dns.Msg) + msg.SetReply(req) + msg.Opcode = dns.OpcodeUpdate + + // dns update re-aporopriates DNS message blocks: + // msg.Question: Zone info for whole request + // msg.Answer: prerequisites + // msg.Ns: the actual update RRs + + // Verify the updates are within the zone we're modifying, since cross + // zone updates are invalid. + for _, rrs := range rrsets { + for _, rr := range rrs { + if dns.CompareDomainName(rr.Header().Name, zone) != dns.CountLabel(zone) { + debugMsg("Domain " + rr.Header().Name + " is not in the " + zone + " zone") + msg.Rcode = dns.RcodeNotZone + return + } + } + } + + // Ensure we recover from any panicking goroutine + defer func() { + if r := recover(); r != nil { + debugMsg("[PANIC] " + fmt.Sprint(r)) + msg.Rcode = dns.RcodeServerFailure + } + }() + + // Attempt to acquire the dns-updates lock key. + // TODO (orls): This means all updates from all running instances are + // applied fully serially; this is less than ideal, the spec says they + // should be serial only when conflicting with one another. But...this is + // easier than building full transactions isolation mgmt :) For a low + // frequency of updates, this should fine. + lock := NewEtcdKeyLock(u.etcd, u.etcdPrefix + "._DISCODNS_UPDATE_LOCK") + defer lock.Abandon() + // block until locked or timed-out + _, err := lock.WaitForAcquire(30) + if err != nil { + debugMsg("Failed to acquire or keep the update lock: ", err) + msg.Rcode = dns.RcodeServerFailure + return + } + + // Validate the prerequisites of the update, returning immediately if they + // are not satisfied. + prereqValidation := validatePrerequisites(req.Answer, u.resolver) + if prereqValidation != dns.RcodeSuccess { + debugMsg("Validation of prerequisites failed") + msg.Rcode = prereqValidation + return + } + + updateValidation := validateUpdates(req.Ns, req.Question[0]) + if updateValidation != dns.RcodeSuccess { + debugMsg("Validation of update instructions failed") + msg.Rcode = updateValidation + return + } + + // Perform the updates to the domain name system + // This is not inside any kind of transaction, so a failure here *can* + // result in a partially updated zone. + // TODO(tarnfeld): Figure out a way of rolling back changes, perhaps make + // use of the etcd indexes? + msg.Rcode = performUpdate(u.etcdPrefix, u.etcd, u.resolver, req.Question[0], req.Ns) + + return +} + +// internal utility struct for making a map of RRsets a bit neater to construct +type matchKey struct { + name string + rrType uint16 +} + +// validatePrerequisites will perform all necessary validation checks against +// update prerequisites and return the relevant status is validation fails, +// otherwise NOERROR(0) will be returned. +// See RFC 2136, section 3.2 +func validatePrerequisites(rr []dns.RR, resolver *Resolver) (rcode int) { + rrSetsToMatch := make(map[matchKey][]dns.RR) + for _, record := range rr { + header := record.Header() + if header.Ttl != 0 { + return dns.RcodeFormatError + } + + if header.Class == dns.ClassANY { + if header.Rdlength != 0 { + return dns.RcodeFormatError + } else if header.Rrtype == dns.TypeANY { + // RFC Meaning: "Name is in use" + exists, err := resolver.NameExists(header.Name) + if err != nil { + return dns.RcodeServerFailure + } + if !exists { + debugMsg("Domain that should exist does not ", header.Name) + return dns.RcodeNameError + } + } else { + // RFC Meaning: "RRset exists (value independent)" + exists, err := resolver.RRSetExists(header.Name, header.Rrtype) + if err != nil { + return dns.RcodeServerFailure + } + if !exists { + debugMsg("RRset that should exist does not ", header.Name, header.Rrtype) + return dns.RcodeNXRrset + } + } + } else if header.Class == dns.ClassNONE { + if header.Rdlength != 0 { + return dns.RcodeFormatError + } else if header.Rrtype == dns.TypeANY { + // RFC Meaning: "Name is not in use" + exists, err := resolver.NameExists(header.Name) + if err != nil { + return dns.RcodeServerFailure + } + if exists { + debugMsg("Domain that should not exist does ", header.Name) + return dns.RcodeYXDomain + } + } else { + // RFC meaning: "RRset does not exist" + exists, err := resolver.RRSetExists(header.Name, header.Rrtype) + if err != nil { + return dns.RcodeServerFailure + } + if exists { + debugMsg("RRset that should not exist does ", header.Name) + return dns.RcodeYXRrset + } + } + } else if header.Class == dns.ClassINET { + if header.Rrtype == dns.TypeANY { + return dns.RcodeFormatError + } else { + // RFC Meaning: "RRset exists (value dependent)" + mKey := matchKey{name: header.Name, rrType: header.Rrtype} + rrSetsToMatch[mKey] = append(rrSetsToMatch[mKey], record) + } + } else { + return dns.RcodeFormatError + } + } + + for matchKey, rrs := range rrSetsToMatch { + matched, err := resolver.RRSetMatches(matchKey.name, matchKey.rrType, rrs) + if err != nil { + return dns.RcodeServerFailure + } + if !matched { + return dns.RcodeNXRrset + } + } + + return dns.RcodeSuccess +} + +// validateUpdates ensures that the given update instructions conform to the RFC +// and are processable, before we begin mutating state +// See RFC 2136, section 3.4.1 +func validateUpdates(rrs []dns.RR, updateZone dns.Question) (rcode int) { + + // name-in-zone checks have already been performed. + + badTypes := map[uint16]bool{ dns.TypeIXFR : true, + dns.TypeAXFR : true, dns.TypeMAILB : true, dns.TypeMAILA : true, + dns.TypeANY : true} + anyClsBadTypes := map[uint16]bool{ dns.TypeIXFR : true, + dns.TypeAXFR : true, dns.TypeMAILB : true,dns.TypeMAILA : true} + + for _, rr := range rrs { + header := rr.Header() + if header.Class == updateZone.Qclass { + if badTypes[header.Rrtype] { + debugMsg("Bad type for class:", dns.ClassToString[header.Class], + header.Name, dns.TypeToString[header.Rrtype]) + return dns.RcodeFormatError + } + } else if header.Class == dns.ClassANY { + if header.Ttl != 0 || header.Rdlength != 0 || anyClsBadTypes[header.Rrtype] { + debugMsg("Bad ttl/length/type for class:", dns.ClassToString[header.Class], + header.Name, header.Ttl, header.Rdlength, dns.TypeToString[header.Rrtype]) + return dns.RcodeFormatError + } + } else if header.Class == dns.ClassNONE { + if header.Ttl != 0 || badTypes[header.Rrtype] { + debugMsg("Bad ttl/type for class:", dns.ClassToString[header.Class], + header.Name, header.Ttl, dns.TypeToString[header.Rrtype]) + return dns.RcodeFormatError + } + } else { + return dns.RcodeFormatError + } + // separately from the RFC validation, fail for RR types we don't understand yet + if _, ok := convertersFromRR[header.Rrtype]; ok != true { + debugMsg("Record converter doesn't exist for " + dns.TypeToString[header.Rrtype]) + return dns.RcodeServerFailure + } + } + return dns.RcodeSuccess +} + +// performUpdate will commit the requested updates to the database +// It is assumed by this point all prerequisites have been validated and all +// domains are locked. +// See RFC 2136, section 3.4.2 +func performUpdate(prefix string, etcdClient *etcd.Client, resolver *Resolver, updateZone dns.Question, records []dns.RR) (rcode int) { + + for _, rr := range records { + header := rr.Header() + if _, ok := convertersFromRR[header.Rrtype]; ok != true { + panic("Record converter doesn't exist for " + dns.TypeToString[header.Rrtype]) + } + + nameDir := nameToKey(header.Name, "") + + // Gather up all deletes + if header.Class == dns.ClassANY { + var typesToDelete []uint16 + if header.Rrtype == dns.TypeANY { + // RFC Meaning: Delete all RRsets from a name + + // Look up all RRs for supported types. Do this 'manually' because we + // don't want standard resolver behaviour about CNAMEs + // TODO: this means it's not parallelized like it normally is + for rrType, _ := range convertersToRR { + if header.Name == updateZone.Name && (rrType == dns.TypeNS || rrType == dns.TypeSOA) { + // RFC explicitly forbids deleting NS/SOA for the zone in this way + continue + } + typesToDelete = append(typesToDelete, rrType) + } + } else { + // RFC Meaning: Delete an RRset (all RRs of type) + typesToDelete = append(typesToDelete, header.Rrtype) + } + for _, rrType := range typesToDelete { + records, err := resolver.GetFromStorage(nameDir + "/." + dns.TypeToString[rrType]) + if err != nil && !missingKeyErr(err) { + debugMsg(err) + panic("Failed to fetch existing records for " + nameDir) + } + for _, toDelete := range records { + deleteKeyAndTtl(etcdClient, toDelete.node.Key) + } + } + } else { + + updateNode, err := convertRRToNode(rr, *header) + if err != nil { + panic("Got error when converting node") + } else if updateNode == nil { + panic("Got NIL after successfully converting node") + } + + existingRecords, err := resolver.GetFromStorage(updateNode.Key) + if err != nil && !missingKeyErr(err) { + debugMsg(err) + panic("Failed to fetch existing RRs for key " + updateNode.Key) + } + + if header.Class == dns.ClassNONE { + // RFC Meaning: Delete an RR from an RRset + for _, existing := range existingRecords { + if existing.node.Value == updateNode.Value { + deleteKeyAndTtl(etcdClient, existing.node.Key) + } + } + } else { + // RFC Meaning: Add to an RRset + + // Ignore certain inserts in presence of CNAMEs, as per RFC. + // Further explanation from http://docs.freebsd.org/doc/8.0-RELEASE/usr/share/doc/bind9/arm/man.nsupdate.html : + // "...cannot conflict with the long-standing rule in RFC1034 that a name must not exist as any other + // record type if it exists as a CNAME. (The rule has been updated for DNSSEC in RFC2535 to allow + // CNAMEs to have RRSIG, DNSKEY and NSEC records.)" + // TODO(orls): add these special cases to the special case for CNAMEs. Yayyyyy standards + + hasCNAME := false + hasNonCNAME := false + neighbouringTypes, err := etcdClient.Get(prefix + nameDir, false, false) + if err == nil { + for _, n := range neighbouringTypes.Node.Nodes { + splits := strings.Split(n.Key, nameDir + "/.") + if len(splits) != 2 || strings.HasSuffix(splits[1], ".ttl") { + continue + } + if splits[1] == "CNAME" { + hasCNAME = true + } else { + hasNonCNAME = true + } + } + } + + if header.Rrtype == dns.TypeCNAME && hasNonCNAME { + debugMsg("Ignoring insert for CNAME due to existing non-CNAME record(s) for " + header.Name) + continue + } else if header.Rrtype != dns.TypeCNAME && hasCNAME { + debugMsg("Ignoring insert for " + dns.TypeToString[header.Rrtype] + " due to existing CNAME record(s) for " + header.Name) + continue + } + + foundExisting := false + var ttlKeys []string + + // Check for existing matching records, in which case just update TTL. + // Otherwise there's risk of duplicates + for _, existing := range existingRecords { + if existing.node.Value == updateNode.Value { + debugMsg("update req matched existing rr with key " + existing.node.Key) + foundExisting = true + ttlKeys = append(ttlKeys, existing.node.Key + ".ttl") + } + } + + if !foundExisting { + // Then we need to add, which means we need a directory. + // Convert any old-style single-keys to directories + if len(existingRecords) == 1 && existingRecords[0].node.Key == "/" + prefix + updateNode.Key { + originalNode := existingRecords[0].node + logger.Printf("[WARNING] ------") + logger.Printf("[WARNING] Converting existing value to a directory!") + logger.Printf("[WARNING] Existing record is old-style single key: " + originalNode.Key) + logger.Printf("[WARNING] ------") + + convertedKey := originalNode.Key + "/" + recordSubkey(originalNode.Value) + _, convertErr := etcdClient.SetDir(originalNode.Key, 0) + if convertErr != nil { + debugMsg(convertErr) + // panic("Failed to insert record into etcd") + } + _, convertErr = etcdClient.Set(convertedKey, originalNode.Value, 0) + if convertErr != nil { + debugMsg(convertErr) + // panic("Failed to insert record into etcd") + } + if existingRecords[0].ttl != 0 { + convertTTL := strconv.FormatInt(int64(existingRecords[0].ttl), 10) + _, convertErr = etcdClient.Set(convertedKey + ".ttl", convertTTL, 0) + if convertErr != nil { + debugMsg(convertErr) + // panic("Failed to insert record into etcd") + } + } + } + + newKey := prefix + updateNode.Key + "/" + recordSubkey(updateNode.Value) + ttlKeys = append(ttlKeys, newKey + ".ttl") + + debugMsg("Inserting new record to " + newKey) + _, err := etcdClient.Set(newKey, updateNode.Value, 0) + if err != nil { + debugMsg(err) + // panic("Failed to insert record into etcd") + } + } + + // Insert the TTL record if one has been requested + if header.Ttl > 0 { + ttl := strconv.FormatInt(int64(header.Ttl), 10) + for _, ttlKey := range ttlKeys { + debugMsg("Inserting/updating TTL key " + ttlKey) + _, err = etcdClient.Set(ttlKey, ttl, 0) + if err != nil { + debugMsg(err) + // panic("Failed to insert ttl into etcd") + } + } + } + } + } + } + + return dns.RcodeSuccess +} + +// Internal boileplate-reducer to delete a specific key (non-recursive) and +// handle it's TTL key too +func deleteKeyAndTtl(etcdClient *etcd.Client, delKey string) { + delTTLKey := delKey + ".ttl" + debugMsg("Deleting RR with key " + delKey) + _, err := etcdClient.Delete(delKey, true) + if err != nil && !missingKeyErr(err) { + debugMsg(err) + panic("Failed to delete RRs with key " + delKey) + } + _, err = etcdClient.Delete(delTTLKey, true) + if err != nil && !missingKeyErr(err) { + debugMsg(err) + panic("Failed to delete RR TTL key " + delTTLKey) + } +} + +// recordSubkey yields the sub-key string to be used for a new RR in a +// directory, based on it's data. The MD5 of the node value is used, making +// duplicates impossible. +func recordSubkey(value string) (subkey string) { + hasher := md5.New() + hasher.Write([]byte(value)) + return hex.EncodeToString(hasher.Sum(nil)) +} + +// internal helper to determine if an error from etcd operations is a 100-code +// error, i.e. that the key is missing. +func missingKeyErr(err error) (ok bool) { + etcdErr, cast := err.(*etcd.EtcdError) + if cast && etcdErr.ErrorCode == 100 { + return true + } + return false +} diff --git a/update_test.go b/update_test.go new file mode 100644 index 0000000..9fb9c3b --- /dev/null +++ b/update_test.go @@ -0,0 +1,502 @@ +package main + +import ( + "github.com/miekg/dns" + "net" + "testing" + "reflect" +) + +func TestInsertNewRecordNoPrerequsites(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestRecordNoPrerequsites/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + record := &dns.A{ + Hdr: dns.RR_Header{Name: "disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 1234}, + A: net.ParseIP("1.2.3.4")} + + msg := &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.Insert([]dns.RR{record}) + + result := manager.Update("disco.net.", msg) + + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("Failed to insert new DNS record") + t.Fatal() + } + + answers, err := resolver.LookupAnswersForType("disco.net.", dns.TypeA) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 1 { + t.Error("Expected exactly one answer for discodns.net.") + t.Fatal() + } + answerHeader := answers[0].Header() + if answerHeader.Ttl != 1234 { + t.Error("Didn't get expected TTL on new record") + t.Fatal() + } +} + +func TestDeleteNameNoPrerequsites(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestDeleteNameNoPrerequsites/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestDeleteNameNoPrerequsites/", true) + client.Set("TestDeleteNameNoPrerequsites/net/disco/foo/.A", "1.1.1.1", 0) + client.Set("TestDeleteNameNoPrerequsites/net/disco/foo/.PTR/a", "a", 0) + client.Set("TestDeleteNameNoPrerequsites/net/disco/foo/.PTR/b", "b", 0) + client.Set("TestDeleteNameNoPrerequsites/net/disco/foo/.PTR/a.ttl", "100", 0) + + record := &dns.ANY{Hdr: dns.RR_Header{Name: "foo.disco.net."}} + + msg := &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.RemoveName([]dns.RR{record}) + + result := manager.Update("disco.net.", msg) + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("Failed to remove DNS record") + t.Fatal() + } + + answers, err := resolver.LookupAnswersForType("foo.disco.net.", dns.TypeANY) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) > 0 { + t.Error("Expected zero answers for foo.disco.net.") + t.Fatal() + } + + // Delete for something that doesn't already exist: + record = &dns.ANY{Hdr: dns.RR_Header{Name: "bar.disco.net."}} + msg = &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.RemoveName([]dns.RR{record}) + + result = manager.Update("disco.net.", msg) + if result.Rcode != dns.RcodeSuccess { + t.Error("Got failure from a no-op delete") + t.Fatal() + } +} + +func TestDeleteRecordsetNoPrerequsites(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestDeleteRecordsetNoPrerequsites/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestDeleteRecordsetNoPrerequsites/", true) + client.Set("TestDeleteRecordsetNoPrerequsites/net/disco/foo/.A", "1.1.1.1", 0) + client.Set("TestDeleteRecordsetNoPrerequsites/net/disco/foo/.PTR/a", "a", 0) + client.Set("TestDeleteRecordsetNoPrerequsites/net/disco/foo/.PTR/b", "b", 0) + client.Set("TestDeleteRecordsetNoPrerequsites/net/disco/foo/.PTR/a.ttl", "100", 0) + + record := &dns.PTR{Hdr: dns.RR_Header{Name: "foo.disco.net.", Rrtype: dns.TypePTR}, Ptr: "whatever"} + + msg := &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.RemoveRRset([]dns.RR{record}) + + result := manager.Update("disco.net.", msg) + if result.Rcode != dns.RcodeSuccess { + t.Error("Failed to remove DNS records") + t.Fatal() + } + + answers, err := resolver.LookupAnswersForType("foo.disco.net.", dns.TypePTR) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) > 0 { + t.Error("Expected zero answers for foo.disco.net. PTR") + t.Fatal() + } + + // Check the A record was left alone: + answers, err = resolver.LookupAnswersForType("foo.disco.net.", dns.TypeA) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 1 { + t.Error("Expected one answer for foo.disco.net. A") + t.Fatal() + } + + // Delete for something that doesn't already exist: + record = &dns.PTR{Hdr: dns.RR_Header{Name: "foo.disco.net.", Rrtype: dns.TypePTR}, Ptr: "whatever"} + msg = &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.RemoveRRset([]dns.RR{record}) + + result = manager.Update("disco.net.", msg) + if result.Rcode != dns.RcodeSuccess { + t.Error("Got failure from a no-op delete") + t.Fatal() + } +} + +func TestInsertMultipleRecords(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestInsertMultipleRecords/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + client.Delete("TestInsertMultipleRecords/", true) + + record1 := &dns.SRV{ + Hdr: dns.RR_Header{Name: "disco.net.", Rrtype: dns.TypeSRV, Class: dns.ClassINET}, + Port: 80, Priority: 100, Weight: 100, Target: "foo.disco.net"} + + record2 := &dns.TXT{ + Hdr: dns.RR_Header{Name: "disco.net.", Rrtype: dns.TypeTXT, Class: dns.ClassINET}, + Txt: []string{"lol"}} + + msg := &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.Insert([]dns.RR{record1, record2}) + + result := manager.Update("disco.net.", msg) + + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("Failed to add DNS records") + t.Fatal() + } + + srvAnswers, err := resolver.LookupAnswersForType("disco.net.", dns.TypeSRV) + if err != nil { + t.Error("Caught error retrieving SRV") + t.Fatal() + } + if len(srvAnswers) != 1 { + t.Error("Expected one SRV response") + t.Fatal() + } + txtAnswers, err := resolver.LookupAnswersForType("disco.net.", dns.TypeTXT) + if err != nil { + t.Error("Caught error retrieving txt") + t.Fatal() + } + if len(txtAnswers) != 1 { + t.Error("Expected one TXT response") + t.Fatal() + } +} + +// Internal utility to save boilerplate. Creates a message with the given +// prereqs and tries to perform an update +func _prereqsTestHelper(t *testing.T, manager *DynamicUpdateManager, prereqMethod string, expected int, prereqs []dns.RR) (pass bool) { + + recordToAdd := &dns.A{ + Hdr: dns.RR_Header{Name: "baz.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET}, + A: net.ParseIP("1.2.3.4")} + + msg := &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net.", Qclass: dns.ClassINET}) + msg.Insert([]dns.RR{recordToAdd}) + reflPrereqs := reflect.ValueOf(prereqs) + v := reflect.ValueOf(msg) + m := v.MethodByName(prereqMethod) + m.Call([]reflect.Value{reflPrereqs}) + + var errorMsg string + if (expected == dns.RcodeSuccess) { + errorMsg = "Failed to add DNS record with `" + prereqMethod +"` prereq, got" + } else { + errorMsg = "Expected update with `" + prereqMethod +"` prereqs to fail with " + dns.RcodeToString[expected] + ", got" + } + + result := manager.Update("disco.net.", msg) + if result.Rcode != expected { + debugMsg(result) + t.Error(errorMsg, dns.RcodeToString[result.Rcode]) + return false + } + return true +} + +func TestPrerequisites_NameInUse(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestPrerequisites_NameInUse/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestPrerequisites_NameInUse/", true) + client.Set("TestPrerequisites_NameInUse/net/disco/foo/.A", "1.1.1.1", 0) + + prereq_fail := &dns.ANY{ Hdr: dns.RR_Header{Name: "foofoo.disco.net.", Rrtype: dns.TypeANY, Class: dns.ClassINET}} + prereq_ok := &dns.ANY{ Hdr: dns.RR_Header{Name: "foo.disco.net.", Rrtype: dns.TypeANY, Class: dns.ClassINET}} + + if ! _prereqsTestHelper(t, manager, "NameUsed", dns.RcodeNameError, []dns.RR{prereq_fail}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "NameUsed", dns.RcodeNameError, []dns.RR{prereq_fail, prereq_ok}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "NameUsed", dns.RcodeSuccess, []dns.RR{prereq_ok}) { t.Fatal() } +} + +func TestPrerequisites_NameNotInUse(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestPrerequisites_NameNotInUse/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestPrerequisites_NameNotInUse/", true) + client.Set("TestPrerequisites_NameNotInUse/net/disco/foo/.A", "1.1.1.1", 0) + + prereq_fail := &dns.ANY{ Hdr: dns.RR_Header{Name: "foo.disco.net.", Rrtype: dns.TypeANY, Class: dns.ClassINET}} + prereq_ok := &dns.ANY{ Hdr: dns.RR_Header{Name: "foofoo.disco.net.", Rrtype: dns.TypeANY, Class: dns.ClassINET}} + + if ! _prereqsTestHelper(t, manager, "NameNotUsed", dns.RcodeYXDomain, []dns.RR{prereq_fail}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "NameNotUsed", dns.RcodeYXDomain, []dns.RR{prereq_fail, prereq_ok}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "NameNotUsed", dns.RcodeSuccess, []dns.RR{prereq_ok}) { t.Fatal() } +} + +func TestPrerequisites_ValueIndependentRRSet(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestPrerequisites_ValueIndependentRRSet/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestPrerequisites_ValueIndependentRRSet/", true) + client.Set("TestPrerequisites_ValueIndependentRRSet/net/disco/foo/.A", "1.1.1.1", 0) + client.Set("TestPrerequisites_ValueIndependentRRSet/net/disco/bar/.A", "1.1.1.1", 0) + client.Set("TestPrerequisites_ValueIndependentRRSet/net/disco/bar/.PTR", "bar.disco.net", 0) + + prereq_foo_a := &dns.ANY{ Hdr: dns.RR_Header{Name: "foo.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET}} + prereq_foo_ptr := &dns.ANY{ Hdr: dns.RR_Header{Name: "foo.disco.net.", Rrtype: dns.TypePTR, Class: dns.ClassINET}} + prereq_bar_a := &dns.ANY{ Hdr: dns.RR_Header{Name: "bar.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET}} + prereq_bar_ptr := &dns.ANY{ Hdr: dns.RR_Header{Name: "bar.disco.net.", Rrtype: dns.TypePTR, Class: dns.ClassINET}} + + if ! _prereqsTestHelper(t, manager, "RRsetNotUsed", dns.RcodeSuccess, []dns.RR{prereq_foo_ptr}) { t.Fatal() } + + if ! _prereqsTestHelper(t, manager, "RRsetUsed", dns.RcodeNXRrset, []dns.RR{prereq_foo_a, prereq_foo_ptr}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "RRsetNotUsed", dns.RcodeYXRrset, []dns.RR{prereq_foo_a, prereq_foo_ptr}) { t.Fatal() } + + if ! _prereqsTestHelper(t, manager, "RRsetUsed", dns.RcodeNXRrset, []dns.RR{prereq_foo_ptr, prereq_bar_ptr}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "RRsetNotUsed", dns.RcodeYXRrset, []dns.RR{prereq_foo_ptr, prereq_bar_ptr}) { t.Fatal() } + + if ! _prereqsTestHelper(t, manager, "RRsetUsed", dns.RcodeSuccess, []dns.RR{prereq_bar_a, prereq_bar_ptr}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "RRsetNotUsed", dns.RcodeYXRrset, []dns.RR{prereq_bar_a, prereq_bar_ptr}) { t.Fatal() } + + if ! _prereqsTestHelper(t, manager, "RRsetUsed", dns.RcodeSuccess, []dns.RR{prereq_foo_a, prereq_bar_a}) { t.Fatal() } +} + +func TestPrerequisites_ValueDependentRRSet(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestPrerequisites_ValueDependentRRSet/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestPrerequisites_ValueDependentRRSet/", true) + client.Set("TestPrerequisites_ValueDependentRRSet/net/disco/foo/.A", "1.1.1.1", 0) + client.Set("TestPrerequisites_ValueDependentRRSet/net/disco/bar/.A", "1.1.1.1", 0) + client.Set("TestPrerequisites_ValueDependentRRSet/net/disco/bar/.PTR", "match.disco.net", 0) + + prereq_foo_a_match, _ := dns.NewRR("foo.disco.net. 0 IN A 1.1.1.1") + prereq_foo_a_miss, _ := dns.NewRR("foo.disco.net. 0 IN A 2.2.2.2") + prereq_bar_a_match, _ := dns.NewRR("bar.disco.net. 0 IN A 1.1.1.1") + prereq_bar_a_miss, _ := dns.NewRR("bar.disco.net. 0 IN A 2.2.2.2") + prereq_bar_ptr_match, _ := dns.NewRR("bar.disco.net. 0 IN PTR match.disco.net") + prereq_bar_ptr_miss, _ := dns.NewRR("bar.disco.net. 0 IN PTR miss.disco.net") + + if ! _prereqsTestHelper(t, manager, "Used", dns.RcodeNXRrset, []dns.RR{prereq_foo_a_miss}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "Used", dns.RcodeNXRrset, []dns.RR{prereq_foo_a_miss, prereq_bar_a_miss}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "Used", dns.RcodeNXRrset, []dns.RR{prereq_foo_a_miss, prereq_bar_ptr_miss}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "Used", dns.RcodeNXRrset, []dns.RR{prereq_foo_a_match, prereq_bar_a_miss}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "Used", dns.RcodeNXRrset, []dns.RR{prereq_foo_a_match, prereq_bar_ptr_miss}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "Used", dns.RcodeSuccess, []dns.RR{prereq_foo_a_match, prereq_bar_a_match}) { t.Fatal() } + if ! _prereqsTestHelper(t, manager, "Used", dns.RcodeSuccess, []dns.RR{prereq_foo_a_match, prereq_bar_a_match, prereq_bar_ptr_match}) { t.Fatal() } +} + +func TestUpsertExisting(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestUpsertExisting/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestUpsertExisting/", true) + client.Set("TestUpsertExisting/net/disco/singlekey/.A", "1.1.1.1", 0) + client.Set("TestUpsertExisting/net/disco/singlekey/.A.ttl", "123", 0) + client.Set("TestUpsertExisting/net/disco/directory/.A/6465ec74397c9126916786bbcd6d7601", "1.1.1.1", 0) + client.Set("TestUpsertExisting/net/disco/directory/.A/6465ec74397c9126916786bbcd6d7601.ttl", "123", 0) + client.Set("TestUpsertExisting/net/disco/directory/.A/nonMd5KeyName", "2.2.2.2", 0) + client.Set("TestUpsertExisting/net/disco/directory/.A/nonMd5KeyName.ttl", "123", 0) + + // Update with same value (to a non-directory key): TTL should change + updateSingle := &dns.A{ + Hdr: dns.RR_Header{Name: "singlekey.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 1234}, + A: net.ParseIP("1.1.1.1")} + + msg := &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net.", Qclass: dns.ClassINET}) + msg.Insert([]dns.RR{updateSingle}) + + result := manager.Update("disco.net.", msg) + + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("Failed to update existing DNS record") + t.Fatal() + } + + answers, err := resolver.LookupAnswersForType("singlekey.disco.net.", dns.TypeA) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 1 { + t.Error("Expected exactly one answer for discodns.net.") + t.Fatal() + } + answerHeader := answers[0].Header() + if answerHeader.Ttl != 1234 { + t.Error("Didn't get expected TTL on new record") + t.Fatal() + } + + // Insert a new one: should auto-convert single-value to directory? + // TODO: not sure what we should consider correct behaviour here. + addNewToSingle := &dns.A{ + Hdr: dns.RR_Header{Name: "singlekey.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 1234}, + A: net.ParseIP("2.2.2.2")} + + msg = &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net.", Qclass: dns.ClassINET}) + msg.Insert([]dns.RR{addNewToSingle}) + + result = manager.Update("disco.net.", msg) + + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("Failed to insert new DNS record to single-value (non-directory) node") + t.Error(" -- (Permitting test to continue for now...) --") + // t.Fatal() + } + + answers, err = resolver.LookupAnswersForType("singlekey.disco.net.", dns.TypeA) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 2 { + t.Error("Expected two answers for singlekey.discodns.net. after update") + t.Error(" -- (Permitting test to continue for now...) --") + // t.Fatal() + } + + // Update with same value (to a directory child key, md5 subkey): TTL should change + updateDirChild := &dns.A{ + Hdr: dns.RR_Header{Name: "directory.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 1234}, + A: net.ParseIP("1.1.1.1")} + + msg = &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net.", Qclass: dns.ClassINET}) + msg.Insert([]dns.RR{updateDirChild}) + + result = manager.Update("disco.net.", msg) + + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("Failed to update existing record (directory child with hashed subkey)") + t.Fatal() + } + + answers, err = resolver.LookupAnswersForType("directory.disco.net.", dns.TypeA) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 2 { + t.Error("Expected two answers for directory.discodns.net.") + t.Fatal() + } + + // Update with same value (to a directory child key, non-md5 subkey): TTL should change + updateDirChildMessyName := &dns.A{ + Hdr: dns.RR_Header{Name: "directory.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 1234}, + A: net.ParseIP("2.2.2.2")} + + msg = &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net.", Qclass: dns.ClassINET}) + msg.Insert([]dns.RR{updateDirChildMessyName}) + + result = manager.Update("disco.net.", msg) + + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("Failed to update existing record (directory child with messy unhashed subkey)") + t.Fatal() + } + + answers, err = resolver.LookupAnswersForType("directory.disco.net.", dns.TypeA) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 2 { + t.Error("Expected two answers for directory.discodns.net.") + t.Fatal() + } + + for _, answer := range answers { + header := answer.Header() + if header.Ttl != 1234 { + t.Error("Didn't get expected TTL on directory.discodns.net record", header.Name) + t.Fatal() + } + } +} + +func TestInsertCname(t *testing.T) { + manager := &DynamicUpdateManager{etcd: client, etcdPrefix: "TestInsertCname/", resolver: resolver} + resolver.etcdPrefix = manager.etcdPrefix + + client.Delete("TestInsertCname/", true) + client.Set("TestInsertCname/net/disco/target/.A", "1.1.1.1", 0) + client.Set("TestInsertCname/net/disco/alias/.CNAME", "target.disco.net.", 0) + + newCname := &dns.CNAME{ + Hdr: dns.RR_Header{Name: "target.disco.net.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 1234}, + Target: "foo.disco.net."} + msg := &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.Insert([]dns.RR{newCname}) + + result := manager.Update("disco.net.", msg) + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("DNS update query failed") + t.Fatal() + } + + // despite success, the cname insert should have been ignored + answers, err := resolver.LookupAnswersForType("target.disco.net.", dns.TypeCNAME) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 0 { + t.Error("Expected no answers for target.discodns.net CNAME") + t.Fatal() + } + + newA := &dns.A{ + Hdr: dns.RR_Header{Name: "alias.disco.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 1234}, + A: net.ParseIP("1.2.3.4")} + msg = &dns.Msg{} + msg.Question = append(msg.Question, dns.Question{Name: "disco.net."}) + msg.Insert([]dns.RR{newA}) + + result = manager.Update("disco.net.", msg) + if result.Rcode != dns.RcodeSuccess { + debugMsg(result) + t.Error("DNS update query failed") + t.Fatal() + } + + // despite success, the A insert should have been ignored + answers, err = resolver.LookupAnswersForType("alias.disco.net.", dns.TypeA) + if err != nil { + t.Error("Caught error resolving domain") + t.Fatal() + } + if len(answers) != 0 { + t.Error("Expected no answers for alias.discodns.net A") + t.Fatal() + } +}