@@ -5,7 +5,7 @@ use gethostname::gethostname;
55use nix:: unistd:: sethostname;
66use socket2:: { Domain , Protocol , Socket , Type as SocketType } ;
77use std:: io:: { self , prelude:: * } ;
8- use std:: net:: { IpAddr , Ipv4Addr , Ipv6Addr , Shutdown , SocketAddr , ToSocketAddrs } ;
8+ use std:: net:: { Ipv4Addr , Ipv6Addr , Shutdown , SocketAddr , ToSocketAddrs } ;
99use std:: time:: Duration ;
1010
1111use crate :: builtins:: bytearray:: PyByteArrayRef ;
@@ -130,7 +130,7 @@ impl PySocket {
130130
131131 #[ pymethod]
132132 fn connect ( & self , address : Address , vm : & VirtualMachine ) -> PyResult < ( ) > {
133- let sock_addr = get_addr ( vm, address) ?;
133+ let sock_addr = get_addr ( vm, address, self . family . load ( ) ) ?;
134134 let res = if let Some ( duration) = self . sock ( ) . read_timeout ( ) . unwrap ( ) {
135135 self . sock ( ) . connect_timeout ( & sock_addr, duration)
136136 } else {
@@ -141,7 +141,7 @@ impl PySocket {
141141
142142 #[ pymethod]
143143 fn bind ( & self , address : Address , vm : & VirtualMachine ) -> PyResult < ( ) > {
144- let sock_addr = get_addr ( vm, address) ?;
144+ let sock_addr = get_addr ( vm, address, self . family . load ( ) ) ?;
145145 self . sock ( )
146146 . bind ( & sock_addr)
147147 . map_err ( |err| convert_sock_error ( vm, err) )
@@ -213,7 +213,7 @@ impl PySocket {
213213
214214 #[ pymethod]
215215 fn sendto ( & self , bytes : PyBytesLike , address : Address , vm : & VirtualMachine ) -> PyResult < ( ) > {
216- let addr = get_addr ( vm, address) ?;
216+ let addr = get_addr ( vm, address, self . family . load ( ) ) ?;
217217 bytes
218218 . with_ref ( |b| self . sock ( ) . send_to ( b, & addr) )
219219 . map_err ( |err| convert_sock_error ( vm, err) ) ?;
@@ -678,7 +678,7 @@ fn socket_getnameinfo(
678678 flags : i32 ,
679679 vm : & VirtualMachine ,
680680) -> PyResult < ( String , String ) > {
681- let addr = get_addr ( vm, address) ?;
681+ let addr = get_addr ( vm, address, IpAddrFmt :: Any ( ) ) ?;
682682 let nameinfo = addr
683683 . as_std ( )
684684 . and_then ( |addr| dns_lookup:: getnameinfo ( & addr, flags) . ok ( ) ) ;
@@ -691,32 +691,51 @@ fn socket_getnameinfo(
691691 } )
692692}
693693
694- fn get_addr ( vm : & VirtualMachine , addr : impl ToSocketAddrs ) -> PyResult < socket2:: SockAddr > {
695- match addr. to_socket_addrs ( ) {
696- Ok ( mut sock_addrs) => {
697- if let Some ( mut addr) = sock_addrs. next ( ) {
698- if option_env ! ( "RUSTPYTHON_NO_IPV6" ) . is_some ( ) {
699- while addr. ip ( ) == IpAddr :: V6 ( Ipv6Addr :: LOCALHOST ) {
700- if let Some ( other) = sock_addrs. next ( ) {
701- addr = other
702- } else {
703- break ;
704- }
705- }
706- }
707-
708- Ok ( addr. into ( ) )
709- } else {
710- let error_type = vm. class ( "_socket" , "gaierror" ) ;
711- Err ( vm. new_exception_msg (
712- error_type,
713- "nodename nor servname provided, or not known" . to_owned ( ) ,
714- ) )
715- }
694+ enum IpAddrFmt {
695+ Ipv4 ( ) ,
696+ Ipv6 ( ) ,
697+ Any ( ) ,
698+ }
699+
700+ impl < T > From < T > for IpAddrFmt
701+ where
702+ T : PartialEq < i32 > + Sized + std:: fmt:: Debug ,
703+ {
704+ fn from ( i : T ) -> Self {
705+ if i == Domain :: ipv4 ( ) . into ( ) {
706+ Self :: Ipv4 ( )
707+ } else if i == Domain :: ipv6 ( ) . into ( ) {
708+ Self :: Ipv6 ( )
709+ } else {
710+ error ! ( "Unknown IP family/domain: {:?}" , i) ;
711+ Self :: Any ( )
716712 }
713+ }
714+ }
715+
716+ fn get_addr < F > ( vm : & VirtualMachine , addr : impl ToSocketAddrs , fmt : F ) -> PyResult < socket2:: SockAddr >
717+ where
718+ F : Into < IpAddrFmt > ,
719+ {
720+ let sock_addr = match addr. to_socket_addrs ( ) {
721+ Ok ( mut sock_addrs) => match fmt. into ( ) {
722+ IpAddrFmt :: Ipv6 ( ) => sock_addrs. find ( |a| a. is_ipv6 ( ) ) ,
723+ IpAddrFmt :: Ipv4 ( ) => sock_addrs. find ( |a| a. is_ipv4 ( ) ) ,
724+ IpAddrFmt :: Any ( ) => sock_addrs. next ( ) ,
725+ } ,
717726 Err ( e) => {
718727 let error_type = vm. class ( "_socket" , "gaierror" ) ;
719- Err ( vm. new_exception_msg ( error_type, e. to_string ( ) ) )
728+ return Err ( vm. new_exception_msg ( error_type, e. to_string ( ) ) ) ;
729+ }
730+ } ;
731+ match sock_addr {
732+ Some ( sock_addr) => Ok ( sock_addr. into ( ) ) ,
733+ None => {
734+ let error_type = vm. class ( "_socket" , "gaierror" ) ;
735+ Err ( vm. new_exception_msg (
736+ error_type,
737+ "nodename nor servname provided, or not known" . to_owned ( ) ,
738+ ) )
720739 }
721740 }
722741}
0 commit comments