diff --git a/dns/client_truncate.go b/dns/client_truncate.go index 19165f99..43ba0a78 100644 --- a/dns/client_truncate.go +++ b/dns/client_truncate.go @@ -6,7 +6,7 @@ import ( "github.com/miekg/dns" ) -func TruncateDNSMessage(request *dns.Msg, response *dns.Msg, headroom int) (*buf.Buffer, error) { +func TruncateDNSMessage(request *dns.Msg, response *dns.Msg, frontHeadroom int, rearHeadroom int) (*buf.Buffer, error) { maxLen := 512 if edns0Option := request.IsEdns0(); edns0Option != nil { if udpSize := int(edns0Option.UDPSize()); udpSize > 512 { @@ -18,8 +18,8 @@ func TruncateDNSMessage(request *dns.Msg, response *dns.Msg, headroom int) (*buf response = response.Copy() response.Truncate(maxLen) } - buffer := buf.NewSize(headroom*2 + 1 + responseLen) - buffer.Resize(headroom, 0) + buffer := buf.NewSize(frontHeadroom + responseLen + 1 + rearHeadroom) + buffer.Resize(frontHeadroom, 0) rawMessage, err := response.PackBuffer(buffer.FreeBytes()) if err != nil { buffer.Release() diff --git a/protocol/dns/handle.go b/protocol/dns/handle.go index 72197140..f47b0243 100644 --- a/protocol/dns/handle.go +++ b/protocol/dns/handle.go @@ -66,6 +66,8 @@ func writeStreamResponse(conn net.Conn, response *mDNS.Msg) { func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn N.PacketConn, cachedPackets []*N.PacketBuffer, metadata adapter.InboundContext) error { metadata.Destination = M.Socksaddr{} + frontHeadroom := N.CalculateFrontHeadroom(conn) + rearHeadroom := N.CalculateRearHeadroom(conn) var reader N.PacketReader = conn var counters []N.CountFunc cachedPackets = common.Reverse(cachedPackets) @@ -131,7 +133,7 @@ func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn return } timeout.Update() - responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, 1024) + responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, frontHeadroom, rearHeadroom) if truncateErr != nil { cancel(truncateErr) return @@ -150,6 +152,8 @@ func NewDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn } func newDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn N.PacketConn, readWaiter N.PacketReadWaiter, readCounters []N.CountFunc, cached []*N.PacketBuffer, metadata adapter.InboundContext) error { + frontHeadroom := N.CalculateFrontHeadroom(conn) + rearHeadroom := N.CalculateRearHeadroom(conn) fastClose, cancel := context.WithCancelCause(ctx) timeout := canceler.New(fastClose, cancel, C.DNSTimeout) var group task.Group @@ -199,7 +203,7 @@ func newDNSPacketConnection(ctx context.Context, router adapter.DNSRouter, conn return } timeout.Update() - responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, 1024) + responseBuffer, truncateErr := dns.TruncateDNSMessage(&message, response, frontHeadroom, rearHeadroom) if truncateErr != nil { cancel(truncateErr) return diff --git a/route/dns.go b/route/dns.go index 58707152..c65778cd 100644 --- a/route/dns.go +++ b/route/dns.go @@ -61,7 +61,7 @@ func (r *Router) HijackDNSPacket(ctx context.Context, payload []byte, writer N.P } func (r *Router) writeDNSPacketResponse(message *mDNS.Msg, response *mDNS.Msg, writer N.PacketWriter, destination M.Socksaddr) error { - responseBuffer, err := dns.TruncateDNSMessage(message, response, 1024) + responseBuffer, err := dns.TruncateDNSMessage(message, response, N.CalculateFrontHeadroom(writer), N.CalculateRearHeadroom(writer)) if err != nil { return err } diff --git a/service/resolved/service.go b/service/resolved/service.go index d14082d7..a634b9b2 100644 --- a/service/resolved/service.go +++ b/service/resolved/service.go @@ -173,7 +173,7 @@ func (i *Service) exchangePacket0(ctx context.Context, buffer *buf.Buffer, oob [ if err != nil { return err } - responseBuffer, err := dns.TruncateDNSMessage(&message, response, 0) + responseBuffer, err := dns.TruncateDNSMessage(&message, response, 0, 0) if err != nil { return err }