package ai.aiflow.mathink import org.junit.Assert.assertEquals import org.junit.Assert.assertFalse import org.junit.Assert.assertTrue import org.junit.Test class AIFlowMathInkTest { /** 필요 변수: 결정적 logit. 작동 원리: native runtime 없이 API와 privacy 계약을 검증한다. */ private class FakeSession : LiteRtSession { override val modelVersion = "test-0.6" override val exactLabelCount = 3 var onlineCalls = 0 var rasterCalls = 0 override fun runOnline(features: FloatArray): FloatArray { assertEquals(128 * 19, features.size) onlineCalls += 1 return floatArrayOf(0.0f, 3.0f, 1.0f) } override fun runRaster(rasterInk: FloatArray): FloatArray { assertEquals(128 * 128, rasterInk.size) rasterCalls += 1 return floatArrayOf(2.0f, 0.0f, 1.0f) } override fun runRasterDebug(rasterInk: FloatArray): RasterInference { assertEquals(128 * 128, rasterInk.size) rasterCalls += 1 return RasterInference( exactLogits = floatArrayOf(2.0f, 0.0f, 1.0f), virtualHypotheses = List(4) { hypothesis -> VirtualStrokeHypothesis( points = FloatArray(128 * 2) { hypothesis.toFloat() }, stateLogits = FloatArray(128 * 3), progress = FloatArray(128) { it / 127.0f }, logProbability = -hypothesis.toFloat(), ) }, ) } override fun close() = Unit } private fun stroke(withTime: Boolean = true) = InkStroke( strokeId = 7, order = 0, points = listOf( InkPoint(10.0f, 10.0f, if (withTime) 0L else null), InkPoint(40.0f, 40.0f, if (withTime) 200L else null), InkPoint(80.0f, 20.0f, if (withTime) 400L else null), ), ) @Test fun canonicalEncoderPreservesRawAndObservedAnchors() { val ink = CanonicalTapEncoder().encodeOnline( listOf(stroke()), InkCanvas(100.0f, 100.0f), ) assertEquals(TimestampMode.OBSERVED, ink.timestampMode) assertEquals(7, ink.rawStrokes.single().strokeId) assertTrue(ink.canonicalTaps.first().strokeStart) assertTrue(ink.canonicalTaps.last().strokeEnd) assertTrue(ink.canonicalTaps.last().penUp) assertEquals(128 * 19, ink.features.size) assertTrue((0 until 128).all { ink.features[it * 19 + 17] == 0.0f }) } @Test fun canonicalFeaturesMatchPythonTrainingPreprocessorWithinFloatTolerance() { val ink = CanonicalTapEncoder().encodeOnline( listOf(stroke()), InkCanvas(100.0f, 100.0f), ) val expected = mapOf( 0 to floatArrayOf( 0.0f, 0.0f, 0.0625f, 0.3125f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 2.3333333f, 0.1f, 0.4f, 0.3f, 0.25f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, ), 1 to floatArrayOf( 0.006602364f, 0.0185185f, 0.06827707f, 0.31827706f, 0.33582273f, 0.94192517f, 0.0f, 0.0f, 0.007874016f, 2.3333333f, 0.1f, 0.4f, 0.3f, 0.25f, 1.0f, 0.0030811024f, 1.3258247f, 0.0f, 0.0f, ), 63 to floatArrayOf( 0.4375688f, 0.9423623f, 0.4453727f, 0.60648084f, 0.804563f, -0.5938673f, -0.000012785f, 0.0f, 0.496063f, 2.3333333f, 0.1f, 0.4f, 0.3f, 0.25f, 1.0f, 0.0033267364f, 1.2279294f, 0.0f, 0.0f, ), 127 to floatArrayOf( 1.0f, 0.40069035f, 0.9375f, 0.4375f, 0.580574f, -0.81420743f, 0.0f, 0.0f, 1.0f, 2.3333333f, 0.1f, 0.4f, 0.3f, 0.25f, 1.0f, 0.002923004f, 1.3975347f, 0.0f, 0.0f, ), ) expected.forEach { (row, values) -> values.forEachIndexed { channel, value -> assertEquals( "row=$row channel=$channel", value, ink.features[row * 19 + channel], 1e-4f, ) } } } @Test fun missingTimestampUsesCanonicalModeAndMask() { val ink = CanonicalTapEncoder().encodeOnline( listOf(stroke(withTime = false)), InkCanvas(100.0f, 100.0f), ) assertEquals(TimestampMode.CANONICAL, ink.timestampMode) assertTrue((0 until 128).all { ink.features[it * 19 + 17] == 1.0f }) } @Test fun recognizeReturnsOnlyPublicFields() { val session = FakeSession() val runtime = AIFlowMathInk(session, listOf("0", "x", "+")) val result = runtime.recognizeOnline( listOf(stroke()), InkCanvas(100.0f, 100.0f), ) assertEquals("x", result.candidates.first().token) assertEquals(result.candidates.first().probability, result.confidence) assertEquals(1, session.onlineCalls) val payload = result.toServerPayload() assertEquals(setOf("candidates", "confidence", "modelVersion", "latencyMs"), payload.keys) assertFalse(payload.containsKey("rawStrokes")) assertFalse(payload.containsKey("canonicalTaps")) assertFalse(payload.containsKey("virtualHypotheses")) } @Test fun rasterUsesLocalSessionAndReturnsTopK() { val session = FakeSession() val runtime = AIFlowMathInk(session, listOf("0", "x", "+"), topK = 2) val result = runtime.recognizeRaster( RasterInput(128, 128, FloatArray(128 * 128)), ) assertEquals(listOf("0", "+"), result.candidates.map { it.token }) assertEquals(1, session.rasterCalls) } @Test fun rasterDebugReturnsFourLocalHypothesesWithoutServerPayload() { val session = FakeSession() val runtime = AIFlowMathInk(session, listOf("0", "x", "+"), topK = 2) val debug = runtime.recognizeRasterDebug( RasterInput(128, 128, FloatArray(128 * 128)), ) assertEquals("0", debug.symbol.candidates.first().token) assertEquals(4, debug.virtualHypotheses.size) assertEquals(128 * 2, debug.virtualHypotheses.first().points.size) assertEquals(1, session.rasterCalls) assertFalse(debug.symbol.toServerPayload().containsKey("virtualHypotheses")) } @Test fun androidBenchmarkUsesNearestRankAndAllReleaseChecks() { class Probe : AndroidBenchmarkProbe { var nanos = 0L var pss = 20L * 1024L * 1024L override fun nowNanos() = nanos override fun totalPssBytes() = pss override fun batteryChargeMicroAh() = 1000 } val probe = Probe() val session = object : LiteRtSession { override val modelVersion = "benchmark-0.6" override val exactLabelCount = 3 override fun runOnline(features: FloatArray): FloatArray { probe.nanos += 10_000_000L return floatArrayOf(1.0f, 0.0f, 0.0f) } override fun runRaster(rasterInk: FloatArray): FloatArray { probe.nanos += 100_000_000L return floatArrayOf(1.0f, 0.0f, 0.0f) } override fun close() = Unit } val report = AndroidBenchmarkRunner(session, probe).run( deviceTier = AndroidDeviceTier.LOW, onlineModelSha256 = "a".repeat(64), rasterModelSha256 = "b".repeat(64), onlineInputs = listOf(FloatArray(128 * 19)), rasterInputs = listOf(FloatArray(128 * 128)), config = AndroidBenchmarkConfig(warmupRuns = 1, measuredRuns = 5), deviceManufacturer = "fixture", deviceModel = "fixture-low", sdkInt = 24, ) assertEquals(10.0, report.online.p95Ms, 0.0) assertEquals(100.0, report.raster.p95Ms, 0.0) assertEquals(20L * 1024L * 1024L, report.peakPssBytes) assertTrue(report.gatePassed) assertEquals(false, report.productValidation) assertEquals("low", report.toMap()["device_tier"]) assertEquals( modelBundleSha256("a".repeat(64), "b".repeat(64)), report.modelBundleSha256, ) } @Test fun androidBenchmarkFailsAnyExceededThreshold() { val summary = summarizeLatency(listOf(1.0, 2.0, 3.0, 100.0)) assertEquals(100.0, summary.p95Ms, 0.0) assertEquals(2.0, summary.p50Ms, 0.0) } }