diff --git a/android/app/src/androidTest/kotlin/com/zarz/spotiflac/ExtensionInputValidationTest.kt b/android/app/src/androidTest/kotlin/com/zarz/spotiflac/ExtensionInputValidationTest.kt new file mode 100644 index 00000000..827ebac3 --- /dev/null +++ b/android/app/src/androidTest/kotlin/com/zarz/spotiflac/ExtensionInputValidationTest.kt @@ -0,0 +1,28 @@ +package com.zarz.spotiflac + +import androidx.test.ext.junit.runners.AndroidJUnit4 +import com.spotiflac.backend.JsExtension +import com.spotiflac.backend.JsExtensionException +import org.json.JSONObject +import org.junit.Assert.assertEquals +import org.junit.Assert.assertTrue +import org.junit.Test +import org.junit.runner.RunWith + +@RunWith(AndroidJUnit4::class) +class ExtensionInputValidationTest { + @Test + fun invalidArgumentsDoNotReachJavaScriptAndValidObjectsKeepTheirValues() { + val source = "let calls=0; registerExtension({echo(value){return {calls:++calls,value}}});" + JsExtension(source, "{}", 2000uL).use { runtime -> + for (input in listOf("{}", "[1e400]", "[\"\\uD800\"]", "[] true")) { + val failure = runCatching { runtime.call("echo", input, null, 2000uL) }.exceptionOrNull() + assertTrue(failure is JsExtensionException.InvalidInput) + } + val result = JSONObject(runtime.call("echo", "[{\"title\":\"音楽 🎵\",\"options\":[true,null,42]}]", null, 2000uL)) + assertEquals(1, result.getInt("calls")) + assertEquals("音楽 🎵", result.getJSONObject("value").getString("title")) + assertEquals(42, result.getJSONObject("value").getJSONArray("options").getInt(2)) + } + } +} diff --git a/rust_backend/crates/extensions/src/runtime.rs b/rust_backend/crates/extensions/src/runtime.rs index 27ddb237..7878b1a9 100644 --- a/rust_backend/crates/extensions/src/runtime.rs +++ b/rust_backend/crates/extensions/src/runtime.rs @@ -6,6 +6,8 @@ use std::sync::{Arc, Condvar, Mutex, OnceLock, Weak}; use std::thread::{self, JoinHandle}; use std::time::{Duration, Instant}; +mod json_input; + const MAX_INPUT_BYTES: usize = 8 * 1024 * 1024; const MAX_TIMEOUT_MS: u64 = 300_000; const QUEUE_CAPACITY: usize = 8; @@ -1010,9 +1012,7 @@ impl Vm { context, _runtime: runtime, }; - let empty_settings = validate_json(settings)? - .as_object() - .is_some_and(|settings| settings.is_empty()); + let empty_settings = validate_json(settings)?.is_empty_object(); if services.load_mode == LoadMode::Initialize && (services.initialize_empty_settings || !empty_settings) { @@ -1218,7 +1218,7 @@ fn validate_size(value: &str) -> Result<(), ExtensionError> { Ok(()) } -fn validate_json(value: &str) -> Result { +fn validate_json(value: &str) -> Result { validate_size(value)?; serde_json::from_str(value).map_err(|error| ExtensionError::InvalidInput(error.to_string())) } @@ -1227,6 +1227,28 @@ fn validate_json(value: &str) -> Result { mod tests { use super::*; + #[test] + fn argument_validation_rejects_before_execution_and_keeps_runtime_usable() { + let runtime = ExtensionRuntime::load( + "let calls=0; registerExtension({echo(value){return {calls:++calls,value}}});", + "{}", + RuntimeLimits::default(), + ) + .unwrap(); + for input in ["{}", "[1e400]", "[\"\\uD800\"]", "[] true"] { + assert!(matches!( + runtime.call("echo", input, None, 0), + Err(ExtensionError::InvalidInput(_)) + )); + } + let input = "[{\"title\":\"音楽 🎵\",\"nested\":[true,null,{\"key\":\"value\"}]}]"; + let output: serde_json::Value = + serde_json::from_str(&runtime.call("echo", input, None, 0).unwrap()).unwrap(); + assert_eq!(output["calls"], 1); + assert_eq!(output["value"]["title"], "音楽 🎵"); + assert_eq!(output["value"]["nested"][2]["key"], "value"); + } + #[test] fn cached_source_still_obeys_each_vms_memory_and_stack_limits() { let source = "registerExtension({allocate(){const a=[];while(true)a.push(new Uint8Array(65536));},recurse(){function f(){return 1+f()}return f()}});"; diff --git a/rust_backend/crates/extensions/src/runtime/json_input.rs b/rust_backend/crates/extensions/src/runtime/json_input.rs new file mode 100644 index 00000000..664f7252 --- /dev/null +++ b/rust_backend/crates/extensions/src/runtime/json_input.rs @@ -0,0 +1,214 @@ +//! Validate caller JSON without retaining a second object tree before QuickJS +//! parses it. Use normal Serde visits rather than IgnoredAny/RawValue so number +//! range, string escape, nesting, and trailing-data errors remain unchanged. + +use serde::de::{Deserialize, Deserializer, MapAccess, SeqAccess, Visitor}; +use std::fmt; + +#[derive(Debug, PartialEq, Eq)] +pub(super) enum Shape { + Scalar, + Array, + Object { empty: bool }, +} + +impl Shape { + pub(super) fn is_array(&self) -> bool { + matches!(self, Self::Array) + } + + pub(super) fn is_object(&self) -> bool { + matches!(self, Self::Object { .. }) + } + + pub(super) fn is_empty_object(&self) -> bool { + matches!(self, Self::Object { empty: true }) + } +} + +impl<'de> Deserialize<'de> for Shape { + fn deserialize>(deserializer: D) -> Result { + struct ShapeVisitor; + + impl<'de> Visitor<'de> for ShapeVisitor { + type Value = Shape; + + fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result { + formatter.write_str("any valid JSON value") + } + + fn visit_bool(self, _: bool) -> Result { + Ok(Shape::Scalar) + } + + fn visit_i64(self, _: i64) -> Result { + Ok(Shape::Scalar) + } + + fn visit_u64(self, _: u64) -> Result { + Ok(Shape::Scalar) + } + + fn visit_f64(self, _: f64) -> Result { + Ok(Shape::Scalar) + } + + fn visit_str(self, _: &str) -> Result { + Ok(Shape::Scalar) + } + + fn visit_unit(self) -> Result { + Ok(Shape::Scalar) + } + + fn visit_seq>(self, mut values: A) -> Result { + while values.next_element::()?.is_some() {} + Ok(Shape::Array) + } + + fn visit_map>(self, mut values: A) -> Result { + let mut empty = true; + while values.next_entry::()?.is_some() { + empty = false; + } + Ok(Shape::Object { empty }) + } + } + + deserializer.deserialize_any(ShapeVisitor) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn compare(input: &str) { + let expected = serde_json::from_str::(input); + let actual = serde_json::from_str::(input); + match (expected, actual) { + (Ok(value), Ok(shape)) => { + assert_eq!(shape.is_array(), value.is_array(), "{input}"); + assert_eq!(shape.is_object(), value.is_object(), "{input}"); + assert_eq!( + shape.is_empty_object(), + value.as_object().is_some_and(|value| value.is_empty()), + "{input}", + ); + } + (Err(expected), Err(actual)) => { + assert_eq!(actual.to_string(), expected.to_string(), "{input}"); + } + (expected, actual) => { + panic!("validation differs for {input:?}: {expected:?} / {actual:?}") + } + } + } + + #[test] + fn validation_matches_value_shapes_numbers_unicode_and_errors() { + for input in [ + "null", + "true", + "false", + "0", + "-0", + "1.25e-100", + "1e400", + "[1e400,0]", + "{\"x\":1e400,\"x\":0}", + "18446744073709551615", + "18446744073709551616", + "-9223372036854775809", + "NaN", + "Infinity", + "[]", + "{}", + "{\"x\":null}", + " { }\n", + "[1,{\"a\":[true,null,\"text\"]}]", + "\"音楽 🎵\"", + "\"\\uD83C\\uDFB5\"", + "\"\\uD800\"", + "\"\\uDC00\"", + "\"\\x20\"", + "\"line\nfeed\"", + "{\"a\\tb\":\"\\u0000\"}", + "{\"a\":1,\"a\":2}", + "{1:true}", + "[01]", + "[1,]", + "[] true", + "[", + "{\"x\":", + "", + ] { + compare(input); + } + for depth in [1, 100, 127, 128, 129] { + compare(&format!("{}0{}", "[".repeat(depth), "]".repeat(depth))); + compare(&format!( + "{}0{}", + "{\"key\":".repeat(depth), + "}".repeat(depth) + )); + } + } + + #[test] + fn validation_matches_value_for_truncated_and_mutated_payloads() { + let input = r#"[{"track":"a\\b\n\uD83C\uDFB5","options":{"count":123.4e-3,"ids":[1,2,3],"ok":true,"nil":null}}]"#; + for index in 0..input.len() { + compare(&input[..index]); + for &replacement in b" \"\\0,]}" { + let mut bytes = input.as_bytes().to_vec(); + bytes[index] = replacement; + compare(std::str::from_utf8(&bytes).unwrap()); + } + } + } + + #[test] + #[ignore = "manual release benchmark: JSON validation only, not device latency"] + fn benchmark_shape_validation() { + use std::hint::black_box; + use std::time::Instant; + + for count in [1, 100, 1000, 10000] { + let row = serde_json::json!({"id":"track-123","title":"音楽 🎵","artists":["Artist"],"duration_ms":240000,"options":{"quality":"lossless","enabled":true}}); + let input = serde_json::json!([vec![row; count]]).to_string(); + let mut old = Vec::new(); + let mut new = Vec::new(); + for round in 0..44 { + for baseline in if round % 2 == 0 { + [true, false] + } else { + [false, true] + } { + let start = Instant::now(); + if baseline { + black_box( + serde_json::from_str::(black_box(&input)).unwrap(), + ); + } else { + black_box(serde_json::from_str::(black_box(&input)).unwrap()); + } + if round >= 4 { + (if baseline { &mut old } else { &mut new }) + .push(start.elapsed().as_nanos()); + } + } + } + old.sort_unstable(); + new.sort_unstable(); + println!( + "validation rows={count} bytes={} samples=40 old_median_ns={} new_median_ns={} old_p95_ns={} new_p95_ns={}", + input.len(), + old[20], + new[20], + old[37], + new[37] + ); + } + } +}