diff --git a/ping/destination.go b/ping/destination.go index 553c7f8..0640b2e 100644 --- a/ping/destination.go +++ b/ping/destination.go @@ -1,6 +1,7 @@ package ping import ( + "context" "errors" "net/netip" "os" @@ -16,12 +17,13 @@ import ( var _ tun.DirectRouteDestination = (*Destination)(nil) type Destination struct { - logger logger.Logger + ctx context.Context + logger logger.ContextLogger routeContext tun.DirectRouteContext conn *Conn } -func ConnectDestination(logger logger.Logger, controlFunc control.Func, address netip.Addr, routeContext tun.DirectRouteContext) (tun.DirectRouteDestination, error) { +func ConnectDestination(ctx context.Context, logger logger.ContextLogger, controlFunc control.Func, address netip.Addr, routeContext tun.DirectRouteContext) (tun.DirectRouteDestination, error) { var ( conn *Conn err error @@ -39,6 +41,7 @@ func ConnectDestination(logger logger.Logger, controlFunc control.Func, address return nil, err } d := &Destination{ + ctx: ctx, logger: logger, routeContext: routeContext, conn: conn, @@ -54,13 +57,13 @@ func (d *Destination) loopRead() { if err != nil { buffer.Release() if !E.IsClosed(err) { - d.logger.Error(E.Cause(err, "receive ICMP echo reply")) + d.logger.ErrorContext(d.ctx, E.Cause(err, "receive ICMP echo reply")) } return } err = d.routeContext.WritePacket(buffer.Bytes()) if err != nil { - d.logger.Error(E.Cause(err, "write ICMP echo reply")) + d.logger.Error(d.ctx, E.Cause(err, "write ICMP echo reply")) } buffer.Release() } diff --git a/ping/socket_unix.go b/ping/socket_unix.go index 11fa7a9..caf0f31 100644 --- a/ping/socket_unix.go +++ b/ping/socket_unix.go @@ -12,6 +12,7 @@ import ( "github.com/sagernet/sing/common/control" E "github.com/sagernet/sing/common/exceptions" M "github.com/sagernet/sing/common/metadata" + "golang.org/x/sys/unix" ) diff --git a/route_nat.go b/route_nat.go index a4f33cc..ccb8ae3 100644 --- a/route_nat.go +++ b/route_nat.go @@ -6,8 +6,14 @@ import ( "github.com/sagernet/sing-tun/internal/gtcpip/checksum" "github.com/sagernet/sing-tun/internal/gtcpip/header" + "github.com/sagernet/sing/common/buf" ) +type DirectRouteDestination interface { + WritePacket(packet *buf.Buffer) error + Close() error +} + type NatMapping struct { access sync.RWMutex sessions map[DirectRouteSession]DirectRouteContext diff --git a/route_nat_gvisor.go b/route_nat_gvisor.go index be8febb..487e86b 100644 --- a/route_nat_gvisor.go +++ b/route_nat_gvisor.go @@ -5,16 +5,9 @@ package tun import ( "github.com/sagernet/gvisor/pkg/tcpip" "github.com/sagernet/gvisor/pkg/tcpip/header" - stack "github.com/sagernet/gvisor/pkg/tcpip/stack" - "github.com/sagernet/sing/common/buf" + "github.com/sagernet/gvisor/pkg/tcpip/stack" ) -type DirectRouteDestination interface { - WritePacket(packet *buf.Buffer) error - WritePacketBuffer(packetBuffer *stack.PacketBuffer) error - Close() error -} - func (w *NatWriter) RewritePacketBuffer(packetBuffer *stack.PacketBuffer) { var bindAddr tcpip.Address if packetBuffer.NetworkProtocolNumber == header.IPv4ProtocolNumber { diff --git a/route_nat_non_gvisor.go b/route_nat_non_gvisor.go deleted file mode 100644 index a0c6cea..0000000 --- a/route_nat_non_gvisor.go +++ /dev/null @@ -1,12 +0,0 @@ -//go:build !with_gvisor - -package tun - -import ( - "github.com/sagernet/sing/common/buf" -) - -type DirectRouteDestination interface { - WritePacket(packet *buf.Buffer) error - Close() error -} diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index a45ba58..6888eac 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -15,6 +15,7 @@ import ( "github.com/sagernet/gvisor/pkg/tcpip/network/ipv4" "github.com/sagernet/gvisor/pkg/tcpip/network/ipv6" "github.com/sagernet/gvisor/pkg/tcpip/stack" + "github.com/sagernet/sing/common/buf" M "github.com/sagernet/sing/common/metadata" N "github.com/sagernet/sing/common/network" ) @@ -62,8 +63,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa } if action != nil { // TODO: handle error - pkt.IncRef() - _ = action.WritePacketBuffer(pkt) + _ = icmpWritePacketBuffer(action, pkt) return true } icmpHdr.SetType(header.ICMPv4EchoReply) @@ -120,7 +120,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa if action != nil { // TODO: handle error pkt.IncRef() - _ = action.WritePacketBuffer(pkt) + _ = icmpWritePacketBuffer(action, pkt) return true } icmpHdr.SetType(header.ICMPv6EchoReply) @@ -207,3 +207,9 @@ func (w *ICMPBackWriter) WritePacket(p []byte) error { } return nil } + +func icmpWritePacketBuffer(action DirectRouteDestination, packetBuffer *stack.PacketBuffer) error { + packetSlice := packetBuffer.TransportHeader().Slice() + packetSlice = append(packetSlice, packetBuffer.Data().AsRange().ToSlice()...) + return action.WritePacket(buf.As(packetSlice).ToOwned()) +}