diff --git a/server_test.go b/server_test.go index d726c50..858426b 100644 --- a/server_test.go +++ b/server_test.go @@ -22,7 +22,51 @@ func TestServer_StartStop(t *testing.T) { } func TestServer_Lookup(t *testing.T) { - serv, err := NewServer(&Config{Zone: makeServiceWithServiceName(t, "_foobar._tcp")}) + entries := make(chan *ServiceEntry, 1) + errCh := make(chan error, 1) + defer close(errCh) + testCases := []struct { + description string + checkup func() + }{ + {description: "normal case", checkup: func() { + select { + case e := <-entries: + if e.Name != "hostname._foobar._tcp.local." { + errCh <- fmt.Errorf("Entry has the wrong name: %+v", e) + return + } + if e.Port != 80 { + errCh <- fmt.Errorf("Entry has the wrong port: %+v", e) + return + } + if e.Info != "Local web server" { + errCh <- fmt.Errorf("Entry as the wrong Info: %+v", e) + return + } + errCh <- nil + case <-time.After(80 * time.Millisecond): + errCh <- fmt.Errorf("Timed out waiting for response") + } + }}, { + description: "change txt", + checkup: func() { + select { + case e := <-entries: + if e.Info != "a=a|b=b" { + errCh <- fmt.Errorf("Entry as the wrong Info: %+v", e) + return + } + errCh <- nil + case <-time.After(80 * time.Millisecond): + errCh <- fmt.Errorf("Timed out waiting for response") + } + }, + }, + } + + svc := makeServiceWithServiceName(t, "_foobar._tcp") + serv, err := NewServer(&Config{Zone: svc}) if err != nil { t.Fatalf("err: %v", err) } @@ -31,45 +75,26 @@ func TestServer_Lookup(t *testing.T) { t.Fatalf("err: %v", err) } }() - - entries := make(chan *ServiceEntry, 1) - errCh := make(chan error, 1) - defer close(errCh) - go func() { - select { - case e := <-entries: - if e.Name != "hostname._foobar._tcp.local." { - errCh <- fmt.Errorf("Entry has the wrong name: %+v", e) - return - } - if e.Port != 80 { - errCh <- fmt.Errorf("Entry has the wrong port: %+v", e) - return - } - if e.Info != "Local web server" { - errCh <- fmt.Errorf("Entry as the wrong Info: %+v", e) - return - } - errCh <- nil - case <-time.After(80 * time.Millisecond): - errCh <- fmt.Errorf("Timed out waiting for response") + for idx, testCase := range testCases { + go testCase.checkup() + if idx == 1 { + svc.UpdateTXT([]string{"a=a", "b=b"}) + } + params := &QueryParam{ + Service: "_foobar._tcp", + Domain: "local", + Timeout: 50 * time.Millisecond, + Entries: entries, + DisableIPv6: true, } - }() - - params := &QueryParam{ - Service: "_foobar._tcp", - Domain: "local", - Timeout: 50 * time.Millisecond, - Entries: entries, - DisableIPv6: true, - } - err = Query(params) - if err != nil { - t.Fatalf("err: %v", err) - } - err = <-errCh - if err != nil { - t.Fatalf("err: %v", err) + err = Query(params) + if err != nil { + t.Fatalf("description: %s, err: %v", testCase.description, err) + } + err = <-errCh + if err != nil { + t.Fatalf("description: %s, err: %v", testCase.description, err) + } } } diff --git a/zone.go b/zone.go index c9e8664..679f1f0 100644 --- a/zone.go +++ b/zone.go @@ -8,6 +8,7 @@ import ( "net" "os" "strings" + "sync" "github.com/miekg/dns" ) @@ -26,6 +27,7 @@ type Zone interface { // MDNSService is used to export a named service by implementing a Zone type MDNSService struct { + mutex sync.Mutex Instance string // Instance name (e.g. "hostService name") Service string // Service name (e.g. "_http._tcp.") Domain string // If blank, assumes "local" @@ -302,9 +304,21 @@ func (m *MDNSService) instanceRecords(q dns.Question) []dns.RR { Class: dns.ClassINET, Ttl: defaultTTL, }, - Txt: m.TXT, + Txt: m.GetTXT(), } return []dns.RR{txt} } return nil } + +func (m *MDNSService) UpdateTXT(txt []string) { + m.mutex.Lock() + defer m.mutex.Unlock() + m.TXT = txt +} + +func (m *MDNSService) GetTXT() []string { + m.mutex.Lock() + defer m.mutex.Unlock() + return m.TXT +}