1use runmat_builtins::{
4 BuiltinCompletionPolicy, BuiltinDescriptor, BuiltinErrorDescriptor, BuiltinExtensionDescriptor,
5 BuiltinExtensionMode, BuiltinIntegerBackendRule, BuiltinIntegerCapabilityDescriptor,
6 BuiltinIntegerComputationDomain, BuiltinIntegerInputAvailability,
7 BuiltinIntegerInputCapability, BuiltinIntegerOutputClassRule, BuiltinIntegerOverflowRule,
8 BuiltinIntegerOverloadKind, BuiltinIntegerScalarDoubleRule, BuiltinOutputMode,
9 BuiltinParamArity, BuiltinParamDescriptor, BuiltinParamType, BuiltinSignatureDescriptor,
10};
11use runmat_macros::runtime_builtin;
12use runmat_value::{ComplexTensor, NumericScalar, Tensor, Value};
13
14use crate::builtins::common::format::{flatten_arguments, format_variadic};
15use crate::builtins::common::gpu_helpers;
16use crate::builtins::common::spec::{
17 BroadcastSemantics, BuiltinFusionSpec, BuiltinGpuSpec, ConstantStrategy, GpuOpKind,
18 ReductionNaN, ResidencyPolicy, ShapeRequirements,
19};
20use crate::builtins::diagnostics::type_resolvers::assert_type;
21use crate::{build_runtime_error, RuntimeError};
22
23const BUILTIN_NAME: &str = "assert";
24
25const ASSERT_OUTPUTS: [BuiltinParamDescriptor; 0] = [];
26
27pub(crate) const ASSERT_COMPLEX_CONDITION_EXTENSION: BuiltinExtensionDescriptor =
28 BuiltinExtensionDescriptor {
29 id: "assert-complex-condition",
30 mode: BuiltinExtensionMode::RunMatOnly,
31 description: "assert with a complex condition is a RunMat extension",
32 error_identifier: Some("RunMat:compatibility:AssertComplexConditionExtension"),
33 };
34
35pub(crate) const ASSERT_UNQUALIFIED_IDENTIFIER_EXTENSION: BuiltinExtensionDescriptor =
36 BuiltinExtensionDescriptor {
37 id: "assert-unqualified-identifier",
38 mode: BuiltinExtensionMode::RunMatOnly,
39 description: "assert with an unqualified custom error identifier is a RunMat extension",
40 error_identifier: Some("RunMat:compatibility:AssertUnqualifiedIdentifierExtension"),
41 };
42
43pub const ASSERT_EXTENSIONS: [BuiltinExtensionDescriptor; 2] = [
44 ASSERT_COMPLEX_CONDITION_EXTENSION,
45 ASSERT_UNQUALIFIED_IDENTIFIER_EXTENSION,
46];
47
48const ASSERT_INTEGER_CONDITION_INPUTS: [BuiltinIntegerInputCapability; 1] =
49 [BuiltinIntegerInputCapability {
50 name: "cond",
51 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
52 availability: BuiltinIntegerInputAvailability::Documented,
53 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
54 notes: "Every real integer scalar or array is convertible to logical; the condition passes only when it is nonempty and every element is nonzero.",
55 }];
56
57const ASSERT_COMPLEX_INTEGER_CONDITION_INPUTS: [BuiltinIntegerInputCapability; 1] =
58 [BuiltinIntegerInputCapability {
59 name: "cond",
60 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
61 availability: BuiltinIntegerInputAvailability::RunMatOnly,
62 scalar_double: BuiltinIntegerScalarDoubleRule::NotApplicable,
63 notes: "RunMat mode additionally accepts paired complex-integer storage and tests each element for a nonzero real or imaginary component.",
64 }];
65
66const ASSERT_INTEGER_FORMAT_INPUTS: [BuiltinIntegerInputCapability; 1] =
67 [BuiltinIntegerInputCapability {
68 name: "A",
69 classes: &crate::builtins::common::integer_capability::ALL_INTEGER_CLASSES,
70 availability: BuiltinIntegerInputAvailability::Documented,
71 scalar_double: BuiltinIntegerScalarDoubleRule::Allowed,
72 notes: "Formatting replacement values accept numeric scalars and preserve exact integer values through integer and string conversion specifiers.",
73 }];
74
75pub const ASSERT_INTEGER_CAPABILITIES: [BuiltinIntegerCapabilityDescriptor; 3] = [
76 BuiltinIntegerCapabilityDescriptor {
77 form: "assert(integer_cond, ...)",
78 inputs: &ASSERT_INTEGER_CONDITION_INPUTS,
79 computation_domain: BuiltinIntegerComputationDomain::Predicate,
80 output_class: BuiltinIntegerOutputClassRule::NotApplicable,
81 overflow: BuiltinIntegerOverflowRule::NotApplicable,
82 backend: BuiltinIntegerBackendRule::GatherFallback,
83 overload: BuiltinIntegerOverloadKind::Multiple,
84 notes: "Truth testing reads exact authoritative integer storage. The public builtin has no output; resident conditions gather to the host because assert accepts gpuArray input but does not execute on the GPU.",
85 },
86 BuiltinIntegerCapabilityDescriptor {
87 form: "assert(complex_integer_cond, ...)",
88 inputs: &ASSERT_COMPLEX_INTEGER_CONDITION_INPUTS,
89 computation_domain: BuiltinIntegerComputationDomain::Predicate,
90 output_class: BuiltinIntegerOutputClassRule::NotApplicable,
91 overflow: BuiltinIntegerOverflowRule::NotApplicable,
92 backend: BuiltinIntegerBackendRule::GatherFallback,
93 overload: BuiltinIntegerOverloadKind::Multiple,
94 notes: "Public logical conversion rejects complex numeric values. RunMat mode retains the pre-existing nonzero-component predicate without floating materialization.",
95 },
96 BuiltinIntegerCapabilityDescriptor {
97 form: "assert(cond, msg, integer_A...)",
98 inputs: &ASSERT_INTEGER_FORMAT_INPUTS,
99 computation_domain: BuiltinIntegerComputationDomain::Structural,
100 output_class: BuiltinIntegerOutputClassRule::NotApplicable,
101 overflow: BuiltinIntegerOverflowRule::NotApplicable,
102 backend: BuiltinIntegerBackendRule::GatherFallback,
103 overload: BuiltinIntegerOverloadKind::StructuralParameter,
104 notes: "Each documented replacement A is a character vector, string scalar, or numeric scalar. Integer scalars remain exact during host formatting, including after resident gather.",
105 },
106];
107
108const ASSERT_INPUTS_CONDITION: [BuiltinParamDescriptor; 1] = [BuiltinParamDescriptor {
109 name: "condition",
110 ty: BuiltinParamType::Any,
111 arity: BuiltinParamArity::Required,
112 default: None,
113 description: "Logical or logically convertible condition that must evaluate to true.",
114}];
115
116const ASSERT_INPUTS_MESSAGE: [BuiltinParamDescriptor; 2] = [
117 BuiltinParamDescriptor {
118 name: "condition",
119 ty: BuiltinParamType::Any,
120 arity: BuiltinParamArity::Required,
121 default: None,
122 description: "Logical or logically convertible condition that must evaluate to true.",
123 },
124 BuiltinParamDescriptor {
125 name: "message",
126 ty: BuiltinParamType::StringScalar,
127 arity: BuiltinParamArity::Required,
128 default: Some("\"Assertion failed.\""),
129 description: "Failure message text.",
130 },
131];
132
133const ASSERT_INPUTS_MESSAGE_VARIADIC: [BuiltinParamDescriptor; 3] = [
134 BuiltinParamDescriptor {
135 name: "condition",
136 ty: BuiltinParamType::Any,
137 arity: BuiltinParamArity::Required,
138 default: None,
139 description: "Logical or logically convertible condition that must evaluate to true.",
140 },
141 BuiltinParamDescriptor {
142 name: "message",
143 ty: BuiltinParamType::StringScalar,
144 arity: BuiltinParamArity::Required,
145 default: Some("\"Assertion failed.\""),
146 description: "Failure message template text.",
147 },
148 BuiltinParamDescriptor {
149 name: "A",
150 ty: BuiltinParamType::Any,
151 arity: BuiltinParamArity::Variadic,
152 default: None,
153 description: "Formatting values for the message template.",
154 },
155];
156
157const ASSERT_INPUTS_IDENTIFIER_MESSAGE: [BuiltinParamDescriptor; 3] = [
158 BuiltinParamDescriptor {
159 name: "condition",
160 ty: BuiltinParamType::Any,
161 arity: BuiltinParamArity::Required,
162 default: None,
163 description: "Logical or logically convertible condition that must evaluate to true.",
164 },
165 BuiltinParamDescriptor {
166 name: "message_id",
167 ty: BuiltinParamType::StringScalar,
168 arity: BuiltinParamArity::Required,
169 default: Some("\"RunMat:assertion:failed\""),
170 description: "Message identifier.",
171 },
172 BuiltinParamDescriptor {
173 name: "message",
174 ty: BuiltinParamType::StringScalar,
175 arity: BuiltinParamArity::Required,
176 default: Some("\"Assertion failed.\""),
177 description: "Failure message text.",
178 },
179];
180
181const ASSERT_INPUTS_IDENTIFIER_MESSAGE_VARIADIC: [BuiltinParamDescriptor; 4] = [
182 BuiltinParamDescriptor {
183 name: "condition",
184 ty: BuiltinParamType::Any,
185 arity: BuiltinParamArity::Required,
186 default: None,
187 description: "Logical or logically convertible condition that must evaluate to true.",
188 },
189 BuiltinParamDescriptor {
190 name: "message_id",
191 ty: BuiltinParamType::StringScalar,
192 arity: BuiltinParamArity::Required,
193 default: Some("\"RunMat:assertion:failed\""),
194 description: "Message identifier.",
195 },
196 BuiltinParamDescriptor {
197 name: "message",
198 ty: BuiltinParamType::StringScalar,
199 arity: BuiltinParamArity::Required,
200 default: Some("\"Assertion failed.\""),
201 description: "Failure message template text.",
202 },
203 BuiltinParamDescriptor {
204 name: "A",
205 ty: BuiltinParamType::Any,
206 arity: BuiltinParamArity::Variadic,
207 default: None,
208 description: "Formatting values for the message template.",
209 },
210];
211
212const ASSERT_SIGNATURES: [BuiltinSignatureDescriptor; 5] = [
213 BuiltinSignatureDescriptor {
214 label: "assert(condition)",
215 inputs: &ASSERT_INPUTS_CONDITION,
216 outputs: &ASSERT_OUTPUTS,
217 },
218 BuiltinSignatureDescriptor {
219 label: "assert(condition, message)",
220 inputs: &ASSERT_INPUTS_MESSAGE,
221 outputs: &ASSERT_OUTPUTS,
222 },
223 BuiltinSignatureDescriptor {
224 label: "assert(condition, message, A...)",
225 inputs: &ASSERT_INPUTS_MESSAGE_VARIADIC,
226 outputs: &ASSERT_OUTPUTS,
227 },
228 BuiltinSignatureDescriptor {
229 label: "assert(condition, message_id, message)",
230 inputs: &ASSERT_INPUTS_IDENTIFIER_MESSAGE,
231 outputs: &ASSERT_OUTPUTS,
232 },
233 BuiltinSignatureDescriptor {
234 label: "assert(condition, message_id, message, A...)",
235 inputs: &ASSERT_INPUTS_IDENTIFIER_MESSAGE_VARIADIC,
236 outputs: &ASSERT_OUTPUTS,
237 },
238];
239
240const ASSERT_ERROR_ASSERTION_FAILED: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
241 code: "RM.ASSERT.ASSERTION_FAILED",
242 identifier: Some("RunMat:assertion:failed"),
243 when: "Condition evaluates to false and no custom identifier/message override is provided.",
244 message: "Assertion failed.",
245};
246
247const ASSERT_ERROR_INVALID_CONDITION: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
248 code: "RM.ASSERT.INVALID_CONDITION",
249 identifier: Some("RunMat:assertion:invalidCondition"),
250 when: "First argument is not logical or convertible to a logical condition.",
251 message: "assert: first input must be logical or convertible to logical.",
252};
253
254const ASSERT_ERROR_INVALID_INPUT: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
255 code: "RM.ASSERT.INVALID_INPUT",
256 identifier: Some("RunMat:assertion:invalidInput"),
257 when: "Message identifier/message text or formatting payload is invalid.",
258 message: "assert: invalid input argument",
259};
260
261const ASSERT_ERROR_NOT_ENOUGH_INPUTS: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
262 code: "RM.ASSERT.NOT_ENOUGH_INPUTS",
263 identifier: Some("RunMat:minrhs"),
264 when: "No condition argument is provided.",
265 message: "Not enough input arguments.",
266};
267
268const ASSERT_ERROR_TOO_MANY_OUTPUTS: BuiltinErrorDescriptor = BuiltinErrorDescriptor {
269 code: "RM.ASSERT.TOO_MANY_OUTPUTS",
270 identifier: Some("RunMat:assertion:TooManyOutputs"),
271 when: "One or more public outputs are requested from assert.",
272 message: "assert: too many output arguments",
273};
274
275const ASSERT_ERRORS: [BuiltinErrorDescriptor; 5] = [
276 ASSERT_ERROR_ASSERTION_FAILED,
277 ASSERT_ERROR_INVALID_CONDITION,
278 ASSERT_ERROR_INVALID_INPUT,
279 ASSERT_ERROR_NOT_ENOUGH_INPUTS,
280 ASSERT_ERROR_TOO_MANY_OUTPUTS,
281];
282
283pub const ASSERT_DESCRIPTOR: BuiltinDescriptor = BuiltinDescriptor {
284 signatures: &ASSERT_SIGNATURES,
285 output_mode: BuiltinOutputMode::Fixed,
286 completion_policy: BuiltinCompletionPolicy::Public,
287 errors: &ASSERT_ERRORS,
288};
289
290#[runmat_macros::register_gpu_spec(builtin_path = "crate::builtins::diagnostics::assert")]
291pub const GPU_SPEC: BuiltinGpuSpec = BuiltinGpuSpec {
292 name: "assert",
293 op_kind: GpuOpKind::Custom("control"),
294 supported_precisions: &[],
295 broadcast: BroadcastSemantics::None,
296 provider_hooks: &[],
297 constant_strategy: ConstantStrategy::InlineLiteral,
298 residency: ResidencyPolicy::GatherImmediately,
299 nan_mode: ReductionNaN::Include,
300 two_pass_threshold: None,
301 workgroup_size: None,
302 accepts_nan_mode: false,
303 notes: "Control-flow builtin; GPU tensors are gathered to host memory before evaluation.",
304};
305
306#[runmat_macros::register_fusion_spec(builtin_path = "crate::builtins::diagnostics::assert")]
307pub const FUSION_SPEC: BuiltinFusionSpec = BuiltinFusionSpec {
308 name: "assert",
309 shape: ShapeRequirements::Any,
310 constant_strategy: ConstantStrategy::InlineLiteral,
311 elementwise: None,
312 reduction: None,
313 emits_nan: false,
314 notes: "Control-flow builtin with no fusion support.",
315};
316
317fn assert_error(error: &'static BuiltinErrorDescriptor) -> RuntimeError {
318 assert_error_with_message(error.message, error)
319}
320
321fn assert_default_identifier() -> &'static str {
322 ASSERT_ERROR_ASSERTION_FAILED
323 .identifier
324 .expect("assert default identifier must be defined")
325}
326
327fn assert_default_message() -> &'static str {
328 ASSERT_ERROR_ASSERTION_FAILED.message
329}
330
331fn assert_error_with_message(
332 message: impl Into<String>,
333 error: &'static BuiltinErrorDescriptor,
334) -> RuntimeError {
335 let mut builder = build_runtime_error(message).with_builtin(BUILTIN_NAME);
336 if let Some(identifier) = error.identifier {
337 builder = builder.with_identifier(normalize_identifier(identifier));
338 }
339 builder.build()
340}
341
342fn assert_flow(identifier: &str, message: impl Into<String>) -> RuntimeError {
343 build_runtime_error(message)
344 .with_builtin(BUILTIN_NAME)
345 .with_identifier(normalize_identifier(identifier))
346 .build()
347}
348
349fn remap_assert_flow<F>(
350 err: RuntimeError,
351 error: &'static BuiltinErrorDescriptor,
352 message: F,
353) -> RuntimeError
354where
355 F: FnOnce(&crate::RuntimeError) -> String,
356{
357 let mut builder = build_runtime_error(message(&err))
358 .with_builtin(BUILTIN_NAME)
359 .with_source(err);
360 if let Some(identifier) = error.identifier {
361 builder = builder.with_identifier(normalize_identifier(identifier));
362 }
363 builder.build()
364}
365
366#[runtime_builtin(
367 name = "assert",
368 category = "diagnostics",
369 summary = "Throw an error when a condition is false, matching MATLAB assert semantics.",
370 keywords = "assert,diagnostics,validation,error",
371 accel = "metadata",
372 type_resolver(assert_type),
373 descriptor(crate::builtins::diagnostics::assert::ASSERT_DESCRIPTOR),
374 extensions(crate::builtins::diagnostics::assert::ASSERT_EXTENSIONS),
375 integer_capabilities(crate::builtins::diagnostics::assert::ASSERT_INTEGER_CAPABILITIES),
376 builtin_path = "crate::builtins::diagnostics::assert"
377)]
378async fn assert_builtin(args: Vec<Value>) -> crate::BuiltinResult<Value> {
379 if matches!(crate::output_count::current_output_count(), Some(count) if count > 0) {
380 return Err(assert_error(&ASSERT_ERROR_TOO_MANY_OUTPUTS));
381 }
382 if args.is_empty() {
383 return Err(assert_error(&ASSERT_ERROR_NOT_ENOUGH_INPUTS));
384 }
385
386 let mut iter = args.into_iter();
387 let condition_raw = iter.next().expect("checked length above");
388 let rest: Vec<Value> = iter.collect();
389
390 let condition = normalize_condition_value(condition_raw).await?;
391 match evaluate_condition(condition)? {
392 ConditionOutcome::Pass => Ok(Value::Num(0.0)),
393 ConditionOutcome::Fail => {
394 let payload = failure_payload(&rest).await?;
395 Err(assert_flow(&payload.identifier, payload.message))
396 }
397 }
398}
399
400async fn normalize_condition_value(condition: Value) -> crate::BuiltinResult<Value> {
401 match condition {
402 Value::GpuTensor(handle) => {
403 let gpu_value = Value::GpuTensor(handle);
404 gpu_helpers::gather_value_async(&gpu_value)
405 .await
406 .map_err(|flow| {
407 remap_assert_flow(flow, &ASSERT_ERROR_INVALID_INPUT, |err| {
408 format!("assert: {}", err.message())
409 })
410 })
411 }
412 other => Ok(other),
413 }
414}
415
416#[derive(Copy, Clone, Debug, PartialEq, Eq)]
417enum ConditionOutcome {
418 Pass,
419 Fail,
420}
421
422fn evaluate_condition(value: Value) -> crate::BuiltinResult<ConditionOutcome> {
423 match value {
424 Value::Bool(flag) => Ok(if flag {
425 ConditionOutcome::Pass
426 } else {
427 ConditionOutcome::Fail
428 }),
429 Value::Int(int_value) => {
430 if !int_value.is_zero() {
431 Ok(ConditionOutcome::Pass)
432 } else {
433 Ok(ConditionOutcome::Fail)
434 }
435 }
436 Value::Num(num) => {
437 if num.is_nan() {
438 Err(assert_error(&ASSERT_ERROR_INVALID_CONDITION))
439 } else if num == 0.0 {
440 Ok(ConditionOutcome::Fail)
441 } else {
442 Ok(ConditionOutcome::Pass)
443 }
444 }
445 Value::Complex(re, im) => {
446 crate::compatibility::ensure_builtin_extension_enabled(
447 &ASSERT_COMPLEX_CONDITION_EXTENSION,
448 BUILTIN_NAME,
449 )?;
450 if complex_element_passes(re, im) {
451 Ok(ConditionOutcome::Pass)
452 } else {
453 Ok(ConditionOutcome::Fail)
454 }
455 }
456 Value::LogicalArray(array) => {
457 if !array.data.is_empty() && array.data.iter().all(|&bit| bit != 0) {
458 Ok(ConditionOutcome::Pass)
459 } else {
460 Ok(ConditionOutcome::Fail)
461 }
462 }
463 Value::Tensor(tensor) => evaluate_tensor_condition(&tensor),
464 Value::ComplexTensor(tensor) => {
465 crate::compatibility::ensure_builtin_extension_enabled(
466 &ASSERT_COMPLEX_CONDITION_EXTENSION,
467 BUILTIN_NAME,
468 )?;
469 evaluate_complex_tensor(&tensor)
470 }
471 Value::CharArray(chars) => {
472 if !chars.data.is_empty() && chars.data.iter().all(|character| *character != '\0') {
473 Ok(ConditionOutcome::Pass)
474 } else {
475 Ok(ConditionOutcome::Fail)
476 }
477 }
478 Value::GpuTensor(_) => {
479 unreachable!("gpu tensors are gathered in normalize_condition_value")
480 }
481 _ => Err(assert_error(&ASSERT_ERROR_INVALID_CONDITION)),
482 }
483}
484
485fn evaluate_tensor_condition(tensor: &Tensor) -> crate::BuiltinResult<ConditionOutcome> {
486 if tensor.is_empty() {
487 return Ok(ConditionOutcome::Fail);
488 }
489 for index in 0..tensor.len() {
490 match tensor
491 .numeric_value_at(index)
492 .ok_or_else(|| assert_error(&ASSERT_ERROR_INVALID_CONDITION))?
493 {
494 NumericScalar::F64(value) => {
495 if value.is_nan() {
496 return Err(assert_error(&ASSERT_ERROR_INVALID_CONDITION));
497 }
498 if value == 0.0 {
499 return Ok(ConditionOutcome::Fail);
500 }
501 }
502 NumericScalar::F32(value) => {
503 if value.is_nan() {
504 return Err(assert_error(&ASSERT_ERROR_INVALID_CONDITION));
505 }
506 if value == 0.0 {
507 return Ok(ConditionOutcome::Fail);
508 }
509 }
510 value => {
511 if value
512 .into_int_value()
513 .is_none_or(|integer| integer.is_zero())
514 {
515 return Ok(ConditionOutcome::Fail);
516 }
517 }
518 }
519 }
520 Ok(ConditionOutcome::Pass)
521}
522
523fn evaluate_complex_tensor(tensor: &ComplexTensor) -> crate::BuiltinResult<ConditionOutcome> {
524 if let Some(storage) = tensor.integer_storage() {
525 if storage.is_empty() {
526 return Ok(ConditionOutcome::Fail);
527 }
528 for idx in 0..storage.len() {
529 let real = storage.real.value_at(idx);
530 let imag = storage.imag.value_at(idx);
531 if real.is_none_or(|value| value.is_zero()) && imag.is_none_or(|value| value.is_zero())
532 {
533 return Ok(ConditionOutcome::Fail);
534 }
535 }
536 return Ok(ConditionOutcome::Pass);
537 }
538
539 if tensor.is_empty() {
540 return Ok(ConditionOutcome::Fail);
541 }
542 for index in 0..tensor.len() {
543 let (real, imag) = tensor
544 .numeric_value_at(index)
545 .ok_or_else(|| assert_error(&ASSERT_ERROR_INVALID_CONDITION))?;
546 let (re, im) = match (real, imag) {
547 (NumericScalar::F64(re), NumericScalar::F64(im)) => (re, im),
548 (NumericScalar::F32(re), NumericScalar::F32(im)) => (f64::from(re), f64::from(im)),
549 _ => return Err(assert_error(&ASSERT_ERROR_INVALID_CONDITION)),
550 };
551 if !complex_element_passes(re, im) {
552 return Ok(ConditionOutcome::Fail);
553 }
554 }
555 Ok(ConditionOutcome::Pass)
556}
557
558fn complex_element_passes(re: f64, im: f64) -> bool {
559 if re.is_nan() || im.is_nan() {
560 return false;
561 }
562 re != 0.0 || im != 0.0
563}
564
565struct FailurePayload {
566 identifier: String,
567 message: String,
568}
569
570async fn failure_payload(args: &[Value]) -> crate::BuiltinResult<FailurePayload> {
571 if args.is_empty() {
572 return Ok(FailurePayload {
573 identifier: assert_default_identifier().to_string(),
574 message: assert_default_message().to_string(),
575 });
576 }
577
578 let candidate = &args[0];
579 let treat_as_identifier = args.len() >= 2 && value_is_identifier(candidate)?;
580
581 if treat_as_identifier {
582 if args.len() < 2 {
583 return Err(assert_flow(
584 ASSERT_ERROR_INVALID_INPUT
585 .identifier
586 .expect("assert invalid-input identifier must be defined"),
587 "assert: message text must follow the message identifier.",
588 ));
589 }
590 let identifier = identifier_from_value(candidate)?;
591 let template = message_from_value(&args[1])?;
592 let formatting_args = normalize_formatting_arguments(&args[2..]).await?;
593 let message = format_message(&template, &formatting_args)?;
594 Ok(FailurePayload {
595 identifier,
596 message,
597 })
598 } else {
599 let template = message_from_value(candidate)?;
600 let formatting_args = normalize_formatting_arguments(&args[1..]).await?;
601 let message = format_message(&template, &formatting_args)?;
602 Ok(FailurePayload {
603 identifier: assert_default_identifier().to_string(),
604 message,
605 })
606 }
607}
608
609async fn normalize_formatting_arguments(args: &[Value]) -> crate::BuiltinResult<Vec<Value>> {
610 let mut normalized = Vec::with_capacity(args.len());
611 for value in args {
612 let mut flattened = flatten_arguments(std::slice::from_ref(value), BUILTIN_NAME)
613 .await
614 .map_err(|flow| {
615 remap_assert_flow(flow, &ASSERT_ERROR_INVALID_INPUT, |err| {
616 format!("assert: {}", err.message())
617 })
618 })?;
619 if flattened.len() != 1 {
620 return Err(assert_error_with_message(
621 "assert: each message replacement value must be a character vector, string scalar, or numeric scalar.",
622 &ASSERT_ERROR_INVALID_INPUT,
623 ));
624 }
625 normalized.push(flattened.remove(0));
626 }
627 Ok(normalized)
628}
629
630fn value_is_identifier(value: &Value) -> crate::BuiltinResult<bool> {
631 if let Some(text) = string_scalar_opt(value) {
632 if text.contains(':') {
633 return Ok(true);
634 }
635 if looks_like_unqualified_identifier(&text)
636 && crate::compatibility::runmat_extensions_enabled()
637 {
638 crate::compatibility::ensure_builtin_extension_enabled(
639 &ASSERT_UNQUALIFIED_IDENTIFIER_EXTENSION,
640 BUILTIN_NAME,
641 )?;
642 return Ok(true);
643 }
644 Ok(false)
645 } else {
646 Ok(false)
647 }
648}
649
650fn identifier_from_value(value: &Value) -> crate::BuiltinResult<String> {
651 let text = string_scalar_from_value(
652 value,
653 "assert: message identifier must be a string scalar or character vector.",
654 )?;
655 if text.trim().is_empty() {
656 return Err(assert_flow(
657 ASSERT_ERROR_INVALID_INPUT
658 .identifier
659 .expect("assert invalid-input identifier must be defined"),
660 "assert: message identifier must be nonempty.",
661 ));
662 }
663 let trimmed = text.trim();
664 if is_message_identifier(trimmed) {
665 return Ok(trimmed.to_string());
666 }
667 if looks_like_unqualified_identifier(trimmed)
668 && crate::compatibility::runmat_extensions_enabled()
669 {
670 crate::compatibility::ensure_builtin_extension_enabled(
671 &ASSERT_UNQUALIFIED_IDENTIFIER_EXTENSION,
672 BUILTIN_NAME,
673 )?;
674 return Ok(normalize_identifier(trimmed));
675 }
676 Err(assert_error_with_message(
677 "assert: error identifier must contain colon-separated fields that each begin with a letter and otherwise contain only letters, digits, or underscores.",
678 &ASSERT_ERROR_INVALID_INPUT,
679 ))
680}
681
682fn message_from_value(value: &Value) -> crate::BuiltinResult<String> {
683 string_scalar_from_value(
684 value,
685 "assert: message text must be a string scalar or character vector.",
686 )
687}
688
689fn format_message(template: &str, args: &[Value]) -> crate::BuiltinResult<String> {
690 format_variadic(template, args).map_err(|flow| {
691 remap_assert_flow(flow, &ASSERT_ERROR_INVALID_INPUT, |err| {
692 format!("assert: {}", err.message())
693 })
694 })
695}
696
697fn normalize_identifier(raw: &str) -> String {
698 let trimmed = raw.trim();
699 if trimmed.is_empty() {
700 assert_default_identifier().to_string()
701 } else if trimmed.contains(':') {
702 trimmed.to_string()
703 } else {
704 format!("RunMat:{trimmed}")
705 }
706}
707
708fn is_message_identifier(text: &str) -> bool {
709 let trimmed = text.trim();
710 let fields: Vec<&str> = trimmed.split(':').collect();
711 if fields.len() < 2 {
712 return false;
713 }
714 fields.into_iter().all(is_identifier_field)
715}
716
717fn looks_like_unqualified_identifier(text: &str) -> bool {
718 let trimmed = text.trim();
719 !trimmed.contains(':') && is_identifier_field(trimmed)
720}
721
722fn is_identifier_field(field: &str) -> bool {
723 let mut chars = field.chars();
724 chars
725 .next()
726 .is_some_and(|first| first.is_ascii_alphabetic())
727 && chars.all(|character| character.is_ascii_alphanumeric() || character == '_')
728}
729
730fn string_scalar_from_value(value: &Value, context: &str) -> crate::BuiltinResult<String> {
731 match value {
732 Value::String(text) => Ok(text.clone()),
733 Value::StringArray(array) if array.data.len() == 1 => Ok(array.data[0].clone()),
734 Value::CharArray(char_array) if char_array.rows == 1 => {
735 Ok(char_array.data.iter().collect::<String>())
736 }
737 _ => Err(assert_error_with_message(
738 context,
739 &ASSERT_ERROR_INVALID_INPUT,
740 )),
741 }
742}
743
744fn string_scalar_opt(value: &Value) -> Option<String> {
745 match value {
746 Value::String(text) => Some(text.clone()),
747 Value::StringArray(array) if array.data.len() == 1 => Some(array.data[0].clone()),
748 Value::CharArray(char_array) if char_array.rows == 1 => {
749 Some(char_array.data.iter().collect())
750 }
751 _ => None,
752 }
753}
754
755#[cfg(test)]
756pub(crate) mod tests {
757 use super::*;
758 use crate::builtins::common::test_support;
759 use futures::executor::block_on;
760 use runmat_builtins::{ResolveContext, Type};
761 use runmat_value::{
762 ComplexTensor, IntValue, IntegerComplexStorage, IntegerStorage, LogicalArray, Tensor,
763 };
764
765 fn assert_builtin(args: Vec<Value>) -> crate::BuiltinResult<Value> {
766 block_on(super::assert_builtin(args))
767 }
768
769 fn unwrap_error(err: crate::RuntimeError) -> crate::RuntimeError {
770 err
771 }
772
773 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
774 #[test]
775 fn assert_true_passes() {
776 let result = assert_builtin(vec![Value::Bool(true)]).expect("assert should pass");
777 assert_eq!(result, Value::Num(0.0));
778 }
779
780 #[test]
781 fn assert_scalar_wide_uint64_passes() {
782 let result =
783 assert_builtin(vec![Value::Int(IntValue::U64(u64::MAX))]).expect("assert should pass");
784 assert_eq!(result, Value::Num(0.0));
785 }
786
787 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
788 #[test]
789 fn assert_empty_tensor_fails() {
790 let tensor = Tensor::new(Vec::new(), vec![0, 3]).unwrap();
791 let err =
792 assert_builtin(vec![Value::Tensor(tensor)]).expect_err("empty condition should fail");
793 assert_eq!(err.identifier(), Some(assert_default_identifier()));
794 }
795
796 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
797 #[test]
798 fn assert_empty_logical_fails() {
799 let logical = LogicalArray::new(Vec::new(), vec![0]).unwrap();
800 let err = assert_builtin(vec![Value::LogicalArray(logical)])
801 .expect_err("empty condition should fail");
802 assert_eq!(err.identifier(), Some(assert_default_identifier()));
803 }
804
805 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
806 #[test]
807 fn assert_false_uses_default_message() {
808 let err =
809 unwrap_error(assert_builtin(vec![Value::Bool(false)]).expect_err("assert should fail"));
810 assert_eq!(err.identifier(), Some(assert_default_identifier()));
811 assert_eq!(err.message(), assert_default_message());
812 }
813
814 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
815 #[test]
816 fn assert_handles_numeric_tensor() {
817 let tensor = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
818 assert_builtin(vec![Value::Tensor(tensor)]).expect("assert should pass");
819 }
820
821 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
822 #[test]
823 fn assert_detects_zero_in_tensor() {
824 let tensor = Tensor::new(vec![1.0, 0.0, 3.0], vec![3, 1]).unwrap();
825 let err = unwrap_error(
826 assert_builtin(vec![Value::Tensor(tensor)]).expect_err("assert should fail"),
827 );
828 assert_eq!(err.identifier(), Some(assert_default_identifier()));
829 }
830
831 #[test]
832 fn assert_reads_typed_integer_tensor_storage_exactly() {
833 let passing =
834 Tensor::new_integer(IntegerStorage::U64(vec![u64::MAX, 1]), vec![2, 1]).unwrap();
835 assert_builtin(vec![Value::Tensor(passing)]).expect("assert should pass");
836
837 let failing =
838 Tensor::new_integer(IntegerStorage::U64(vec![u64::MAX, 0]), vec![2, 1]).unwrap();
839 let err = unwrap_error(
840 assert_builtin(vec![Value::Tensor(failing)]).expect_err("assert should fail"),
841 );
842 assert_eq!(err.identifier(), Some(assert_default_identifier()));
843 }
844
845 #[test]
846 fn assert_tests_every_real_integer_class_exactly() {
847 for (passing, failing) in [
848 (
849 IntegerStorage::I8(vec![i8::MIN, i8::MAX]),
850 IntegerStorage::I8(vec![i8::MIN, 0]),
851 ),
852 (
853 IntegerStorage::I16(vec![i16::MIN, i16::MAX]),
854 IntegerStorage::I16(vec![i16::MIN, 0]),
855 ),
856 (
857 IntegerStorage::I32(vec![i32::MIN, i32::MAX]),
858 IntegerStorage::I32(vec![i32::MIN, 0]),
859 ),
860 (
861 IntegerStorage::I64(vec![i64::MIN, i64::MAX]),
862 IntegerStorage::I64(vec![i64::MIN, 0]),
863 ),
864 (
865 IntegerStorage::U8(vec![1, u8::MAX]),
866 IntegerStorage::U8(vec![u8::MAX, 0]),
867 ),
868 (
869 IntegerStorage::U16(vec![1, u16::MAX]),
870 IntegerStorage::U16(vec![u16::MAX, 0]),
871 ),
872 (
873 IntegerStorage::U32(vec![1, u32::MAX]),
874 IntegerStorage::U32(vec![u32::MAX, 0]),
875 ),
876 (
877 IntegerStorage::U64(vec![9_007_199_254_740_993, u64::MAX]),
878 IntegerStorage::U64(vec![u64::MAX, 0]),
879 ),
880 ] {
881 assert_builtin(vec![Value::Tensor(
882 Tensor::new_integer(passing, vec![2, 1]).expect("passing integer condition"),
883 )])
884 .expect("all nonzero integers pass");
885 let err = assert_builtin(vec![Value::Tensor(
886 Tensor::new_integer(failing, vec![2, 1]).expect("failing integer condition"),
887 )])
888 .expect_err("zero integer fails");
889 assert_eq!(err.identifier(), Some(assert_default_identifier()));
890 }
891 }
892
893 #[test]
894 fn assert_formats_every_integer_scalar_class_exactly() {
895 for (value, format, expected) in [
896 (IntValue::I8(i8::MIN), "%d", i8::MIN.to_string()),
897 (IntValue::I16(i16::MIN), "%d", i16::MIN.to_string()),
898 (IntValue::I32(i32::MIN), "%d", i32::MIN.to_string()),
899 (IntValue::I64(i64::MIN), "%d", i64::MIN.to_string()),
900 (IntValue::U8(u8::MAX), "%u", u8::MAX.to_string()),
901 (IntValue::U16(u16::MAX), "%u", u16::MAX.to_string()),
902 (IntValue::U32(u32::MAX), "%u", u32::MAX.to_string()),
903 (IntValue::U64(u64::MAX), "%u", u64::MAX.to_string()),
904 ] {
905 let err = assert_builtin(vec![
906 Value::Bool(false),
907 Value::String(format.to_string()),
908 Value::Int(value),
909 ])
910 .expect_err("formatted assertion should fail");
911 assert_eq!(err.message(), expected);
912 }
913
914 let scalar =
915 Tensor::new_integer(IntegerStorage::U64(vec![u64::MAX]), vec![1, 1]).expect("scalar");
916 let err = assert_builtin(vec![
917 Value::Bool(false),
918 Value::String("%u".to_string()),
919 Value::Tensor(scalar),
920 ])
921 .expect_err("formatted assertion should fail");
922 assert_eq!(err.message(), u64::MAX.to_string());
923
924 let nonscalar =
925 Tensor::new_integer(IntegerStorage::U8(vec![1, 2]), vec![1, 2]).expect("array");
926 let err = assert_builtin(vec![
927 Value::Bool(false),
928 Value::String("%u".to_string()),
929 Value::Tensor(nonscalar),
930 ])
931 .expect_err("format replacement arrays reject");
932 assert_eq!(
933 err.identifier(),
934 Some(ASSERT_ERROR_INVALID_INPUT.identifier.unwrap())
935 );
936 }
937
938 #[test]
939 fn assert_real_condition_conversion_rejects_nan_and_accepts_character_vectors() {
940 let chars = runmat_value::CharArray::new(vec!['o', 'k'], 1, 2).expect("chars");
941 assert_builtin(vec![Value::CharArray(chars)]).expect("nonzero character codes pass");
942
943 for chars in [
944 runmat_value::CharArray::new(Vec::new(), 1, 0).expect("empty"),
945 runmat_value::CharArray::new(vec!['o', '\0'], 1, 2).expect("zero character"),
946 ] {
947 let err =
948 assert_builtin(vec![Value::CharArray(chars)]).expect_err("condition should fail");
949 assert_eq!(err.identifier(), Some(assert_default_identifier()));
950 }
951
952 for value in [
953 Value::Num(f64::NAN),
954 Value::Tensor(Tensor::new(vec![1.0, f64::NAN], vec![2, 1]).expect("double")),
955 Value::Tensor(Tensor::from_f32(vec![1.0, f32::NAN], vec![2, 1]).expect("single")),
956 ] {
957 let err = assert_builtin(vec![value]).expect_err("NaN cannot convert to logical");
958 assert_eq!(
959 err.identifier(),
960 Some(ASSERT_ERROR_INVALID_CONDITION.identifier.unwrap())
961 );
962 }
963 }
964
965 #[test]
966 fn assert_complex_conditions_are_mode_gated() {
967 {
968 let _compat = crate::compatibility::push_runmat_extensions_enabled(false);
969 let err = assert_builtin(vec![Value::Complex(1.0, 0.0)])
970 .expect_err("MATLAB mode rejects complex condition");
971 assert_eq!(
972 err.identifier(),
973 ASSERT_COMPLEX_CONDITION_EXTENSION.error_identifier
974 );
975 }
976 {
977 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
978 assert_builtin(vec![Value::Complex(1.0, 0.0)])
979 .expect("RunMat mode admits complex condition");
980 }
981 }
982
983 #[test]
984 fn assert_identifier_grammar_and_unqualified_extension_are_explicit() {
985 {
986 let _compat = crate::compatibility::push_runmat_extensions_enabled(false);
987 let err = assert_builtin(vec![
988 Value::Bool(false),
989 Value::String("plainMessage".to_string()),
990 Value::Int(IntValue::I32(7)),
991 ])
992 .expect_err("plain text is the message form");
993 assert_eq!(err.identifier(), Some(assert_default_identifier()));
994 assert_eq!(err.message(), "plainMessage");
995 }
996 {
997 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
998 let err = assert_builtin(vec![
999 Value::Bool(false),
1000 Value::String("customFailure".to_string()),
1001 Value::String("failed".to_string()),
1002 ])
1003 .expect_err("RunMat mode admits unqualified identifier");
1004 assert_eq!(err.identifier(), Some("RunMat:customFailure"));
1005 }
1006 for identifier in ["bad.segment:mnemonic", "component:9bad", "component::bad"] {
1007 let err = assert_builtin(vec![
1008 Value::Bool(false),
1009 Value::String(identifier.to_string()),
1010 Value::String("failed".to_string()),
1011 ])
1012 .expect_err("invalid qualified identifier rejects");
1013 assert_eq!(
1014 err.identifier(),
1015 Some(ASSERT_ERROR_INVALID_INPUT.identifier.unwrap())
1016 );
1017 }
1018 }
1019
1020 #[test]
1021 fn assert_reads_native_single_tensor_storage() {
1022 let passing = Tensor::from_f32(vec![f32::MIN_POSITIVE, -2.0], vec![2, 1]).unwrap();
1023 assert_builtin(vec![Value::Tensor(passing)]).expect("assert should pass");
1024
1025 let zero = Tensor::from_f32(vec![1.0_f32, 0.0], vec![2, 1]).unwrap();
1026 let err =
1027 assert_builtin(vec![Value::Tensor(zero)]).expect_err("zero condition should fail");
1028 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1029
1030 let nan = Tensor::from_f32(vec![1.0_f32, f32::NAN], vec![2, 1]).unwrap();
1031 let err =
1032 assert_builtin(vec![Value::Tensor(nan)]).expect_err("NaN condition should reject");
1033 assert_eq!(
1034 err.identifier(),
1035 Some(ASSERT_ERROR_INVALID_CONDITION.identifier.unwrap())
1036 );
1037 }
1038
1039 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1040 #[test]
1041 fn assert_detects_nan() {
1042 let err = unwrap_error(
1043 assert_builtin(vec![Value::Num(f64::NAN)]).expect_err("assert should reject NaN"),
1044 );
1045 assert_eq!(
1046 err.identifier(),
1047 Some(ASSERT_ERROR_INVALID_CONDITION.identifier.unwrap())
1048 );
1049 }
1050
1051 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1052 #[test]
1053 fn assert_complex_scalar_passes() {
1054 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
1055 assert_builtin(vec![Value::Complex(0.0, 2.0)]).expect("assert should pass");
1056 }
1057
1058 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1059 #[test]
1060 fn assert_complex_scalar_failure() {
1061 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
1062 let err = unwrap_error(
1063 assert_builtin(vec![Value::Complex(0.0, 0.0)]).expect_err("assert should fail"),
1064 );
1065 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1066 }
1067
1068 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1069 #[test]
1070 fn assert_complex_tensor_failure() {
1071 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
1072 let tensor = ComplexTensor::new(vec![(1.0, 0.0), (0.0, 0.0)], vec![2, 1]).expect("tensor");
1073 let err = unwrap_error(
1074 assert_builtin(vec![Value::ComplexTensor(tensor)]).expect_err("assert should fail"),
1075 );
1076 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1077 }
1078
1079 #[test]
1080 fn assert_reads_typed_complex_integer_tensor_storage_exactly() {
1081 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
1082 let storage = IntegerComplexStorage::new(
1083 IntegerStorage::U64(vec![0, u64::MAX]),
1084 IntegerStorage::U64(vec![5, 0]),
1085 )
1086 .expect("complex integer storage");
1087 let passing = ComplexTensor::new_integer(storage, vec![2, 1]).unwrap();
1088 assert_builtin(vec![Value::ComplexTensor(passing)]).expect("assert should pass");
1089
1090 let storage = IntegerComplexStorage::new(
1091 IntegerStorage::U64(vec![u64::MAX, 0]),
1092 IntegerStorage::U64(vec![0, 0]),
1093 )
1094 .expect("complex integer storage");
1095 let failing = ComplexTensor::new_integer(storage, vec![2, 1]).unwrap();
1096 let err = unwrap_error(
1097 assert_builtin(vec![Value::ComplexTensor(failing)]).expect_err("assert should fail"),
1098 );
1099 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1100 }
1101
1102 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1103 #[test]
1104 fn assert_accepts_custom_message() {
1105 let err = unwrap_error(
1106 assert_builtin(vec![
1107 Value::Bool(false),
1108 Value::from("Vector length must be positive."),
1109 ])
1110 .expect_err("assert should fail"),
1111 );
1112 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1113 assert!(err.message().contains("Vector length must be positive."));
1114 }
1115
1116 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1117 #[test]
1118 fn assert_supports_message_formatting() {
1119 let err = unwrap_error(
1120 assert_builtin(vec![
1121 Value::Bool(false),
1122 Value::from("Expected positive value, got %d."),
1123 Value::Int(IntValue::I32(-4)),
1124 ])
1125 .expect_err("assert should fail"),
1126 );
1127 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1128 assert!(err.message().contains("Expected positive value, got -4."));
1129 }
1130
1131 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1132 #[test]
1133 fn assert_supports_custom_identifier() {
1134 let err = unwrap_error(
1135 assert_builtin(vec![
1136 Value::Bool(false),
1137 Value::from("runmat:tests:failed"),
1138 Value::from("Failure %d occurred."),
1139 Value::Int(IntValue::I32(3)),
1140 ])
1141 .expect_err("assert should fail"),
1142 );
1143 assert_eq!(err.identifier(), Some("runmat:tests:failed"));
1144 assert!(err.message().contains("Failure 3 occurred."));
1145 }
1146
1147 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1148 #[test]
1149 fn assert_unqualified_identifier_prefixed() {
1150 let _compat = crate::compatibility::push_runmat_extensions_enabled(true);
1151 let err = unwrap_error(
1152 assert_builtin(vec![
1153 Value::Bool(false),
1154 Value::from("customAssertionFailed"),
1155 Value::from("runtime failure"),
1156 ])
1157 .expect_err("assert should fail"),
1158 );
1159 assert_eq!(err.identifier(), Some("RunMat:customAssertionFailed"));
1160 }
1161
1162 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1163 #[test]
1164 fn assert_rejects_invalid_condition_type() {
1165 let err = unwrap_error(
1166 assert_builtin(vec![Value::from("invalid")]).expect_err("assert should error"),
1167 );
1168 assert_eq!(
1169 err.identifier(),
1170 Some(ASSERT_ERROR_INVALID_CONDITION.identifier.unwrap())
1171 );
1172 }
1173
1174 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1175 #[test]
1176 fn assert_gpu_tensor_passes() {
1177 test_support::with_test_provider(|provider| {
1178 let tensor = Tensor::new(vec![1.0, 2.0, 3.0], vec![3, 1]).unwrap();
1179 let view = runmat_accelerate_api::HostTensorView {
1180 data: &tensor.materialize_f64(),
1181 shape: &tensor.shape,
1182 };
1183 let handle = provider.upload(&view).expect("upload");
1184 let result = assert_builtin(vec![Value::GpuTensor(handle)]).expect("assert");
1185 assert_eq!(result, Value::Num(0.0));
1186 });
1187 }
1188
1189 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1190 #[test]
1191 fn assert_invalid_message_type_errors() {
1192 let err = unwrap_error(
1193 assert_builtin(vec![Value::Bool(false), Value::Num(5.0)])
1194 .expect_err("assert should error"),
1195 );
1196 assert_eq!(
1197 err.identifier(),
1198 Some(ASSERT_ERROR_INVALID_INPUT.identifier.unwrap())
1199 );
1200 }
1201
1202 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1203 #[test]
1204 fn assert_formatting_error_propagates() {
1205 let err = unwrap_error(
1206 assert_builtin(vec![
1207 Value::Bool(false),
1208 Value::from("number %d must be > 0"),
1209 ])
1210 .expect_err("assert should fail"),
1211 );
1212 assert_eq!(
1213 err.identifier(),
1214 Some(ASSERT_ERROR_INVALID_INPUT.identifier.unwrap())
1215 );
1216 assert!(err.message().contains("sprintf"));
1217 }
1218
1219 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1220 #[test]
1221 fn assert_gpu_tensor_failure() {
1222 test_support::with_test_provider(|provider| {
1223 let tensor = Tensor::new(vec![1.0, 0.0, 3.0], vec![3, 1]).unwrap();
1224 let view = runmat_accelerate_api::HostTensorView {
1225 data: &tensor.materialize_f64(),
1226 shape: &tensor.shape,
1227 };
1228 let handle = provider.upload(&view).expect("upload");
1229 let err =
1230 unwrap_error(assert_builtin(vec![Value::GpuTensor(handle)]).expect_err("assert"));
1231 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1232 });
1233 }
1234
1235 #[test]
1236 fn assert_provider_gather_tests_every_integer_class_and_formats_wide_scalar() {
1237 test_support::with_test_provider(|provider| {
1238 for storage in [
1239 IntegerStorage::I8(vec![i8::MIN, 0]),
1240 IntegerStorage::I16(vec![i16::MIN, 0]),
1241 IntegerStorage::I32(vec![i32::MIN, 0]),
1242 IntegerStorage::I64(vec![i64::MIN, 0]),
1243 IntegerStorage::U8(vec![u8::MAX, 0]),
1244 IntegerStorage::U16(vec![u16::MAX, 0]),
1245 IntegerStorage::U32(vec![u32::MAX, 0]),
1246 IntegerStorage::U64(vec![u64::MAX, 0]),
1247 ] {
1248 let handle = gpu_helpers::upload_tensor(
1249 provider,
1250 &Tensor::new_integer(storage, vec![2, 1]).expect("condition"),
1251 )
1252 .expect("upload");
1253 let err = assert_builtin(vec![Value::GpuTensor(handle.clone())])
1254 .expect_err("resident zero fails");
1255 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1256 let _ = provider.free(&handle);
1257 }
1258
1259 let handle = gpu_helpers::upload_tensor(
1260 provider,
1261 &Tensor::new_integer(IntegerStorage::U64(vec![u64::MAX]), vec![1, 1])
1262 .expect("format value"),
1263 )
1264 .expect("upload");
1265 let err = assert_builtin(vec![
1266 Value::Bool(false),
1267 Value::String("%u".to_string()),
1268 Value::GpuTensor(handle.clone()),
1269 ])
1270 .expect_err("resident scalar formats");
1271 assert_eq!(err.message(), u64::MAX.to_string());
1272 let _ = provider.free(&handle);
1273 });
1274 }
1275
1276 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1277 #[test]
1278 fn assert_logical_array_failure() {
1279 let logical = LogicalArray::new(vec![1, 0], vec![2]).unwrap();
1280 let err = unwrap_error(
1281 assert_builtin(vec![Value::LogicalArray(logical)]).expect_err("assert should fail"),
1282 );
1283 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1284 }
1285
1286 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1287 #[test]
1288 fn assert_requires_condition_argument() {
1289 let err = unwrap_error(assert_builtin(Vec::new()).expect_err("assert should error"));
1290 assert_eq!(
1291 err.identifier(),
1292 Some(ASSERT_ERROR_NOT_ENOUGH_INPUTS.identifier.unwrap())
1293 );
1294 assert_eq!(err.message(), ASSERT_ERROR_NOT_ENOUGH_INPUTS.message);
1295 }
1296
1297 #[test]
1298 fn assert_rejects_requested_public_output() {
1299 let _outputs = crate::output_count::push_output_count(Some(1));
1300 let err = assert_builtin(vec![Value::Bool(true)]).expect_err("assert has no output");
1301 assert_eq!(err.identifier(), ASSERT_ERROR_TOO_MANY_OUTPUTS.identifier);
1302 }
1303
1304 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
1305 #[test]
1306 #[cfg(feature = "wgpu")]
1307 fn assert_wgpu_tensor_failure_matches_cpu() {
1308 use runmat_accelerate::backend::wgpu::provider::{
1309 register_wgpu_provider, WgpuProviderOptions,
1310 };
1311
1312 if register_wgpu_provider(WgpuProviderOptions::default()).is_err() {
1313 return;
1314 }
1315 let Some(provider) = runmat_accelerate_api::provider() else {
1316 return;
1317 };
1318
1319 let tensor = Tensor::new(vec![1.0, 0.0], vec![2, 1]).unwrap();
1320 let view = runmat_accelerate_api::HostTensorView {
1321 data: &tensor.materialize_f64(),
1322 shape: &tensor.shape,
1323 };
1324 let handle = provider.upload(&view).expect("upload");
1325 let err = unwrap_error(
1326 assert_builtin(vec![Value::GpuTensor(handle)]).expect_err("assert should fail"),
1327 );
1328 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1329 }
1330
1331 #[test]
1332 #[cfg(feature = "wgpu")]
1333 fn assert_wgpu_integer_conditions_and_formatting_remain_exact() {
1334 use runmat_accelerate::backend::wgpu::provider::{
1335 register_wgpu_provider, WgpuProviderOptions,
1336 };
1337
1338 let _guard = test_support::accel_test_lock();
1339 if register_wgpu_provider(WgpuProviderOptions::default()).is_err() {
1340 return;
1341 }
1342 let Some(provider) = runmat_accelerate_api::provider() else {
1343 return;
1344 };
1345 for storage in [
1346 IntegerStorage::I8(vec![i8::MIN, 0]),
1347 IntegerStorage::I16(vec![i16::MIN, 0]),
1348 IntegerStorage::I32(vec![i32::MIN, 0]),
1349 IntegerStorage::I64(vec![i64::MIN, 0]),
1350 IntegerStorage::U8(vec![u8::MAX, 0]),
1351 IntegerStorage::U16(vec![u16::MAX, 0]),
1352 IntegerStorage::U32(vec![u32::MAX, 0]),
1353 IntegerStorage::U64(vec![u64::MAX, 0]),
1354 ] {
1355 let handle = gpu_helpers::upload_tensor(
1356 provider,
1357 &Tensor::new_integer(storage, vec![2, 1]).expect("condition"),
1358 )
1359 .expect("upload");
1360 let err = assert_builtin(vec![Value::GpuTensor(handle.clone())])
1361 .expect_err("resident zero fails");
1362 assert_eq!(err.identifier(), Some(assert_default_identifier()));
1363 let _ = provider.free(&handle);
1364 }
1365
1366 let handle = gpu_helpers::upload_tensor(
1367 provider,
1368 &Tensor::new_integer(IntegerStorage::U64(vec![u64::MAX]), vec![1, 1])
1369 .expect("format value"),
1370 )
1371 .expect("upload");
1372 let err = assert_builtin(vec![
1373 Value::Bool(false),
1374 Value::String("%u".to_string()),
1375 Value::GpuTensor(handle.clone()),
1376 ])
1377 .expect_err("resident scalar formats");
1378 assert_eq!(err.message(), u64::MAX.to_string());
1379 let _ = provider.free(&handle);
1380 }
1381
1382 #[test]
1383 fn assert_has_no_public_output_type() {
1384 assert_eq!(
1385 assert_type(&[Type::Bool], &ResolveContext::new(Vec::new())),
1386 Type::Unknown
1387 );
1388 }
1389
1390 #[test]
1391 fn assert_metadata_classifies_integer_and_extension_forms() {
1392 assert_eq!(ASSERT_INTEGER_CAPABILITIES.len(), 3);
1393 assert_eq!(
1394 ASSERT_INTEGER_CAPABILITIES[0].output_class,
1395 BuiltinIntegerOutputClassRule::NotApplicable
1396 );
1397 assert_eq!(
1398 ASSERT_INTEGER_CAPABILITIES[1].inputs[0].availability,
1399 BuiltinIntegerInputAvailability::RunMatOnly
1400 );
1401 assert_eq!(
1402 ASSERT_EXTENSIONS,
1403 [
1404 ASSERT_COMPLEX_CONDITION_EXTENSION,
1405 ASSERT_UNQUALIFIED_IDENTIFIER_EXTENSION
1406 ]
1407 );
1408 assert!(ASSERT_SIGNATURES
1409 .iter()
1410 .all(|signature| signature.outputs.is_empty()));
1411 }
1412}