@@ -7,7 +7,7 @@ use std::collections::HashMap;
77use std:: hash:: { Hash , Hasher } ;
88use std:: io;
99use std:: path:: { Path , PathBuf } ;
10- use tokio_util:: bytes:: { Buf , BytesMut } ;
10+ use tokio_util:: bytes:: { Buf , BufMut , BytesMut } ;
1111use tokio_util:: codec:: { Decoder , Encoder } ;
1212
1313use crate :: client:: log:: { fmt_bin, trace} ;
@@ -454,17 +454,29 @@ impl Encoder<CommandResponse<'_>> for CommandCodec {
454454 type Error = io:: Error ;
455455
456456 fn encode ( & mut self , item : CommandResponse < ' _ > , dst : & mut BytesMut ) -> Result < ( ) , Self :: Error > {
457- let mut buf = Vec :: new ( ) ;
458- let mut serializer = rmp_serde:: Serializer :: new ( & mut buf) ;
457+ let start = dst. len ( ) ;
458+ let header_len = std:: mem:: size_of :: < Header > ( ) ;
459+ let header = Header {
460+ marker : Header :: VALID_MARKER ,
461+ size : 0 ,
462+ } ;
463+
464+ dst. extend_from_slice ( header. as_slice ( ) ) ;
459465
460466 // The protocol supports responding with several messages, but actually
461467 // only one message is ever sent (see command_helpers.c)
462- [ item]
463- . serialize ( & mut serializer)
464- . map_err ( |e| io:: Error :: new ( io:: ErrorKind :: InvalidData , e) ) ?;
468+ {
469+ let mut writer = ( & mut * dst) . writer ( ) ;
470+ let mut serializer = rmp_serde:: Serializer :: new ( & mut writer) ;
471+ if let Err ( e) = [ item] . serialize ( & mut serializer) {
472+ dst. truncate ( start) ;
473+ return Err ( io:: Error :: new ( io:: ErrorKind :: InvalidData , e) ) ;
474+ }
475+ }
465476
466- let size = buf . len ( ) ;
477+ let size = dst . len ( ) - start - header_len ;
467478 if size > MAX_MESSAGE_SIZE as usize {
479+ dst. truncate ( start) ;
468480 return Err ( io:: Error :: new (
469481 io:: ErrorKind :: InvalidData ,
470482 format ! (
@@ -479,12 +491,13 @@ impl Encoder<CommandResponse<'_>> for CommandCodec {
479491 marker : Header :: VALID_MARKER ,
480492 size,
481493 } ;
494+ dst[ start..start + header_len] . copy_from_slice ( header. as_slice ( ) ) ;
482495
483- trace ! ( "Encoding message with size {}: {:?}" , size , fmt_bin ( & buf ) ) ;
484-
485- dst . extend_from_slice ( header . as_slice ( ) ) ;
486- dst . reserve ( size as usize ) ;
487- dst . extend_from_slice ( & buf ) ;
496+ trace ! (
497+ "Encoding message with size {}: {:?}" ,
498+ size ,
499+ fmt_bin ( & dst [ start + header_len.. ] )
500+ ) ;
488501
489502 Ok ( ( ) )
490503 }
@@ -561,6 +574,31 @@ mod tests {
561574 encoder. encode ( resp, & mut buf) . unwrap ( ) ;
562575 }
563576
577+ #[ test]
578+ fn test_encode_oversized_truncates_and_preserves_prior_frame ( ) {
579+ let mut buf = BytesMut :: new ( ) ;
580+ let mut encoder = CommandCodec ;
581+ encoder
582+ . encode ( CommandResponse :: ConfigSync , & mut buf)
583+ . unwrap ( ) ;
584+ let good_frame = buf. clone ( ) ;
585+
586+ let huge = CommandResponse :: ClientInit ( ClientInitResp {
587+ status : "x" . repeat ( 5 * 1024 * 1024 ) ,
588+ version : "1.0.0" ,
589+ errors : vec ! [ ] ,
590+ meta : HashMap :: new ( ) ,
591+ metrics : HashMap :: new ( ) ,
592+ helper_runtime : None ,
593+ } ) ;
594+ let res = encoder. encode ( huge, & mut buf) ;
595+ assert ! ( res. is_err( ) ) ;
596+ assert_eq ! (
597+ buf, good_frame,
598+ "oversized-message error must not corrupt already-buffered data"
599+ ) ;
600+ }
601+
564602 fn serialize_message < T : serde:: Serialize > ( command : & T ) -> Vec < u8 > {
565603 let mut buf = Vec :: new ( ) ;
566604 let mut serializer = Serializer :: new ( & mut buf) ;
0 commit comments