1515
1616import os
1717import socket
18+ import stat
1819import struct
20+ import tempfile
1921import threading
2022import time
2123from array import array
2628DEFAULT_LAYERWISE_EVENTFD_SOCKET = "/tmp/flexkv_layerwise_eventfd.sock"
2729
2830
31+ def create_layerwise_eventfd_socket_path () -> str :
32+ socket_directory = tempfile .mkdtemp (
33+ prefix = f"recsys-flexkv-layerwise-{ os .geteuid ()} -"
34+ )
35+ os .chmod (socket_directory , 0o700 )
36+ return os .path .join (socket_directory , "eventfd.sock" )
37+
38+
2939class FlexKVLayerwiseEventfdSender :
30- """Creates mock layerwise eventfds and sends them to FlexKV's worker."""
40+ """Creates layerwise eventfds and sends them to FlexKV's worker."""
3141
3242 def __init__ (
3343 self ,
@@ -42,6 +52,8 @@ def __init__(
4252 self .timeout_s = float (timeout_s )
4353 self ._eventfds : Optional [List [List [int ]]] = None
4454 self ._thread : Optional [threading .Thread ] = None
55+ self ._handoff_done = threading .Event ()
56+ self ._handoff_error : Optional [BaseException ] = None
4557
4658 @staticmethod
4759 def _create_eventfd () -> int :
@@ -76,20 +88,87 @@ def start(self) -> None:
7688 return
7789 self .create_eventfds ()
7890 self ._thread = threading .Thread (
79- target = self ._send_eventfds ,
91+ target = self ._run_sender ,
8092 name = "flexkv-layerwise-eventfd-sender" ,
8193 daemon = True ,
8294 )
8395 self ._thread .start ()
8496
97+ def wait_until_ready (self , timeout_s : Optional [float ] = None ) -> None :
98+ timeout = self .timeout_s + 1.0 if timeout_s is None else float (timeout_s )
99+ if not self ._handoff_done .wait (timeout ):
100+ raise TimeoutError (
101+ "Timed out waiting for FlexKV layerwise eventfd handoff "
102+ f"on socket { self .socket_path } "
103+ )
104+ if self ._handoff_error is not None :
105+ raise RuntimeError (
106+ "FlexKV layerwise eventfd handoff failed"
107+ ) from self ._handoff_error
108+
109+ def _run_sender (self ) -> None :
110+ try :
111+ self ._send_eventfds ()
112+ except BaseException as error :
113+ self ._handoff_error = error
114+ finally :
115+ self ._handoff_done .set ()
116+
117+ @staticmethod
118+ def _is_process_descendant (process_id : int , ancestor_id : int ) -> bool :
119+ current_id = int (process_id )
120+ ancestor_id = int (ancestor_id )
121+ visited = set ()
122+ while current_id > 1 and current_id not in visited :
123+ if current_id == ancestor_id :
124+ return True
125+ visited .add (current_id )
126+ try :
127+ with open (f"/proc/{ current_id } /stat" , encoding = "utf-8" ) as stat_file :
128+ process_stat = stat_file .read ()
129+ fields_after_name = process_stat [process_stat .rfind (")" ) + 1 :].split ()
130+ current_id = int (fields_after_name [1 ])
131+ except (OSError , IndexError , ValueError ):
132+ return False
133+ return current_id == ancestor_id
134+
135+ def _authenticate_peer (self , sock : socket .socket ) -> None :
136+ socket_stat = os .stat (self .socket_path , follow_symlinks = False )
137+ if not stat .S_ISSOCK (socket_stat .st_mode ):
138+ raise PermissionError (
139+ f"Layerwise eventfd path is not a socket: { self .socket_path } "
140+ )
141+ if socket_stat .st_uid != os .geteuid ():
142+ raise PermissionError (
143+ "FlexKV layerwise socket owner does not match the current user"
144+ )
145+
146+ if not hasattr (socket , "SO_PEERCRED" ):
147+ raise RuntimeError ("SO_PEERCRED is required for layerwise eventfd handoff" )
148+ credentials_size = struct .calcsize ("3i" )
149+ peer_pid , peer_uid , peer_gid = struct .unpack (
150+ "3i" ,
151+ sock .getsockopt (socket .SOL_SOCKET , socket .SO_PEERCRED , credentials_size ),
152+ )
153+ if peer_uid != os .geteuid () or peer_gid != os .getegid ():
154+ raise PermissionError (
155+ "FlexKV layerwise peer credentials do not match the current process"
156+ )
157+ if not self ._is_process_descendant (peer_pid , os .getpid ()):
158+ raise PermissionError (
159+ f"FlexKV layerwise peer pid { peer_pid } is outside the current process tree"
160+ )
161+
85162 def _send_eventfds (self ) -> None :
86163 eventfds = self .create_eventfds ()
87164 deadline = time .time () + self .timeout_s
88165 last_error : Optional [Exception ] = None
89166 while time .time () < deadline :
90167 sock = socket .socket (socket .AF_UNIX , socket .SOCK_STREAM )
91168 try :
169+ sock .settimeout (min (1.0 , max (deadline - time .time (), 0.01 )))
92170 sock .connect (self .socket_path )
171+ self ._authenticate_peer (sock )
93172 metadata = struct .pack (
94173 "iiii" ,
95174 0 ,
@@ -116,12 +195,12 @@ def _send_eventfds(self) -> None:
116195 f"FlexKV layerwise eventfd receiver returned ack={ ack !r} "
117196 )
118197 return
119- except (FileNotFoundError , ConnectionRefusedError , socket . timeout ) as e :
198+ except (OSError , RuntimeError ) as e :
120199 last_error = e
121200 time .sleep (0.05 )
122201 finally :
123202 sock .close ()
124203 raise RuntimeError (
125- "Timed out sending mock layerwise eventfds to FlexKV "
204+ "Timed out sending layerwise eventfds to FlexKV "
126205 f"socket { self .socket_path } : { last_error } "
127206 )
0 commit comments