@@ -46,7 +46,8 @@ use crate::{
4646 } ,
4747} ;
4848
49- use arrow:: array:: { Array , ArrayRef , Int32Array , UInt32Array , UInt64Array } ;
49+ use arrow:: array:: { Array , ArrayRef , BooleanArray , Int32Array , UInt32Array , UInt64Array } ;
50+ use arrow:: compute:: filter_record_batch;
5051use arrow:: datatypes:: { Schema , SchemaRef } ;
5152use arrow:: record_batch:: RecordBatch ;
5253use arrow_schema:: DataType ;
@@ -727,44 +728,30 @@ impl HashJoinStream {
727728 Map :: RoaringMap ( bitmap) => {
728729 let key_col = & state. values [ 0 ] ;
729730 let is_semi = matches ! ( self . join_type, JoinType :: RightSemi ) ;
730- let right_indices = match key_col. data_type ( ) {
731+ let mask : BooleanArray = match key_col. data_type ( ) {
731732 DataType :: Int32 => {
732733 let arr = key_col. as_any ( ) . downcast_ref :: < Int32Array > ( ) . unwrap ( ) ;
733- arr. values ( )
734- . iter ( )
735- . enumerate ( )
736- . filter_map ( |( i, v) | {
737- let contains = bitmap. contains ( * v as u32 ) ;
738- let emit = if is_semi { contains } else { !contains } ;
739- emit. then_some ( i as u32 )
734+ arr. iter ( )
735+ . map ( |v| match v {
736+ Some ( v) => bitmap. contains ( v as u32 ) == is_semi,
737+ None => !is_semi,
740738 } )
741- . collect :: < Vec < u32 > > ( )
739+ . collect ( )
742740 }
743741 DataType :: UInt32 => {
744742 let arr = key_col. as_any ( ) . downcast_ref :: < UInt32Array > ( ) . unwrap ( ) ;
745- arr. values ( )
746- . iter ( )
747- . enumerate ( )
748- . filter_map ( |( i, v) | {
749- let contains = bitmap. contains ( * v) ;
750- let emit = if is_semi { contains } else { !contains } ;
751- emit. then_some ( i as u32 )
743+ arr. iter ( )
744+ . map ( |v| match v {
745+ Some ( v) => bitmap. contains ( v) == is_semi,
746+ None => !is_semi,
752747 } )
753- . collect :: < Vec < u32 > > ( )
748+ . collect ( )
754749 }
755750 _ => {
756751 return internal_err ! ( "unsupported data type for roaring bitmap" ) ;
757752 }
758753 } ;
759- let indices = UInt32Array :: from ( right_indices) ;
760- let columns: Vec < ArrayRef > = state
761- . batch
762- . columns ( )
763- . iter ( )
764- . map ( |col| arrow:: compute:: take ( col, & indices, None ) )
765- . collect :: < Result < Vec < _ > , _ > > ( ) ?;
766-
767- let batch = RecordBatch :: try_new ( self . schema . clone ( ) , columns) ?;
754+ let batch = filter_record_batch ( & state. batch , & mask) ?;
768755 self . output_buffer . push_batch ( batch) ?;
769756 self . state = HashJoinStreamState :: FetchProbeBatch ;
770757 return Ok ( StatefulStreamResult :: Continue ) ;
0 commit comments