dns: Add neighbor-based hostname resolution to local server

This commit is contained in:
世界
2026-08-30 17:41:40 +08:00
parent 35540b4ff3
commit f8a623e0c6
16 changed files with 290 additions and 71 deletions
+27 -7
View File
@@ -14,6 +14,7 @@ import (
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing/service"
mDNS "github.com/miekg/dns"
)
@@ -26,12 +27,14 @@ var _ adapter.DNSTransport = (*Transport)(nil)
type Transport struct {
dns.TransportAdapter
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
preferGo bool
resolved ResolvedResolver
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
preferGo bool
resolved ResolvedResolver
neighborResolver adapter.NeighborResolver
neighborSuffixes []string
}
func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, options option.LocalDNSServerOptions) (adapter.DNSTransport, error) {
@@ -39,13 +42,17 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
if err != nil {
return nil, err
}
suffixes, err := buildNeighborMatchers(options.NeighborDomain)
if err != nil {
return nil, err
}
return &Transport{
TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeLocal, tag, options),
ctx: ctx,
logger: logger,
dialer: transportDialer,
preferGo: options.PreferGo,
neighborSuffixes: suffixes,
}, nil
}
@@ -71,6 +78,11 @@ func (t *Transport) Start(stage adapter.StartStage) error {
}
}
}
case adapter.StartStateStart:
router := service.FromContext[adapter.Router](t.ctx)
if router != nil {
t.neighborResolver = router.NeighborResolver()
}
}
return nil
}
@@ -87,6 +99,10 @@ func (t *Transport) Reset() {
func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) {
if t.resolved != nil {
response := t.lookupNeighbor(message)
if response != nil {
return response, nil
}
return t.resolved.Exchange(ctx, message)
}
question := message.Question[0]
@@ -96,5 +112,9 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL), nil
}
}
response := t.lookupNeighbor(message)
if response != nil {
return response, nil
}
return t.exchange(ctx, message, question.Name)
}
+38 -27
View File
@@ -28,12 +28,14 @@ var _ adapter.DNSTransport = (*Transport)(nil)
type Transport struct {
dns.TransportAdapter
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
fallback bool
dhcpTransport dhcpTransport
ctx context.Context
logger logger.ContextLogger
hosts *hosts.File
dialer N.Dialer
fallback bool
dhcpTransport dhcpTransport
neighborResolver adapter.NeighborResolver
neighborSuffixes []string
}
type dhcpTransport interface {
@@ -47,39 +49,48 @@ func NewTransport(ctx context.Context, logger log.ContextLogger, tag string, opt
if err != nil {
return nil, err
}
suffixes, err := buildNeighborMatchers(options.NeighborDomain)
if err != nil {
return nil, err
}
return &Transport{
TransportAdapter: dns.NewTransportAdapterWithLocalOptions(C.DNSTypeLocal, tag, options),
ctx: ctx,
logger: logger,
dialer: transportDialer,
neighborSuffixes: suffixes,
}, nil
}
func (t *Transport) Start(stage adapter.StartStage) error {
if stage != adapter.StartStateStart {
return nil
}
defaultHosts, err := hosts.NewDefault()
if err != nil {
t.logger.Warn(err)
} else {
t.hosts = defaultHosts
}
inboundManager := service.FromContext[adapter.InboundManager](t.ctx)
for _, inbound := range inboundManager.Inbounds() {
if inbound.Type() == C.TypeTun {
t.fallback = true
break
switch stage {
case adapter.StartStateStart:
defaultHosts, err := hosts.NewDefault()
if err != nil {
t.logger.Warn(err)
} else {
t.hosts = defaultHosts
}
}
if t.fallback {
t.dhcpTransport = newDHCPTransport(t.TransportAdapter, log.ContextWithOverrideLevel(t.ctx, log.LevelDebug), t.dialer, t.logger)
if t.dhcpTransport != nil {
err := t.dhcpTransport.Start(stage)
if err != nil {
return err
inboundManager := service.FromContext[adapter.InboundManager](t.ctx)
for _, inbound := range inboundManager.Inbounds() {
if inbound.Type() == C.TypeTun {
t.fallback = true
break
}
}
if t.fallback {
t.dhcpTransport = newDHCPTransport(t.TransportAdapter, log.ContextWithOverrideLevel(t.ctx, log.LevelDebug), t.dialer, t.logger)
if t.dhcpTransport != nil {
err = t.dhcpTransport.Start(stage)
if err != nil {
return err
}
}
}
router := service.FromContext[adapter.Router](t.ctx)
if router != nil {
t.neighborResolver = router.NeighborResolver()
}
}
return nil
}
+5 -1
View File
@@ -86,7 +86,11 @@ func (t *Transport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg,
return dns.FixedResponse(message.Id, question, addresses, boxC.DefaultDNSTTL), nil
}
}
if t.fallback && t.dhcpTransport != nil {
response := t.lookupNeighbor(message)
if response != nil {
return response, nil
}
if t.dhcpTransport != nil {
dhcpServers := t.dhcpTransport.Fetch()
if len(dhcpServers) > 0 {
return t.dhcpTransport.Exchange0(ctx, message, dhcpServers)
+57
View File
@@ -0,0 +1,57 @@
package local
import (
"strings"
C "github.com/sagernet/sing-box/constant"
"github.com/sagernet/sing-box/dns"
E "github.com/sagernet/sing/common/exceptions"
mDNS "github.com/miekg/dns"
)
func buildNeighborMatchers(domains []string) ([]string, error) {
if len(domains) == 0 {
return nil, nil
}
var suffixes []string
for _, domain := range domains {
if !strings.HasPrefix(domain, ".") {
return nil, E.New("neighbor_domain entry must start with '.': ", domain)
}
suffixes = append(suffixes, mDNS.CanonicalName(domain))
}
return suffixes, nil
}
func (t *Transport) lookupNeighbor(message *mDNS.Msg) *mDNS.Msg {
if t.neighborResolver == nil {
return nil
}
question := message.Question[0]
if question.Qtype != mDNS.TypeA && question.Qtype != mDNS.TypeAAAA {
return nil
}
host := extractNeighborHost(mDNS.CanonicalName(question.Name), t.neighborSuffixes)
if host == "" {
return nil
}
addresses := t.neighborResolver.LookupAddresses(host)
if len(addresses) == 0 {
return nil
}
return dns.FixedResponse(message.Id, question, addresses, C.DefaultDNSTTL)
}
func extractNeighborHost(canonical string, suffixes []string) string {
for _, suffix := range suffixes {
if !strings.HasSuffix(canonical, suffix) || len(canonical) <= len(suffix) {
continue
}
host := canonical[:len(canonical)-len(suffix)]
if !strings.ContainsRune(host, '.') {
return host
}
}
return ""
}