@@ -22,7 +22,6 @@ use std::{
2222 thread:: JoinHandle ,
2323} ;
2424
25- use nix:: libc:: iovec;
2625use vmm_sys_util:: eventfd:: { EventFd , EFD_SEMAPHORE } ;
2726use xen_bindings:: bindings:: xs_watch_type;
2827
@@ -34,6 +33,25 @@ pub const XS_WATCH: u32 = 4;
3433pub const XS_WRITE : u32 = 11 ;
3534pub const XS_WATCH_EVENT : u32 = 15 ;
3635
36+ fn message_bytes ( message : & XenSocketMessage ) -> & [ u8 ] {
37+ // SAFETY: XenSocketMessage is #[repr(C)] with four u32 fields and no padding.
38+ unsafe {
39+ std:: slice:: from_raw_parts (
40+ std:: ptr:: addr_of!( * message) . cast ( ) ,
41+ mem:: size_of :: < XenSocketMessage > ( ) ,
42+ )
43+ }
44+ }
45+
46+ fn write_request < W : Write > (
47+ writer : & mut W ,
48+ message : & XenSocketMessage ,
49+ payload : & [ u8 ] ,
50+ ) -> Result < ( ) , std:: io:: Error > {
51+ writer. write_all ( message_bytes ( message) ) ?;
52+ writer. write_all ( payload)
53+ }
54+
3755fn queue_message (
3856 condvar : & Arc < (
3957 Mutex < VecDeque < Result < XenStoreMessage , std:: io:: Error > > > ,
@@ -188,44 +206,14 @@ impl XenStoreHandle {
188206 } )
189207 }
190208
191- fn xs_transaction (
192- & self ,
193- r#type : u32 ,
194- iovec_buffers : & mut Vec < iovec > ,
195- ) -> Result < String , std:: io:: Error > {
196- let mut xen_socket_msg = XenSocketMessage :: new ( r#type, iovec_buffers) ?;
209+ fn xs_transaction ( & self , r#type : u32 , payload : & [ u8 ] ) -> Result < String , std:: io:: Error > {
210+ let xen_socket_msg = XenSocketMessage :: new ( r#type, payload. len ( ) ) ?;
197211 let ( lock, cvar) = & * self . reply_condvar ;
198212
199213 let mut tx_socket = self . tx_socket . lock ( ) . unwrap ( ) ;
200- {
201- // SAFETY: `xen_socket_msg` is `XenSocketMessage` bytes sized.
202- let xen_socket_msg_slice: & [ u8 ] = unsafe {
203- std:: slice:: from_raw_parts (
204- std:: ptr:: addr_of_mut!( xen_socket_msg) . cast ( ) ,
205- mem:: size_of :: < XenSocketMessage > ( ) ,
206- )
207- } ;
208-
209- /*
210- * Grabbing the mutex guarantees there will only be
211- * one active transcation at a time.
212- */
213- tx_socket. write_all ( xen_socket_msg_slice) ?;
214- }
215214
216- // SAFETY: tx_socket is a valid file descriptor and the pointer/length we pass
217- // are valid allocated values.
218- let ret = unsafe {
219- libc:: writev (
220- tx_socket. as_raw_fd ( ) ,
221- iovec_buffers. as_ptr ( ) ,
222- iovec_buffers. len ( ) as i32 ,
223- )
224- } ;
225-
226- if ret < 0 {
227- return Err ( Error :: last_os_error ( ) ) ;
228- }
215+ // Serialize complete request/response transactions.
216+ write_request ( & mut * tx_socket, & xen_socket_msg, payload) ?;
229217
230218 let mut reply_vec = lock. lock ( ) . unwrap ( ) ;
231219 while reply_vec. is_empty ( ) {
@@ -247,49 +235,23 @@ impl XenStoreHandle {
247235 }
248236
249237 pub fn read_str ( & self , path : & str ) -> Result < String , std:: io:: Error > {
250- let c_path = CString :: new ( path) ?;
251- let mut iovec_buffers = vec ! [ iovec {
252- iov_base: c_path. as_ptr( ) as * mut _,
253- iov_len: path. len( ) + 1 ,
254- } ] ;
238+ let payload = CString :: new ( path) ?. into_bytes_with_nul ( ) ;
255239
256- self . xs_transaction ( XS_READ , & mut iovec_buffers )
240+ self . xs_transaction ( XS_READ , & payload )
257241 }
258242
259243 pub fn write_str ( & self , path : & str , val : & str ) -> Result < ( ) , std:: io:: Error > {
260- let cpath = CString :: new ( path) ?;
261- let cval = CString :: new ( val) ?;
262- let mut iovec_buffers = vec ! [
263- iovec {
264- iov_base: cpath. as_ptr( ) as * mut _,
265- iov_len: path. len( ) + 1 ,
266- } ,
267- iovec {
268- iov_base: cval. as_ptr( ) as * mut _,
269- iov_len: val. len( ) ,
270- } ,
271- ] ;
244+ let mut payload = CString :: new ( path) ?. into_bytes_with_nul ( ) ;
245+ payload. extend ( CString :: new ( val) ?. into_bytes ( ) ) ;
272246
273- self . xs_transaction ( XS_WRITE , & mut iovec_buffers)
274- . map ( |_| ( ) )
247+ self . xs_transaction ( XS_WRITE , & payload) . map ( |_| ( ) )
275248 }
276249
277250 pub fn create_watch ( & self , path : & str , token : & str ) -> Result < ( ) , std:: io:: Error > {
278- let cpath = CString :: new ( path) ?;
279- let ctoken = CString :: new ( token) ?;
280- let mut iovec_buffers = vec ! [
281- iovec {
282- iov_base: cpath. as_ptr( ) as * mut _,
283- iov_len: path. len( ) + 1 ,
284- } ,
285- iovec {
286- iov_base: ctoken. as_ptr( ) as * mut _,
287- iov_len: token. len( ) + 1 ,
288- } ,
289- ] ;
251+ let mut payload = CString :: new ( path) ?. into_bytes_with_nul ( ) ;
252+ payload. extend ( CString :: new ( token) ?. into_bytes_with_nul ( ) ) ;
290253
291- self . xs_transaction ( XS_WATCH , & mut iovec_buffers)
292- . map ( |_| ( ) )
254+ self . xs_transaction ( XS_WATCH , & payload) . map ( |_| ( ) )
293255 }
294256
295257 pub fn read_watch ( & self , index : xs_watch_type ) -> Result < String , std:: io:: Error > {
@@ -329,13 +291,9 @@ impl XenStoreHandle {
329291 }
330292
331293 pub fn directory ( & self , path : & str ) -> Result < Vec < i32 > , std:: io:: Error > {
332- let c_path = CString :: new ( path) ?;
333- let mut iovec_buffers = vec ! [ iovec {
334- iov_base: c_path. as_ptr( ) as * mut _,
335- iov_len: path. len( ) + 1 ,
336- } ] ;
294+ let payload = CString :: new ( path) ?. into_bytes_with_nul ( ) ;
337295
338- match self . xs_transaction ( XS_DIRECTORY , & mut iovec_buffers ) {
296+ match self . xs_transaction ( XS_DIRECTORY , & payload ) {
339297 Ok ( res) => Ok ( res
340298 . as_str ( )
341299 . split ( '\0' )
@@ -367,3 +325,49 @@ impl Drop for XenStoreHandle {
367325 let _ = self . handler . take ( ) . unwrap ( ) . join ( ) ;
368326 }
369327}
328+
329+ #[ cfg( test) ]
330+ mod tests {
331+ use super :: * ;
332+
333+ #[ derive( Default ) ]
334+ struct ShortWriter {
335+ bytes : Vec < u8 > ,
336+ write_calls : usize ,
337+ }
338+
339+ impl Write for ShortWriter {
340+ fn write ( & mut self , buffer : & [ u8 ] ) -> Result < usize , std:: io:: Error > {
341+ self . write_calls += 1 ;
342+
343+ // Write the header, interrupt the payload, then force short writes.
344+ let written = match self . write_calls {
345+ 1 => buffer. len ( ) ,
346+ 2 => return Err ( Error :: from ( ErrorKind :: Interrupted ) ) ,
347+ _ => buffer. len ( ) . min ( 3 ) ,
348+ } ;
349+
350+ self . bytes . extend_from_slice ( & buffer[ ..written] ) ;
351+ Ok ( written)
352+ }
353+
354+ fn flush ( & mut self ) -> Result < ( ) , std:: io:: Error > {
355+ Ok ( ( ) )
356+ }
357+ }
358+
359+ #[ test]
360+ fn writes_complete_request_after_interrupted_and_short_writes ( ) -> Result < ( ) , std:: io:: Error > {
361+ let payload = b"domid\0 " ;
362+ let message = XenSocketMessage :: new ( XS_READ , payload. len ( ) ) ?;
363+ let mut writer = ShortWriter :: default ( ) ;
364+
365+ write_request ( & mut writer, & message, payload) ?;
366+
367+ let mut expected = message_bytes ( & message) . to_vec ( ) ;
368+ expected. extend_from_slice ( payload) ;
369+ assert_eq ! ( writer. bytes, expected) ;
370+ assert ! ( writer. write_calls > 2 ) ;
371+ Ok ( ( ) )
372+ }
373+ }
0 commit comments