Skip to main content

runmat_runtime/builtins/array/sorting_sets/
intersect.rs

1//! MATLAB-compatible `intersect` builtin with GPU-aware semantics for RunMat.
2//!
3//! Supports element-wise and row-wise intersections with optional stable ordering,
4//! and index outputs that mirror MathWorks MATLAB semantics. GPU tensors use
5//! typed host fallback and their public outputs are restored to the owning
6//! provider.
7
8use std::cmp::Ordering;
9use std::collections::{HashMap, HashSet};
10
11use runmat_accelerate_api::GpuTensorHandle;
12use runmat_builtins::{
13    BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinIntegerBackendRule,
14    BuiltinIntegerCapabilityDescriptor, BuiltinIntegerComputationDomain,
15    BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule, BuiltinIntegerOverloadKind,
16    BuiltinOutputMode, BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType,
17    BuiltinSignatureDescriptor,
18};
19use runmat_macros::runtime_builtin;
20use runmat_value::{
21    CharArray, ComplexStorage, ComplexTensor, IntValue, IntegerStorage, NumericDType,
22    NumericStorage, StringArray, Tensor, Value,
23};
24
25use super::{float_order::SetFloat, integer_order, type_resolvers::set_values_output_type};
26use crate::build_runtime_error;
27use crate::builtins::common::arg_tokens::tokens_from_values;
28use crate::builtins::common::gpu_helpers;
29use crate::builtins::common::random_args::complex_tensor_into_value;
30use crate::builtins::common::spec::{
31    BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
32    ReductionNaN, ResidencyPolicy, ScalarType, ShapeRequirements,
33};
34use crate::builtins::common::tensor;
35use crate::builtins::math::elementwise::integer_cast::IntegerTarget;
36
37#[runmat_macros::register_gpu_spec(
38    builtin_path = "crate::builtins::array::sorting_sets::intersect"
39)]
40pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
41    name: "intersect",
42    op_kind: GpuOpKind::Custom("intersect"),
43    supported_precisions: &[ScalarType::F32, ScalarType::F64],
44    broadcast: BroadcastSemantics::None,
45    provider_hooks: &[],
46    constant_strategy: ConstantStrategy::InlineLiteral,
47    residency: ResidencyPolicy::NewHandle,
48    nan_mode: ReductionNaN::Include,
49    two_pass_threshold: None,
50    workgroup_size: None,
51    accepts_nan_mode: true,
52    notes: "Exact typed fallback gathers when needed and restores intersection values plus double indices to the input owner.",
53};
54
55#[runmat_macros::register_fusion_spec(
56    builtin_path = "crate::builtins::array::sorting_sets::intersect"
57)]
58pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
59    name: "intersect",
60    shape: ShapeRequirements::Any,
61    constant_strategy: ConstantStrategy::InlineLiteral,
62    elementwise: None,
63    reduction: None,
64    emits_nan: true,
65    notes: "`intersect` materialises its inputs and terminates fusion chains; upstream GPU tensors are gathered when necessary.",
66};
67
68const BUILTIN_NAME: &str = "intersect";
69
70const INTERSECT_OUTPUT_C: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
71    name: "C",
72    ty: BuiltinParamType::Any,
73    arity: BuiltinParamArity::Required,
74    default: None,
75    description: "Intersection values or rows.",
76}];
77
78const INTERSECT_OUTPUT_C_IA: [BuiltinParamDescriptor; 2] = [
79    BuiltinParamDescriptor {
80        name: "C",
81        ty: BuiltinParamType::Any,
82        arity: BuiltinParamArity::Required,
83        default: None,
84        description: "Intersection values or rows.",
85    },
86    BuiltinParamDescriptor {
87        name: "ia",
88        ty: BuiltinParamType::NumericArray,
89        arity: BuiltinParamArity::Required,
90        default: None,
91        description: "Indices selecting matching elements/rows in A.",
92    },
93];
94
95const INTERSECT_OUTPUT_C_IA_IB: [BuiltinParamDescriptor; 3] = [
96    BuiltinParamDescriptor {
97        name: "C",
98        ty: BuiltinParamType::Any,
99        arity: BuiltinParamArity::Required,
100        default: None,
101        description: "Intersection values or rows.",
102    },
103    BuiltinParamDescriptor {
104        name: "ia",
105        ty: BuiltinParamType::NumericArray,
106        arity: BuiltinParamArity::Required,
107        default: None,
108        description: "Indices selecting matching elements/rows in A.",
109    },
110    BuiltinParamDescriptor {
111        name: "ib",
112        ty: BuiltinParamType::NumericArray,
113        arity: BuiltinParamArity::Required,
114        default: None,
115        description: "Indices selecting matching elements/rows in B.",
116    },
117];
118
119const INTERSECT_INPUTS_A_B: [BuiltinParamDescriptor; 2] = [
120    BuiltinParamDescriptor {
121        name: "A",
122        ty: BuiltinParamType::Any,
123        arity: BuiltinParamArity::Required,
124        default: None,
125        description: "First input array.",
126    },
127    BuiltinParamDescriptor {
128        name: "B",
129        ty: BuiltinParamType::Any,
130        arity: BuiltinParamArity::Required,
131        default: None,
132        description: "Second input array.",
133    },
134];
135
136const INTERSECT_INPUTS_A_B_OPTIONS: [BuiltinParamDescriptor; 3] = [
137    BuiltinParamDescriptor {
138        name: "A",
139        ty: BuiltinParamType::Any,
140        arity: BuiltinParamArity::Required,
141        default: None,
142        description: "First input array.",
143    },
144    BuiltinParamDescriptor {
145        name: "B",
146        ty: BuiltinParamType::Any,
147        arity: BuiltinParamArity::Required,
148        default: None,
149        description: "Second input array.",
150    },
151    BuiltinParamDescriptor {
152        name: "option",
153        ty: BuiltinParamType::StringScalar,
154        arity: BuiltinParamArity::Variadic,
155        default: None,
156        description: "Option tokens: 'rows'|'sorted'|'stable'.",
157    },
158];
159
160const INTERSECT_SIGNATURES: [BuiltinSignatureDescriptor; 6] = [
161    BuiltinSignatureDescriptor {
162        label: "C = intersect(A, B)",
163        inputs: &INTERSECT_INPUTS_A_B,
164        outputs: &INTERSECT_OUTPUT_C,
165    },
166    BuiltinSignatureDescriptor {
167        label: "C = intersect(A, B, option...)",
168        inputs: &INTERSECT_INPUTS_A_B_OPTIONS,
169        outputs: &INTERSECT_OUTPUT_C,
170    },
171    BuiltinSignatureDescriptor {
172        label: "[C, ia] = intersect(A, B)",
173        inputs: &INTERSECT_INPUTS_A_B,
174        outputs: &INTERSECT_OUTPUT_C_IA,
175    },
176    BuiltinSignatureDescriptor {
177        label: "[C, ia] = intersect(A, B, option...)",
178        inputs: &INTERSECT_INPUTS_A_B_OPTIONS,
179        outputs: &INTERSECT_OUTPUT_C_IA,
180    },
181    BuiltinSignatureDescriptor {
182        label: "[C, ia, ib] = intersect(A, B)",
183        inputs: &INTERSECT_INPUTS_A_B,
184        outputs: &INTERSECT_OUTPUT_C_IA_IB,
185    },
186    BuiltinSignatureDescriptor {
187        label: "[C, ia, ib] = intersect(A, B, option...)",
188        inputs: &INTERSECT_INPUTS_A_B_OPTIONS,
189        outputs: &INTERSECT_OUTPUT_C_IA_IB,
190    },
191];
192
193const INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
194    code: "RM.INTERSECT.LEGACY_OPTION_UNSUPPORTED",
195    identifier: Some("RunMat:intersect:LegacyOptionUnsupported"),
196    when: "Legacy compatibility options are requested.",
197    message: "intersect: the 'legacy' behaviour is not supported",
198};
199
200const INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
201    code: "RM.INTERSECT.CONFLICTING_ORDER_OPTIONS",
202    identifier: Some("RunMat:intersect:ConflictingOrderOptions"),
203    when: "Both 'sorted' and 'stable' options are provided.",
204    message: "intersect: cannot combine 'sorted' with 'stable'",
205};
206
207const INTERSECT_ERROR_UNKNOWN_OPTION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
208    code: "RM.INTERSECT.UNKNOWN_OPTION",
209    identifier: Some("RunMat:intersect:UnknownOption"),
210    when: "An unsupported option token is provided.",
211    message: "intersect: unrecognised option",
212};
213
214const INTERSECT_ERROR_ROWS_COLUMN_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
215    code: "RM.INTERSECT.ROWS_COLUMN_MISMATCH",
216    identifier: Some("RunMat:intersect:RowsColumnMismatch"),
217    when: "'rows' mode is used and column counts differ.",
218    message: "intersect: inputs must have the same number of columns when using 'rows'",
219};
220
221const INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
222    code: "RM.INTERSECT.UNSUPPORTED_INPUT_TYPE",
223    identifier: Some("RunMat:intersect:UnsupportedInputType"),
224    when: "Input values cannot be converted into supported intersect domains.",
225    message: "intersect: unsupported input type",
226};
227
228const INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
229    code: "RM.INTERSECT.NUMERIC_CLASS_MISMATCH",
230    identifier: Some("RunMat:intersect:NumericClassMismatch"),
231    when: "Numeric inputs have incompatible nondouble classes.",
232    message: "intersect: numeric inputs must have the same class, except double may be combined with one nondouble class",
233};
234
235const INTERSECT_ERROR_INVALID_ARGUMENT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
236    code: "RM.INTERSECT.INVALID_ARGUMENT",
237    identifier: Some("RunMat:intersect:InvalidArgument"),
238    when: "Option arguments are not string-like where required.",
239    message: "intersect: expected string option arguments",
240};
241
242const INTERSECT_ERROR_INTERNAL: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
243    code: "RM.INTERSECT.INTERNAL",
244    identifier: Some("RunMat:intersect:Internal"),
245    when: "Internal conversion/allocation/provider decode fails.",
246    message: "intersect: internal operation failed",
247};
248
249const INTERSECT_ERRORS: [BuiltinErrorDescriptor; 8] = [
250    INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED,
251    INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS,
252    INTERSECT_ERROR_UNKNOWN_OPTION,
253    INTERSECT_ERROR_ROWS_COLUMN_MISMATCH,
254    INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE,
255    INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH,
256    INTERSECT_ERROR_INVALID_ARGUMENT,
257    INTERSECT_ERROR_INTERNAL,
258];
259
260const INTERSECT_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 1] =
261    [BuiltinIntegerCapabilityDescriptor {
262        form: "[C, ia, ib] = intersect(integer_A, integer_B, options)",
263        inputs: &super::BINARY_SET_INTEGER_INPUTS,
264        computation_domain: BuiltinIntegerComputationDomain::ExactInteger,
265        output_class: BuiltinIntegerOutputClassRule::FunctionSpecific,
266        overflow: BuiltinIntegerOverflowRule::NotApplicable,
267        backend: BuiltinIntegerBackendRule::GpuRestricted,
268        overload: BuiltinIntegerOverloadKind::Multiple,
269        notes: "C preserves the common nondouble integer class, including when paired with double; ia and ib are one-based double. GPU supports integer classes through 32 bits and restores outputs after typed fallback.",
270    }];
271
272pub const INTERSECT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
273    signatures: &INTERSECT_SIGNATURES,
274    output_mode: BuiltinOutputMode::ByRequestedOutputCount,
275    completion_policy: BuiltinCompletionPolicy::Public,
276    errors: &INTERSECT_ERRORS,
277};
278
279fn intersect_error_with(
280    error: &'static BuiltinErrorDescriptor,
281    message: impl Into<String>,
282) -> crate::RuntimeError {
283    let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
284    if let Some(identifier) = error.identifier {
285        builder = builder.with_identifier(identifier);
286    }
287    builder.build()
288}
289
290fn intersect_error(error: &'static BuiltinErrorDescriptor) -> crate::RuntimeError {
291    intersect_error_with(error, error.message)
292}
293
294fn intersect_internal_error(message: impl Into<String>) -> crate::RuntimeError {
295    intersect_error_with(&INTERSECT_ERROR_INTERNAL, message)
296}
297
298#[runtime_builtin(
299    name = "intersect",
300    category = "array/sorting_sets",
301    summary = "Return common elements or rows across arrays with index outputs.",
302    keywords = "intersect,set,stable,rows,indices,gpu",
303    accel = "array_construct",
304    sink = true,
305    type_resolver(set_values_output_type),
306    descriptor(crate::builtins::array::sorting_sets::intersect::INTERSECT_DESCRIPTOR),
307    integer_capabilities(INTERSECT_INTEGER_CAPABILITIES),
308    builtin_path = "crate::builtins::array::sorting_sets::intersect"
309)]
310async fn intersect_builtin(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
311    if matches!(crate::output_count::current_output_count(), Some(n) if n > 3) {
312        return Err(intersect_error_with(
313            &INTERSECT_ERROR_INVALID_ARGUMENT,
314            "intersect: too many output arguments; maximum is 3",
315        ));
316    }
317    let provider = super::set_output_provider(&a, &b);
318    let eval = evaluate(a, b, &rest).await?;
319    if let Some(out_count) = crate::output_count::current_output_count() {
320        if out_count == 0 {
321            return Ok(Value::OutputList(Vec::new()));
322        }
323        if out_count == 1 {
324            let outputs = super::restore_set_outputs(
325                provider,
326                BUILTIN_NAME,
327                vec![eval.into_values_value()],
328                intersect_internal_error,
329            )?;
330            return Ok(Value::OutputList(outputs));
331        }
332        if out_count == 2 {
333            let (values, ia) = eval.into_pair();
334            let outputs = super::restore_set_outputs(
335                provider,
336                BUILTIN_NAME,
337                vec![values, ia],
338                intersect_internal_error,
339            )?;
340            return Ok(Value::OutputList(outputs));
341        }
342        let (values, ia, ib) = eval.into_triple();
343        let outputs = super::restore_set_outputs(
344            provider,
345            BUILTIN_NAME,
346            vec![values, ia, ib],
347            intersect_internal_error,
348        )?;
349        return Ok(Value::OutputList(outputs));
350    }
351    let mut outputs = super::restore_set_outputs(
352        provider,
353        BUILTIN_NAME,
354        vec![eval.into_values_value()],
355        intersect_internal_error,
356    )?;
357    Ok(outputs.pop().expect("intersect output"))
358}
359
360/// Evaluate the `intersect` builtin once and expose all outputs.
361pub async fn evaluate(
362    a: Value,
363    b: Value,
364    rest: &[Value],
365) -> crate::BuiltinResult<IntersectEvaluation> {
366    crate::builtins::common::validation::reject_typed_complex_integer(&a, "intersect")?;
367    crate::builtins::common::validation::reject_typed_complex_integer(&b, "intersect")?;
368    let opts = parse_options(rest)?;
369    for value in [&a, &b] {
370        if let Value::GpuTensor(handle) = value {
371            if super::is_unsupported_set_gpu_integer(handle) {
372                return Err(intersect_error_with(
373                    &INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE,
374                    "intersect: resident 64-bit integer inputs are not supported",
375                ));
376            }
377        }
378    }
379    match (a, b) {
380        (Value::GpuTensor(handle_a), Value::GpuTensor(handle_b)) => {
381            intersect_gpu_pair(handle_a, handle_b, &opts).await
382        }
383        (Value::GpuTensor(handle_a), other) => {
384            intersect_gpu_mixed(handle_a, other, &opts, true).await
385        }
386        (other, Value::GpuTensor(handle_b)) => {
387            intersect_gpu_mixed(handle_b, other, &opts, false).await
388        }
389        (left, right) => intersect_host(left, right, &opts),
390    }
391}
392
393#[derive(Debug, Clone, Copy, PartialEq, Eq)]
394enum IntersectOrder {
395    Sorted,
396    Stable,
397}
398
399#[derive(Debug, Clone)]
400struct IntersectOptions {
401    rows: bool,
402    order: IntersectOrder,
403}
404
405fn parse_options(rest: &[Value]) -> crate::BuiltinResult<IntersectOptions> {
406    let mut opts = IntersectOptions {
407        rows: false,
408        order: IntersectOrder::Sorted,
409    };
410    let mut seen_order: Option<IntersectOrder> = None;
411
412    let tokens = tokens_from_values(rest);
413    for (arg, token) in rest.iter().zip(tokens.iter()) {
414        let text = match token {
415            crate::builtins::common::arg_tokens::ArgToken::String(text) => text.as_str(),
416            _ => {
417                let text = tensor::value_to_string(arg)
418                    .ok_or_else(|| intersect_error(&INTERSECT_ERROR_INVALID_ARGUMENT))?;
419                let lowered = text.trim().to_ascii_lowercase();
420                parse_intersect_option(&mut opts, &mut seen_order, &lowered)?;
421                continue;
422            }
423        };
424        parse_intersect_option(&mut opts, &mut seen_order, text)?;
425    }
426
427    Ok(opts)
428}
429
430fn parse_intersect_option(
431    opts: &mut IntersectOptions,
432    seen_order: &mut Option<IntersectOrder>,
433    lowered: &str,
434) -> crate::BuiltinResult<()> {
435    match lowered {
436        "rows" => opts.rows = true,
437        "sorted" => {
438            if let Some(prev) = seen_order {
439                if *prev != IntersectOrder::Sorted {
440                    return Err(intersect_error(&INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS));
441                }
442            }
443            *seen_order = Some(IntersectOrder::Sorted);
444            opts.order = IntersectOrder::Sorted;
445        }
446        "stable" => {
447            if let Some(prev) = seen_order {
448                if *prev != IntersectOrder::Stable {
449                    return Err(intersect_error(&INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS));
450                }
451            }
452            *seen_order = Some(IntersectOrder::Stable);
453            opts.order = IntersectOrder::Stable;
454        }
455        "legacy" | "r2012a" => {
456            return Err(intersect_error(&INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED));
457        }
458        other => {
459            return Err(intersect_error_with(
460                &INTERSECT_ERROR_UNKNOWN_OPTION,
461                format!("intersect: unrecognised option '{other}'"),
462            ))
463        }
464    }
465    Ok(())
466}
467
468async fn intersect_gpu_pair(
469    handle_a: GpuTensorHandle,
470    handle_b: GpuTensorHandle,
471    opts: &IntersectOptions,
472) -> crate::BuiltinResult<IntersectEvaluation> {
473    let tensor_a = gpu_helpers::gather_tensor_async(&handle_a).await?;
474    let tensor_b = gpu_helpers::gather_tensor_async(&handle_b).await?;
475    intersect_numeric(tensor_a, tensor_b, opts)
476}
477
478async fn intersect_gpu_mixed(
479    handle_gpu: GpuTensorHandle,
480    other: Value,
481    opts: &IntersectOptions,
482    gpu_is_a: bool,
483) -> crate::BuiltinResult<IntersectEvaluation> {
484    let tensor_gpu = gpu_helpers::gather_tensor_async(&handle_gpu).await?;
485    let tensor_other = tensor::value_into_tensor_for("intersect", other)
486        .map_err(|e| intersect_internal_error(e))?;
487    if gpu_is_a {
488        intersect_numeric(tensor_gpu, tensor_other, opts)
489    } else {
490        intersect_numeric(tensor_other, tensor_gpu, opts)
491    }
492}
493
494fn intersect_host(
495    a: Value,
496    b: Value,
497    opts: &IntersectOptions,
498) -> crate::BuiltinResult<IntersectEvaluation> {
499    match (a, b) {
500        (Value::ComplexTensor(at), Value::ComplexTensor(bt)) => intersect_complex(at, bt, opts),
501        (Value::ComplexTensor(at), Value::Complex(re, im)) => {
502            let bt = scalar_complex_tensor(re, im)?;
503            intersect_complex(at, bt, opts)
504        }
505        (Value::Complex(re, im), Value::ComplexTensor(bt)) => {
506            let at = scalar_complex_tensor(re, im)?;
507            intersect_complex(at, bt, opts)
508        }
509        (Value::Complex(a_re, a_im), Value::Complex(b_re, b_im)) => {
510            let at = scalar_complex_tensor(a_re, a_im)?;
511            let bt = scalar_complex_tensor(b_re, b_im)?;
512            intersect_complex(at, bt, opts)
513        }
514        (Value::ComplexTensor(at), other) => {
515            let bt = value_into_complex_tensor(other)?;
516            intersect_complex(at, bt, opts)
517        }
518        (other, Value::ComplexTensor(bt)) => {
519            let at = value_into_complex_tensor(other)?;
520            intersect_complex(at, bt, opts)
521        }
522        (Value::Complex(re, im), other) => {
523            let at = scalar_complex_tensor(re, im)?;
524            let bt = value_into_complex_tensor(other)?;
525            intersect_complex(at, bt, opts)
526        }
527        (other, Value::Complex(re, im)) => {
528            let at = value_into_complex_tensor(other)?;
529            let bt = scalar_complex_tensor(re, im)?;
530            intersect_complex(at, bt, opts)
531        }
532
533        (Value::CharArray(ac), Value::CharArray(bc)) => intersect_char(ac, bc, opts),
534
535        (Value::StringArray(astring), Value::StringArray(bstring)) => {
536            intersect_string(astring, bstring, opts)
537        }
538        (Value::StringArray(astring), Value::String(b)) => {
539            let bstring = StringArray::new(vec![b], vec![1, 1])
540                .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
541            intersect_string(astring, bstring, opts)
542        }
543        (Value::String(a), Value::StringArray(bstring)) => {
544            let astring = StringArray::new(vec![a], vec![1, 1])
545                .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
546            intersect_string(astring, bstring, opts)
547        }
548        (Value::String(a), Value::String(b)) => {
549            let astring = StringArray::new(vec![a], vec![1, 1])
550                .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
551            let bstring = StringArray::new(vec![b], vec![1, 1])
552                .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
553            intersect_string(astring, bstring, opts)
554        }
555
556        (left, right) => {
557            let tensor_a = tensor::value_into_tensor_for("intersect", left)
558                .map_err(|e| intersect_error_with(&INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
559            let tensor_b = tensor::value_into_tensor_for("intersect", right)
560                .map_err(|e| intersect_error_with(&INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE, e))?;
561            intersect_numeric(tensor_a, tensor_b, opts)
562        }
563    }
564}
565
566fn intersect_numeric(
567    a: Tensor,
568    b: Tensor,
569    opts: &IntersectOptions,
570) -> crate::BuiltinResult<IntersectEvaluation> {
571    let a_dtype = a.numeric_dtype();
572    let b_dtype = b.numeric_dtype();
573    if let (Some(a_storage), Some(b_storage)) = (a.integer_storage(), b.integer_storage()) {
574        if a_storage.class_name() == b_storage.class_name() {
575            return if opts.rows {
576                intersect_integer_rows(a_storage, a.shape.clone(), b_storage, b.shape.clone(), opts)
577            } else {
578                intersect_integer_elements(a_storage, b_storage, opts)
579            };
580        }
581        return Err(intersect_error(&INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH));
582    }
583    match (a.integer_storage(), b.integer_storage()) {
584        (Some(storage), None) if b_dtype == NumericDType::F64 => {
585            let target = IntegerTarget::from_storage(storage);
586            let b = target.cast_tensor(b).map_err(intersect_internal_error)?;
587            return intersect_numeric(a, b, opts);
588        }
589        (None, Some(storage)) if a_dtype == NumericDType::F64 => {
590            let target = IntegerTarget::from_storage(storage);
591            let a = target.cast_tensor(a).map_err(intersect_internal_error)?;
592            return intersect_numeric(a, b, opts);
593        }
594        _ => {}
595    }
596    if a_dtype != b_dtype && a_dtype != NumericDType::F64 && b_dtype != NumericDType::F64 {
597        return Err(intersect_error(&INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH));
598    }
599    let a_shape = a.shape.clone();
600    let b_shape = b.shape.clone();
601    let a_storage = a
602        .into_numeric_storage()
603        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
604    let b_storage = b
605        .into_numeric_storage()
606        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
607    match (a_storage, b_storage) {
608        (NumericStorage::F64(a), NumericStorage::F64(b)) => {
609            intersect_floating(a, a_shape, b, b_shape, opts)
610        }
611        (NumericStorage::F32(a), NumericStorage::F32(b)) => {
612            intersect_floating(a, a_shape, b, b_shape, opts)
613        }
614        (a, b) => intersect_promoted_f64(a, a_shape, b, b_shape, opts),
615    }
616}
617
618fn intersect_promoted_f64(
619    a: NumericStorage,
620    a_shape: Vec<usize>,
621    b: NumericStorage,
622    b_shape: Vec<usize>,
623    opts: &IntersectOptions,
624) -> crate::BuiltinResult<IntersectEvaluation> {
625    intersect_floating(
626        a.materialize_f64(),
627        a_shape,
628        b.materialize_f64(),
629        b_shape,
630        opts,
631    )
632}
633
634fn intersect_floating<T: SetFloat>(
635    a: Vec<T>,
636    a_shape: Vec<usize>,
637    b: Vec<T>,
638    b_shape: Vec<usize>,
639    opts: &IntersectOptions,
640) -> crate::BuiltinResult<IntersectEvaluation> {
641    if opts.rows {
642        intersect_floating_rows(a, a_shape, b, b_shape, opts)
643    } else {
644        intersect_floating_elements(a, b, opts)
645    }
646}
647
648fn intersect_integer_elements(
649    a: &IntegerStorage,
650    b: &IntegerStorage,
651    opts: &IntersectOptions,
652) -> crate::BuiltinResult<IntersectEvaluation> {
653    let mut b_map = HashMap::<IntValue, usize>::new();
654    for (index, value) in b.exact_values().into_iter().enumerate() {
655        b_map.entry(value).or_insert(index);
656    }
657    let mut seen = HashSet::<IntValue>::new();
658    let mut entries = Vec::<IntegerIntersectEntry>::new();
659    for (a_index, value) in a.exact_values().into_iter().enumerate() {
660        if seen.contains(&value) {
661            continue;
662        }
663        if let Some(&b_index) = b_map.get(&value) {
664            let order_rank = entries.len();
665            entries.push(IntegerIntersectEntry {
666                value: value.clone(),
667                a_index,
668                b_index,
669                order_rank,
670            });
671            seen.insert(value);
672        }
673    }
674    assemble_integer_intersect(entries, a, opts)
675}
676
677fn intersect_integer_rows(
678    a_storage: &IntegerStorage,
679    a_shape: Vec<usize>,
680    b_storage: &IntegerStorage,
681    b_shape: Vec<usize>,
682    opts: &IntersectOptions,
683) -> crate::BuiltinResult<IntersectEvaluation> {
684    if a_shape.len() != 2 || b_shape.len() != 2 {
685        return Err(intersect_internal_error(
686            "intersect: 'rows' option requires 2-D numeric matrices",
687        ));
688    }
689    if a_shape[1] != b_shape[1] {
690        return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
691    }
692    let (rows_a, rows_b, cols) = (a_shape[0], b_shape[0], a_shape[1]);
693    let a_values = a_storage.exact_values();
694    let b_values = b_storage.exact_values();
695    let mut b_map = HashMap::<Vec<IntValue>, usize>::new();
696    for row in 0..rows_b {
697        let key: Vec<_> = (0..cols)
698            .map(|col| b_values[row + col * rows_b].clone())
699            .collect();
700        b_map.entry(key).or_insert(row);
701    }
702    let mut seen = HashSet::<Vec<IntValue>>::new();
703    let mut entries = Vec::<IntegerRowIntersectEntry>::new();
704    for row in 0..rows_a {
705        let values: Vec<_> = (0..cols)
706            .map(|col| a_values[row + col * rows_a].clone())
707            .collect();
708        if seen.contains(&values) {
709            continue;
710        }
711        if let Some(&b_row) = b_map.get(&values) {
712            let order_rank = entries.len();
713            entries.push(IntegerRowIntersectEntry {
714                row_data: values.clone(),
715                a_row: row,
716                b_row,
717                order_rank,
718            });
719            seen.insert(values);
720        }
721    }
722    assemble_integer_row_intersect(entries, a_storage, opts, cols)
723}
724
725fn intersect_floating_elements<T: SetFloat>(
726    a_values: Vec<T>,
727    b_values: Vec<T>,
728    opts: &IntersectOptions,
729) -> crate::BuiltinResult<IntersectEvaluation> {
730    let mut b_map: HashMap<u64, usize> = HashMap::new();
731    for (idx, &value) in b_values.iter().enumerate() {
732        let key = value.canonical_key();
733        b_map.entry(key).or_insert(idx);
734    }
735
736    let mut seen: HashSet<u64> = HashSet::new();
737    let mut entries = Vec::<FloatingIntersectEntry<T>>::new();
738    let mut order_counter = 0usize;
739
740    for (idx, &value) in a_values.iter().enumerate() {
741        let key = value.canonical_key();
742        if seen.contains(&key) {
743            continue;
744        }
745        if let Some(&b_idx) = b_map.get(&key) {
746            entries.push(FloatingIntersectEntry {
747                value,
748                a_index: idx,
749                b_index: b_idx,
750                order_rank: order_counter,
751            });
752            seen.insert(key);
753            order_counter += 1;
754        }
755    }
756
757    assemble_floating_intersect(entries, opts)
758}
759
760fn intersect_floating_rows<T: SetFloat>(
761    a_values: Vec<T>,
762    a_shape: Vec<usize>,
763    b_values: Vec<T>,
764    b_shape: Vec<usize>,
765    opts: &IntersectOptions,
766) -> crate::BuiltinResult<IntersectEvaluation> {
767    if a_shape.len() != 2 || b_shape.len() != 2 {
768        return Err(intersect_internal_error(
769            "intersect: 'rows' option requires 2-D numeric matrices",
770        ));
771    }
772    if a_shape[1] != b_shape[1] {
773        return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
774    }
775    let rows_a = a_shape[0];
776    let cols = a_shape[1];
777    let rows_b = b_shape[0];
778
779    let mut b_map: HashMap<FloatingRowKey, usize> = HashMap::new();
780    for r in 0..rows_b {
781        let mut row_values = Vec::with_capacity(cols);
782        for c in 0..cols {
783            let idx = r + c * rows_b;
784            row_values.push(b_values[idx]);
785        }
786        let key = FloatingRowKey::from_slice(&row_values);
787        b_map.entry(key).or_insert(r);
788    }
789
790    let mut seen: HashSet<FloatingRowKey> = HashSet::new();
791    let mut entries = Vec::<FloatingRowIntersectEntry<T>>::new();
792    let mut order_counter = 0usize;
793
794    for r in 0..rows_a {
795        let mut row_values = Vec::with_capacity(cols);
796        for c in 0..cols {
797            let idx = r + c * rows_a;
798            row_values.push(a_values[idx]);
799        }
800        let key = FloatingRowKey::from_slice(&row_values);
801        if seen.contains(&key) {
802            continue;
803        }
804        if let Some(&b_row) = b_map.get(&key) {
805            entries.push(FloatingRowIntersectEntry {
806                row_data: row_values,
807                a_row: r,
808                b_row,
809                order_rank: order_counter,
810            });
811            seen.insert(key);
812            order_counter += 1;
813        }
814    }
815
816    assemble_floating_row_intersect(entries, opts, cols)
817}
818
819#[cfg(test)]
820fn intersect_numeric_elements(
821    a: Tensor,
822    b: Tensor,
823    opts: &IntersectOptions,
824) -> crate::BuiltinResult<IntersectEvaluation> {
825    intersect_numeric(a, b, opts)
826}
827
828#[cfg(test)]
829fn intersect_numeric_rows(
830    a: Tensor,
831    b: Tensor,
832    opts: &IntersectOptions,
833) -> crate::BuiltinResult<IntersectEvaluation> {
834    intersect_numeric(a, b, opts)
835}
836
837fn intersect_complex(
838    a: ComplexTensor,
839    b: ComplexTensor,
840    opts: &IntersectOptions,
841) -> crate::BuiltinResult<IntersectEvaluation> {
842    let a_shape = a.shape.clone();
843    let b_shape = b.shape.clone();
844    match (a.into_complex_storage(), b.into_complex_storage()) {
845        (ComplexStorage::F64(a), ComplexStorage::F64(b)) => {
846            intersect_floating_complex(a, a_shape, b, b_shape, opts)
847        }
848        (ComplexStorage::F32(a), ComplexStorage::F32(b)) => {
849            intersect_floating_complex(a, a_shape, b, b_shape, opts)
850        }
851        (a, b) => intersect_promoted_complex_f64(a, a_shape, b, b_shape, opts),
852    }
853}
854
855fn intersect_promoted_complex_f64(
856    a: ComplexStorage,
857    a_shape: Vec<usize>,
858    b: ComplexStorage,
859    b_shape: Vec<usize>,
860    opts: &IntersectOptions,
861) -> crate::BuiltinResult<IntersectEvaluation> {
862    intersect_floating_complex(
863        a.materialize_f64(),
864        a_shape,
865        b.materialize_f64(),
866        b_shape,
867        opts,
868    )
869}
870
871fn intersect_floating_complex<T: SetFloat>(
872    a: Vec<(T, T)>,
873    a_shape: Vec<usize>,
874    b: Vec<(T, T)>,
875    b_shape: Vec<usize>,
876    opts: &IntersectOptions,
877) -> crate::BuiltinResult<IntersectEvaluation> {
878    if opts.rows {
879        intersect_complex_rows(a, a_shape, b, b_shape, opts)
880    } else {
881        intersect_complex_elements(a, b, opts)
882    }
883}
884
885fn intersect_complex_elements<T: SetFloat>(
886    a: Vec<(T, T)>,
887    b: Vec<(T, T)>,
888    opts: &IntersectOptions,
889) -> crate::BuiltinResult<IntersectEvaluation> {
890    let mut b_map: HashMap<ComplexKey, usize> = HashMap::new();
891    for (idx, &value) in b.iter().enumerate() {
892        let key = ComplexKey::new(value);
893        b_map.entry(key).or_insert(idx);
894    }
895
896    let mut seen: HashSet<ComplexKey> = HashSet::new();
897    let mut entries = Vec::<ComplexIntersectEntry<T>>::new();
898    let mut order_counter = 0usize;
899
900    for (idx, &value) in a.iter().enumerate() {
901        let key = ComplexKey::new(value);
902        if seen.contains(&key) {
903            continue;
904        }
905        if let Some(&b_idx) = b_map.get(&key) {
906            entries.push(ComplexIntersectEntry {
907                value,
908                a_index: idx,
909                b_index: b_idx,
910                order_rank: order_counter,
911            });
912            seen.insert(key);
913            order_counter += 1;
914        }
915    }
916
917    assemble_complex_intersect(entries, opts)
918}
919
920fn intersect_complex_rows<T: SetFloat>(
921    a: Vec<(T, T)>,
922    a_shape: Vec<usize>,
923    b: Vec<(T, T)>,
924    b_shape: Vec<usize>,
925    opts: &IntersectOptions,
926) -> crate::BuiltinResult<IntersectEvaluation> {
927    if a_shape.len() != 2 || b_shape.len() != 2 {
928        return Err(intersect_internal_error(
929            "intersect: 'rows' option requires 2-D complex matrices",
930        ));
931    }
932    if a_shape[1] != b_shape[1] {
933        return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
934    }
935    let rows_a = a_shape[0];
936    let cols = a_shape[1];
937    let rows_b = b_shape[0];
938
939    let mut b_map: HashMap<Vec<ComplexKey>, usize> = HashMap::new();
940    for r in 0..rows_b {
941        let mut row_keys = Vec::with_capacity(cols);
942        for c in 0..cols {
943            let idx = r + c * rows_b;
944            row_keys.push(ComplexKey::new(b[idx]));
945        }
946        b_map.entry(row_keys).or_insert(r);
947    }
948
949    let mut seen: HashSet<Vec<ComplexKey>> = HashSet::new();
950    let mut entries = Vec::<ComplexRowIntersectEntry<T>>::new();
951    let mut order_counter = 0usize;
952
953    for r in 0..rows_a {
954        let mut row_values = Vec::with_capacity(cols);
955        let mut row_keys = Vec::with_capacity(cols);
956        for c in 0..cols {
957            let idx = r + c * rows_a;
958            let value = a[idx];
959            row_values.push(value);
960            row_keys.push(ComplexKey::new(value));
961        }
962        if seen.contains(&row_keys) {
963            continue;
964        }
965        if let Some(&b_row) = b_map.get(&row_keys) {
966            entries.push(ComplexRowIntersectEntry {
967                row_data: row_values,
968                a_row: r,
969                b_row,
970                order_rank: order_counter,
971            });
972            seen.insert(row_keys);
973            order_counter += 1;
974        }
975    }
976
977    assemble_complex_row_intersect(entries, opts, cols)
978}
979
980fn intersect_char(
981    a: CharArray,
982    b: CharArray,
983    opts: &IntersectOptions,
984) -> crate::BuiltinResult<IntersectEvaluation> {
985    if opts.rows {
986        intersect_char_rows(a, b, opts)
987    } else {
988        intersect_char_elements(a, b, opts)
989    }
990}
991
992fn intersect_char_elements(
993    a: CharArray,
994    b: CharArray,
995    opts: &IntersectOptions,
996) -> crate::BuiltinResult<IntersectEvaluation> {
997    let mut seen: HashSet<u32> = HashSet::new();
998    let mut entries = Vec::<CharIntersectEntry>::new();
999    let mut order_counter = 0usize;
1000
1001    for col in 0..a.cols {
1002        for row in 0..a.rows {
1003            let linear_idx = row + col * a.rows;
1004            let data_idx = row * a.cols + col;
1005            let ch = a.data[data_idx];
1006            let key = ch as u32;
1007            if seen.contains(&key) {
1008                continue;
1009            }
1010            if let Some(b_idx) = find_char_index(&b, ch) {
1011                entries.push(CharIntersectEntry {
1012                    ch,
1013                    a_index: linear_idx,
1014                    b_index: b_idx,
1015                    order_rank: order_counter,
1016                });
1017                seen.insert(key);
1018                order_counter += 1;
1019            }
1020        }
1021    }
1022
1023    assemble_char_intersect(entries, opts, &b)
1024}
1025
1026fn intersect_char_rows(
1027    a: CharArray,
1028    b: CharArray,
1029    opts: &IntersectOptions,
1030) -> crate::BuiltinResult<IntersectEvaluation> {
1031    if a.cols != b.cols {
1032        return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
1033    }
1034    let rows_a = a.rows;
1035    let rows_b = b.rows;
1036    let cols = a.cols;
1037
1038    let mut b_map: HashMap<RowCharKey, usize> = HashMap::new();
1039    for r in 0..rows_b {
1040        let mut row_values = Vec::with_capacity(cols);
1041        for c in 0..cols {
1042            let idx = r * cols + c;
1043            row_values.push(b.data[idx]);
1044        }
1045        let key = RowCharKey::from_slice(&row_values);
1046        b_map.entry(key).or_insert(r);
1047    }
1048
1049    let mut seen: HashSet<RowCharKey> = HashSet::new();
1050    let mut entries = Vec::<CharRowIntersectEntry>::new();
1051    let mut order_counter = 0usize;
1052
1053    for r in 0..rows_a {
1054        let mut row_values = Vec::with_capacity(cols);
1055        for c in 0..cols {
1056            let idx = r * cols + c;
1057            row_values.push(a.data[idx]);
1058        }
1059        let key = RowCharKey::from_slice(&row_values);
1060        if seen.contains(&key) {
1061            continue;
1062        }
1063        if let Some(&b_row) = b_map.get(&key) {
1064            entries.push(CharRowIntersectEntry {
1065                row_data: row_values,
1066                a_row: r,
1067                b_row,
1068                order_rank: order_counter,
1069            });
1070            seen.insert(key);
1071            order_counter += 1;
1072        }
1073    }
1074
1075    assemble_char_row_intersect(entries, opts, cols)
1076}
1077
1078fn find_char_index(array: &CharArray, target: char) -> Option<usize> {
1079    for col in 0..array.cols {
1080        for row in 0..array.rows {
1081            let data_idx = row * array.cols + col;
1082            if array.data[data_idx] == target {
1083                return Some(row + col * array.rows);
1084            }
1085        }
1086    }
1087    None
1088}
1089
1090fn intersect_string(
1091    a: StringArray,
1092    b: StringArray,
1093    opts: &IntersectOptions,
1094) -> crate::BuiltinResult<IntersectEvaluation> {
1095    if opts.rows {
1096        intersect_string_rows(a, b, opts)
1097    } else {
1098        intersect_string_elements(a, b, opts)
1099    }
1100}
1101
1102fn intersect_string_elements(
1103    a: StringArray,
1104    b: StringArray,
1105    opts: &IntersectOptions,
1106) -> crate::BuiltinResult<IntersectEvaluation> {
1107    let mut b_map: HashMap<String, usize> = HashMap::new();
1108    for (idx, value) in b.data.iter().enumerate() {
1109        b_map.entry(value.clone()).or_insert(idx);
1110    }
1111
1112    let mut seen: HashSet<String> = HashSet::new();
1113    let mut entries = Vec::<StringIntersectEntry>::new();
1114    let mut order_counter = 0usize;
1115
1116    for (idx, value) in a.data.iter().enumerate() {
1117        if seen.contains(value) {
1118            continue;
1119        }
1120        if let Some(&b_idx) = b_map.get(value) {
1121            entries.push(StringIntersectEntry {
1122                value: value.clone(),
1123                a_index: idx,
1124                b_index: b_idx,
1125                order_rank: order_counter,
1126            });
1127            seen.insert(value.clone());
1128            order_counter += 1;
1129        }
1130    }
1131
1132    assemble_string_intersect(entries, opts)
1133}
1134
1135fn intersect_string_rows(
1136    a: StringArray,
1137    b: StringArray,
1138    opts: &IntersectOptions,
1139) -> crate::BuiltinResult<IntersectEvaluation> {
1140    if a.shape.len() != 2 || b.shape.len() != 2 {
1141        return Err(intersect_internal_error(
1142            "intersect: 'rows' option requires 2-D string arrays",
1143        ));
1144    }
1145    if a.shape[1] != b.shape[1] {
1146        return Err(intersect_error(&INTERSECT_ERROR_ROWS_COLUMN_MISMATCH));
1147    }
1148    let rows_a = a.shape[0];
1149    let cols = a.shape[1];
1150    let rows_b = b.shape[0];
1151
1152    let mut b_map: HashMap<RowStringKey, usize> = HashMap::new();
1153    for r in 0..rows_b {
1154        let mut row_values = Vec::with_capacity(cols);
1155        for c in 0..cols {
1156            let idx = r + c * rows_b;
1157            row_values.push(b.data[idx].clone());
1158        }
1159        let key = RowStringKey::from_slice(&row_values);
1160        b_map.entry(key).or_insert(r);
1161    }
1162
1163    let mut seen: HashSet<RowStringKey> = HashSet::new();
1164    let mut entries = Vec::<StringRowIntersectEntry>::new();
1165    let mut order_counter = 0usize;
1166
1167    for r in 0..rows_a {
1168        let mut row_values = Vec::with_capacity(cols);
1169        for c in 0..cols {
1170            let idx = r + c * rows_a;
1171            row_values.push(a.data[idx].clone());
1172        }
1173        let key = RowStringKey::from_slice(&row_values);
1174        if seen.contains(&key) {
1175            continue;
1176        }
1177        if let Some(&b_row) = b_map.get(&key) {
1178            entries.push(StringRowIntersectEntry {
1179                row_data: row_values,
1180                a_row: r,
1181                b_row,
1182                order_rank: order_counter,
1183            });
1184            seen.insert(key);
1185            order_counter += 1;
1186        }
1187    }
1188
1189    assemble_string_row_intersect(entries, opts, cols)
1190}
1191
1192#[derive(Debug, Clone)]
1193pub struct IntersectEvaluation {
1194    values: Value,
1195    ia: Tensor,
1196    ib: Tensor,
1197}
1198
1199impl IntersectEvaluation {
1200    fn new(values: Value, ia: Tensor, ib: Tensor) -> Self {
1201        Self { values, ia, ib }
1202    }
1203
1204    pub fn into_values_value(self) -> Value {
1205        self.values
1206    }
1207
1208    pub fn into_pair(self) -> (Value, Value) {
1209        let ia = tensor::tensor_into_value(self.ia);
1210        (self.values, ia)
1211    }
1212
1213    pub fn into_triple(self) -> (Value, Value, Value) {
1214        let ia = tensor::tensor_into_value(self.ia);
1215        let ib = tensor::tensor_into_value(self.ib);
1216        (self.values, ia, ib)
1217    }
1218
1219    pub fn values_value(&self) -> Value {
1220        self.values.clone()
1221    }
1222
1223    pub fn ia_value(&self) -> Value {
1224        tensor::tensor_into_value(self.ia.clone())
1225    }
1226
1227    pub fn ib_value(&self) -> Value {
1228        tensor::tensor_into_value(self.ib.clone())
1229    }
1230}
1231
1232#[derive(Debug)]
1233struct FloatingIntersectEntry<T> {
1234    value: T,
1235    a_index: usize,
1236    b_index: usize,
1237    order_rank: usize,
1238}
1239
1240#[derive(Debug)]
1241struct IntegerIntersectEntry {
1242    value: IntValue,
1243    a_index: usize,
1244    b_index: usize,
1245    order_rank: usize,
1246}
1247
1248#[derive(Debug)]
1249struct FloatingRowIntersectEntry<T> {
1250    row_data: Vec<T>,
1251    a_row: usize,
1252    b_row: usize,
1253    order_rank: usize,
1254}
1255
1256#[derive(Debug)]
1257struct IntegerRowIntersectEntry {
1258    row_data: Vec<IntValue>,
1259    a_row: usize,
1260    b_row: usize,
1261    order_rank: usize,
1262}
1263
1264#[derive(Debug)]
1265struct ComplexIntersectEntry<T> {
1266    value: (T, T),
1267    a_index: usize,
1268    b_index: usize,
1269    order_rank: usize,
1270}
1271
1272#[derive(Debug)]
1273struct ComplexRowIntersectEntry<T> {
1274    row_data: Vec<(T, T)>,
1275    a_row: usize,
1276    b_row: usize,
1277    order_rank: usize,
1278}
1279
1280#[derive(Debug)]
1281struct CharIntersectEntry {
1282    ch: char,
1283    a_index: usize,
1284    b_index: usize,
1285    order_rank: usize,
1286}
1287
1288#[derive(Debug)]
1289struct CharRowIntersectEntry {
1290    row_data: Vec<char>,
1291    a_row: usize,
1292    b_row: usize,
1293    order_rank: usize,
1294}
1295
1296#[derive(Debug)]
1297struct StringIntersectEntry {
1298    value: String,
1299    a_index: usize,
1300    b_index: usize,
1301    order_rank: usize,
1302}
1303
1304#[derive(Debug)]
1305struct StringRowIntersectEntry {
1306    row_data: Vec<String>,
1307    a_row: usize,
1308    b_row: usize,
1309    order_rank: usize,
1310}
1311
1312fn assemble_floating_intersect<T: SetFloat>(
1313    entries: Vec<FloatingIntersectEntry<T>>,
1314    opts: &IntersectOptions,
1315) -> crate::BuiltinResult<IntersectEvaluation> {
1316    let mut order: Vec<usize> = (0..entries.len()).collect();
1317    match opts.order {
1318        IntersectOrder::Sorted => {
1319            order.sort_by(|&lhs, &rhs| entries[lhs].value.compare(entries[rhs].value));
1320        }
1321        IntersectOrder::Stable => {
1322            order.sort_by_key(|&idx| entries[idx].order_rank);
1323        }
1324    }
1325
1326    let mut values = Vec::with_capacity(order.len());
1327    let mut ia = Vec::with_capacity(order.len());
1328    let mut ib = Vec::with_capacity(order.len());
1329    for &idx in &order {
1330        let entry = &entries[idx];
1331        values.push(entry.value);
1332        ia.push((entry.a_index + 1) as f64);
1333        ib.push((entry.b_index + 1) as f64);
1334    }
1335
1336    let value_tensor =
1337        Tensor::from_numeric_storage(T::numeric_storage(values), vec![order.len(), 1])
1338            .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1339    let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1340        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1341    let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1342        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1343
1344    Ok(IntersectEvaluation::new(
1345        tensor::tensor_into_value(value_tensor),
1346        ia_tensor,
1347        ib_tensor,
1348    ))
1349}
1350
1351fn assemble_integer_intersect(
1352    entries: Vec<IntegerIntersectEntry>,
1353    storage: &IntegerStorage,
1354    opts: &IntersectOptions,
1355) -> crate::BuiltinResult<IntersectEvaluation> {
1356    let mut order: Vec<_> = (0..entries.len()).collect();
1357    match opts.order {
1358        IntersectOrder::Sorted => order.sort_by(|&a, &b| {
1359            integer_order::compare(&entries[a].value, &entries[b].value, false, false)
1360        }),
1361        IntersectOrder::Stable => order.sort_by_key(|&index| entries[index].order_rank),
1362    }
1363    let values: Vec<_> = order
1364        .iter()
1365        .map(|&index| entries[index].value.clone())
1366        .collect();
1367    let ia: Vec<_> = order
1368        .iter()
1369        .map(|&index| (entries[index].a_index + 1) as f64)
1370        .collect();
1371    let ib: Vec<_> = order
1372        .iter()
1373        .map(|&index| (entries[index].b_index + 1) as f64)
1374        .collect();
1375    let values = Tensor::new_integer(
1376        storage
1377            .from_exact_values_like(values)
1378            .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?,
1379        vec![order.len(), 1],
1380    )
1381    .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1382    let ia = Tensor::new(ia, vec![order.len(), 1])
1383        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1384    let ib = Tensor::new(ib, vec![order.len(), 1])
1385        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1386    Ok(IntersectEvaluation::new(Value::Tensor(values), ia, ib))
1387}
1388
1389fn assemble_floating_row_intersect<T: SetFloat>(
1390    entries: Vec<FloatingRowIntersectEntry<T>>,
1391    opts: &IntersectOptions,
1392    cols: usize,
1393) -> crate::BuiltinResult<IntersectEvaluation> {
1394    let mut order: Vec<usize> = (0..entries.len()).collect();
1395    match opts.order {
1396        IntersectOrder::Sorted => {
1397            order.sort_by(|&lhs, &rhs| {
1398                compare_floating_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1399            });
1400        }
1401        IntersectOrder::Stable => {
1402            order.sort_by_key(|&idx| entries[idx].order_rank);
1403        }
1404    }
1405
1406    let rows_out = order.len();
1407    let mut values = vec![T::default(); rows_out * cols];
1408    let mut ia = Vec::with_capacity(rows_out);
1409    let mut ib = Vec::with_capacity(rows_out);
1410
1411    for (row_pos, &entry_idx) in order.iter().enumerate() {
1412        let entry = &entries[entry_idx];
1413        for col in 0..cols {
1414            let dest = row_pos + col * rows_out;
1415            values[dest] = entry.row_data[col];
1416        }
1417        ia.push((entry.a_row + 1) as f64);
1418        ib.push((entry.b_row + 1) as f64);
1419    }
1420
1421    let value_tensor =
1422        Tensor::from_numeric_storage(T::numeric_storage(values), vec![rows_out, cols])
1423            .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1424    let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1425        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1426    let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1427        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1428
1429    Ok(IntersectEvaluation::new(
1430        tensor::tensor_into_value(value_tensor),
1431        ia_tensor,
1432        ib_tensor,
1433    ))
1434}
1435
1436fn assemble_integer_row_intersect(
1437    entries: Vec<IntegerRowIntersectEntry>,
1438    storage: &IntegerStorage,
1439    opts: &IntersectOptions,
1440    cols: usize,
1441) -> crate::BuiltinResult<IntersectEvaluation> {
1442    let mut order: Vec<_> = (0..entries.len()).collect();
1443    match opts.order {
1444        IntersectOrder::Sorted => order.sort_by(|&a, &b| {
1445            for (left, right) in entries[a].row_data.iter().zip(&entries[b].row_data) {
1446                let ordering = integer_order::compare(left, right, false, false);
1447                if ordering != Ordering::Equal {
1448                    return ordering;
1449                }
1450            }
1451            Ordering::Equal
1452        }),
1453        IntersectOrder::Stable => order.sort_by_key(|&index| entries[index].order_rank),
1454    }
1455    let rows = order.len();
1456    let mut values = Vec::with_capacity(rows * cols);
1457    for col in 0..cols {
1458        for &index in &order {
1459            values.push(entries[index].row_data[col].clone());
1460        }
1461    }
1462    let ia: Vec<_> = order
1463        .iter()
1464        .map(|&index| (entries[index].a_row + 1) as f64)
1465        .collect();
1466    let ib: Vec<_> = order
1467        .iter()
1468        .map(|&index| (entries[index].b_row + 1) as f64)
1469        .collect();
1470    let values = Tensor::new_integer(
1471        storage
1472            .from_exact_values_like(values)
1473            .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?,
1474        vec![rows, cols],
1475    )
1476    .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1477    let ia = Tensor::new(ia, vec![rows, 1])
1478        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1479    let ib = Tensor::new(ib, vec![rows, 1])
1480        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1481    Ok(IntersectEvaluation::new(Value::Tensor(values), ia, ib))
1482}
1483
1484fn assemble_complex_intersect<T: SetFloat>(
1485    entries: Vec<ComplexIntersectEntry<T>>,
1486    opts: &IntersectOptions,
1487) -> crate::BuiltinResult<IntersectEvaluation> {
1488    let mut order: Vec<usize> = (0..entries.len()).collect();
1489    match opts.order {
1490        IntersectOrder::Sorted => {
1491            order.sort_by(|&lhs, &rhs| compare_complex(entries[lhs].value, entries[rhs].value));
1492        }
1493        IntersectOrder::Stable => {
1494            order.sort_by_key(|&idx| entries[idx].order_rank);
1495        }
1496    }
1497
1498    let mut values = Vec::with_capacity(order.len());
1499    let mut ia = Vec::with_capacity(order.len());
1500    let mut ib = Vec::with_capacity(order.len());
1501    for &idx in &order {
1502        let entry = &entries[idx];
1503        values.push(entry.value);
1504        ia.push((entry.a_index + 1) as f64);
1505        ib.push((entry.b_index + 1) as f64);
1506    }
1507
1508    let value_tensor =
1509        ComplexTensor::from_complex_storage(T::complex_storage(values), vec![order.len(), 1])
1510            .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1511    let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1512        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1513    let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1514        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1515
1516    let value = if value_tensor.as_f32_slice().is_some() {
1517        Value::ComplexTensor(value_tensor)
1518    } else {
1519        complex_tensor_into_value(value_tensor)
1520    };
1521    Ok(IntersectEvaluation::new(value, ia_tensor, ib_tensor))
1522}
1523
1524fn assemble_complex_row_intersect<T: SetFloat>(
1525    entries: Vec<ComplexRowIntersectEntry<T>>,
1526    opts: &IntersectOptions,
1527    cols: usize,
1528) -> crate::BuiltinResult<IntersectEvaluation> {
1529    let mut order: Vec<usize> = (0..entries.len()).collect();
1530    match opts.order {
1531        IntersectOrder::Sorted => {
1532            order.sort_by(|&lhs, &rhs| {
1533                compare_complex_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1534            });
1535        }
1536        IntersectOrder::Stable => {
1537            order.sort_by_key(|&idx| entries[idx].order_rank);
1538        }
1539    }
1540
1541    let rows_out = order.len();
1542    let mut values = vec![(T::default(), T::default()); rows_out * cols];
1543    let mut ia = Vec::with_capacity(rows_out);
1544    let mut ib = Vec::with_capacity(rows_out);
1545
1546    for (row_pos, &entry_idx) in order.iter().enumerate() {
1547        let entry = &entries[entry_idx];
1548        for col in 0..cols {
1549            let dest = row_pos + col * rows_out;
1550            values[dest] = entry.row_data[col];
1551        }
1552        ia.push((entry.a_row + 1) as f64);
1553        ib.push((entry.b_row + 1) as f64);
1554    }
1555
1556    let value_tensor =
1557        ComplexTensor::from_complex_storage(T::complex_storage(values), vec![rows_out, cols])
1558            .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1559    let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1560        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1561    let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1562        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1563
1564    let value = if value_tensor.as_f32_slice().is_some() {
1565        Value::ComplexTensor(value_tensor)
1566    } else {
1567        complex_tensor_into_value(value_tensor)
1568    };
1569    Ok(IntersectEvaluation::new(value, ia_tensor, ib_tensor))
1570}
1571
1572fn assemble_char_intersect(
1573    entries: Vec<CharIntersectEntry>,
1574    opts: &IntersectOptions,
1575    b: &CharArray,
1576) -> crate::BuiltinResult<IntersectEvaluation> {
1577    let mut order: Vec<usize> = (0..entries.len()).collect();
1578    match opts.order {
1579        IntersectOrder::Sorted => {
1580            order.sort_by(|&lhs, &rhs| entries[lhs].ch.cmp(&entries[rhs].ch));
1581        }
1582        IntersectOrder::Stable => {
1583            order.sort_by_key(|&idx| entries[idx].order_rank);
1584        }
1585    }
1586
1587    let mut values = Vec::with_capacity(order.len());
1588    let mut ia = Vec::with_capacity(order.len());
1589    let mut ib = Vec::with_capacity(order.len());
1590    for &idx in &order {
1591        let entry = &entries[idx];
1592        values.push(entry.ch);
1593        ia.push((entry.a_index + 1) as f64);
1594        let b_idx = find_char_index(b, entry.ch).unwrap_or(entry.b_index);
1595        ib.push((b_idx + 1) as f64);
1596    }
1597
1598    let value_array = CharArray::new(values, order.len(), 1)
1599        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1600    let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1601        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1602    let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1603        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1604
1605    Ok(IntersectEvaluation::new(
1606        Value::CharArray(value_array),
1607        ia_tensor,
1608        ib_tensor,
1609    ))
1610}
1611
1612fn assemble_char_row_intersect(
1613    entries: Vec<CharRowIntersectEntry>,
1614    opts: &IntersectOptions,
1615    cols: usize,
1616) -> crate::BuiltinResult<IntersectEvaluation> {
1617    let mut order: Vec<usize> = (0..entries.len()).collect();
1618    match opts.order {
1619        IntersectOrder::Sorted => {
1620            order.sort_by(|&lhs, &rhs| {
1621                compare_char_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1622            });
1623        }
1624        IntersectOrder::Stable => {
1625            order.sort_by_key(|&idx| entries[idx].order_rank);
1626        }
1627    }
1628
1629    let rows_out = order.len();
1630    let mut values = vec!['\0'; rows_out * cols];
1631    let mut ia = Vec::with_capacity(rows_out);
1632    let mut ib = Vec::with_capacity(rows_out);
1633
1634    for (row_pos, &entry_idx) in order.iter().enumerate() {
1635        let entry = &entries[entry_idx];
1636        for col in 0..cols {
1637            let dest = row_pos * cols + col;
1638            values[dest] = entry.row_data[col];
1639        }
1640        ia.push((entry.a_row + 1) as f64);
1641        ib.push((entry.b_row + 1) as f64);
1642    }
1643
1644    let value_array = CharArray::new(values, rows_out, cols)
1645        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1646    let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1647        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1648    let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1649        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1650
1651    Ok(IntersectEvaluation::new(
1652        Value::CharArray(value_array),
1653        ia_tensor,
1654        ib_tensor,
1655    ))
1656}
1657
1658fn assemble_string_intersect(
1659    entries: Vec<StringIntersectEntry>,
1660    opts: &IntersectOptions,
1661) -> crate::BuiltinResult<IntersectEvaluation> {
1662    let mut order: Vec<usize> = (0..entries.len()).collect();
1663    match opts.order {
1664        IntersectOrder::Sorted => {
1665            order.sort_by(|&lhs, &rhs| entries[lhs].value.cmp(&entries[rhs].value));
1666        }
1667        IntersectOrder::Stable => {
1668            order.sort_by_key(|&idx| entries[idx].order_rank);
1669        }
1670    }
1671
1672    let mut values = Vec::with_capacity(order.len());
1673    let mut ia = Vec::with_capacity(order.len());
1674    let mut ib = Vec::with_capacity(order.len());
1675    for &idx in &order {
1676        let entry = &entries[idx];
1677        values.push(entry.value.clone());
1678        ia.push((entry.a_index + 1) as f64);
1679        ib.push((entry.b_index + 1) as f64);
1680    }
1681
1682    let value_array = StringArray::new(values, vec![order.len(), 1])
1683        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1684    let ia_tensor = Tensor::new(ia, vec![order.len(), 1])
1685        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1686    let ib_tensor = Tensor::new(ib, vec![order.len(), 1])
1687        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1688
1689    Ok(IntersectEvaluation::new(
1690        Value::StringArray(value_array),
1691        ia_tensor,
1692        ib_tensor,
1693    ))
1694}
1695
1696fn assemble_string_row_intersect(
1697    entries: Vec<StringRowIntersectEntry>,
1698    opts: &IntersectOptions,
1699    cols: usize,
1700) -> crate::BuiltinResult<IntersectEvaluation> {
1701    let mut order: Vec<usize> = (0..entries.len()).collect();
1702    match opts.order {
1703        IntersectOrder::Sorted => {
1704            order.sort_by(|&lhs, &rhs| {
1705                compare_string_rows(&entries[lhs].row_data, &entries[rhs].row_data)
1706            });
1707        }
1708        IntersectOrder::Stable => {
1709            order.sort_by_key(|&idx| entries[idx].order_rank);
1710        }
1711    }
1712
1713    let rows_out = order.len();
1714    let mut values = vec![String::new(); rows_out * cols];
1715    let mut ia = Vec::with_capacity(rows_out);
1716    let mut ib = Vec::with_capacity(rows_out);
1717
1718    for (row_pos, &entry_idx) in order.iter().enumerate() {
1719        let entry = &entries[entry_idx];
1720        for col in 0..cols {
1721            let dest = row_pos + col * rows_out;
1722            values[dest] = entry.row_data[col].clone();
1723        }
1724        ia.push((entry.a_row + 1) as f64);
1725        ib.push((entry.b_row + 1) as f64);
1726    }
1727
1728    let value_array = StringArray::new(values, vec![rows_out, cols])
1729        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1730    let ia_tensor = Tensor::new(ia, vec![rows_out, 1])
1731        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1732    let ib_tensor = Tensor::new(ib, vec![rows_out, 1])
1733        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))?;
1734
1735    Ok(IntersectEvaluation::new(
1736        Value::StringArray(value_array),
1737        ia_tensor,
1738        ib_tensor,
1739    ))
1740}
1741
1742#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1743struct FloatingRowKey(Vec<u64>);
1744
1745impl FloatingRowKey {
1746    fn from_slice<T: SetFloat>(values: &[T]) -> Self {
1747        Self(values.iter().map(|&value| value.canonical_key()).collect())
1748    }
1749}
1750
1751#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
1752struct ComplexKey {
1753    re: u64,
1754    im: u64,
1755}
1756
1757impl ComplexKey {
1758    fn new<T: SetFloat>(value: (T, T)) -> Self {
1759        Self {
1760            re: value.0.canonical_key(),
1761            im: value.1.canonical_key(),
1762        }
1763    }
1764}
1765
1766#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1767struct RowCharKey(Vec<u32>);
1768
1769impl RowCharKey {
1770    fn from_slice(values: &[char]) -> Self {
1771        RowCharKey(values.iter().map(|&ch| ch as u32).collect())
1772    }
1773}
1774
1775#[derive(Debug, Clone, PartialEq, Eq, Hash)]
1776struct RowStringKey(Vec<String>);
1777
1778impl RowStringKey {
1779    fn from_slice(values: &[String]) -> Self {
1780        RowStringKey(values.to_vec())
1781    }
1782}
1783
1784fn scalar_complex_tensor(re: f64, im: f64) -> crate::BuiltinResult<ComplexTensor> {
1785    ComplexTensor::new(vec![(re, im)], vec![1, 1])
1786        .map_err(|e| intersect_internal_error(format!("intersect: {e}")))
1787}
1788
1789fn tensor_to_complex_owned(name: &str, tensor: Tensor) -> crate::BuiltinResult<ComplexTensor> {
1790    let shape = tensor.shape.clone();
1791    let complex = tensor
1792        .into_numeric_storage()
1793        .map_err(|e| intersect_internal_error(format!("{name}: {e}")))?
1794        .materialize_f64()
1795        .into_iter()
1796        .map(|real| (real, 0.0))
1797        .collect();
1798    ComplexTensor::new(complex, shape).map_err(|e| intersect_internal_error(format!("{name}: {e}")))
1799}
1800
1801fn value_into_complex_tensor(value: Value) -> crate::BuiltinResult<ComplexTensor> {
1802    match value {
1803        Value::ComplexTensor(tensor) => Ok(tensor),
1804        Value::Complex(re, im) => scalar_complex_tensor(re, im),
1805        other => {
1806            let tensor = tensor::value_into_tensor_for("intersect", other)
1807                .map_err(|e| intersect_internal_error(e))?;
1808            tensor_to_complex_owned("intersect", tensor)
1809        }
1810    }
1811}
1812
1813fn compare_floating_rows<T: SetFloat>(a: &[T], b: &[T]) -> Ordering {
1814    for (lhs, rhs) in a.iter().zip(b.iter()) {
1815        let ord = lhs.compare(*rhs);
1816        if ord != Ordering::Equal {
1817            return ord;
1818        }
1819    }
1820    Ordering::Equal
1821}
1822
1823fn complex_is_nan<T: SetFloat>(value: (T, T)) -> bool {
1824    value.0.is_nan() || value.1.is_nan()
1825}
1826
1827fn compare_complex<T: SetFloat>(a: (T, T), b: (T, T)) -> Ordering {
1828    match (complex_is_nan(a), complex_is_nan(b)) {
1829        (true, true) => Ordering::Equal,
1830        (true, false) => Ordering::Greater,
1831        (false, true) => Ordering::Less,
1832        (false, false) => {
1833            let mag_a = a.0.hypot(a.1);
1834            let mag_b = b.0.hypot(b.1);
1835            let mag_cmp = mag_a.compare(mag_b);
1836            if mag_cmp != Ordering::Equal {
1837                return mag_cmp;
1838            }
1839            let re_cmp = a.0.compare(b.0);
1840            if re_cmp != Ordering::Equal {
1841                return re_cmp;
1842            }
1843            a.1.compare(b.1)
1844        }
1845    }
1846}
1847
1848fn compare_complex_rows<T: SetFloat>(a: &[(T, T)], b: &[(T, T)]) -> Ordering {
1849    for (lhs, rhs) in a.iter().zip(b.iter()) {
1850        let ord = compare_complex(*lhs, *rhs);
1851        if ord != Ordering::Equal {
1852            return ord;
1853        }
1854    }
1855    Ordering::Equal
1856}
1857
1858fn compare_char_rows(a: &[char], b: &[char]) -> Ordering {
1859    for (lhs, rhs) in a.iter().zip(b.iter()) {
1860        let ord = lhs.cmp(rhs);
1861        if ord != Ordering::Equal {
1862            return ord;
1863        }
1864    }
1865    Ordering::Equal
1866}
1867
1868fn compare_string_rows(a: &[String], b: &[String]) -> Ordering {
1869    for (lhs, rhs) in a.iter().zip(b.iter()) {
1870        let ord = lhs.cmp(rhs);
1871        if ord != Ordering::Equal {
1872            return ord;
1873        }
1874    }
1875    Ordering::Equal
1876}
1877
1878#[cfg(test)]
1879pub(crate) mod tests {
1880    use super::*;
1881    use crate::builtins::common::test_support;
1882    use runmat_accelerate_api::HostTensorView;
1883    use runmat_builtins::{ResolveContext, Type};
1884
1885    fn evaluate_sync(
1886        a: Value,
1887        b: Value,
1888        rest: &[Value],
1889    ) -> crate::BuiltinResult<IntersectEvaluation> {
1890        futures::executor::block_on(evaluate(a, b, rest))
1891    }
1892
1893    fn builtin_sync(a: Value, b: Value, rest: Vec<Value>) -> crate::BuiltinResult<Value> {
1894        futures::executor::block_on(intersect_builtin(a, b, rest))
1895    }
1896
1897    #[test]
1898    fn registered_builtin_restores_resident_outputs_and_rejects_excess_arity() {
1899        test_support::with_test_provider(|provider| {
1900            let left = Tensor::new_integer(IntegerStorage::I32(vec![7, 2, 9]), vec![3, 1]).unwrap();
1901            let right = Tensor::new_integer(IntegerStorage::I32(vec![2, 7]), vec![2, 1]).unwrap();
1902            let left =
1903                Value::GpuTensor(gpu_helpers::upload_tensor(provider, &left).expect("upload left"));
1904            let right = Value::GpuTensor(
1905                gpu_helpers::upload_tensor(provider, &right).expect("upload right"),
1906            );
1907
1908            {
1909                let _guard = crate::output_count::push_output_count(Some(3));
1910                let Value::OutputList(outputs) =
1911                    builtin_sync(left, right, Vec::new()).expect("resident intersect")
1912                else {
1913                    panic!("expected output list");
1914                };
1915                assert_eq!(outputs.len(), 3);
1916                assert!(outputs
1917                    .iter()
1918                    .all(|output| matches!(output, Value::GpuTensor(_))));
1919                assert_eq!(
1920                    test_support::gather(outputs[0].clone())
1921                        .expect("gather values")
1922                        .integer_storage(),
1923                    Some(&IntegerStorage::I32(vec![2, 7]))
1924                );
1925            }
1926
1927            let _guard = crate::output_count::push_output_count(Some(4));
1928            let err = builtin_sync(Value::Num(1.0), Value::Num(1.0), Vec::new())
1929                .expect_err("excess outputs must fail");
1930            assert_eq!(
1931                err.identifier(),
1932                INTERSECT_ERROR_INVALID_ARGUMENT.identifier
1933            );
1934        });
1935    }
1936
1937    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1938    #[test]
1939    fn intersect_numeric_sorted() {
1940        let a = Tensor::new(vec![5.0, 7.0, 5.0, 1.0], vec![4, 1]).unwrap();
1941        let b = Tensor::new(vec![7.0, 1.0, 3.0], vec![3, 1]).unwrap();
1942        let eval = intersect_numeric_elements(
1943            a,
1944            b,
1945            &IntersectOptions {
1946                rows: false,
1947                order: IntersectOrder::Sorted,
1948            },
1949        )
1950        .expect("intersect");
1951        let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
1952        assert_eq!(values.materialize_f64(), vec![1.0, 7.0]);
1953        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
1954        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
1955        assert_eq!(ia.materialize_f64(), vec![4.0, 2.0]);
1956        assert_eq!(ib.materialize_f64(), vec![2.0, 1.0]);
1957    }
1958
1959    #[test]
1960    fn intersect_preserves_native_single_elements_and_rows() {
1961        let a = Tensor::from_f32(vec![5.0, 7.0, 5.0, 1.0], vec![4, 1]).unwrap();
1962        let b = Tensor::from_f32(vec![7.0, 1.0, 3.0], vec![3, 1]).unwrap();
1963        let values = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[])
1964            .expect("single intersect")
1965            .values_value();
1966        let Value::Tensor(values) = values else {
1967            panic!("expected native single values");
1968        };
1969        assert_eq!(
1970            values.into_numeric_storage().unwrap(),
1971            NumericStorage::F32(vec![1.0, 7.0])
1972        );
1973
1974        let a = Tensor::from_f32(vec![1.0, 3.0, 1.0, 2.0, 4.0, 2.0], vec![3, 2]).unwrap();
1975        let b = Tensor::from_f32(vec![1.0, 5.0, 2.0, 6.0], vec![2, 2]).unwrap();
1976        let values = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[Value::from("rows")])
1977            .expect("single row intersect")
1978            .values_value();
1979        let Value::Tensor(values) = values else {
1980            panic!("expected native single rows");
1981        };
1982        assert_eq!(values.shape, vec![1, 2]);
1983        assert_eq!(
1984            values.into_numeric_storage().unwrap(),
1985            NumericStorage::F32(vec![1.0, 2.0])
1986        );
1987    }
1988
1989    #[test]
1990    fn intersect_preserves_native_complex_single_elements_and_rows() {
1991        let a =
1992            ComplexTensor::from_f32(vec![(1.0, 1.0), (0.0, 2.0), (1.0, -1.0)], vec![3, 1]).unwrap();
1993        let b = ComplexTensor::from_f32(vec![(0.0, 2.0), (4.0, 0.0)], vec![2, 1]).unwrap();
1994        let values = evaluate_sync(Value::ComplexTensor(a), Value::ComplexTensor(b), &[])
1995            .expect("complex single intersect")
1996            .values_value();
1997        let Value::ComplexTensor(values) = values else {
1998            panic!("expected native complex single value");
1999        };
2000        assert_eq!(values.as_f32_slice(), Some(&[(0.0, 2.0)][..]));
2001
2002        let a = ComplexTensor::from_f32(
2003            vec![
2004                (1.0, 0.0),
2005                (3.0, 0.0),
2006                (1.0, 0.0),
2007                (2.0, 1.0),
2008                (4.0, 1.0),
2009                (2.0, 1.0),
2010            ],
2011            vec![3, 2],
2012        )
2013        .unwrap();
2014        let b = ComplexTensor::from_f32(
2015            vec![(1.0, 0.0), (5.0, 0.0), (2.0, 1.0), (6.0, 1.0)],
2016            vec![2, 2],
2017        )
2018        .unwrap();
2019        let values = evaluate_sync(
2020            Value::ComplexTensor(a),
2021            Value::ComplexTensor(b),
2022            &[Value::from("rows")],
2023        )
2024        .expect("complex single row intersect")
2025        .values_value();
2026        let Value::ComplexTensor(values) = values else {
2027            panic!("expected native complex single rows");
2028        };
2029        assert_eq!(values.shape, vec![1, 2]);
2030        assert_eq!(values.as_f32_slice(), Some(&[(1.0, 0.0), (2.0, 1.0)][..]));
2031    }
2032
2033    #[test]
2034    fn intersect_preserves_exact_integer_elements_and_rows() {
2035        let a = Tensor::new_integer(
2036            runmat_value::IntegerStorage::U64(vec![u64::MAX, 0, 9_007_199_254_740_993]),
2037            vec![3, 1],
2038        )
2039        .expect("input");
2040        let b = Tensor::new_integer(
2041            runmat_value::IntegerStorage::U64(vec![0, u64::MAX]),
2042            vec![2, 1],
2043        )
2044        .expect("input");
2045        let (values, ia, ib) = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[])
2046            .expect("intersect")
2047            .into_triple();
2048        let Value::Tensor(values) = values else {
2049            panic!("exact values");
2050        };
2051        assert_eq!(
2052            values.integer_storage(),
2053            Some(&runmat_value::IntegerStorage::U64(vec![0, u64::MAX]))
2054        );
2055        let Value::Tensor(ia) = ia else {
2056            panic!("indices");
2057        };
2058        assert_eq!(ia.materialize_f64(), vec![2.0, 1.0]);
2059        let Value::Tensor(ib) = ib else {
2060            panic!("indices");
2061        };
2062        assert_eq!(ib.materialize_f64(), vec![1.0, 2.0]);
2063    }
2064
2065    #[test]
2066    fn intersect_rejects_mixed_nondouble_integer_classes() {
2067        let a = Tensor::new_integer(IntegerStorage::U16(vec![7, 2, 9, 7]), vec![4, 1]).unwrap();
2068        let b = Tensor::new_integer(IntegerStorage::I32(vec![2, 7]), vec![2, 1]).unwrap();
2069
2070        let error = evaluate_sync(Value::Tensor(a), Value::Tensor(b), &[])
2071            .expect_err("mixed integer classes must reject");
2072        assert_eq!(
2073            error.identifier(),
2074            INTERSECT_ERROR_NUMERIC_CLASS_MISMATCH.identifier
2075        );
2076    }
2077
2078    #[test]
2079    fn intersect_type_resolver_numeric() {
2080        assert_eq!(
2081            set_values_output_type(&[Type::tensor()], &ResolveContext::new(Vec::new())),
2082            Type::tensor()
2083        );
2084    }
2085
2086    #[test]
2087    fn intersect_type_resolver_string_array() {
2088        assert_eq!(
2089            set_values_output_type(
2090                &[Type::cell_of(Type::String)],
2091                &ResolveContext::new(Vec::new()),
2092            ),
2093            Type::cell_of(Type::String)
2094        );
2095    }
2096
2097    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2098    #[test]
2099    fn intersect_numeric_stable() {
2100        let a = Tensor::new(vec![4.0, 2.0, 4.0, 1.0, 3.0], vec![5, 1]).unwrap();
2101        let b = Tensor::new(vec![3.0, 4.0, 5.0, 1.0], vec![4, 1]).unwrap();
2102        let eval = intersect_numeric_elements(
2103            a,
2104            b,
2105            &IntersectOptions {
2106                rows: false,
2107                order: IntersectOrder::Stable,
2108            },
2109        )
2110        .expect("intersect");
2111        let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2112        assert_eq!(values.materialize_f64(), vec![4.0, 1.0, 3.0]);
2113        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2114        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2115        assert_eq!(ia.materialize_f64(), vec![1.0, 4.0, 5.0]);
2116        assert_eq!(ib.materialize_f64(), vec![2.0, 4.0, 1.0]);
2117    }
2118
2119    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2120    #[test]
2121    fn intersect_numeric_handles_nan() {
2122        let a = Tensor::new(vec![f64::NAN, 1.0, f64::NAN], vec![3, 1]).unwrap();
2123        let b = Tensor::new(vec![2.0, f64::NAN], vec![2, 1]).unwrap();
2124        let eval = intersect_numeric_elements(
2125            a,
2126            b,
2127            &IntersectOptions {
2128                rows: false,
2129                order: IntersectOrder::Sorted,
2130            },
2131        )
2132        .expect("intersect");
2133        let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2134        assert_eq!(values.materialize_f64().len(), 1);
2135        assert!(values.materialize_f64()[0].is_nan());
2136        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2137        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2138        assert_eq!(ia.materialize_f64(), vec![1.0]);
2139        assert_eq!(ib.materialize_f64(), vec![2.0]);
2140    }
2141
2142    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2143    #[test]
2144    fn intersect_complex_with_real_inputs() {
2145        let complex =
2146            ComplexTensor::new(vec![(1.0, 0.0), (2.0, 0.0), (3.0, 1.0)], vec![3, 1]).unwrap();
2147        let real = Tensor::new(vec![2.0, 4.0, 1.0], vec![3, 1]).unwrap();
2148        let real_complex = tensor_to_complex_owned("intersect", real).unwrap();
2149        let eval = intersect_complex(
2150            complex,
2151            real_complex,
2152            &IntersectOptions {
2153                rows: false,
2154                order: IntersectOrder::Sorted,
2155            },
2156        )
2157        .expect("intersect complex");
2158        match eval.values_value() {
2159            Value::ComplexTensor(t) => {
2160                assert_eq!(t.materialize_f64(), vec![(1.0, 0.0), (2.0, 0.0)]);
2161            }
2162            other => panic!("expected complex tensor, got {other:?}"),
2163        }
2164        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2165        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2166        assert_eq!(ia.materialize_f64(), vec![1.0, 2.0]);
2167        assert_eq!(ib.materialize_f64(), vec![3.0, 1.0]);
2168    }
2169
2170    #[test]
2171    fn intersect_complex_real_alignment_reads_typed_integer_storage_exactly() {
2172        let real =
2173            Tensor::new_integer(IntegerStorage::I64(vec![i64::MIN, -7, 3]), vec![3, 1]).unwrap();
2174        let complex = ComplexTensor::new(vec![(-7.0, 0.0), (4.0, 0.0)], vec![2, 1]).unwrap();
2175
2176        let eval = evaluate_sync(Value::Tensor(real), Value::ComplexTensor(complex), &[])
2177            .expect("intersect");
2178        let Value::Complex(re, im) = eval.values_value() else {
2179            panic!("expected complex scalar");
2180        };
2181        assert_eq!(re, -7.0);
2182        assert_eq!(im, 0.0);
2183        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2184        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2185        assert_eq!(ia.materialize_f64(), vec![2.0]);
2186        assert_eq!(ib.materialize_f64(), vec![1.0]);
2187    }
2188
2189    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2190    #[test]
2191    fn intersect_numeric_rows_default() {
2192        let a = Tensor::new(vec![1.0, 3.0, 1.0, 2.0, 4.0, 2.0], vec![3, 2]).unwrap();
2193        let b = Tensor::new(vec![1.0, 5.0, 2.0, 6.0], vec![2, 2]).unwrap();
2194        let eval = intersect_numeric_rows(
2195            a,
2196            b,
2197            &IntersectOptions {
2198                rows: true,
2199                order: IntersectOrder::Sorted,
2200            },
2201        )
2202        .expect("intersect rows");
2203        let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2204        assert_eq!(values.shape, vec![1, 2]);
2205        assert_eq!(values.materialize_f64(), vec![1.0, 2.0]);
2206        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2207        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2208        assert_eq!(ia.materialize_f64(), vec![1.0]);
2209        assert_eq!(ib.materialize_f64(), vec![1.0]);
2210    }
2211
2212    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2213    #[test]
2214    fn intersect_char_elements_basic() {
2215        let a = CharArray::new("cab".chars().collect(), 1, 3).unwrap();
2216        let b = CharArray::new("bcd".chars().collect(), 1, 3).unwrap();
2217        assert_eq!(find_char_index(&b, 'b'), Some(0));
2218        assert_eq!(find_char_index(&b, 'c'), Some(1));
2219        let b_for_eval = CharArray::new("bcd".chars().collect(), 1, 3).unwrap();
2220        let eval = intersect_char_elements(
2221            a,
2222            b_for_eval,
2223            &IntersectOptions {
2224                rows: false,
2225                order: IntersectOrder::Sorted,
2226            },
2227        )
2228        .expect("intersect char");
2229        match eval.values_value() {
2230            Value::CharArray(arr) => {
2231                assert_eq!(arr.rows, 2);
2232                assert_eq!(arr.cols, 1);
2233                assert_eq!(arr.data, vec!['b', 'c']);
2234            }
2235            other => panic!("expected char array, got {other:?}"),
2236        }
2237        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2238        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2239        assert_eq!(ia.materialize_f64(), vec![3.0, 1.0]);
2240        assert_eq!(ib.materialize_f64(), vec![1.0, 2.0]);
2241    }
2242
2243    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2244    #[test]
2245    fn intersect_string_elements_stable() {
2246        let a = StringArray::new(
2247            vec!["apple".into(), "orange".into(), "pear".into()],
2248            vec![3, 1],
2249        )
2250        .unwrap();
2251        let b = StringArray::new(
2252            vec!["pear".into(), "grape".into(), "orange".into()],
2253            vec![3, 1],
2254        )
2255        .unwrap();
2256        let eval = intersect_string_elements(
2257            a,
2258            b,
2259            &IntersectOptions {
2260                rows: false,
2261                order: IntersectOrder::Stable,
2262            },
2263        )
2264        .expect("intersect string");
2265        match eval.values_value() {
2266            Value::StringArray(arr) => {
2267                assert_eq!(arr.shape, vec![2, 1]);
2268                assert_eq!(arr.data, vec!["orange".to_string(), "pear".to_string()]);
2269            }
2270            other => panic!("expected string array, got {other:?}"),
2271        }
2272        let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2273        let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2274        assert_eq!(ia.materialize_f64(), vec![2.0, 3.0]);
2275        assert_eq!(ib.materialize_f64(), vec![3.0, 1.0]);
2276    }
2277
2278    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2279    #[test]
2280    fn intersect_rejects_legacy_option() {
2281        let tensor = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
2282        let err = evaluate_sync(
2283            Value::Tensor(tensor.clone()),
2284            Value::Tensor(tensor),
2285            &[Value::from("legacy")],
2286        )
2287        .unwrap_err();
2288        assert_eq!(
2289            err.identifier(),
2290            INTERSECT_ERROR_LEGACY_OPTION_UNSUPPORTED.identifier
2291        );
2292    }
2293
2294    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2295    #[test]
2296    fn intersect_rejects_conflicting_order_options() {
2297        let tensor = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2298        let err = evaluate_sync(
2299            Value::Tensor(tensor.clone()),
2300            Value::Tensor(tensor),
2301            &[Value::from("stable"), Value::from("sorted")],
2302        )
2303        .unwrap_err();
2304        assert_eq!(
2305            err.identifier(),
2306            INTERSECT_ERROR_CONFLICTING_ORDER_OPTIONS.identifier
2307        );
2308    }
2309
2310    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2311    #[test]
2312    fn intersect_rejects_unknown_option() {
2313        let tensor = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2314        let err = evaluate_sync(
2315            Value::Tensor(tensor.clone()),
2316            Value::Tensor(tensor),
2317            &[Value::from("bogus")],
2318        )
2319        .unwrap_err();
2320        assert_eq!(err.identifier(), INTERSECT_ERROR_UNKNOWN_OPTION.identifier);
2321    }
2322
2323    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2324    #[test]
2325    fn intersect_rows_dimension_mismatch() {
2326        let a = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
2327        let b = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
2328        let err = intersect_numeric_rows(
2329            a,
2330            b,
2331            &IntersectOptions {
2332                rows: true,
2333                order: IntersectOrder::Sorted,
2334            },
2335        )
2336        .unwrap_err();
2337        assert_eq!(
2338            err.identifier(),
2339            INTERSECT_ERROR_ROWS_COLUMN_MISMATCH.identifier
2340        );
2341    }
2342
2343    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2344    #[test]
2345    fn intersect_mixed_types_error() {
2346        let a = Tensor::new(vec![1.0, 2.0], vec![2, 1]).unwrap();
2347        let b = CharArray::new(vec!['a', 'b'], 1, 2).unwrap();
2348        let err = intersect_host(
2349            Value::Tensor(a),
2350            Value::CharArray(b),
2351            &IntersectOptions {
2352                rows: false,
2353                order: IntersectOrder::Sorted,
2354            },
2355        )
2356        .unwrap_err();
2357        assert_eq!(
2358            err.identifier(),
2359            INTERSECT_ERROR_UNSUPPORTED_INPUT_TYPE.identifier
2360        );
2361    }
2362
2363    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2364    #[test]
2365    fn intersect_gpu_roundtrip() {
2366        test_support::with_test_provider(|provider| {
2367            let a = Tensor::new(vec![4.0, 1.0, 2.0, 1.0], vec![4, 1]).unwrap();
2368            let b = Tensor::new(vec![2.0, 5.0, 1.0], vec![3, 1]).unwrap();
2369            let view_a = HostTensorView {
2370                data: &a.materialize_f64(),
2371                shape: &a.shape,
2372            };
2373            let view_b = HostTensorView {
2374                data: &b.materialize_f64(),
2375                shape: &b.shape,
2376            };
2377            let handle_a = provider.upload(&view_a).expect("upload A");
2378            let handle_b = provider.upload(&view_b).expect("upload B");
2379            let eval = evaluate_sync(Value::GpuTensor(handle_a), Value::GpuTensor(handle_b), &[])
2380                .expect("intersect");
2381            let values = tensor::value_into_tensor_for("intersect", eval.values_value()).unwrap();
2382            assert_eq!(values.materialize_f64(), vec![1.0, 2.0]);
2383            let ia = tensor::value_into_tensor_for("intersect", eval.ia_value()).unwrap();
2384            let ib = tensor::value_into_tensor_for("intersect", eval.ib_value()).unwrap();
2385            assert_eq!(ia.materialize_f64(), vec![2.0, 3.0]);
2386            assert_eq!(ib.materialize_f64(), vec![3.0, 1.0]);
2387        });
2388    }
2389
2390    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2391    #[test]
2392    fn intersect_two_outputs_from_evaluate() {
2393        let a = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
2394        let b = Tensor::new(vec![3.0, 1.0], vec![2, 1]).unwrap();
2395        let eval = intersect_numeric_elements(
2396            a,
2397            b,
2398            &IntersectOptions {
2399                rows: false,
2400                order: IntersectOrder::Sorted,
2401            },
2402        )
2403        .unwrap();
2404        let (_c, ia) = eval.clone().into_pair();
2405        let ia_tensor = tensor::value_into_tensor_for("intersect", ia).unwrap();
2406        assert_eq!(ia_tensor.materialize_f64(), vec![1.0, 3.0]);
2407        let (_c, ia2, ib2) = eval.into_triple();
2408        let ia_tensor2 = tensor::value_into_tensor_for("intersect", ia2).unwrap();
2409        let ib_tensor2 = tensor::value_into_tensor_for("intersect", ib2).unwrap();
2410        assert_eq!(ia_tensor2.materialize_f64(), vec![1.0, 3.0]);
2411        assert_eq!(ib_tensor2.materialize_f64(), vec![2.0, 1.0]);
2412    }
2413
2414    #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
2415    #[test]
2416    #[cfg(feature = "wgpu")]
2417    fn intersect_wgpu_matches_cpu() {
2418        let _ = runmat_accelerate::backend::wgpu::provider::register_wgpu_provider(
2419            runmat_accelerate::backend::wgpu::provider::WgpuProviderOptions::default(),
2420        );
2421        let a = Tensor::new(vec![4.0, 1.0, 2.0, 3.0], vec![4, 1]).unwrap();
2422        let b = Tensor::new(vec![2.0, 6.0, 3.0], vec![3, 1]).unwrap();
2423
2424        let cpu_eval = intersect_numeric_elements(
2425            a.clone(),
2426            b.clone(),
2427            &IntersectOptions {
2428                rows: false,
2429                order: IntersectOrder::Sorted,
2430            },
2431        )
2432        .unwrap();
2433        let cpu_values =
2434            tensor::value_into_tensor_for("intersect", cpu_eval.values_value()).unwrap();
2435        let cpu_ia = tensor::value_into_tensor_for("intersect", cpu_eval.ia_value()).unwrap();
2436        let cpu_ib = tensor::value_into_tensor_for("intersect", cpu_eval.ib_value()).unwrap();
2437
2438        let provider = runmat_accelerate_api::provider().expect("provider");
2439        let view_a = HostTensorView {
2440            data: &a.materialize_f64(),
2441            shape: &a.shape,
2442        };
2443        let view_b = HostTensorView {
2444            data: &b.materialize_f64(),
2445            shape: &b.shape,
2446        };
2447        let handle_a = provider.upload(&view_a).expect("upload A");
2448        let handle_b = provider.upload(&view_b).expect("upload B");
2449        let gpu_eval = evaluate_sync(Value::GpuTensor(handle_a), Value::GpuTensor(handle_b), &[])
2450            .expect("intersect");
2451        let gpu_values =
2452            tensor::value_into_tensor_for("intersect", gpu_eval.values_value()).unwrap();
2453        let gpu_ia = tensor::value_into_tensor_for("intersect", gpu_eval.ia_value()).unwrap();
2454        let gpu_ib = tensor::value_into_tensor_for("intersect", gpu_eval.ib_value()).unwrap();
2455
2456        assert_eq!(gpu_values.materialize_f64(), cpu_values.materialize_f64());
2457        assert_eq!(gpu_ia.materialize_f64(), cpu_ia.materialize_f64());
2458        assert_eq!(gpu_ib.materialize_f64(), cpu_ib.materialize_f64());
2459    }
2460}