Skip to main content

runmat_runtime/builtins/array/sorting_sets/
sortrows.rs

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