@@ -948,8 +948,9 @@ use socket2::{Domain, Protocol, SockAddr, Socket, Type};
948948use std:: net:: SocketAddr ;
949949use std:: sync:: Mutex ;
950950
951- enum HostSocket {
952- Socket ( Socket , i32 ) ,
951+ struct HostSocket {
952+ socket : Socket ,
953+ sock_type : i32 ,
953954}
954955
955956const MAX_SOCKETS : usize = 1024 ;
@@ -989,15 +990,11 @@ impl SocketTable {
989990 }
990991
991992 fn get_socket ( & self , fd : u64 ) -> Result < & Socket > {
992- match self . get ( fd) ? {
993- HostSocket :: Socket ( s, _) => Ok ( s) ,
994- }
993+ Ok ( & self . get ( fd) ?. socket )
995994 }
996995
997996 fn get_sock_type ( & self , fd : u64 ) -> Result < i32 > {
998- match self . get ( fd) ? {
999- HostSocket :: Socket ( _, t) => Ok ( * t) ,
1000- }
997+ Ok ( self . get ( fd) ?. sock_type )
1001998 }
1002999
10031000 fn remove ( & mut self , fd : u64 ) -> Result < ( ) > {
@@ -1181,10 +1178,10 @@ fn register_net_tools(
11811178 Some ( Protocol :: from ( protocol) )
11821179 } ;
11831180 let sock = Socket :: new ( domain, stype, proto) ?;
1184- let fd = t
1185- . lock ( )
1186- . unwrap ( )
1187- . insert ( HostSocket :: Socket ( sock , sock_type ) ) ?;
1181+ let fd = t. lock ( ) . unwrap ( ) . insert ( HostSocket {
1182+ socket : sock ,
1183+ sock_type ,
1184+ } ) ?;
11881185 Ok ( json ! ( { "fd" : fd } ) )
11891186 } ) ;
11901187
@@ -1242,10 +1239,10 @@ fn register_net_tools(
12421239 ( s, p, st)
12431240 } ;
12441241 let peer_addr: Option < SocketAddr > = peer. as_socket ( ) ;
1245- let new_fd = t
1246- . lock ( )
1247- . unwrap ( )
1248- . insert ( HostSocket :: Socket ( new_sock , parent_type ) ) ?;
1242+ let new_fd = t. lock ( ) . unwrap ( ) . insert ( HostSocket {
1243+ socket : new_sock ,
1244+ sock_type : parent_type ,
1245+ } ) ?;
12491246 let mut resp = json ! ( { "fd" : new_fd } ) ;
12501247 if let Some ( pa) = peer_addr {
12511248 resp[ "addr" ] = json ! ( pa. ip( ) . to_string( ) ) ;
@@ -3015,10 +3012,14 @@ mod tests {
30153012 let req = br#"{"name":"net_socket","args":{"family":2,"type":1}}"# ;
30163013 let resp = tools. dispatch ( req) ;
30173014 let s = std:: str:: from_utf8 ( & resp) . unwrap ( ) ;
3018- assert ! ( s. contains( "\" fd\" " ) , "net_socket should work: {s}" ) ;
3015+ let v: serde_json:: Value = serde_json:: from_str ( s) . unwrap ( ) ;
3016+ let fd = v[ "result" ] [ "fd" ] . as_u64 ( ) . unwrap ( ) ;
30193017 // Try to bind — should fail because no listen_ports
3020- let req = br#"{"name":"net_bind","args":{"fd":0,"addr":"127.0.0.1","port":8080}}"# ;
3021- let resp = tools. dispatch ( req) ;
3018+ let req = format ! (
3019+ r#"{{"name":"net_bind","args":{{"fd":{},"addr":"127.0.0.1","port":8080}}}}"# ,
3020+ fd
3021+ ) ;
3022+ let resp = tools. dispatch ( req. as_bytes ( ) ) ;
30223023 let s = std:: str:: from_utf8 ( & resp) . unwrap ( ) ;
30233024 assert ! ( s. contains( "\" error\" " ) , "net_bind should be denied: {s}" ) ;
30243025 assert ! ( s. contains( "no --port" ) , "{s}" ) ;
@@ -3133,17 +3134,32 @@ mod tests {
31333134 let mut table = SocketTable :: new ( ) ;
31343135 for _ in 0 ..MAX_SOCKETS {
31353136 let sock = Socket :: new ( Domain :: IPV4 , Type :: STREAM , None ) . unwrap ( ) ;
3136- table. insert ( HostSocket :: Socket ( sock, 1 ) ) . unwrap ( ) ;
3137+ table
3138+ . insert ( HostSocket {
3139+ socket : sock,
3140+ sock_type : 1 ,
3141+ } )
3142+ . unwrap ( ) ;
31373143 }
31383144 let sock = Socket :: new ( Domain :: IPV4 , Type :: STREAM , None ) . unwrap ( ) ;
3139- assert ! ( table. insert( HostSocket :: Socket ( sock, 1 ) ) . is_err( ) ) ;
3145+ assert ! ( table
3146+ . insert( HostSocket {
3147+ socket: sock,
3148+ sock_type: 1
3149+ } )
3150+ . is_err( ) ) ;
31403151 }
31413152
31423153 #[ test]
31433154 fn socket_table_clear ( ) {
31443155 let mut table = SocketTable :: new ( ) ;
31453156 let sock = Socket :: new ( Domain :: IPV4 , Type :: STREAM , None ) . unwrap ( ) ;
3146- table. insert ( HostSocket :: Socket ( sock, 1 ) ) . unwrap ( ) ;
3157+ table
3158+ . insert ( HostSocket {
3159+ socket : sock,
3160+ sock_type : 1 ,
3161+ } )
3162+ . unwrap ( ) ;
31473163 assert_eq ! ( table. sockets. len( ) , 1 ) ;
31483164 table. clear ( ) ;
31493165 assert_eq ! ( table. sockets. len( ) , 0 ) ;
0 commit comments