@@ -316,6 +316,187 @@ SearchGroupBy(milvus::OpContext* op_ctx,
316316 }
317317}
318318
319+ template <typename T>
320+ static void
321+ PopulateGroupByValuesByType (const std::shared_ptr<DataGetter<T>>& data_getter,
322+ std::vector<GroupByValueType>& group_by_values,
323+ const std::vector<int64_t >& seg_offsets) {
324+ group_by_values.reserve (seg_offsets.size ());
325+ for (const auto offset : seg_offsets) {
326+ if (offset == INVALID_SEG_OFFSET ) {
327+ group_by_values.emplace_back (std::nullopt );
328+ continue ;
329+ }
330+ group_by_values.emplace_back (data_getter->Get (offset));
331+ }
332+ }
333+
334+ void
335+ PopulateGroupByValues (milvus::OpContext* op_ctx,
336+ const SearchInfo& search_info,
337+ std::vector<GroupByValueType>& group_by_values,
338+ const segcore::SegmentInternalInterface& segment,
339+ const std::vector<int64_t >& seg_offsets) {
340+ FieldId group_by_field_id = search_info.group_by_field_id_ .value ();
341+ auto data_type = segment.GetFieldDataType (group_by_field_id);
342+ switch (data_type) {
343+ case DataType::INT8 : {
344+ PopulateGroupByValuesByType<int8_t >(
345+ GetDataGetter<int8_t >(op_ctx, segment, group_by_field_id),
346+ group_by_values,
347+ seg_offsets);
348+ break ;
349+ }
350+ case DataType::INT16 : {
351+ PopulateGroupByValuesByType<int16_t >(
352+ GetDataGetter<int16_t >(op_ctx, segment, group_by_field_id),
353+ group_by_values,
354+ seg_offsets);
355+ break ;
356+ }
357+ case DataType::INT32 : {
358+ PopulateGroupByValuesByType<int32_t >(
359+ GetDataGetter<int32_t >(op_ctx, segment, group_by_field_id),
360+ group_by_values,
361+ seg_offsets);
362+ break ;
363+ }
364+ case DataType::INT64 :
365+ case DataType::TIMESTAMPTZ : {
366+ PopulateGroupByValuesByType<int64_t >(
367+ GetDataGetter<int64_t >(op_ctx, segment, group_by_field_id),
368+ group_by_values,
369+ seg_offsets);
370+ break ;
371+ }
372+ case DataType::BOOL : {
373+ PopulateGroupByValuesByType<bool >(
374+ GetDataGetter<bool >(op_ctx, segment, group_by_field_id),
375+ group_by_values,
376+ seg_offsets);
377+ break ;
378+ }
379+ case DataType::VARCHAR : {
380+ PopulateGroupByValuesByType<std::string>(
381+ GetDataGetter<std::string>(op_ctx, segment, group_by_field_id),
382+ group_by_values,
383+ seg_offsets);
384+ break ;
385+ }
386+ case DataType::JSON : {
387+ AssertInfo (search_info.json_path_ .has_value (),
388+ " json_path is required for json field when doing "
389+ " search_group_by" );
390+ if (search_info.json_type_ .has_value ()) {
391+ switch (search_info.json_type_ .value ()) {
392+ case DataType::BOOL : {
393+ PopulateGroupByValuesByType<bool >(
394+ GetDataGetter<bool , milvus::Json>(
395+ op_ctx,
396+ segment,
397+ group_by_field_id,
398+ search_info.json_path_ ,
399+ search_info.json_type_ ,
400+ search_info.strict_cast_ ),
401+ group_by_values,
402+ seg_offsets);
403+ break ;
404+ }
405+ case DataType::INT8 : {
406+ PopulateGroupByValuesByType<int8_t >(
407+ GetDataGetter<int8_t , milvus::Json>(
408+ op_ctx,
409+ segment,
410+ group_by_field_id,
411+ search_info.json_path_ ,
412+ search_info.json_type_ ,
413+ search_info.strict_cast_ ),
414+ group_by_values,
415+ seg_offsets);
416+ break ;
417+ }
418+ case DataType::INT16 : {
419+ PopulateGroupByValuesByType<int16_t >(
420+ GetDataGetter<int16_t , milvus::Json>(
421+ op_ctx,
422+ segment,
423+ group_by_field_id,
424+ search_info.json_path_ ,
425+ search_info.json_type_ ,
426+ search_info.strict_cast_ ),
427+ group_by_values,
428+ seg_offsets);
429+ break ;
430+ }
431+ case DataType::INT32 : {
432+ PopulateGroupByValuesByType<int32_t >(
433+ GetDataGetter<int32_t , milvus::Json>(
434+ op_ctx,
435+ segment,
436+ group_by_field_id,
437+ search_info.json_path_ ,
438+ search_info.json_type_ ,
439+ search_info.strict_cast_ ),
440+ group_by_values,
441+ seg_offsets);
442+ break ;
443+ }
444+ case DataType::INT64 : {
445+ PopulateGroupByValuesByType<int64_t >(
446+ GetDataGetter<int64_t , milvus::Json>(
447+ op_ctx,
448+ segment,
449+ group_by_field_id,
450+ search_info.json_path_ ,
451+ search_info.json_type_ ,
452+ search_info.strict_cast_ ),
453+ group_by_values,
454+ seg_offsets);
455+ break ;
456+ }
457+ case DataType::VARCHAR : {
458+ PopulateGroupByValuesByType<std::string>(
459+ GetDataGetter<std::string, milvus::Json>(
460+ op_ctx,
461+ segment,
462+ group_by_field_id,
463+ search_info.json_path_ ,
464+ search_info.json_type_ ,
465+ search_info.strict_cast_ ),
466+ group_by_values,
467+ seg_offsets);
468+ break ;
469+ }
470+ default : {
471+ ThrowInfo (Unsupported,
472+ fmt::format (" unsupported data type {} for "
473+ " group by operator" ,
474+ data_type));
475+ }
476+ }
477+ } else {
478+ PopulateGroupByValuesByType<std::string>(
479+ GetDataGetter<std::string, milvus::Json>(
480+ op_ctx,
481+ segment,
482+ group_by_field_id,
483+ search_info.json_path_ ,
484+ search_info.json_type_ ,
485+ search_info.strict_cast_ ),
486+ group_by_values,
487+ seg_offsets);
488+ }
489+ break ;
490+ }
491+ default : {
492+ ThrowInfo (
493+ Unsupported,
494+ fmt::format (" unsupported data type {} for group by operator" ,
495+ data_type));
496+ }
497+ }
498+ }
499+
319500template <typename T>
320501void
321502GroupIteratorsByType (
@@ -392,4 +573,4 @@ GroupIteratorResult(const std::shared_ptr<VectorIterator>& iterator,
392573}
393574
394575} // namespace exec
395- } // namespace milvus
576+ } // namespace milvus
0 commit comments