From e166e2ecd79dd84150f7f5293067a86eff4bc073 Mon Sep 17 00:00:00 2001 From: nerpa Date: Mon, 14 Sep 2026 15:59:53 +0300 Subject: [PATCH] fix(amnezia): adopt DNS field fix from wireguard Adopt fix from https://github.com/XTLS/Xray-core/commit/c7e569b0377724600af1ea2a05eb8f4c7c3e0609 for Amnezia protocol, now user's DNS is work properly --- infra/conf/amnezia.go | 2 + proxy/amnezia/client.go | 75 +++++++++++++++++++++++++++--------- proxy/amnezia/config.pb.go | 79 +++++++++++++++++++++----------------- proxy/amnezia/config.proto | 33 ++++++++-------- proxy/amnezia/netstack.go | 21 +++++----- 5 files changed, 131 insertions(+), 79 deletions(-) diff --git a/infra/conf/amnezia.go b/infra/conf/amnezia.go index 58ea0fa5..08d1f6ec 100644 --- a/infra/conf/amnezia.go +++ b/infra/conf/amnezia.go @@ -66,6 +66,7 @@ type AmneziaConfig struct { MTU int32 `json:"mtu"` Reserved []byte `json:"reserved"` DomainStrategy string `json:"domainStrategy"` + DNS []string `json:"remoteDNS"` // Amnezia stuff (version 2.0) JunkCount int32 `json:"jc"` JunkMin int32 `json:"jmin"` @@ -161,6 +162,7 @@ func (c *AmneziaConfig) Build() (proto.Message, error) { config.IsClient = c.IsClient config.NoKernelTun = c.NoKernelTun + config.DNS = c.DNS config.Jc = c.JunkCount config.JMin = c.JunkMin diff --git a/proxy/amnezia/client.go b/proxy/amnezia/client.go index 2835a15a..d7edf33a 100644 --- a/proxy/amnezia/client.go +++ b/proxy/amnezia/client.go @@ -5,9 +5,10 @@ import ( "fmt" gonet "net" "net/netip" - reflect "reflect" + "reflect" "strings" "sync" + "time" "github.com/amnezia-vpn/amneziawg-go/tun" @@ -30,6 +31,11 @@ import ( "github.com/xtls/xray-core/transport/internet" ) +type entry struct { + got []net.IP + time time.Time +} + type Handler struct { conf *DeviceConfig policyManager policy.Manager @@ -43,6 +49,11 @@ type Handler struct { tnet *Net dev *device.Device mu sync.Mutex + + // TODO: cache cleanup loop + local bool + cache map[string]entry + cacheMu sync.Mutex } func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) { @@ -98,6 +109,20 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) { return nil, err } + local := false + dns := conf.DNS + if len(dns) == 0 { + dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"} + } + if len(dns) == 1 && dns[0] == "local" { + local = true + dns = nil + } + dnses := make([]netip.Addr, 0, len(dns)) + for _, dns := range dns { + dnses = append(dnses, netip.MustParseAddr(dns)) + } + kernelTunSupported, err := KernelTunSupported() if err != nil { errors.LogWarningInner(context.Background(), err, "Failed to check kernel TUN support") @@ -106,10 +131,10 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) { var tnet *Net if !conf.NoKernelTun && kernelTunSupported { errors.LogWarning(context.Background(), "Using kernel TUN") - tun, tnet, err = createKernelTun(localAddresses, []netip.Addr{netip.MustParseAddr("1.1.1.1"), netip.MustParseAddr("1.0.0.1"), netip.MustParseAddr("2606:4700:4700::1111"), netip.MustParseAddr("2606:4700:4700::1001")}, int(conf.Mtu)) + tun, tnet, err = createKernelTun(localAddresses, dnses, int(conf.Mtu)) } else { errors.LogWarning(context.Background(), "Using gVisor TUN") - tun, tnet, _, err = CreateNetTUN(localAddresses, []netip.Addr{netip.MustParseAddr("1.1.1.1"), netip.MustParseAddr("1.0.0.1"), netip.MustParseAddr("2606:4700:4700::1111"), netip.MustParseAddr("2606:4700:4700::1001")}, int(conf.Mtu), true) + tun, tnet, _, err = CreateNetTUN(localAddresses, dnses, int(conf.Mtu), true) } if err != nil { return nil, err @@ -126,6 +151,9 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) { tun: tun, tnet: tnet, + + local: local, + cache: make(map[string]entry), }, nil } @@ -138,6 +166,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte } ob.Name = "amnezia" ob.CanSpliceCopy = 3 + dialer.SetOutboundGateway(ctx, ob) if err := h.init(ctx); err != nil { return err @@ -335,7 +364,6 @@ func (h *Handler) init(ctx context.Context) error { cfg.WriteString("i3=" + fmt.Sprint(h.conf.I3) + "\n") cfg.WriteString("i4=" + fmt.Sprint(h.conf.I4) + "\n") cfg.WriteString("i5=" + fmt.Sprint(h.conf.I5) + "\n") - for _, peer := range h.conf.Peers { cfg.WriteString("public_key=" + peer.PublicKey + "\n") if peer.PreSharedKey != "" { @@ -362,31 +390,34 @@ func (h *Handler) init(ctx context.Context) error { } func (h *Handler) resolveLocal(host string) (net.IP, error) { - return resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, error) { - ips, _, err := h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true}) - return ips, err + return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) { + return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true}) }) } func (h *Handler) resolveRemote(host string) (net.IP, error) { - return resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, error) { - addrs, err := h.tnet.LookupHost(host) - if err != nil { - return nil, err + return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) { + if h.local { + return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true}) } - ips := make([]net.IP, 0, len(addrs)) - for _, addr := range addrs { - ips = append(ips, net.ParseIP(addr)) - } - return ips, nil + return h.tnet.LookupHost(host) }) } -func resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP func(host string) ([]net.IP, error)) (net.IP, error) { +func (h *Handler) resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP func(host string) ([]net.IP, uint32, error)) (net.IP, error) { if ip := net.ParseIP(host); ip != nil { return ip, nil } - ips, err := lookupIP(host) + h.cacheMu.Lock() + if entry, ok := h.cache[host]; ok { + if time.Now().Before(entry.time) { + h.cacheMu.Unlock() + return entry.got[dice.Roll(len(entry.got))], nil + } + delete(h.cache, host) + } + h.cacheMu.Unlock() + ips, ttl, err := lookupIP(host) if err != nil { return nil, err } @@ -401,7 +432,6 @@ func resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP f got6 = append(got6, ip) } } - var got []net.IP switch strategy { case DeviceConfig_FORCE_IP: @@ -427,6 +457,13 @@ func resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP f if len(got) == 0 { return nil, dns.ErrEmptyResponse } + entry := entry{ + got: got, + time: time.Now().Add(time.Duration(ttl) * time.Second), + } + h.cacheMu.Lock() + h.cache[host] = entry + h.cacheMu.Unlock() return got[dice.Roll(len(got))], nil } diff --git a/proxy/amnezia/config.pb.go b/proxy/amnezia/config.pb.go index 15e251a3..fe159bec 100644 --- a/proxy/amnezia/config.pb.go +++ b/proxy/amnezia/config.pb.go @@ -164,22 +164,23 @@ type DeviceConfig struct { DomainStrategy DeviceConfig_DomainStrategy `protobuf:"varint,7,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.proxy.amnezia.DeviceConfig_DomainStrategy" json:"domain_strategy,omitempty"` IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"` NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"` - Jc int32 `protobuf:"varint,10,opt,name=jc,proto3" json:"jc,omitempty"` - JMin int32 `protobuf:"varint,11,opt,name=j_min,json=jMin,proto3" json:"j_min,omitempty"` - JMax int32 `protobuf:"varint,12,opt,name=j_max,json=jMax,proto3" json:"j_max,omitempty"` - S1 int32 `protobuf:"varint,13,opt,name=s1,proto3" json:"s1,omitempty"` // InitPadding - S2 int32 `protobuf:"varint,14,opt,name=s2,proto3" json:"s2,omitempty"` // ResponsePadding - S3 int32 `protobuf:"varint,15,opt,name=s3,proto3" json:"s3,omitempty"` // CookiePadding - S4 int32 `protobuf:"varint,16,opt,name=s4,proto3" json:"s4,omitempty"` // TransportPadding - H1 string `protobuf:"bytes,17,opt,name=h1,proto3" json:"h1,omitempty"` // InitHeader - H2 string `protobuf:"bytes,18,opt,name=h2,proto3" json:"h2,omitempty"` // ResponseHeader - H3 string `protobuf:"bytes,19,opt,name=h3,proto3" json:"h3,omitempty"` // CookieHeader - H4 string `protobuf:"bytes,20,opt,name=h4,proto3" json:"h4,omitempty"` // TransportHeader - I1 string `protobuf:"bytes,21,opt,name=i1,proto3" json:"i1,omitempty"` // Signature (1-5) - I2 string `protobuf:"bytes,22,opt,name=i2,proto3" json:"i2,omitempty"` - I3 string `protobuf:"bytes,23,opt,name=i3,proto3" json:"i3,omitempty"` - I4 string `protobuf:"bytes,24,opt,name=i4,proto3" json:"i4,omitempty"` - I5 string `protobuf:"bytes,25,opt,name=i5,proto3" json:"i5,omitempty"` + DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"` + Jc int32 `protobuf:"varint,11,opt,name=jc,proto3" json:"jc,omitempty"` + JMin int32 `protobuf:"varint,12,opt,name=j_min,json=jMin,proto3" json:"j_min,omitempty"` + JMax int32 `protobuf:"varint,13,opt,name=j_max,json=jMax,proto3" json:"j_max,omitempty"` + S1 int32 `protobuf:"varint,14,opt,name=s1,proto3" json:"s1,omitempty"` // InitPadding + S2 int32 `protobuf:"varint,15,opt,name=s2,proto3" json:"s2,omitempty"` // ResponsePadding + S3 int32 `protobuf:"varint,16,opt,name=s3,proto3" json:"s3,omitempty"` // CookiePadding + S4 int32 `protobuf:"varint,17,opt,name=s4,proto3" json:"s4,omitempty"` // TransportPadding + H1 string `protobuf:"bytes,18,opt,name=h1,proto3" json:"h1,omitempty"` // InitHeader + H2 string `protobuf:"bytes,19,opt,name=h2,proto3" json:"h2,omitempty"` // ResponseHeader + H3 string `protobuf:"bytes,20,opt,name=h3,proto3" json:"h3,omitempty"` // CookieHeader + H4 string `protobuf:"bytes,21,opt,name=h4,proto3" json:"h4,omitempty"` // TransportHeader + I1 string `protobuf:"bytes,22,opt,name=i1,proto3" json:"i1,omitempty"` // Signature (1-5) + I2 string `protobuf:"bytes,23,opt,name=i2,proto3" json:"i2,omitempty"` + I3 string `protobuf:"bytes,24,opt,name=i3,proto3" json:"i3,omitempty"` + I4 string `protobuf:"bytes,25,opt,name=i4,proto3" json:"i4,omitempty"` + I5 string `protobuf:"bytes,26,opt,name=i5,proto3" json:"i5,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -277,6 +278,13 @@ func (x *DeviceConfig) GetNoKernelTun() bool { return false } +func (x *DeviceConfig) GetDNS() []string { + if x != nil { + return x.DNS + } + return nil +} + func (x *DeviceConfig) GetJc() int32 { if x != nil { return x.Jc @@ -403,7 +411,7 @@ const file_proxy_amnezia_config_proto_rawDesc = "" + "\n" + "keep_alive\x18\x04 \x01(\tR\tkeepAlive\x12\x1f\n" + "\vallowed_ips\x18\x05 \x03(\tR\n" + - "allowedIps\"\xe2\x05\n" + + "allowedIps\"\xf4\x05\n" + "\fDeviceConfig\x12\x1d\n" + "\n" + "secret_key\x18\x01 \x01(\tR\tsecretKey\x12\x1a\n" + @@ -414,24 +422,25 @@ const file_proxy_amnezia_config_proto_rawDesc = "" + "\breserved\x18\x06 \x01(\fR\breserved\x12X\n" + "\x0fdomain_strategy\x18\a \x01(\x0e2/.xray.proxy.amnezia.DeviceConfig.DomainStrategyR\x0edomainStrategy\x12\x1b\n" + "\tis_client\x18\b \x01(\bR\bisClient\x12\"\n" + - "\rno_kernel_tun\x18\t \x01(\bR\vnoKernelTun\x12\x0e\n" + - "\x02jc\x18\n" + - " \x01(\x05R\x02jc\x12\x13\n" + - "\x05j_min\x18\v \x01(\x05R\x04jMin\x12\x13\n" + - "\x05j_max\x18\f \x01(\x05R\x04jMax\x12\x0e\n" + - "\x02s1\x18\r \x01(\x05R\x02s1\x12\x0e\n" + - "\x02s2\x18\x0e \x01(\x05R\x02s2\x12\x0e\n" + - "\x02s3\x18\x0f \x01(\x05R\x02s3\x12\x0e\n" + - "\x02s4\x18\x10 \x01(\x05R\x02s4\x12\x0e\n" + - "\x02h1\x18\x11 \x01(\tR\x02h1\x12\x0e\n" + - "\x02h2\x18\x12 \x01(\tR\x02h2\x12\x0e\n" + - "\x02h3\x18\x13 \x01(\tR\x02h3\x12\x0e\n" + - "\x02h4\x18\x14 \x01(\tR\x02h4\x12\x0e\n" + - "\x02i1\x18\x15 \x01(\tR\x02i1\x12\x0e\n" + - "\x02i2\x18\x16 \x01(\tR\x02i2\x12\x0e\n" + - "\x02i3\x18\x17 \x01(\tR\x02i3\x12\x0e\n" + - "\x02i4\x18\x18 \x01(\tR\x02i4\x12\x0e\n" + - "\x02i5\x18\x19 \x01(\tR\x02i5\"\\\n" + + "\rno_kernel_tun\x18\t \x01(\bR\vnoKernelTun\x12\x10\n" + + "\x03DNS\x18\n" + + " \x03(\tR\x03DNS\x12\x0e\n" + + "\x02jc\x18\v \x01(\x05R\x02jc\x12\x13\n" + + "\x05j_min\x18\f \x01(\x05R\x04jMin\x12\x13\n" + + "\x05j_max\x18\r \x01(\x05R\x04jMax\x12\x0e\n" + + "\x02s1\x18\x0e \x01(\x05R\x02s1\x12\x0e\n" + + "\x02s2\x18\x0f \x01(\x05R\x02s2\x12\x0e\n" + + "\x02s3\x18\x10 \x01(\x05R\x02s3\x12\x0e\n" + + "\x02s4\x18\x11 \x01(\x05R\x02s4\x12\x0e\n" + + "\x02h1\x18\x12 \x01(\tR\x02h1\x12\x0e\n" + + "\x02h2\x18\x13 \x01(\tR\x02h2\x12\x0e\n" + + "\x02h3\x18\x14 \x01(\tR\x02h3\x12\x0e\n" + + "\x02h4\x18\x15 \x01(\tR\x02h4\x12\x0e\n" + + "\x02i1\x18\x16 \x01(\tR\x02i1\x12\x0e\n" + + "\x02i2\x18\x17 \x01(\tR\x02i2\x12\x0e\n" + + "\x02i3\x18\x18 \x01(\tR\x02i3\x12\x0e\n" + + "\x02i4\x18\x19 \x01(\tR\x02i4\x12\x0e\n" + + "\x02i5\x18\x1a \x01(\tR\x02i5\"\\\n" + "\x0eDomainStrategy\x12\f\n" + "\bFORCE_IP\x10\x00\x12\r\n" + "\tFORCE_IP4\x10\x01\x12\r\n" + diff --git a/proxy/amnezia/config.proto b/proxy/amnezia/config.proto index 0e45d7a6..1e28f644 100644 --- a/proxy/amnezia/config.proto +++ b/proxy/amnezia/config.proto @@ -34,21 +34,22 @@ message DeviceConfig { DomainStrategy domain_strategy = 7; bool is_client = 8; bool no_kernel_tun = 9; + repeated string DNS = 10; - int32 jc = 10; - int32 j_min = 11; - int32 j_max = 12; - int32 s1 = 13; // InitPadding - int32 s2 = 14; // ResponsePadding - int32 s3 = 15; // CookiePadding - int32 s4 = 16; // TransportPadding - string h1 = 17; // InitHeader - string h2 = 18; // ResponseHeader - string h3 = 19; // CookieHeader - string h4 = 20; // TransportHeader - string i1 = 21; // Signature (1-5) - string i2 = 22; - string i3 = 23; - string i4 = 24; - string i5 = 25; + int32 jc = 11; + int32 j_min = 12; + int32 j_max = 13; + int32 s1 = 14; // InitPadding + int32 s2 = 15; // ResponsePadding + int32 s3 = 16; // CookiePadding + int32 s4 = 17; // TransportPadding + string h1 = 18; // InitHeader + string h2 = 19; // ResponseHeader + string h3 = 20; // CookieHeader + string h4 = 21; // TransportHeader + string i1 = 22; // Signature (1-5) + string i2 = 23; + string i3 = 24; + string i4 = 25; + string i5 = 26; } diff --git a/proxy/amnezia/netstack.go b/proxy/amnezia/netstack.go index bba72b51..a2429545 100644 --- a/proxy/amnezia/netstack.go +++ b/proxy/amnezia/netstack.go @@ -248,7 +248,7 @@ var ( errTimeout = errors.New("i/o timeout") ) -func (net *Net) LookupHost(host string) (addrs []string, err error) { +func (net *Net) LookupHost(host string) (addrs []net.IP, ttl uint32, err error) { return net.LookupContextHost(context.Background(), host) } @@ -567,9 +567,9 @@ func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.T return dnsmessage.Parser{}, "", lastErr } -func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) { +func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP, uint32, error) { if host == "" || (!tnet.hasV6 && !tnet.hasV4) { - return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} } zlen := len(host) if strings.IndexByte(host, ':') != -1 { @@ -578,11 +578,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, } } if ip, err := netip.ParseAddr(host[:zlen]); err == nil { - return []string{ip.String()}, nil + return []net.IP{ip.AsSlice()}, 0, nil } if !isDomainName(host) { - return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} + return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true} } type result struct { p dnsmessage.Parser @@ -611,6 +611,7 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, lane <- result{p, server, err} }() } + ttl := uint32(300) for l := 0; l < lanes; l++ { result := <-lane if result.error != nil { @@ -644,6 +645,7 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, } break loop } + ttl = min(ttl, h.TTL) addrsV4 = append(addrsV4, netip.AddrFrom4(a.A)) case dnsmessage.TypeAAAA: @@ -656,6 +658,7 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, } break loop } + ttl = min(ttl, h.TTL) addrsV6 = append(addrsV6, netip.AddrFrom16(aaaa.AAAA)) default: @@ -680,11 +683,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, } if len(addrs) == 0 && lastErr != nil { - return nil, lastErr + return nil, 0, lastErr } - saddrs := make([]string, 0, len(addrs)) + ips := make([]net.IP, 0, len(addrs)) for _, ip := range addrs { - saddrs = append(saddrs, ip.String()) + ips = append(ips, ip.AsSlice()) } - return saddrs, nil + return ips, ttl, nil }