Skip to main content

runmat_runtime/builtins/array/sorting_sets/
sort.rs

1//! MATLAB-compatible `sort` builtin with multi-output and GPU-aware semantics.
2
3use std::cmp::Ordering;
4
5use runmat_accelerate_api::{
6    GpuTensorHandle, SortComparison as ProviderSortComparison, SortOrder as ProviderSortOrder,
7};
8use runmat_builtins::{
9    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinIntegerBackendRule,
10    BuiltinIntegerCapabilityDescriptor, BuiltinIntegerComputationDomain,
11    BuiltinIntegerInputAvailability, BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule,
12    BuiltinIntegerOverflowRule, BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule,
13    BuiltinOutputMode, BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType,
14    BuiltinSignatureDescriptor,
15};
16use runmat_macros::runtime_builtin;
17use runmat_value::{
18    ComplexStorage, ComplexTensor, IntValue, IntegerStorage, LogicalArray, NumericStorage, Tensor,
19    Value,
20};
21
22use super::{float_order::SetFloat, integer_order, type_resolvers::tensor_output_type};
23use crate::build_runtime_error;
24use crate::builtins::common::arg_tokens::{tokens_from_values, ArgToken};
25use crate::builtins::common::gpu_helpers;
26use crate::builtins::common::spec::{
27    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
28    ProviderHook, ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
29};
30use crate::builtins::common::tensor;
31
32#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::array::sorting_sets::sort")]
33pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
34    name: "sort",
35    op_kind: GpuOpKind::Custom("sort"),
36    supported_precisions: &[ScalarType::F32, ScalarType::F64],
37    broadcast: BroadcastSemantics::None,
38    provider_hooks: &[ProviderHook::Custom("sort_dim")],
39    constant_strategy: ConstantStrategy::InlineLiteral,
40    residency: ResidencyPolicy::NewHandle,
41    nan_mode: ReductionNaN::Include,
42    two_pass_threshold: None,
43    workgroup_size: None,
44    accepts_nan_mode: true,
45    notes: "Plain real tensors may use the provider sort hook; typed integer, logical, complex, explicit missing-placement, and unsupported-provider paths gather through authoritative host storage and restore both outputs to the owning provider.",
46};
47
48#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::array::sorting_sets::sort")]
49pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
50    name: "sort",
51    shape: ShapeRequirements::Any,
52    constant_strategy: ConstantStrategy::InlineLiteral,
53    elementwise: None,
54    reduction: None,
55    emits_nan: true,
56    notes: "Sorting breaks fusion chains; host fallback may gather upstream tensors before restoring new resident output handles.",
57};
58
59const BUILTIN_NAME: &str = "sort";
60
61const SORT_OUTPUT_B: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
62    name: "B",
63    ty: BuiltinParamType::Any,
64    arity: BuiltinParamArity::Required,
65    default: None,
66    description: "Sorted values.",
67}];
68
69const SORT_OUTPUT_BI: [BuiltinParamDescriptor; 2] = [
70    BuiltinParamDescriptor {
71        name: "B",
72        ty: BuiltinParamType::Any,
73        arity: BuiltinParamArity::Required,
74        default: None,
75        description: "Sorted values.",
76    },
77    BuiltinParamDescriptor {
78        name: "I",
79        ty: BuiltinParamType::NumericArray,
80        arity: BuiltinParamArity::Required,
81        default: None,
82        description: "One-based index permutation for each sorted slice.",
83    },
84];
85
86const SORT_INPUTS_A: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
87    name: "A",
88    ty: BuiltinParamType::Any,
89    arity: BuiltinParamArity::Required,
90    default: None,
91    description: "Input array.",
92}];
93
94const SORT_INPUTS_A_ARG1: [BuiltinParamDescriptor; 2] = [
95    BuiltinParamDescriptor {
96        name: "A",
97        ty: BuiltinParamType::Any,
98        arity: BuiltinParamArity::Required,
99        default: None,
100        description: "Input array.",
101    },
102    BuiltinParamDescriptor {
103        name: "arg1",
104        ty: BuiltinParamType::Any,
105        arity: BuiltinParamArity::Required,
106        default: None,
107        description: "Dimension selector or direction token ('ascend'/'descend').",
108    },
109];
110
111const SORT_INPUTS_A_ARG1_ARG2: [BuiltinParamDescriptor; 3] = [
112    BuiltinParamDescriptor {
113        name: "A",
114        ty: BuiltinParamType::Any,
115        arity: BuiltinParamArity::Required,
116        default: None,
117        description: "Input array.",
118    },
119    BuiltinParamDescriptor {
120        name: "arg1",
121        ty: BuiltinParamType::Any,
122        arity: BuiltinParamArity::Required,
123        default: None,
124        description: "Dimension selector, placeholder, or direction token.",
125    },
126    BuiltinParamDescriptor {
127        name: "arg2",
128        ty: BuiltinParamType::Any,
129        arity: BuiltinParamArity::Required,
130        default: None,
131        description: "Dimension selector or direction token.",
132    },
133];
134
135const SORT_INPUTS_COMPARISON_METHOD: [BuiltinParamDescriptor; 4] = [
136    BuiltinParamDescriptor {
137        name: "A",
138        ty: BuiltinParamType::Any,
139        arity: BuiltinParamArity::Required,
140        default: None,
141        description: "Input array.",
142    },
143    BuiltinParamDescriptor {
144        name: "arg",
145        ty: BuiltinParamType::Any,
146        arity: BuiltinParamArity::Variadic,
147        default: None,
148        description: "Optional dimension/direction arguments.",
149    },
150    BuiltinParamDescriptor {
151        name: "name",
152        ty: BuiltinParamType::StringScalar,
153        arity: BuiltinParamArity::Required,
154        default: Some("\"ComparisonMethod\""),
155        description: "Name-value option key.",
156    },
157    BuiltinParamDescriptor {
158        name: "method",
159        ty: BuiltinParamType::StringScalar,
160        arity: BuiltinParamArity::Required,
161        default: Some("\"auto\""),
162        description: "Comparison method: 'auto', 'real', or 'abs'.",
163    },
164];
165
166const SORT_INPUTS_MISSING_PLACEMENT: [BuiltinParamDescriptor; 4] = [
167    BuiltinParamDescriptor {
168        name: "A",
169        ty: BuiltinParamType::Any,
170        arity: BuiltinParamArity::Required,
171        default: None,
172        description: "Input array.",
173    },
174    BuiltinParamDescriptor {
175        name: "arg",
176        ty: BuiltinParamType::Any,
177        arity: BuiltinParamArity::Variadic,
178        default: None,
179        description: "Optional dimension/direction arguments.",
180    },
181    BuiltinParamDescriptor {
182        name: "name",
183        ty: BuiltinParamType::StringScalar,
184        arity: BuiltinParamArity::Required,
185        default: Some("\"MissingPlacement\""),
186        description: "Name-value option key.",
187    },
188    BuiltinParamDescriptor {
189        name: "placement",
190        ty: BuiltinParamType::StringScalar,
191        arity: BuiltinParamArity::Required,
192        default: Some("\"auto\""),
193        description: "Missing-value placement: 'auto', 'first', or 'last'.",
194    },
195];
196
197const SORT_SIGNATURES: [BuiltinSignatureDescriptor; 10] = [
198    BuiltinSignatureDescriptor {
199        label: "B = sort(A)",
200        inputs: &SORT_INPUTS_A,
201        outputs: &SORT_OUTPUT_B,
202    },
203    BuiltinSignatureDescriptor {
204        label: "B = sort(A, arg1)",
205        inputs: &SORT_INPUTS_A_ARG1,
206        outputs: &SORT_OUTPUT_B,
207    },
208    BuiltinSignatureDescriptor {
209        label: "B = sort(A, arg1, arg2)",
210        inputs: &SORT_INPUTS_A_ARG1_ARG2,
211        outputs: &SORT_OUTPUT_B,
212    },
213    BuiltinSignatureDescriptor {
214        label: "B = sort(A, ..., \"ComparisonMethod\", method)",
215        inputs: &SORT_INPUTS_COMPARISON_METHOD,
216        outputs: &SORT_OUTPUT_B,
217    },
218    BuiltinSignatureDescriptor {
219        label: "B = sort(A, ..., \"MissingPlacement\", placement)",
220        inputs: &SORT_INPUTS_MISSING_PLACEMENT,
221        outputs: &SORT_OUTPUT_B,
222    },
223    BuiltinSignatureDescriptor {
224        label: "[B, I] = sort(A)",
225        inputs: &SORT_INPUTS_A,
226        outputs: &SORT_OUTPUT_BI,
227    },
228    BuiltinSignatureDescriptor {
229        label: "[B, I] = sort(A, arg1)",
230        inputs: &SORT_INPUTS_A_ARG1,
231        outputs: &SORT_OUTPUT_BI,
232    },
233    BuiltinSignatureDescriptor {
234        label: "[B, I] = sort(A, arg1, arg2)",
235        inputs: &SORT_INPUTS_A_ARG1_ARG2,
236        outputs: &SORT_OUTPUT_BI,
237    },
238    BuiltinSignatureDescriptor {
239        label: "[B, I] = sort(A, ..., \"ComparisonMethod\", method)",
240        inputs: &SORT_INPUTS_COMPARISON_METHOD,
241        outputs: &SORT_OUTPUT_BI,
242    },
243    BuiltinSignatureDescriptor {
244        label: "[B, I] = sort(A, ..., \"MissingPlacement\", placement)",
245        inputs: &SORT_INPUTS_MISSING_PLACEMENT,
246        outputs: &SORT_OUTPUT_BI,
247    },
248];
249
250const SORT_ERROR_INVALID_DIMENSION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
251    code: "RM.SORT.INVALID_DIMENSION",
252    identifier: Some("RunMat:sort:InvalidDimension"),
253    when: "Dimension argument is non-positive, non-integer, or otherwise invalid.",
254    message: "sort: invalid dimension argument",
255};
256
257const SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING: BuiltinErrorDescriptor =
258    BuiltinErrorDescriptor {
259        code: "RM.SORT.COMPARISON_METHOD_REQUIRES_STRING",
260        identifier: Some("RunMat:sort:ComparisonMethodRequiresString"),
261        when: "ComparisonMethod option value is not string-like.",
262        message: "sort: 'ComparisonMethod' requires a string value",
263    };
264
265const SORT_ERROR_COMPARISON_METHOD_UNKNOWN: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
266    code: "RM.SORT.COMPARISON_METHOD_UNKNOWN",
267    identifier: Some("RunMat:sort:ComparisonMethodUnknown"),
268    when: "ComparisonMethod option value is not one of 'auto'/'real'/'abs'.",
269    message: "sort: unsupported ComparisonMethod",
270};
271
272const SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING: BuiltinErrorDescriptor =
273    BuiltinErrorDescriptor {
274        code: "RM.SORT.MISSINGPLACEMENT_REQUIRES_STRING",
275        identifier: Some("RunMat:sort:MissingPlacementRequiresString"),
276        when: "MissingPlacement option value is not string-like.",
277        message: "sort: 'MissingPlacement' requires a string value",
278    };
279
280const SORT_ERROR_MISSINGPLACEMENT_UNKNOWN: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
281    code: "RM.SORT.MISSINGPLACEMENT_UNKNOWN",
282    identifier: Some("RunMat:sort:MissingPlacementUnknown"),
283    when: "MissingPlacement option value is not one of 'auto'/'first'/'last'.",
284    message: "sort: unsupported MissingPlacement",
285};
286
287const SORT_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
288    code: "RM.SORT.INVALID_ARGUMENT",
289    identifier: Some("RunMat:sort:InvalidArgument"),
290    when: "Parser encounters invalid or unrecognized option/value arguments.",
291    message: "sort: invalid argument sequence",
292};
293
294const SORT_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
295    code: "RM.SORT.INTERNAL",
296    identifier: Some("RunMat:sort:Internal"),
297    when: "Internal conversion, allocation, or provider result construction fails.",
298    message: "sort: internal operation failed",
299};
300
301const SORT_ERRORS: [BuiltinErrorDescriptor; 7] = [
302    SORT_ERROR_INVALID_DIMENSION,
303    SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING,
304    SORT_ERROR_COMPARISON_METHOD_UNKNOWN,
305    SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING,
306    SORT_ERROR_MISSINGPLACEMENT_UNKNOWN,
307    SORT_ERROR_INVALID_ARGUMENT,
308    SORT_ERROR_INTERNAL,
309];
310
311const SORT_INTEGER_INPUTS: [BuiltinIntegerInputCapability; 2] = [
312    BuiltinIntegerInputCapability {
313        name: "A",
314        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
315        availability: BuiltinIntegerInputAvailability::Documented,
316        scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
317        notes: "The documented sortable input domain includes all eight real integer classes.",
318    },
319    BuiltinIntegerInputCapability {
320        name: "dim",
321        classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
322        availability: BuiltinIntegerInputAvailability::Documented,
323        scalar_double: BuiltinIntegerScalarDoubleRule::Allowed,
324        notes: "The documented positive integer-scalar dimension accepts every integer class and integer-valued scalar double.",
325    },
326];
327
328const SORT_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
329    [BuiltinIntegerCapabilityDescriptor {
330        form: "[B, I] = sort(integer_A, integer_dim, direction, options)",
331        inputs: &SORT_INTEGER_INPUTS,
332        computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
333        output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
334        overflow: BuiltinIntegerOverflowRule::NotApplicable,
335        backend: BuiltinIntegerBackendRule::HostAndGpu,
336        overload: BuiltinIntegerOverloadKind::Multiple,
337        notes: "B preserves A's exact integer class and stable equal-value order; optional I is one-based double. Resident integer input uses exact typed gather fallback and restores both outputs to the owning provider.",
338    }];
339
340pub const SORT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
341    signatures: &SORT_SIGNATURES,
342    output_mode: BuiltinOutputMode::ByRequestedOutputCount,
343    completion_policy: BuiltinCompletionPolicy::Public,
344    errors: &SORT_ERRORS,
345};
346
347fn sort_error(
348    error: &'static BuiltinErrorDescriptor,
349    message: impl Into<String>,
350) -> crate::RuntimeError {
351    let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
352    if let Some(identifier) = error.identifier {
353        builder = builder.with_identifier(identifier);
354    }
355    builder.build()
356}
357
358fn sort_internal(message: impl Into<String>) -> crate::RuntimeError {
359    sort_error(&SORT_ERROR_INTERNAL, message)
360}
361
362fn sort_invalid_argument(message: impl Into<String>) -> crate::RuntimeError {
363    sort_error(&SORT_ERROR_INVALID_ARGUMENT, message)
364}
365
366#[runtime_builtin(
367    name = "sort",
368    category = "array/sorting_sets",
369    summary = "Sort array elements along a dimension with optional index outputs.",
370    keywords = "sort,ascending,descending,indices,comparisonmethod,gpu",
371    accel = "sink",
372    sink = true,
373    type_resolver(tensor_output_type),
374    descriptor(crate::builtins::array::sorting_sets::sort::SORT_DESCRIPTOR),
375    integer_capabilities(SORT_INTEGER_CAPABILITIES),
376    builtin_path = "crate::builtins::array::sorting_sets::sort"
377)]
378async fn sort_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
379    let eval = evaluate(value, &rest).await?;
380    if let Some(out_count) = crate::output_count::current_output_count() {
381        if out_count == 0 {
382            return Ok(Value::OutputList(Vec::new()));
383        }
384        let (sorted, indices) = eval.into_values();
385        if out_count == 1 {
386            return Ok(Value::OutputList(vec![sorted]));
387        }
388        if out_count == 2 {
389            return Ok(Value::OutputList(vec![sorted, indices]));
390        }
391        return Err(sort_invalid_argument(
392            "sort: too many output arguments; maximum is 2",
393        ));
394    }
395    Ok(eval.into_sorted_value())
396}
397
398/// Evaluate the `sort` builtin once and expose both outputs.
399pub async fn evaluate(value: Value, rest: &[Value]) -> crate::BuiltinResult<SortEvaluation> {
400    let args = SortArgs::parse(rest)?;
401    match value {
402        Value::GpuTensor(handle) => sort_gpu(handle, &args).await,
403        other => sort_host(other, &args),
404    }
405}
406
407async fn sort_gpu(
408    handle: GpuTensorHandle,
409    args: &SortArgs,
410) -> crate::BuiltinResult<SortEvaluation> {
411    let shape = handle.shape.clone();
412    let provider = runmat_accelerate_api::provider_for_handle(&handle)
413        .or_else(runmat_accelerate_api::provider);
414    let dim = args.dimension.unwrap_or_else(|| default_dimension(&shape));
415    if dim == 0 {
416        return Err(sort_error(
417            &SORT_ERROR_INVALID_DIMENSION,
418            "sort: dimension must be >= 1",
419        ));
420    }
421    let dim_len = dimension_length(&shape, dim);
422    let plain_real = runmat_accelerate_api::handle_integer_type(&handle).is_none()
423        && !runmat_accelerate_api::handle_is_logical(&handle)
424        && runmat_accelerate_api::handle_storage(&handle)
425            == runmat_accelerate_api::GpuTensorStorage::Real;
426    if dim_len > 1 && plain_real && matches!(args.missing, MissingPlacement::Auto) {
427        if let Some(provider) = provider {
428            let order = args.direction.to_provider();
429            let comparison = args.comparison.to_provider();
430            let zero_based = dim - 1;
431            if let Ok(result) = provider
432                .sort_dim(&handle, zero_based, order, comparison)
433                .await
434            {
435                let sorted_tensor = Tensor::new(result.values.data, result.values.shape)
436                    .map_err(|e| sort_internal(format!("sort: {e}")))?;
437                let sorted_value = tensor::tensor_into_value(sorted_tensor);
438                let indices_tensor = Tensor::new(result.indices.data, result.indices.shape)
439                    .map_err(|e| sort_internal(format!("sort: {e}")))?;
440                return upload_sort_evaluation(
441                    provider,
442                    SortEvaluation {
443                        sorted: sorted_value,
444                        indices: tensor::tensor_into_value(indices_tensor),
445                    },
446                );
447            }
448        }
449    }
450    let host = gpu_helpers::gather_value_async(&Value::GpuTensor(handle)).await?;
451    let evaluation = sort_host(host, args)?;
452    match provider {
453        Some(provider) => upload_sort_evaluation(provider, evaluation),
454        None => Ok(evaluation),
455    }
456}
457
458fn sort_host(value: Value, args: &SortArgs) -> crate::BuiltinResult<SortEvaluation> {
459    match value {
460        Value::ComplexTensor(ct) => sort_complex_tensor(ct, args),
461        Value::Complex(re, im) => {
462            let tensor = ComplexTensor::new(vec![(re, im)], vec![1, 1])
463                .map_err(|e| sort_internal(format!("sort: {e}")))?;
464            sort_complex_tensor(tensor, args)
465        }
466        Value::Int(value) => {
467            let tensor = Tensor::new_integer(IntegerStorage::from_scalar(value), vec![1, 1])
468                .map_err(|e| sort_internal(format!("sort: {e}")))?;
469            sort_real_tensor(tensor, args)
470        }
471        Value::LogicalArray(logical) => sort_logical(logical, args),
472        Value::Bool(value) => {
473            let logical = LogicalArray::new(vec![u8::from(value)], vec![1, 1])
474                .map_err(|e| sort_internal(format!("sort: {e}")))?;
475            sort_logical(logical, args)
476        }
477        other => {
478            let tensor =
479                tensor::value_into_tensor_for("sort", other).map_err(sort_invalid_argument)?;
480            sort_real_tensor(tensor, args)
481        }
482    }
483}
484
485fn sort_logical(logical: LogicalArray, args: &SortArgs) -> crate::BuiltinResult<SortEvaluation> {
486    let shape = logical.shape;
487    let evaluation = sort_real_tensor(
488        Tensor::new_integer(IntegerStorage::U8(logical.data), shape)
489            .map_err(|e| sort_internal(format!("sort: {e}")))?,
490        args,
491    )?;
492    let sorted = match evaluation.sorted {
493        Value::Int(IntValue::U8(value)) => Value::Bool(value != 0),
494        Value::Tensor(tensor) => {
495            let shape = tensor.shape.clone();
496            let storage = tensor
497                .into_numeric_storage()
498                .map_err(|e| sort_internal(format!("sort: {e}")))?
499                .into_integer_storage()
500                .map_err(|_| sort_internal("sort: logical output lost integer storage"))?;
501            let IntegerStorage::U8(values) = storage else {
502                return Err(sort_internal("sort: logical output changed storage class"));
503            };
504            if values.len() == 1 {
505                Value::Bool(values[0] != 0)
506            } else {
507                Value::LogicalArray(
508                    LogicalArray::new(values, shape)
509                        .map_err(|e| sort_internal(format!("sort: {e}")))?,
510                )
511            }
512        }
513        other => {
514            return Err(sort_internal(format!(
515                "sort: unexpected logical output {other:?}"
516            )))
517        }
518    };
519    Ok(SortEvaluation {
520        sorted,
521        indices: evaluation.indices,
522    })
523}
524
525fn upload_sort_evaluation(
526    provider: &'static dyn runmat_accelerate_api::AccelProvider,
527    evaluation: SortEvaluation,
528) -> crate::BuiltinResult<SortEvaluation> {
529    Ok(SortEvaluation {
530        sorted: upload_sort_value(provider, evaluation.sorted)?,
531        indices: upload_sort_value(provider, evaluation.indices)?,
532    })
533}
534
535fn upload_sort_value(
536    provider: &'static dyn runmat_accelerate_api::AccelProvider,
537    value: Value,
538) -> crate::BuiltinResult<Value> {
539    let upload_tensor = |tensor: Tensor, logical: bool| -> crate::BuiltinResult<Value> {
540        let handle = gpu_helpers::upload_tensor(provider, &tensor)
541            .map_err(|error| sort_internal(format!("sort: GPU upload failed: {error}")))?;
542        Ok(if logical {
543            gpu_helpers::logical_gpu_value(handle)
544        } else {
545            gpu_helpers::resident_gpu_value(handle)
546        })
547    };
548    match value {
549        Value::Tensor(tensor) => upload_tensor(tensor, false),
550        Value::Num(number) => upload_tensor(
551            Tensor::new(vec![number], vec![1, 1])
552                .map_err(|error| sort_internal(format!("sort: {error}")))?,
553            false,
554        ),
555        Value::Int(integer) => upload_tensor(
556            Tensor::new_integer(IntegerStorage::from_scalar(integer), vec![1, 1])
557                .map_err(|error| sort_internal(format!("sort: {error}")))?,
558            false,
559        ),
560        Value::Bool(logical) => upload_tensor(
561            Tensor::new(vec![if logical { 1.0 } else { 0.0 }], vec![1, 1])
562                .map_err(|error| sort_internal(format!("sort: {error}")))?,
563            true,
564        ),
565        Value::LogicalArray(logical) => {
566            let tensor = tensor::logical_to_tensor(&logical).map_err(sort_internal)?;
567            upload_tensor(tensor, true)
568        }
569        Value::Complex(real, imag) => {
570            let tensor = ComplexTensor::new(vec![(real, imag)], vec![1, 1])
571                .map_err(|error| sort_internal(format!("sort: {error}")))?;
572            let handle = gpu_helpers::upload_complex_tensor(provider, &tensor)
573                .map_err(|error| sort_internal(format!("sort: {}", error.message())))?;
574            Ok(gpu_helpers::complex_gpu_value(handle))
575        }
576        Value::ComplexTensor(tensor) => {
577            let handle = gpu_helpers::upload_complex_tensor(provider, &tensor)
578                .map_err(|error| sort_internal(format!("sort: {}", error.message())))?;
579            Ok(gpu_helpers::complex_gpu_value(handle))
580        }
581        other => Err(sort_internal(format!(
582            "sort: cannot upload unexpected output {other:?}"
583        ))),
584    }
585}
586
587fn sort_real_tensor(tensor: Tensor, args: &SortArgs) -> crate::BuiltinResult<SortEvaluation> {
588    let shape = tensor.shape.clone();
589    let dim = args.dimension.unwrap_or_else(|| default_dimension(&shape));
590    if dim == 0 {
591        return Err(sort_error(
592            &SORT_ERROR_INVALID_DIMENSION,
593            "sort: dimension must be >= 1",
594        ));
595    }
596
597    let storage = tensor
598        .into_numeric_storage()
599        .map_err(|e| sort_internal(format!("sort: {e}")))?;
600    match storage {
601        NumericStorage::F64(values) => {
602            sort_floating_tensor(values, shape, dim, args, NumericStorage::F64)
603        }
604        NumericStorage::F32(values) => {
605            sort_floating_tensor(values, shape, dim, args, NumericStorage::F32)
606        }
607        storage => sort_integer_tensor(
608            storage
609                .into_integer_storage()
610                .expect("non-floating numeric storage is integer"),
611            shape,
612            dim,
613            args,
614        ),
615    }
616}
617
618fn sort_floating_tensor<T>(
619    values: Vec<T>,
620    shape: Vec<usize>,
621    dim: usize,
622    args: &SortArgs,
623    wrap: fn(Vec<T>) -> NumericStorage,
624) -> crate::BuiltinResult<SortEvaluation>
625where
626    T: SetFloat,
627{
628    let dim_len = dimension_length(&shape, dim);
629    if values.is_empty() || dim_len <= 1 {
630        let indices = vec![1.0; values.len()];
631        let index_tensor =
632            Tensor::new(indices, shape.clone()).map_err(|e| sort_internal(format!("sort: {e}")))?;
633        let sorted_tensor = Tensor::from_numeric_storage(wrap(values), shape)
634            .map_err(|e| sort_internal(format!("sort: {e}")))?;
635        return Ok(SortEvaluation {
636            sorted: tensor::tensor_into_value(sorted_tensor),
637            indices: tensor::tensor_into_value(index_tensor),
638        });
639    }
640
641    let stride_before = stride_before(&shape, dim);
642    let stride_after = stride_after(&shape, dim);
643    let mut sorted = values;
644    let mut indices = vec![0.0f64; sorted.len()];
645    let mut buffer: Vec<(usize, T)> = Vec::with_capacity(dim_len);
646
647    for after in 0..stride_after {
648        for before in 0..stride_before {
649            buffer.clear();
650            for k in 0..dim_len {
651                let idx = before + k * stride_before + after * stride_before * dim_len;
652                let value = sorted[idx];
653                buffer.push((k, value));
654            }
655            buffer.sort_by(|a, b| compare_real_values(a.1, b.1, args));
656            for (pos, (original_index, value)) in buffer.iter().enumerate() {
657                let target = before + pos * stride_before + after * stride_before * dim_len;
658                sorted[target] = *value;
659                indices[target] = (*original_index + 1) as f64;
660            }
661        }
662    }
663
664    let sorted_tensor = Tensor::from_numeric_storage(wrap(sorted), shape.clone())
665        .map_err(|e| sort_internal(format!("sort: {e}")))?;
666    let index_tensor =
667        Tensor::new(indices, shape).map_err(|e| sort_internal(format!("sort: {e}")))?;
668
669    Ok(SortEvaluation {
670        sorted: tensor::tensor_into_value(sorted_tensor),
671        indices: tensor::tensor_into_value(index_tensor),
672    })
673}
674
675fn sort_integer_tensor(
676    storage: IntegerStorage,
677    shape: Vec<usize>,
678    dim: usize,
679    args: &SortArgs,
680) -> crate::BuiltinResult<SortEvaluation> {
681    let stride_before = stride_before(&shape, dim);
682    let stride_after = stride_after(&shape, dim);
683    let dim_len = dimension_length(&shape, dim);
684    let mut sorted = storage.exact_values();
685    let mut indices = vec![0.0f64; sorted.len()];
686    let mut buffer: Vec<(usize, IntValue)> = Vec::with_capacity(dim_len);
687
688    for after in 0..stride_after {
689        for before in 0..stride_before {
690            buffer.clear();
691            for k in 0..dim_len {
692                let idx = before + k * stride_before + after * stride_before * dim_len;
693                buffer.push((k, sorted[idx].clone()));
694            }
695            buffer.sort_by(|a, b| {
696                integer_order::compare(
697                    &a.1,
698                    &b.1,
699                    matches!(args.direction, SortDirection::Descend),
700                    matches!(args.comparison, ComparisonMethod::Abs),
701                )
702            });
703            for (pos, (original_index, value)) in buffer.iter().enumerate() {
704                let target = before + pos * stride_before + after * stride_before * dim_len;
705                sorted[target] = value.clone();
706                indices[target] = (*original_index + 1) as f64;
707            }
708        }
709    }
710
711    let sorted_storage = storage
712        .from_exact_values_like(sorted)
713        .map_err(|e| sort_internal(format!("sort: {e}")))?;
714    let sorted_tensor = Tensor::new_integer(sorted_storage, shape.clone())
715        .map_err(|e| sort_internal(format!("sort: {e}")))?;
716    let index_tensor =
717        Tensor::new(indices, shape).map_err(|e| sort_internal(format!("sort: {e}")))?;
718    Ok(SortEvaluation {
719        sorted: Value::Tensor(sorted_tensor),
720        indices: tensor::tensor_into_value(index_tensor),
721    })
722}
723
724fn sort_complex_tensor(
725    tensor: ComplexTensor,
726    args: &SortArgs,
727) -> crate::BuiltinResult<SortEvaluation> {
728    let shape = tensor.shape.clone();
729    let dim = args.dimension.unwrap_or_else(|| default_dimension(&shape));
730    if dim == 0 {
731        return Err(sort_error(
732            &SORT_ERROR_INVALID_DIMENSION,
733            "sort: dimension must be >= 1",
734        ));
735    }
736
737    let dim_len = dimension_length(&shape, dim);
738    let storage = tensor.into_complex_storage();
739    if storage.is_empty() || dim_len <= 1 {
740        let indices = vec![1.0; storage.len()];
741        let index_tensor =
742            Tensor::new(indices, shape.clone()).map_err(|e| sort_internal(format!("sort: {e}")))?;
743        let sorted_tensor = ComplexTensor::from_complex_storage(storage, shape)
744            .map_err(|e| sort_internal(format!("sort: {e}")))?;
745        return Ok(SortEvaluation {
746            sorted: complex_tensor_into_value(sorted_tensor),
747            indices: tensor::tensor_into_value(index_tensor),
748        });
749    }
750
751    let (source_indices, indices) = match &storage {
752        ComplexStorage::F64(values) => complex_sort_permutation(values, &shape, dim, args),
753        ComplexStorage::F32(values) => complex_sort_permutation(values, &shape, dim, args),
754        ComplexStorage::Integer(values) => {
755            complex_integer_sort_permutation(values, &shape, dim, args)
756        }
757    };
758    let sorted_storage = storage
759        .gather(&source_indices)
760        .map_err(|e| sort_internal(format!("sort: {e}")))?;
761    let sorted_tensor = ComplexTensor::from_complex_storage(sorted_storage, shape.clone())
762        .map_err(|e| sort_internal(format!("sort: {e}")))?;
763    let index_tensor =
764        Tensor::new(indices, shape).map_err(|e| sort_internal(format!("sort: {e}")))?;
765
766    Ok(SortEvaluation {
767        sorted: complex_tensor_into_value(sorted_tensor),
768        indices: tensor::tensor_into_value(index_tensor),
769    })
770}
771
772fn complex_integer_sort_permutation(
773    storage: &runmat_value::IntegerComplexStorage,
774    shape: &[usize],
775    dim: usize,
776    args: &SortArgs,
777) -> (Vec<usize>, Vec<f64>) {
778    let stride_before = stride_before(shape, dim);
779    let stride_after = stride_after(shape, dim);
780    let dim_len = dimension_length(shape, dim);
781    let mut source_indices = vec![0usize; storage.len()];
782    let mut indices = vec![0.0f64; storage.len()];
783    let mut buffer: Vec<(usize, usize)> = Vec::with_capacity(dim_len);
784
785    for after in 0..stride_after {
786        for before in 0..stride_before {
787            buffer.clear();
788            for k in 0..dim_len {
789                let source = before + k * stride_before + after * stride_before * dim_len;
790                buffer.push((k, source));
791            }
792            buffer.sort_by(|a, b| {
793                let a_real = storage.real.value_at(a.1).expect("validated complex index");
794                let a_imag = storage.imag.value_at(a.1).expect("validated complex index");
795                let b_real = storage.real.value_at(b.1).expect("validated complex index");
796                let b_imag = storage.imag.value_at(b.1).expect("validated complex index");
797                integer_order::compare_complex(
798                    (&a_real, &a_imag),
799                    (&b_real, &b_imag),
800                    matches!(args.direction, SortDirection::Descend),
801                    matches!(args.comparison, ComparisonMethod::Real),
802                )
803            });
804            for (position, (original_index, source)) in buffer.iter().enumerate() {
805                let target = before + position * stride_before + after * stride_before * dim_len;
806                source_indices[target] = *source;
807                indices[target] = (*original_index + 1) as f64;
808            }
809        }
810    }
811
812    (source_indices, indices)
813}
814
815fn complex_sort_permutation<T: SetFloat>(
816    values: &[(T, T)],
817    shape: &[usize],
818    dim: usize,
819    args: &SortArgs,
820) -> (Vec<usize>, Vec<f64>) {
821    let stride_before = stride_before(shape, dim);
822    let stride_after = stride_after(shape, dim);
823    let dim_len = dimension_length(shape, dim);
824    let mut source_indices = vec![0usize; values.len()];
825    let mut indices = vec![0.0f64; values.len()];
826    let mut buffer: Vec<(usize, usize, (T, T))> = Vec::with_capacity(dim_len);
827
828    for after in 0..stride_after {
829        for before in 0..stride_before {
830            buffer.clear();
831            for k in 0..dim_len {
832                let source = before + k * stride_before + after * stride_before * dim_len;
833                buffer.push((k, source, values[source]));
834            }
835            buffer.sort_by(|a, b| compare_complex_values(a.2, b.2, args));
836            for (position, (original_index, source, _)) in buffer.iter().enumerate() {
837                let target = before + position * stride_before + after * stride_before * dim_len;
838                source_indices[target] = *source;
839                indices[target] = (*original_index + 1) as f64;
840            }
841        }
842    }
843
844    (source_indices, indices)
845}
846
847fn complex_tensor_into_value(tensor: ComplexTensor) -> Value {
848    if let Some([value]) = tensor.as_f64_slice() {
849        Value::Complex(value.0, value.1)
850    } else {
851        Value::ComplexTensor(tensor)
852    }
853}
854
855fn compare_real_values<T: SetFloat>(a: T, b: T, args: &SortArgs) -> Ordering {
856    match (a.is_nan(), b.is_nan()) {
857        (true, true) => Ordering::Equal,
858        (true, false) => match args.missing.resolve(args.direction) {
859            MissingPlacementResolved::First => Ordering::Less,
860            MissingPlacementResolved::Last => Ordering::Greater,
861        },
862        (false, true) => match args.missing.resolve(args.direction) {
863            MissingPlacementResolved::First => Ordering::Greater,
864            MissingPlacementResolved::Last => Ordering::Less,
865        },
866        (false, false) => compare_real_finite(a, b, args),
867    }
868}
869
870fn compare_real_finite<T: SetFloat>(a: T, b: T, args: &SortArgs) -> Ordering {
871    let primary = match args.comparison {
872        ComparisonMethod::Abs => {
873            let abs_cmp = a.abs().compare(b.abs());
874            if abs_cmp != Ordering::Equal {
875                return match args.direction {
876                    SortDirection::Ascend => abs_cmp,
877                    SortDirection::Descend => abs_cmp.reverse(),
878                };
879            }
880            Ordering::Equal
881        }
882        ComparisonMethod::Auto | ComparisonMethod::Real => Ordering::Equal,
883    };
884    if primary != Ordering::Equal {
885        return primary;
886    }
887    let ordering = if matches!(args.comparison, ComparisonMethod::Abs) {
888        b.compare(a)
889    } else {
890        a.compare(b)
891    };
892    match args.direction {
893        SortDirection::Ascend => ordering,
894        SortDirection::Descend => ordering.reverse(),
895    }
896}
897
898fn compare_complex_values<T: SetFloat>(a: (T, T), b: (T, T), args: &SortArgs) -> Ordering {
899    match (complex_is_nan(a), complex_is_nan(b)) {
900        (true, true) => Ordering::Equal,
901        (true, false) => match args.missing.resolve(args.direction) {
902            MissingPlacementResolved::First => Ordering::Less,
903            MissingPlacementResolved::Last => Ordering::Greater,
904        },
905        (false, true) => match args.missing.resolve(args.direction) {
906            MissingPlacementResolved::First => Ordering::Greater,
907            MissingPlacementResolved::Last => Ordering::Less,
908        },
909        (false, false) => compare_complex_finite(a, b, args),
910    }
911}
912
913fn compare_complex_finite<T: SetFloat>(a: (T, T), b: (T, T), args: &SortArgs) -> Ordering {
914    match args.comparison {
915        ComparisonMethod::Real => compare_complex_real_imag(a, b, args.direction),
916        ComparisonMethod::Abs | ComparisonMethod::Auto => {
917            let abs_cmp = complex_abs(a).compare(complex_abs(b));
918            if abs_cmp != Ordering::Equal {
919                return match args.direction {
920                    SortDirection::Ascend => abs_cmp,
921                    SortDirection::Descend => abs_cmp.reverse(),
922                };
923            }
924            compare_complex_phase(a, b, args.direction)
925        }
926    }
927}
928
929fn compare_complex_phase<T: SetFloat>(a: (T, T), b: (T, T), direction: SortDirection) -> Ordering {
930    let ordering = complex_phase(a).compare(complex_phase(b));
931    match direction {
932        SortDirection::Ascend => ordering,
933        SortDirection::Descend => ordering.reverse(),
934    }
935}
936
937fn complex_phase<T: SetFloat>((real, imaginary): (T, T)) -> T {
938    let imaginary = if imaginary == T::default() {
939        T::default()
940    } else {
941        imaginary
942    };
943    imaginary.atan2(real)
944}
945
946fn compare_complex_real_imag<T: SetFloat>(
947    a: (T, T),
948    b: (T, T),
949    direction: SortDirection,
950) -> Ordering {
951    let real_cmp = match direction {
952        SortDirection::Ascend => a.0.compare(b.0),
953        SortDirection::Descend => b.0.compare(a.0),
954    };
955    if real_cmp != Ordering::Equal {
956        return real_cmp;
957    }
958    match direction {
959        SortDirection::Ascend => a.1.compare(b.1),
960        SortDirection::Descend => b.1.compare(a.1),
961    }
962}
963
964fn complex_is_nan<T: SetFloat>(value: (T, T)) -> bool {
965    value.0.is_nan() || value.1.is_nan()
966}
967
968fn complex_abs<T: SetFloat>(value: (T, T)) -> T {
969    value.0.hypot(value.1)
970}
971
972fn stride_before(shape: &[usize], dim: usize) -> usize {
973    if dim <= 1 {
974        return 1;
975    }
976    let mut product = 1usize;
977    for i in 0..(dim - 1) {
978        product = product.saturating_mul(*shape.get(i).unwrap_or(&1));
979    }
980    product
981}
982
983fn stride_after(shape: &[usize], dim: usize) -> usize {
984    if dim >= shape.len() {
985        return 1;
986    }
987    let mut product = 1usize;
988    for extent in shape.iter().skip(dim) {
989        product = product.saturating_mul(*extent);
990    }
991    product
992}
993
994fn dimension_length(shape: &[usize], dim: usize) -> usize {
995    shape.get(dim - 1).copied().unwrap_or(1)
996}
997
998fn default_dimension(shape: &[usize]) -> usize {
999    shape
1000        .iter()
1001        .position(|&extent| extent > 1)
1002        .map(|idx| idx + 1)
1003        .unwrap_or(1)
1004}
1005
1006#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
1007enum SortDirection {
1008    #[default]
1009    Ascend,
1010    Descend,
1011}
1012
1013impl SortDirection {
1014    fn to_provider(self) -> ProviderSortOrder {
1015        match self {
1016            SortDirection::Ascend => ProviderSortOrder::Ascend,
1017            SortDirection::Descend => ProviderSortOrder::Descend,
1018        }
1019    }
1020}
1021
1022#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
1023enum ComparisonMethod {
1024    #[default]
1025    Auto,
1026    Real,
1027    Abs,
1028}
1029
1030#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
1031enum MissingPlacement {
1032    #[default]
1033    Auto,
1034    First,
1035    Last,
1036}
1037
1038#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1039enum MissingPlacementResolved {
1040    First,
1041    Last,
1042}
1043
1044impl MissingPlacement {
1045    fn resolve(self, direction: SortDirection) -> MissingPlacementResolved {
1046        match self {
1047            Self::First => MissingPlacementResolved::First,
1048            Self::Last => MissingPlacementResolved::Last,
1049            Self::Auto => match direction {
1050                SortDirection::Ascend => MissingPlacementResolved::Last,
1051                SortDirection::Descend => MissingPlacementResolved::First,
1052            },
1053        }
1054    }
1055}
1056
1057impl ComparisonMethod {
1058    fn to_provider(self) -> ProviderSortComparison {
1059        match self {
1060            ComparisonMethod::Auto => ProviderSortComparison::Auto,
1061            ComparisonMethod::Real => ProviderSortComparison::Real,
1062            ComparisonMethod::Abs => ProviderSortComparison::Abs,
1063        }
1064    }
1065}
1066
1067#[derive(Debug, Clone, Default)]
1068struct SortArgs {
1069    dimension: Option<usize>,
1070    direction: SortDirection,
1071    comparison: ComparisonMethod,
1072    missing: MissingPlacement,
1073}
1074
1075impl SortArgs {
1076    fn parse(rest: &[Value]) -> crate::BuiltinResult<Self> {
1077        let mut args = SortArgs::default();
1078        let tokens = tokens_from_values(rest);
1079        let mut i = 0usize;
1080        while i < rest.len() {
1081            if args.dimension.is_none() {
1082                if is_dimension_placeholder(&rest[i]) {
1083                    i += 1;
1084                    continue;
1085                }
1086                match tensor::parse_dimension(&rest[i], "sort") {
1087                    Ok(dim) => {
1088                        args.dimension = Some(dim);
1089                        i += 1;
1090                        continue;
1091                    }
1092                    Err(err) => {
1093                        if matches!(rest[i], Value::Int(_) | Value::Num(_)) {
1094                            return Err(sort_error(&SORT_ERROR_INVALID_DIMENSION, err));
1095                        }
1096                    }
1097                }
1098            }
1099            if let Some(ArgToken::String(text)) = tokens.get(i) {
1100                match text.as_str() {
1101                    "ascend" | "ascending" => {
1102                        args.direction = SortDirection::Ascend;
1103                        i += 1;
1104                        continue;
1105                    }
1106                    "descend" | "descending" => {
1107                        args.direction = SortDirection::Descend;
1108                        i += 1;
1109                        continue;
1110                    }
1111                    "comparisonmethod" => {
1112                        i += 1;
1113                        if i >= rest.len() {
1114                            return Err(sort_invalid_argument(
1115                                "sort: expected a value for 'ComparisonMethod'",
1116                            ));
1117                        }
1118                        let value = match tokens.get(i) {
1119                            Some(ArgToken::String(value)) => value.as_str(),
1120                            _ => {
1121                                return Err(sort_error(
1122                                    &SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING,
1123                                    SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING.message,
1124                                ))
1125                            }
1126                        };
1127                        args.comparison = match value {
1128                            "auto" => ComparisonMethod::Auto,
1129                            "real" => ComparisonMethod::Real,
1130                            "abs" | "magnitude" => ComparisonMethod::Abs,
1131                            other => {
1132                                return Err(sort_error(
1133                                    &SORT_ERROR_COMPARISON_METHOD_UNKNOWN,
1134                                    format!("sort: unsupported ComparisonMethod '{other}'"),
1135                                )
1136                                .into())
1137                            }
1138                        };
1139                        i += 1;
1140                        continue;
1141                    }
1142                    "missingplacement" => {
1143                        i += 1;
1144                        if i >= rest.len() {
1145                            return Err(sort_invalid_argument(
1146                                "sort: expected a value for 'MissingPlacement'",
1147                            ));
1148                        }
1149                        let value = match tokens.get(i) {
1150                            Some(ArgToken::String(value)) => value.as_str(),
1151                            _ => {
1152                                return Err(sort_error(
1153                                    &SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING,
1154                                    SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING.message,
1155                                ))
1156                            }
1157                        };
1158                        args.missing = match value {
1159                            "auto" => MissingPlacement::Auto,
1160                            "first" => MissingPlacement::First,
1161                            "last" => MissingPlacement::Last,
1162                            other => {
1163                                return Err(sort_error(
1164                                    &SORT_ERROR_MISSINGPLACEMENT_UNKNOWN,
1165                                    format!("sort: unsupported MissingPlacement '{other}'"),
1166                                ))
1167                            }
1168                        };
1169                        i += 1;
1170                        continue;
1171                    }
1172                    _ => {}
1173                }
1174            }
1175            if let Some(keyword) = tensor::value_to_string(&rest[i]) {
1176                let lowered = keyword.trim().to_ascii_lowercase();
1177                match lowered.as_str() {
1178                    "ascend" | "ascending" => {
1179                        args.direction = SortDirection::Ascend;
1180                        i += 1;
1181                        continue;
1182                    }
1183                    "descend" | "descending" => {
1184                        args.direction = SortDirection::Descend;
1185                        i += 1;
1186                        continue;
1187                    }
1188                    "comparisonmethod" => {
1189                        i += 1;
1190                        if i >= rest.len() {
1191                            return Err(sort_invalid_argument(
1192                                "sort: expected a value for 'ComparisonMethod'",
1193                            ));
1194                        }
1195                        let raw = &rest[i];
1196                        let value = match raw {
1197                            Value::String(s) => s.clone(),
1198                            Value::StringArray(sa) if sa.data.len() == 1 => sa.data[0].clone(),
1199                            Value::CharArray(ca) if ca.rows == 1 => {
1200                                ca.data.iter().copied().collect()
1201                            }
1202                            _ => {
1203                                return Err(sort_error(
1204                                    &SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING,
1205                                    SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING.message,
1206                                ))
1207                            }
1208                        };
1209                        let lowered_value = value.trim().to_ascii_lowercase();
1210                        args.comparison = match lowered_value.as_str() {
1211                            "auto" => ComparisonMethod::Auto,
1212                            "real" => ComparisonMethod::Real,
1213                            "abs" | "magnitude" => ComparisonMethod::Abs,
1214                            other => {
1215                                return Err(sort_error(
1216                                    &SORT_ERROR_COMPARISON_METHOD_UNKNOWN,
1217                                    format!("sort: unsupported ComparisonMethod '{other}'"),
1218                                )
1219                                .into())
1220                            }
1221                        };
1222                        i += 1;
1223                        continue;
1224                    }
1225                    "missingplacement" => {
1226                        i += 1;
1227                        if i >= rest.len() {
1228                            return Err(sort_invalid_argument(
1229                                "sort: expected a value for 'MissingPlacement'",
1230                            ));
1231                        }
1232                        let value = tensor::value_to_string(&rest[i]).ok_or_else(|| {
1233                            sort_error(
1234                                &SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING,
1235                                SORT_ERROR_MISSINGPLACEMENT_REQUIRES_STRING.message,
1236                            )
1237                        })?;
1238                        args.missing = match value.trim().to_ascii_lowercase().as_str() {
1239                            "auto" => MissingPlacement::Auto,
1240                            "first" => MissingPlacement::First,
1241                            "last" => MissingPlacement::Last,
1242                            other => {
1243                                return Err(sort_error(
1244                                    &SORT_ERROR_MISSINGPLACEMENT_UNKNOWN,
1245                                    format!("sort: unsupported MissingPlacement '{other}'"),
1246                                ))
1247                            }
1248                        };
1249                        i += 1;
1250                        continue;
1251                    }
1252                    _ => {}
1253                }
1254            }
1255            return Err(sort_invalid_argument(format!(
1256                "sort: unrecognised argument {:?}",
1257                rest[i]
1258            )));
1259        }
1260        Ok(args)
1261    }
1262}
1263
1264fn is_dimension_placeholder(value: &Value) -> bool {
1265    match value {
1266        Value::Tensor(t) => tensor::tensor_element_len(t) == 0,
1267        Value::LogicalArray(logical) => logical.data.is_empty(),
1268        _ => false,
1269    }
1270}
1271
1272pub struct SortEvaluation {
1273    sorted: Value,
1274    indices: Value,
1275}
1276
1277impl SortEvaluation {
1278    pub fn into_sorted_value(self) -> Value {
1279        self.sorted
1280    }
1281
1282    pub fn into_values(self) -> (Value, Value) {
1283        (self.sorted, self.indices)
1284    }
1285
1286    pub fn indices_value(&self) -> Value {
1287        self.indices.clone()
1288    }
1289}
1290
1291#[cfg(test)]
1292pub(crate) mod tests {
1293    use super::*;
1294    use crate::builtins::common::test_support;
1295    use futures::executor::block_on;
1296    use runmat_builtins::{ResolveContext, Type};
1297    use runmat_value::{
1298        ComplexTensor, IntValue, IntegerComplexStorage, IntegerStorage, NumericStorage, Tensor,
1299        Value,
1300    };
1301
1302    fn sort_builtin(value: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
1303        block_on(super::sort_builtin(value, rest))
1304    }
1305
1306    fn evaluate(value: Value, rest: &[Value]) -> crate::BuiltinResult<SortEvaluation> {
1307        block_on(super::evaluate(value, rest))
1308    }
1309
1310    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1311    #[test]
1312    fn sort_vector_default() {
1313        let tensor = Tensor::new(vec![3.0, 1.0, 2.0], vec![3, 1]).unwrap();
1314        let result = sort_builtin(Value::Tensor(tensor), Vec::new()).expect("sort");
1315        match result {
1316            Value::Tensor(t) => {
1317                assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 3.0]);
1318                assert_eq!(t.shape, vec![3, 1]);
1319            }
1320            other => panic!("expected tensor result, got {other:?}"),
1321        }
1322    }
1323
1324    #[test]
1325    fn sort_preserves_native_single_storage_and_indices() {
1326        let tensor = Tensor::from_f32(vec![3.0, 1.0, 2.0], vec![3, 1]).expect("input");
1327        let (sorted, indices) = evaluate(Value::Tensor(tensor), &[])
1328            .expect("sort")
1329            .into_values();
1330        let Value::Tensor(sorted) = sorted else {
1331            panic!("expected sorted tensor");
1332        };
1333        assert_eq!(
1334            sorted.into_numeric_storage().expect("storage"),
1335            NumericStorage::F32(vec![1.0, 2.0, 3.0])
1336        );
1337        let Value::Tensor(indices) = indices else {
1338            panic!("expected index tensor");
1339        };
1340        assert_eq!(indices.materialize_f64(), vec![2.0, 3.0, 1.0]);
1341    }
1342
1343    #[test]
1344    fn sort_type_resolver_tensor() {
1345        assert_eq!(
1346            tensor_output_type(&[Type::tensor()], &ResolveContext::new(Vec::new())),
1347            Type::tensor()
1348        );
1349    }
1350
1351    #[test]
1352    fn sort_preserves_all_exact_real_integer_classes_and_indices() {
1353        let cases = [
1354            (
1355                IntegerStorage::I8(vec![i8::MAX, i8::MIN, 0, 7]),
1356                IntegerStorage::I8(vec![i8::MIN, 0, 7, i8::MAX]),
1357            ),
1358            (
1359                IntegerStorage::I16(vec![i16::MAX, i16::MIN, 0, 7]),
1360                IntegerStorage::I16(vec![i16::MIN, 0, 7, i16::MAX]),
1361            ),
1362            (
1363                IntegerStorage::I32(vec![i32::MAX, i32::MIN, 0, 7]),
1364                IntegerStorage::I32(vec![i32::MIN, 0, 7, i32::MAX]),
1365            ),
1366            (
1367                IntegerStorage::I64(vec![i64::MAX, i64::MIN, 0, 7]),
1368                IntegerStorage::I64(vec![i64::MIN, 0, 7, i64::MAX]),
1369            ),
1370            (
1371                IntegerStorage::U8(vec![u8::MAX, 0, 7, 9]),
1372                IntegerStorage::U8(vec![0, 7, 9, u8::MAX]),
1373            ),
1374            (
1375                IntegerStorage::U16(vec![u16::MAX, 0, 7, 700]),
1376                IntegerStorage::U16(vec![0, 7, 700, u16::MAX]),
1377            ),
1378            (
1379                IntegerStorage::U32(vec![u32::MAX, 0, 7, 9_007_199]),
1380                IntegerStorage::U32(vec![0, 7, 9_007_199, u32::MAX]),
1381            ),
1382            (
1383                IntegerStorage::U64(vec![u64::MAX, 0, 7, 9_007_199_254_740_993]),
1384                IntegerStorage::U64(vec![0, 7, 9_007_199_254_740_993, u64::MAX]),
1385            ),
1386        ];
1387
1388        for (input, expected) in cases {
1389            let tensor = Tensor::new_integer(input, vec![4, 1]).expect("input");
1390            let (sorted, indices) = evaluate(Value::Tensor(tensor), &[])
1391                .expect("sort")
1392                .into_values();
1393            let Value::Tensor(sorted) = sorted else {
1394                panic!("expected exact integer sorted values");
1395            };
1396            assert_eq!(sorted.integer_storage(), Some(&expected));
1397            let Value::Tensor(indices) = indices else {
1398                panic!("expected index tensor");
1399            };
1400            assert_eq!(indices.materialize_f64(), vec![2.0, 3.0, 4.0, 1.0]);
1401        }
1402    }
1403
1404    #[test]
1405    fn sort_reads_exact_integer_values_without_mirror() {
1406        let tensor = Tensor::new_integer(
1407            IntegerStorage::U64(vec![u64::MAX, 0, 9_007_199_254_740_993, 7]),
1408            vec![4, 1],
1409        )
1410        .expect("input");
1411
1412        let (sorted, indices) = evaluate(Value::Tensor(tensor), &[])
1413            .expect("sort")
1414            .into_values();
1415        let Value::Tensor(sorted) = sorted else {
1416            panic!("expected exact integer sorted values");
1417        };
1418        assert_eq!(
1419            sorted.integer_storage(),
1420            Some(&IntegerStorage::U64(vec![
1421                0,
1422                7,
1423                9_007_199_254_740_993,
1424                u64::MAX,
1425            ]))
1426        );
1427        let Value::Tensor(indices) = indices else {
1428            panic!("expected index tensor");
1429        };
1430        assert_eq!(indices.materialize_f64(), vec![2.0, 4.0, 3.0, 1.0]);
1431    }
1432
1433    #[test]
1434    fn sort_exact_integer_honors_descend_abs_and_dimension() {
1435        let input = Tensor::new_integer(
1436            IntegerStorage::U64(vec![u64::MAX, 0, 7, 9_007_199_254_740_993]),
1437            vec![4, 1],
1438        )
1439        .expect("input");
1440        let (sorted, indices) = evaluate(Value::Tensor(input), &[Value::from("descend")])
1441            .expect("descending sort")
1442            .into_values();
1443        let Value::Tensor(sorted) = sorted else {
1444            panic!("expected exact integer sorted values");
1445        };
1446        assert_eq!(
1447            sorted.integer_storage(),
1448            Some(&IntegerStorage::U64(vec![
1449                u64::MAX,
1450                9_007_199_254_740_993,
1451                7,
1452                0,
1453            ]))
1454        );
1455        let Value::Tensor(indices) = indices else {
1456            panic!("expected index tensor");
1457        };
1458        assert_eq!(indices.materialize_f64(), vec![1.0, 4.0, 3.0, 2.0]);
1459
1460        let input = Tensor::new_integer(
1461            IntegerStorage::I64(vec![i64::MIN, i64::MAX, 2, -1]),
1462            vec![4, 1],
1463        )
1464        .expect("input");
1465        let (sorted, indices) = evaluate(
1466            Value::Tensor(input),
1467            &[Value::from("ComparisonMethod"), Value::from("abs")],
1468        )
1469        .expect("absolute sort")
1470        .into_values();
1471        let Value::Tensor(sorted) = sorted else {
1472            panic!("expected exact integer sorted values");
1473        };
1474        assert_eq!(
1475            sorted.integer_storage(),
1476            Some(&IntegerStorage::I64(vec![-1, 2, i64::MAX, i64::MIN]))
1477        );
1478        let Value::Tensor(indices) = indices else {
1479            panic!("expected index tensor");
1480        };
1481        assert_eq!(indices.materialize_f64(), vec![4.0, 3.0, 2.0, 1.0]);
1482
1483        let input = Tensor::new_integer(
1484            IntegerStorage::U64(vec![u64::MAX, 0, 7, 9_007_199_254_740_993]),
1485            vec![2, 2],
1486        )
1487        .expect("matrix input");
1488        let (sorted, indices) = evaluate(Value::Tensor(input), &[Value::Int(IntValue::I32(2))])
1489            .expect("dimension sort")
1490            .into_values();
1491        let Value::Tensor(sorted) = sorted else {
1492            panic!("expected exact integer matrix values");
1493        };
1494        assert_eq!(
1495            sorted.integer_storage(),
1496            Some(&IntegerStorage::U64(vec![
1497                7,
1498                0,
1499                u64::MAX,
1500                9_007_199_254_740_993,
1501            ]))
1502        );
1503        let Value::Tensor(indices) = indices else {
1504            panic!("expected index tensor");
1505        };
1506        assert_eq!(indices.materialize_f64(), vec![2.0, 1.0, 1.0, 2.0]);
1507    }
1508
1509    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1510    #[test]
1511    fn sort_descend_direction() {
1512        let tensor = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![4, 1]).unwrap();
1513        let result =
1514            sort_builtin(Value::Tensor(tensor), vec![Value::from("descend")]).expect("sort");
1515        match result {
1516            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![4.0, 3.0, 2.0, 1.0]),
1517            other => panic!("expected tensor, got {other:?}"),
1518        }
1519    }
1520
1521    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1522    #[test]
1523    fn sort_matrix_default_dim1() {
1524        let tensor = Tensor::new(vec![4.0, 2.0, 1.0, 5.0, 6.0, 3.0], vec![2, 3]).unwrap();
1525        let eval = evaluate(Value::Tensor(tensor), &[]).expect("evaluate");
1526        let (sorted, indices) = eval.into_values();
1527        match sorted {
1528            Value::Tensor(t) => {
1529                assert_eq!(t.materialize_f64(), vec![2.0, 4.0, 1.0, 5.0, 3.0, 6.0]);
1530                assert_eq!(t.shape, vec![2, 3]);
1531            }
1532            other => panic!("expected tensor result, got {other:?}"),
1533        }
1534        match indices {
1535            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![2.0, 1.0, 1.0, 2.0, 2.0, 1.0]),
1536            other => panic!("expected tensor indices, got {other:?}"),
1537        }
1538    }
1539
1540    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1541    #[test]
1542    fn sort_matrix_along_dimension_two() {
1543        let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0, 2.0, 5.0], vec![2, 3]).unwrap();
1544        let eval =
1545            evaluate(Value::Tensor(tensor), &[Value::Int(IntValue::I32(2))]).expect("evaluate");
1546        let (sorted, indices) = eval.into_values();
1547        match sorted {
1548            Value::Tensor(t) => {
1549                assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 2.0, 3.0, 4.0, 5.0]);
1550                assert_eq!(t.shape, vec![2, 3]);
1551            }
1552            other => panic!("expected tensor result, got {other:?}"),
1553        }
1554        match indices {
1555            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0]),
1556            other => panic!("expected tensor indices, got {other:?}"),
1557        }
1558    }
1559
1560    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1561    #[test]
1562    fn sort_dimension_placeholder_then_dim() {
1563        let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0], vec![2, 2]).unwrap();
1564        let placeholder = Tensor::new(Vec::new(), vec![0, 0]).unwrap();
1565        let eval = evaluate(
1566            Value::Tensor(tensor),
1567            &[
1568                Value::Tensor(placeholder),
1569                Value::Int(IntValue::I32(2)),
1570                Value::from("descend"),
1571            ],
1572        )
1573        .expect("evaluate");
1574        let (sorted, _) = eval.into_values();
1575        match sorted {
1576            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![4.0, 3.0, 1.0, 2.0]),
1577            other => panic!("expected tensor result, got {other:?}"),
1578        }
1579    }
1580
1581    #[test]
1582    fn sort_typed_integer_dimension_argument_is_not_empty_placeholder_without_mirror() {
1583        let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0, 2.0, 5.0], vec![2, 3]).unwrap();
1584        let dim = Tensor::new_integer(IntegerStorage::U16(vec![2]), vec![1, 1]).expect("dimension");
1585
1586        let eval = evaluate(Value::Tensor(tensor), &[Value::Tensor(dim)]).expect("evaluate");
1587        let (sorted, indices) = eval.into_values();
1588        match sorted {
1589            Value::Tensor(t) => {
1590                assert_eq!(t.shape, vec![2, 3]);
1591                assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 2.0, 3.0, 4.0, 5.0]);
1592            }
1593            other => panic!("expected tensor result, got {other:?}"),
1594        }
1595        match indices {
1596            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0]),
1597            other => panic!("expected tensor indices, got {other:?}"),
1598        }
1599    }
1600
1601    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1602    #[test]
1603    fn sort_descend_then_dimension() {
1604        let tensor = Tensor::new(vec![1.0, 3.0, 4.0, 2.0, 2.0, 5.0], vec![2, 3]).unwrap();
1605        let eval = evaluate(
1606            Value::Tensor(tensor),
1607            &[Value::from("descend"), Value::Int(IntValue::I32(1))],
1608        )
1609        .expect("evaluate");
1610        let (sorted, _) = eval.into_values();
1611        match sorted {
1612            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![3.0, 1.0, 4.0, 2.0, 5.0, 2.0]),
1613            other => panic!("expected tensor result, got {other:?}"),
1614        }
1615    }
1616
1617    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1618    #[test]
1619    fn sort_returns_indices() {
1620        let tensor = Tensor::new(vec![4.0, 1.0, 9.0, 2.0], vec![4, 1]).unwrap();
1621        let eval = evaluate(Value::Tensor(tensor), &[]).expect("evaluate");
1622        let (sorted, indices) = eval.into_values();
1623        match sorted {
1624            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 4.0, 9.0]),
1625            other => panic!("expected tensor, got {other:?}"),
1626        }
1627        match indices {
1628            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![2.0, 4.0, 1.0, 3.0]),
1629            other => panic!("expected tensor, got {other:?}"),
1630        }
1631    }
1632
1633    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1634    #[test]
1635    fn sort_with_nan_handling() {
1636        let tensor = Tensor::new(vec![f64::NAN, 4.0, 1.0, 2.0], vec![4, 1]).unwrap();
1637        let eval = evaluate(Value::Tensor(tensor.clone()), &[]).expect("evaluate");
1638        let (sorted, _) = eval.into_values();
1639        match sorted {
1640            Value::Tensor(t) => {
1641                assert!(t.materialize_f64()[3].is_nan());
1642                assert_eq!(&t.materialize_f64()[0..3], &[1.0, 2.0, 4.0]);
1643            }
1644            other => panic!("expected tensor, got {other:?}"),
1645        }
1646
1647        let eval_desc =
1648            evaluate(Value::Tensor(tensor), &[Value::from("descend")]).expect("evaluate");
1649        let (sorted_desc, _) = eval_desc.into_values();
1650        match sorted_desc {
1651            Value::Tensor(t) => {
1652                assert!(t.materialize_f64()[0].is_nan());
1653                assert_eq!(&t.materialize_f64()[1..], &[4.0, 2.0, 1.0]);
1654            }
1655            other => panic!("expected tensor, got {other:?}"),
1656        }
1657    }
1658
1659    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1660    #[test]
1661    fn sort_by_absolute_value() {
1662        let tensor = Tensor::new(vec![-8.0, -1.0, 3.0, -2.0], vec![4, 1]).unwrap();
1663        let eval = evaluate(
1664            Value::Tensor(tensor),
1665            &[Value::from("ComparisonMethod"), Value::from("abs")],
1666        )
1667        .expect("evaluate");
1668        let (sorted, _) = eval.into_values();
1669        match sorted {
1670            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![-1.0, -2.0, 3.0, -8.0]),
1671            other => panic!("expected tensor, got {other:?}"),
1672        }
1673    }
1674
1675    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1676    #[test]
1677    fn sort_by_absolute_value_descend() {
1678        let tensor = Tensor::new(vec![-1.0, 2.0, -3.0, 4.0], vec![4, 1]).unwrap();
1679        let eval = evaluate(
1680            Value::Tensor(tensor),
1681            &[
1682                Value::from("descend"),
1683                Value::from("ComparisonMethod"),
1684                Value::from("abs"),
1685            ],
1686        )
1687        .expect("evaluate");
1688        let (sorted, _) = eval.into_values();
1689        match sorted {
1690            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![4.0, -3.0, 2.0, -1.0]),
1691            other => panic!("expected tensor, got {other:?}"),
1692        }
1693    }
1694
1695    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1696    #[test]
1697    fn sort_complex_auto_abs() {
1698        let tensor =
1699            ComplexTensor::new(vec![(1.0, 2.0), (-3.0, 0.5), (0.0, -1.0)], vec![3, 1]).unwrap();
1700        let eval = evaluate(Value::ComplexTensor(tensor), &[]).expect("evaluate");
1701        let (sorted, indices) = eval.into_values();
1702        match sorted {
1703            Value::ComplexTensor(t) => {
1704                assert_eq!(
1705                    t.materialize_f64(),
1706                    vec![(0.0, -1.0), (1.0, 2.0), (-3.0, 0.5)]
1707                )
1708            }
1709            other => panic!("expected complex tensor, got {other:?}"),
1710        }
1711        match indices {
1712            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![3.0, 1.0, 2.0]),
1713            other => panic!("expected tensor indices, got {other:?}"),
1714        }
1715    }
1716
1717    #[test]
1718    fn sort_preserves_native_complex_single_storage_and_indices() {
1719        let tensor =
1720            ComplexTensor::from_f32(vec![(3.0, 4.0), (1.0, 0.0), (0.0, 2.0)], vec![3, 1]).unwrap();
1721        let (sorted, indices) = evaluate(Value::ComplexTensor(tensor), &[])
1722            .expect("sort")
1723            .into_values();
1724        let Value::ComplexTensor(sorted) = sorted else {
1725            panic!("expected complex single tensor");
1726        };
1727        assert_eq!(
1728            sorted.as_f32_slice(),
1729            Some(&[(1.0, 0.0), (0.0, 2.0), (3.0, 4.0)][..])
1730        );
1731        let Value::Tensor(indices) = indices else {
1732            panic!("expected index tensor");
1733        };
1734        assert_eq!(indices.as_f64_slice(), Some(&[2.0, 3.0, 1.0][..]));
1735    }
1736
1737    #[test]
1738    fn sort_preserves_one_element_complex_single_tensor() {
1739        let tensor = ComplexTensor::from_f32(vec![(1.25, -2.5)], vec![1, 1]).unwrap();
1740        let sorted = evaluate(Value::ComplexTensor(tensor), &[])
1741            .expect("sort")
1742            .into_sorted_value();
1743        let Value::ComplexTensor(sorted) = sorted else {
1744            panic!("expected complex single tensor");
1745        };
1746        assert_eq!(sorted.as_f32_slice(), Some(&[(1.25, -2.5)][..]));
1747    }
1748
1749    #[test]
1750    fn sort_orders_typed_complex_uint64_exactly_and_preserves_storage() {
1751        for input in [
1752            (
1753                IntegerStorage::I8(vec![3, 1]),
1754                IntegerStorage::I8(vec![4, 0]),
1755            ),
1756            (
1757                IntegerStorage::I16(vec![3, 1]),
1758                IntegerStorage::I16(vec![4, 0]),
1759            ),
1760            (
1761                IntegerStorage::I32(vec![3, 1]),
1762                IntegerStorage::I32(vec![4, 0]),
1763            ),
1764            (
1765                IntegerStorage::I64(vec![3, 1]),
1766                IntegerStorage::I64(vec![4, 0]),
1767            ),
1768            (
1769                IntegerStorage::U8(vec![3, 1]),
1770                IntegerStorage::U8(vec![4, 0]),
1771            ),
1772            (
1773                IntegerStorage::U16(vec![3, 1]),
1774                IntegerStorage::U16(vec![4, 0]),
1775            ),
1776            (
1777                IntegerStorage::U32(vec![3, 1]),
1778                IntegerStorage::U32(vec![4, 0]),
1779            ),
1780            (
1781                IntegerStorage::U64(vec![3, 1]),
1782                IntegerStorage::U64(vec![4, 0]),
1783            ),
1784        ] {
1785            let expected_real = input
1786                .0
1787                .from_exact_values_like(vec![
1788                    input.0.value_at(1).unwrap(),
1789                    input.0.value_at(0).unwrap(),
1790                ])
1791                .unwrap();
1792            let expected_imag = input
1793                .1
1794                .from_exact_values_like(vec![
1795                    input.1.value_at(1).unwrap(),
1796                    input.1.value_at(0).unwrap(),
1797                ])
1798                .unwrap();
1799            let tensor = ComplexTensor::new_integer(
1800                IntegerComplexStorage::new(input.0, input.1).unwrap(),
1801                vec![2, 1],
1802            )
1803            .unwrap();
1804            let sorted = evaluate(Value::ComplexTensor(tensor), &[])
1805                .expect("typed complex integer sort")
1806                .into_sorted_value();
1807            let Value::ComplexTensor(sorted) = sorted else {
1808                panic!("expected typed complex integer tensor");
1809            };
1810            assert_eq!(
1811                sorted.integer_storage(),
1812                Some(&IntegerComplexStorage::new(expected_real, expected_imag).unwrap())
1813            );
1814        }
1815
1816        let input = IntegerComplexStorage::new(
1817            IntegerStorage::U64(vec![u64::MAX, u64::MAX - 1, 0]),
1818            IntegerStorage::U64(vec![0, 1, u64::MAX]),
1819        )
1820        .unwrap();
1821        let tensor = ComplexTensor::new_integer(input, vec![3, 1]).unwrap();
1822        let (sorted, indices) = evaluate(Value::ComplexTensor(tensor), &[])
1823            .expect("typed complex integer sort")
1824            .into_values();
1825        let Value::ComplexTensor(sorted) = sorted else {
1826            panic!("expected typed complex integer tensor");
1827        };
1828        assert_eq!(
1829            sorted.integer_storage(),
1830            Some(
1831                &IntegerComplexStorage::new(
1832                    IntegerStorage::U64(vec![u64::MAX - 1, u64::MAX, 0]),
1833                    IntegerStorage::U64(vec![1, 0, u64::MAX]),
1834                )
1835                .unwrap()
1836            )
1837        );
1838        let Value::Tensor(indices) = indices else {
1839            panic!("expected double indices");
1840        };
1841        assert_eq!(indices.as_f64_slice(), Some(&[2.0, 1.0, 3.0][..]));
1842    }
1843
1844    #[test]
1845    fn sort_restores_exact_typed_complex_integer_results_to_the_input_provider() {
1846        test_support::with_test_provider(|provider| {
1847            let input = ComplexTensor::new_integer(
1848                IntegerComplexStorage::new(
1849                    IntegerStorage::U64(vec![u64::MAX, u64::MAX - 1]),
1850                    IntegerStorage::U64(vec![0, 1]),
1851                )
1852                .unwrap(),
1853                vec![2, 1],
1854            )
1855            .unwrap();
1856            let handle =
1857                crate::builtins::common::gpu_helpers::upload_complex_tensor(provider, &input)
1858                    .expect("typed complex upload");
1859            let (sorted, indices) = evaluate(Value::GpuTensor(handle), &[])
1860                .expect("resident typed complex sort")
1861                .into_values();
1862            let sorted = block_on(crate::builtins::common::gpu_helpers::gather_value_async(
1863                &sorted,
1864            ))
1865            .expect("gather sorted values");
1866            let indices = test_support::gather(indices).expect("gather sorted indices");
1867            let Value::ComplexTensor(sorted) = sorted else {
1868                panic!("expected typed complex integer tensor");
1869            };
1870            assert_eq!(
1871                sorted.integer_storage(),
1872                Some(
1873                    &IntegerComplexStorage::new(
1874                        IntegerStorage::U64(vec![u64::MAX - 1, u64::MAX]),
1875                        IntegerStorage::U64(vec![1, 0]),
1876                    )
1877                    .unwrap()
1878                )
1879            );
1880            assert_eq!(indices.as_f64_slice(), Some(&[2.0, 1.0][..]));
1881        });
1882    }
1883
1884    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1885    #[test]
1886    fn sort_complex_real_descend() {
1887        let tensor =
1888            ComplexTensor::new(vec![(1.0, 2.0), (-3.0, 0.0), (1.0, -1.0)], vec![3, 1]).unwrap();
1889        let eval = evaluate(
1890            Value::ComplexTensor(tensor),
1891            &[
1892                Value::from("descend"),
1893                Value::from("ComparisonMethod"),
1894                Value::from("real"),
1895            ],
1896        )
1897        .expect("evaluate");
1898        let (sorted, _) = eval.into_values();
1899        match sorted {
1900            Value::ComplexTensor(t) => {
1901                assert_eq!(
1902                    t.materialize_f64(),
1903                    vec![(1.0, 2.0), (1.0, -1.0), (-3.0, 0.0)]
1904                );
1905            }
1906            other => panic!("expected complex tensor, got {other:?}"),
1907        }
1908    }
1909
1910    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1911    #[test]
1912    fn sort_stable_with_duplicates() {
1913        let tensor = Tensor::new(vec![2.0, 2.0, 1.0, 2.0], vec![4, 1]).unwrap();
1914        let eval = evaluate(Value::Tensor(tensor), &[]).expect("evaluate");
1915        let (sorted, indices) = eval.into_values();
1916        match sorted {
1917            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![1.0, 2.0, 2.0, 2.0]),
1918            other => panic!("expected tensor, got {other:?}"),
1919        }
1920        match indices {
1921            Value::Tensor(t) => assert_eq!(t.materialize_f64(), vec![3.0, 1.0, 2.0, 4.0]),
1922            other => panic!("expected tensor indices, got {other:?}"),
1923        }
1924    }
1925
1926    #[test]
1927    fn sort_abs_ties_follow_phase_for_real_integer_and_complex_values() {
1928        let args = [Value::from("ComparisonMethod"), Value::from("abs")];
1929        let (real, real_indices) = evaluate(
1930            Value::Tensor(Tensor::new(vec![-3.0, 3.0], vec![2, 1]).expect("real")),
1931            &args,
1932        )
1933        .expect("real sort")
1934        .into_values();
1935        let Value::Tensor(real) = real else {
1936            panic!("expected real tensor");
1937        };
1938        assert_eq!(real.materialize_f64(), vec![3.0, -3.0]);
1939        let Value::Tensor(real_indices) = real_indices else {
1940            panic!("expected real indices");
1941        };
1942        assert_eq!(real_indices.materialize_f64(), vec![2.0, 1.0]);
1943
1944        let (integer, _) = evaluate(
1945            Value::Tensor(
1946                Tensor::new_integer(IntegerStorage::I64(vec![-3, 3]), vec![2, 1]).expect("integer"),
1947            ),
1948            &args,
1949        )
1950        .expect("integer sort")
1951        .into_values();
1952        let Value::Tensor(integer) = integer else {
1953            panic!("expected integer tensor");
1954        };
1955        assert_eq!(
1956            integer.integer_storage(),
1957            Some(&IntegerStorage::I64(vec![3, -3]))
1958        );
1959
1960        let (complex, complex_indices) = evaluate(
1961            Value::ComplexTensor(
1962                ComplexTensor::new(vec![(-1.0, 0.0), (0.0, 1.0), (0.0, -1.0)], vec![3, 1])
1963                    .expect("complex"),
1964            ),
1965            &[],
1966        )
1967        .expect("complex sort")
1968        .into_values();
1969        let Value::ComplexTensor(complex) = complex else {
1970            panic!("expected complex tensor");
1971        };
1972        assert_eq!(
1973            complex.materialize_f64(),
1974            vec![(0.0, -1.0), (0.0, 1.0), (-1.0, 0.0)]
1975        );
1976        let Value::Tensor(complex_indices) = complex_indices else {
1977            panic!("expected complex indices");
1978        };
1979        assert_eq!(complex_indices.materialize_f64(), vec![3.0, 2.0, 1.0]);
1980
1981        let (phase_boundary, phase_boundary_indices) = evaluate(
1982            Value::ComplexTensor(
1983                ComplexTensor::new(vec![(-1.0, 0.0), (-1.0, -0.0)], vec![2, 1])
1984                    .expect("phase boundary"),
1985            ),
1986            &[],
1987        )
1988        .expect("phase-boundary sort")
1989        .into_values();
1990        let Value::ComplexTensor(phase_boundary) = phase_boundary else {
1991            panic!("expected complex phase-boundary tensor");
1992        };
1993        assert_eq!(
1994            phase_boundary.materialize_f64(),
1995            vec![(-1.0, 0.0), (-1.0, -0.0)]
1996        );
1997        let Value::Tensor(phase_boundary_indices) = phase_boundary_indices else {
1998            panic!("expected complex phase-boundary indices");
1999        };
2000        assert_eq!(phase_boundary_indices.materialize_f64(), vec![1.0, 2.0]);
2001    }
2002
2003    #[test]
2004    fn sort_preserves_logical_class_and_accepts_every_integer_dimension_class() {
2005        let (logical, logical_indices) = evaluate(
2006            Value::LogicalArray(LogicalArray::new(vec![1, 0, 1], vec![3, 1]).expect("logical")),
2007            &[],
2008        )
2009        .expect("logical sort")
2010        .into_values();
2011        let Value::LogicalArray(logical) = logical else {
2012            panic!("expected logical array");
2013        };
2014        assert_eq!(logical.data, vec![0, 1, 1]);
2015        let Value::Tensor(logical_indices) = logical_indices else {
2016            panic!("expected logical indices");
2017        };
2018        assert_eq!(logical_indices.materialize_f64(), vec![2.0, 1.0, 3.0]);
2019
2020        for dimension in [
2021            IntegerStorage::I8(vec![2]),
2022            IntegerStorage::I16(vec![2]),
2023            IntegerStorage::I32(vec![2]),
2024            IntegerStorage::I64(vec![2]),
2025            IntegerStorage::U8(vec![2]),
2026            IntegerStorage::U16(vec![2]),
2027            IntegerStorage::U32(vec![2]),
2028            IntegerStorage::U64(vec![2]),
2029        ] {
2030            let input = Tensor::new(vec![4.0, 3.0, 2.0, 1.0], vec![2, 2]).expect("input");
2031            let dim = Value::Tensor(Tensor::new_integer(dimension, vec![1, 1]).expect("dimension"));
2032            let sorted = evaluate(Value::Tensor(input), &[dim])
2033                .expect("typed dimension")
2034                .into_sorted_value();
2035            let Value::Tensor(sorted) = sorted else {
2036                panic!("expected tensor");
2037            };
2038            assert_eq!(sorted.materialize_f64(), vec![2.0, 1.0, 4.0, 3.0]);
2039        }
2040    }
2041
2042    #[test]
2043    fn sort_rejects_more_than_two_outputs() {
2044        let _outputs = crate::output_count::push_output_count(Some(3));
2045        let error = sort_builtin(
2046            Value::Tensor(Tensor::new(vec![2.0, 1.0], vec![2, 1]).expect("input")),
2047            Vec::new(),
2048        )
2049        .expect_err("too many outputs");
2050        assert!(error.message().contains("maximum is 2"));
2051    }
2052
2053    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2054    #[test]
2055    fn sort_empty_tensor() {
2056        let tensor = Tensor::new(Vec::new(), vec![0, 3]).unwrap();
2057        let eval = evaluate(Value::Tensor(tensor.clone()), &[]).expect("evaluate");
2058        let (sorted, indices) = eval.into_values();
2059        match sorted {
2060            Value::Tensor(t) => {
2061                assert!(t.materialize_f64().is_empty());
2062                assert_eq!(t.shape, tensor.shape);
2063            }
2064            other => panic!("expected tensor, got {other:?}"),
2065        }
2066        match indices {
2067            Value::Tensor(t) => assert!(t.materialize_f64().is_empty()),
2068            other => panic!("expected tensor, got {other:?}"),
2069        }
2070    }
2071
2072    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2073    #[test]
2074    fn sort_dim_greater_than_ndims() {
2075        let tensor = Tensor::new(vec![4.0, 2.0, 3.0, 1.0], vec![2, 2]).unwrap();
2076        let eval = evaluate(
2077            Value::Tensor(tensor.clone()),
2078            &[Value::Int(IntValue::I32(3))],
2079        )
2080        .expect("evaluate");
2081        let (sorted, indices) = eval.into_values();
2082        match sorted {
2083            Value::Tensor(t) => assert_eq!(t.materialize_f64(), tensor.materialize_f64()),
2084            other => panic!("expected tensor, got {other:?}"),
2085        }
2086        match indices {
2087            Value::Tensor(t) => assert!(t
2088                .materialize_f64()
2089                .iter()
2090                .all(|v| (*v - 1.0).abs() < f64::EPSILON)),
2091            other => panic!("expected tensor, got {other:?}"),
2092        }
2093    }
2094
2095    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2096    #[test]
2097    fn sort_supports_missing_placement_and_rejects_unknown_values() {
2098        let input =
2099            Value::Tensor(Tensor::new(vec![2.0, f64::NAN, 1.0], vec![3, 1]).expect("input"));
2100        let first = sort_builtin(
2101            input.clone(),
2102            vec![Value::from("MissingPlacement"), Value::from("first")],
2103        )
2104        .expect("missing first");
2105        let Value::Tensor(first) = first else {
2106            panic!("expected tensor");
2107        };
2108        assert!(first.materialize_f64()[0].is_nan());
2109        assert_eq!(&first.materialize_f64()[1..], &[1.0, 2.0]);
2110
2111        let last_descending = sort_builtin(
2112            input,
2113            vec![
2114                Value::from("descend"),
2115                Value::from("MissingPlacement"),
2116                Value::from("last"),
2117            ],
2118        )
2119        .expect("missing last");
2120        let Value::Tensor(last_descending) = last_descending else {
2121            panic!("expected tensor");
2122        };
2123        assert_eq!(&last_descending.materialize_f64()[..2], &[2.0, 1.0]);
2124        assert!(last_descending.materialize_f64()[2].is_nan());
2125
2126        let err = sort_builtin(
2127            Value::Tensor(Tensor::new(vec![1.0], vec![1, 1]).unwrap()),
2128            vec![Value::from("MissingPlacement"), Value::from("middle")],
2129        )
2130        .unwrap_err();
2131        assert_eq!(
2132            err.identifier(),
2133            SORT_ERROR_MISSINGPLACEMENT_UNKNOWN.identifier
2134        );
2135    }
2136
2137    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2138    #[test]
2139    fn sort_invalid_comparison_method_errors() {
2140        let err = sort_builtin(
2141            Value::Tensor(Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap()),
2142            vec![Value::from("ComparisonMethod"), Value::from("unknown")],
2143        )
2144        .unwrap_err();
2145        assert_eq!(
2146            err.identifier(),
2147            SORT_ERROR_COMPARISON_METHOD_UNKNOWN.identifier
2148        );
2149    }
2150
2151    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2152    #[test]
2153    fn sort_invalid_comparison_method_value_errors() {
2154        let err = sort_builtin(
2155            Value::Tensor(Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap()),
2156            vec![
2157                Value::from("ComparisonMethod"),
2158                Value::Int(IntValue::I32(1)),
2159            ],
2160        )
2161        .unwrap_err();
2162        assert_eq!(
2163            err.identifier(),
2164            SORT_ERROR_COMPARISON_METHOD_REQUIRES_STRING.identifier
2165        );
2166    }
2167
2168    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2169    #[test]
2170    fn sort_dimension_zero_errors() {
2171        let err = sort_builtin(
2172            Value::Tensor(Tensor::new(vec![1.0], vec![1, 1]).unwrap()),
2173            vec![Value::Num(0.0)],
2174        )
2175        .unwrap_err();
2176        assert_eq!(err.identifier(), SORT_ERROR_INVALID_DIMENSION.identifier);
2177    }
2178
2179    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2180    #[test]
2181    fn sort_gpu_round_trip() {
2182        test_support::with_test_provider(|provider| {
2183            let tensor = Tensor::new(vec![3.0, 1.0, 2.0], vec![3, 1]).unwrap();
2184            let view = runmat_accelerate_api::HostTensorView {
2185                data: &tensor.materialize_f64(),
2186                shape: &tensor.shape,
2187            };
2188            let handle = provider.upload(&view).expect("upload");
2189            let eval = evaluate(Value::GpuTensor(handle), &[]).expect("evaluate");
2190            let (sorted, indices) = eval.into_values();
2191            assert!(matches!(sorted, Value::GpuTensor(_)));
2192            assert!(matches!(indices, Value::GpuTensor(_)));
2193            let sorted = test_support::gather(sorted).expect("gather sorted");
2194            assert_eq!(sorted.materialize_f64(), vec![1.0, 2.0, 3.0]);
2195            let indices = test_support::gather(indices).expect("gather indices");
2196            assert_eq!(indices.materialize_f64(), vec![2.0, 3.0, 1.0]);
2197        });
2198    }
2199
2200    #[test]
2201    fn sort_gpu_preserves_exact_wide_integer_and_logical_residency() {
2202        test_support::with_test_provider(|provider| {
2203            let integer = Tensor::new_integer(
2204                IntegerStorage::U64(vec![u64::MAX, 0, 9_007_199_254_740_993]),
2205                vec![3, 1],
2206            )
2207            .expect("integer");
2208            let handle = gpu_helpers::upload_tensor(provider, &integer).expect("upload integer");
2209            let (sorted, indices) = evaluate(Value::GpuTensor(handle), &[])
2210                .expect("integer GPU sort")
2211                .into_values();
2212            assert!(matches!(sorted, Value::GpuTensor(_)));
2213            assert!(matches!(indices, Value::GpuTensor(_)));
2214            let sorted = test_support::gather(sorted).expect("gather integer");
2215            assert_eq!(
2216                sorted.into_numeric_storage().expect("integer storage"),
2217                NumericStorage::U64(vec![0, 9_007_199_254_740_993, u64::MAX])
2218            );
2219
2220            let logical = provider
2221                .upload(&runmat_accelerate_api::HostTensorView {
2222                    data: &[1.0, 0.0, 1.0],
2223                    shape: &[3, 1],
2224                })
2225                .expect("upload logical");
2226            runmat_accelerate_api::set_handle_logical(&logical, true);
2227            let (sorted, _) = evaluate(Value::GpuTensor(logical), &[])
2228                .expect("logical GPU sort")
2229                .into_values();
2230            let Value::GpuTensor(sorted_handle) = &sorted else {
2231                panic!("expected resident logical output");
2232            };
2233            assert!(runmat_accelerate_api::handle_is_logical(sorted_handle));
2234            let sorted = test_support::gather(sorted).expect("gather logical");
2235            assert_eq!(sorted.materialize_f64(), vec![0.0, 1.0, 1.0]);
2236        });
2237    }
2238
2239    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2240    #[test]
2241    #[cfg(feature = "wgpu")]
2242    fn sort_wgpu_matches_cpu() {
2243        let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2244            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2245        );
2246        let tensor = Tensor::new(vec![4.0, 1.0, 3.0, 2.0], vec![4, 1]).unwrap();
2247        let cpu_eval = evaluate(Value::Tensor(tensor.clone()), &[]).expect("cpu sort");
2248        let (cpu_sorted, cpu_indices) = cpu_eval.into_values();
2249
2250        let gpu_view = runmat_accelerate_api::HostTensorView {
2251            data: &tensor.materialize_f64(),
2252            shape: &tensor.shape,
2253        };
2254        let provider = runmat_accelerate_api::provider().expect("wgpu provider");
2255        let handle = provider.upload(&gpu_view).expect("upload");
2256        let gpu_eval = evaluate(Value::GpuTensor(handle), &[]).expect("gpu sort");
2257        let (gpu_sorted, gpu_indices) = gpu_eval.into_values();
2258
2259        let cpu_sorted_tensor = match cpu_sorted {
2260            Value::Tensor(t) => t,
2261            Value::Num(n) => Tensor::new(vec![n], vec![1, 1]).unwrap(),
2262            other => panic!("unexpected CPU sorted value {other:?}"),
2263        };
2264        let cpu_indices_tensor = match cpu_indices {
2265            Value::Tensor(t) => t,
2266            Value::Num(n) => Tensor::new(vec![n], vec![1, 1]).unwrap(),
2267            other => panic!("unexpected CPU indices value {other:?}"),
2268        };
2269        assert!(matches!(gpu_sorted, Value::GpuTensor(_)));
2270        assert!(matches!(gpu_indices, Value::GpuTensor(_)));
2271        let gpu_sorted_tensor = test_support::gather(gpu_sorted).expect("gather GPU sorted");
2272        let gpu_indices_tensor = test_support::gather(gpu_indices).expect("gather GPU indices");
2273
2274        assert_eq!(
2275            gpu_sorted_tensor.materialize_f64(),
2276            cpu_sorted_tensor.materialize_f64()
2277        );
2278        assert_eq!(
2279            gpu_indices_tensor.materialize_f64(),
2280            cpu_indices_tensor.materialize_f64()
2281        );
2282    }
2283
2284    #[test]
2285    #[cfg(feature = "wgpu")]
2286    fn sort_wgpu_abs_ties_follow_phase() {
2287        let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2288            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2289        );
2290        let provider = runmat_accelerate_api::provider().expect("wgpu provider");
2291        let input = provider
2292            .upload(&runmat_accelerate_api::HostTensorView {
2293                data: &[-3.0, 3.0],
2294                shape: &[2, 1],
2295            })
2296            .expect("upload");
2297        let (sorted, indices) = evaluate(
2298            Value::GpuTensor(input),
2299            &[Value::from("ComparisonMethod"), Value::from("abs")],
2300        )
2301        .expect("wgpu sort")
2302        .into_values();
2303        let sorted = test_support::gather(sorted).expect("gather sorted");
2304        assert_eq!(sorted.materialize_f64(), vec![3.0, -3.0]);
2305        let indices = test_support::gather(indices).expect("gather indices");
2306        assert_eq!(indices.materialize_f64(), vec![2.0, 1.0]);
2307    }
2308}