diff --git a/protocol/cloudflare/inbound.go b/protocol/cloudflare/inbound.go index f445ab956..2ad3fb296 100644 --- a/protocol/cloudflare/inbound.go +++ b/protocol/cloudflare/inbound.go @@ -9,18 +9,19 @@ import ( "github.com/sagernet/sing-box/adapter" "github.com/sagernet/sing-box/adapter/inbound" - boxDialer "github.com/sagernet/sing-box/common/dialer" + "github.com/sagernet/sing-box/common/dialer" C "github.com/sagernet/sing-box/constant" "github.com/sagernet/sing-box/log" "github.com/sagernet/sing-box/option" "github.com/sagernet/sing-box/route/rule" - cloudflared "github.com/sagernet/sing-cloudflared" - tun "github.com/sagernet/sing-tun" + "github.com/sagernet/sing-cloudflared" + "github.com/sagernet/sing-tun" "github.com/sagernet/sing/common/bufio" E "github.com/sagernet/sing/common/exceptions" M "github.com/sagernet/sing/common/metadata" N "github.com/sagernet/sing/common/network" "github.com/sagernet/sing/common/pipe" + "github.com/sagernet/sing/service" ) func RegisterInbound(registry *inbound.Registry) { @@ -28,28 +29,35 @@ func RegisterInbound(registry *inbound.Registry) { } func NewInbound(ctx context.Context, router adapter.Router, logger log.ContextLogger, tag string, options option.CloudflaredInboundOptions) (adapter.Inbound, error) { - controlDialer, err := boxDialer.NewWithOptions(boxDialer.Options{ - Context: ctx, - Options: options.ControlDialer, - RemoteIsDomain: true, + controlDialer, err := dialer.NewWithOptions(dialer.Options{ + Context: ctx, + Options: options.ControlDialer, + RemoteIsDomain: true, + ResolverOnDetour: true, }) if err != nil { return nil, E.Cause(err, "build cloudflared control dialer") } - tunnelDialer, err := boxDialer.NewWithOptions(boxDialer.Options{ - Context: ctx, - Options: options.TunnelDialer, - RemoteIsDomain: true, + tunnelDialer, err := dialer.NewWithOptions(dialer.Options{ + Context: ctx, + Options: options.TunnelDialer, + RemoteIsDomain: true, + ResolverOnDetour: true, }) if err != nil { return nil, E.Cause(err, "build cloudflared tunnel dialer") } + dnsRouter := service.FromContext[adapter.DNSRouter](ctx) + controlResolver := newRouterResolver(dnsRouter, controlDialer.(dialer.ResolveDialer).QueryOptions()) + tunnelResolver := newRouterResolver(dnsRouter, tunnelDialer.(dialer.ResolveDialer).QueryOptions()) service, err := cloudflared.NewService(cloudflared.ServiceOptions{ Logger: logger, ConnectionDialer: &routerDialer{router: router, tag: tag}, ControlDialer: controlDialer, TunnelDialer: tunnelDialer, + ControlResolver: controlResolver, + TunnelResolver: tunnelResolver, ICMPHandler: &icmpRouterHandler{router: router, logger: logger, tag: tag}, ConnContext: func(connCtx context.Context) context.Context { return adapter.WithContext(connCtx, &adapter.InboundContext{ diff --git a/protocol/cloudflare/resolver.go b/protocol/cloudflare/resolver.go new file mode 100644 index 000000000..7253d030d --- /dev/null +++ b/protocol/cloudflare/resolver.go @@ -0,0 +1,57 @@ +//go:build with_cloudflared + +package cloudflare + +import ( + "context" + "net" + "net/netip" + "sort" + "strings" + + "github.com/sagernet/sing-box/adapter" + + mDNS "github.com/miekg/dns" +) + +type routerResolver struct { + dnsRouter adapter.DNSRouter + queryOptions adapter.DNSQueryOptions +} + +func newRouterResolver(dnsRouter adapter.DNSRouter, queryOptions adapter.DNSQueryOptions) *routerResolver { + return &routerResolver{dnsRouter: dnsRouter, queryOptions: queryOptions} +} + +func (r *routerResolver) LookupNetIP(ctx context.Context, host string) ([]netip.Addr, error) { + return r.dnsRouter.Lookup(ctx, strings.TrimSuffix(host, "."), r.queryOptions) +} + +func (r *routerResolver) LookupSRV(ctx context.Context, service, proto, name string) ([]*net.SRV, error) { + message := &mDNS.Msg{} + message.SetQuestion(mDNS.Fqdn("_"+service+"._"+proto+"."+name), mDNS.TypeSRV) + response, err := r.dnsRouter.Exchange(ctx, message, r.queryOptions) + if err != nil { + return nil, err + } + var records []*net.SRV + for _, answer := range response.Answer { + record, isSRV := answer.(*mDNS.SRV) + if !isSRV { + continue + } + records = append(records, &net.SRV{ + Target: record.Target, + Port: record.Port, + Priority: record.Priority, + Weight: record.Weight, + }) + } + sort.SliceStable(records, func(i, j int) bool { + if records[i].Priority != records[j].Priority { + return records[i].Priority < records[j].Priority + } + return records[i].Weight > records[j].Weight + }) + return records, nil +}