package socks5 import ( "context" "encoding/binary" "fmt" "io" "net" "strconv" "time" ) type Dialer struct { ProxyAddr string Timeout time.Duration } func (d Dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { if network != "tcp" { return nil, fmt.Errorf("socks5 only supports tcp, got %s", network) } host, portText, err := net.SplitHostPort(address) if err != nil { return nil, fmt.Errorf("split target address: %w", err) } port, err := strconv.Atoi(portText) if err != nil || port < 1 || port > 65535 { return nil, fmt.Errorf("invalid target port") } if len(host) == 0 || len(host) > 255 { return nil, fmt.Errorf("invalid target host length") } timeout := d.Timeout if timeout <= 0 { timeout = 90 * time.Second } conn, err := (&net.Dialer{Timeout: timeout}).DialContext(ctx, "tcp", d.ProxyAddr) if err != nil { return nil, fmt.Errorf("connect socks proxy: %w", err) } if deadline, ok := ctx.Deadline(); ok { _ = conn.SetDeadline(deadline) } else { _ = conn.SetDeadline(time.Now().Add(timeout)) } if err := handshake(conn, host, uint16(port)); err != nil { _ = conn.Close() return nil, err } _ = conn.SetDeadline(time.Time{}) return conn, nil } func handshake(conn net.Conn, host string, port uint16) error { if _, err := conn.Write([]byte{0x05, 0x01, 0x00}); err != nil { return fmt.Errorf("write socks greeting: %w", err) } var greeting [2]byte if _, err := io.ReadFull(conn, greeting[:]); err != nil { return fmt.Errorf("read socks greeting: %w", err) } if greeting[0] != 0x05 || greeting[1] != 0x00 { return fmt.Errorf("socks proxy rejected no-auth method") } req := make([]byte, 0, 7+len(host)) req = append(req, 0x05, 0x01, 0x00, 0x03, byte(len(host))) req = append(req, []byte(host)...) var p [2]byte binary.BigEndian.PutUint16(p[:], port) req = append(req, p[:]...) if _, err := conn.Write(req); err != nil { return fmt.Errorf("write socks connect: %w", err) } var head [4]byte if _, err := io.ReadFull(conn, head[:]); err != nil { return fmt.Errorf("read socks response: %w", err) } if head[0] != 0x05 { return fmt.Errorf("invalid socks response version") } if head[1] != 0x00 { return fmt.Errorf("socks connect failed: %s", replyText(head[1])) } var skip int switch head[3] { case 0x01: skip = 4 + 2 case 0x03: var l [1]byte if _, err := io.ReadFull(conn, l[:]); err != nil { return fmt.Errorf("read socks bind host length: %w", err) } skip = int(l[0]) + 2 case 0x04: skip = 16 + 2 default: return fmt.Errorf("unsupported socks bind address type") } if skip > 0 { buf := make([]byte, skip) if _, err := io.ReadFull(conn, buf); err != nil { return fmt.Errorf("read socks bind address: %w", err) } } return nil } func replyText(code byte) string { switch code { case 0x01: return "general failure" case 0x02: return "connection not allowed" case 0x03: return "network unreachable" case 0x04: return "host unreachable" case 0x05: return "connection refused" case 0x06: return "ttl expired" case 0x07: return "command not supported" case 0x08: return "address type not supported" default: return fmt.Sprintf("unknown code %d", code) } }