diff --git a/src/CalcManager/CalcManager.vcxproj b/src/CalcManager/CalcManager.vcxproj index 462e97a9c3..7c31549282 100644 --- a/src/CalcManager/CalcManager.vcxproj +++ b/src/CalcManager/CalcManager.vcxproj @@ -131,6 +131,7 @@ Use false true + Guard /Zm250 /std:c++20 /permissive- /Zc:twoPhase- /utf-8 /w44242 %(AdditionalOptions) $(SolutionDir)..\src\;%(AdditionalIncludeDirectories) Level4 @@ -167,6 +168,7 @@ Use false true + Guard /Zm250 /std:c++20 /permissive- /Zc:twoPhase- /utf-8 /w44242 %(AdditionalOptions) $(SolutionDir)..\src\;%(AdditionalIncludeDirectories) Level4 @@ -202,6 +204,7 @@ Use false true + Guard /Zm250 /std:c++20 /permissive- /Zc:twoPhase- /utf-8 /w44242 %(AdditionalOptions) $(SolutionDir)..\src\;%(AdditionalIncludeDirectories) Level4 diff --git a/src/CalcManager/Ratpack/conv.cpp b/src/CalcManager/Ratpack/conv.cpp index 4f7609de09..27c745a90b 100644 --- a/src/CalcManager/Ratpack/conv.cpp +++ b/src/CalcManager/Ratpack/conv.cpp @@ -1500,34 +1500,30 @@ void numpowi32(_Inout_ PNUMBER* proot, int32_t power, uint32_t radix, int32_t pr void ratpowi32(_Inout_ PRAT* proot, int32_t power, int32_t precision) { - if (power < 0) - { - // Take the positive power and invert answer. - PNUMBER pnumtemp = nullptr; - ratpowi32(proot, -power, precision); - pnumtemp = (*proot)->pp; - (*proot)->pp = (*proot)->pq; - (*proot)->pq = pnumtemp; - } - else - { - PRAT lret = nullptr; + const bool invertResult = power < 0; + uint32_t magnitude = + invertResult ? static_cast(-static_cast(power)) : static_cast(power); - lret = i32torat(1); + PRAT lret = i32torat(1); - while (power > 0) + while (magnitude > 0) + { + if (magnitude & 1) { - if (power & 1) - { - mulnumx(&(lret->pp), (*proot)->pp); - mulnumx(&(lret->pq), (*proot)->pq); - } - mulrat(proot, *proot, precision); - trimit(&lret, precision); - trimit(proot, precision); - power >>= 1; + mulnumx(&(lret->pp), (*proot)->pp); + mulnumx(&(lret->pq), (*proot)->pq); } - destroyrat(*proot); - *proot = lret; + mulrat(proot, *proot, precision); + trimit(&lret, precision); + trimit(proot, precision); + magnitude >>= 1; + } + destroyrat(*proot); + *proot = lret; + + if (invertResult) + { + // Invert after exponentiation so INT32_MIN never needs to be negated as an int32_t. + std::swap((*proot)->pp, (*proot)->pq); } } diff --git a/src/Calculator.Tests/StandardCalculatorViewModelTests.cs b/src/Calculator.Tests/StandardCalculatorViewModelTests.cs index 88c46ed245..4a3d570996 100644 --- a/src/Calculator.Tests/StandardCalculatorViewModelTests.cs +++ b/src/Calculator.Tests/StandardCalculatorViewModelTests.cs @@ -5,8 +5,10 @@ using System.Collections.Generic; using System.Linq; using Microsoft.VisualStudio.TestTools.UnitTesting; +using CalcManager.Interop; using CalculatorApp.ViewModel; using CalculatorApp.ViewModel.Common; +using CalculatorApp.ViewModel.Snapshot; namespace Calculator.Tests { @@ -152,6 +154,58 @@ private void ValidateViewModelValueAndSecondaryExpression(string value, string e } } + private void ValidateSnapshotReplayRejectsInvalidCommands(bool includeExpression) + { + int[] invalidCommands = + { + (int)CalculatorCommand.CommandEQU, + (int)CalculatorCommand.CommandCLEAR, + (int)CalculatorCommand.CommandCENTR, + (int)CalculatorCommand.CommandBACK, + (int)CalculatorCommand.CommandMCLEAR, + -1, + int.MaxValue + }; + + foreach (int invalidCommand in invalidCommands) + { + var snapshot = new StandardCalculatorSnapshot(); + if (includeExpression) + { + snapshot.ExpressionDisplay = new ExpressionDisplaySnapshot(); + } + AppendSnapshotOperand(snapshot, (int)CalculatorCommand.Command1); + AppendSnapshotBinaryCommand(snapshot, invalidCommand); + AppendSnapshotOperand(snapshot, (int)CalculatorCommand.Command2); + + _viewModel.Snapshot = snapshot; + + ValidateViewModelValue("12"); + } + } + + private static void AppendSnapshotOperand(StandardCalculatorSnapshot snapshot, int command) + { + snapshot.DisplayCommands.Add(new ExpressionCommandWrapper( + CommandType.OperandCommand, + 0, + new[] { command }, + false, + false, + false)); + } + + private static void AppendSnapshotBinaryCommand(StandardCalculatorSnapshot snapshot, int command) + { + snapshot.DisplayCommands.Add(new ExpressionCommandWrapper( + CommandType.BinaryCommand, + command, + Array.Empty(), + false, + false, + false)); + } + #endregion #region Constructor Tests @@ -173,6 +227,56 @@ public void ViewModelConstructorButtonPressedInitializedTest() Assert.IsNotNull(vm.ButtonPressed); } + [TestMethod] + public void SnapshotReplayRejectsInvalidCommandsWithoutExpression() + { + ValidateSnapshotReplayRejectsInvalidCommands(false); + } + + [TestMethod] + public void SnapshotReplayRejectsInvalidCommandsWithExpression() + { + ValidateSnapshotReplayRejectsInvalidCommands(true); + } + + [TestMethod] + public void SnapshotReplayRejectsModeCommands() + { + int[] modeCommands = + { + (int)NumbersAndOperatorsEnum.Degree, + (int)NumbersAndOperatorsEnum.Radians, + (int)NumbersAndOperatorsEnum.Grads + }; + + foreach (int modeCommand in modeCommands) + { + var snapshot = new StandardCalculatorSnapshot(); + AppendSnapshotBinaryCommand(snapshot, modeCommand); + AppendSnapshotOperand(snapshot, (int)CalculatorCommand.Command9); + AppendSnapshotOperand(snapshot, (int)CalculatorCommand.Command0); + AppendSnapshotBinaryCommand(snapshot, (int)CalculatorCommand.CommandSIN); + + _viewModel.Snapshot = snapshot; + + ValidateViewModelValue("1"); + } + } + + [TestMethod] + public void SnapshotReplayRejectsMemoryCommands() + { + var snapshot = new StandardCalculatorSnapshot(); + AppendSnapshotOperand(snapshot, (int)CalculatorCommand.Command1); + AppendSnapshotBinaryCommand(snapshot, (int)CalculatorCommand.CommandMPLUS); + AppendSnapshotOperand(snapshot, (int)CalculatorCommand.Command2); + + _viewModel.Snapshot = snapshot; + + ValidateViewModelValue("12"); + Assert.IsTrue(_viewModel.IsMemoryEmpty); + } + #endregion #region Basic Arithmetic Tests diff --git a/src/Calculator.ViewModels/StandardCalculatorViewModel.cs b/src/Calculator.ViewModels/StandardCalculatorViewModel.cs index e715cfed77..2f5cba2ed2 100644 --- a/src/Calculator.ViewModels/StandardCalculatorViewModel.cs +++ b/src/Calculator.ViewModels/StandardCalculatorViewModel.cs @@ -708,10 +708,7 @@ public Snapshot.StandardCalculatorSnapshot Snapshot { // Expression was not evaluated before, or it was an error. var displayCommands = GetCommandsFromExpressionCommands(snapshot.DisplayCommands); - foreach (var cmd in displayCommands) - { - _standardCalculatorManager?.SendCommand((CalculatorCommand)cmd); - } + ReplaySnapshotCommands(displayCommands); if (snapshot.PrimaryDisplay.IsError) { SetPrimaryDisplay(snapshot.PrimaryDisplay.DisplayValue, true); @@ -727,10 +724,7 @@ public Snapshot.StandardCalculatorSnapshot Snapshot else { var commands = GetCommandsFromExpressionCommands(snapshot.DisplayCommands); - foreach (var cmd in commands) - { - _standardCalculatorManager?.SendCommand((CalculatorCommand)cmd); - } + ReplaySnapshotCommands(commands); } } } @@ -2055,6 +2049,34 @@ private static List GetCommandsFromExpressionCommands(IList commands) + { + foreach (int command in commands) + { + if (IsValidSnapshotReplayCommand(command)) + { + _standardCalculatorManager?.SendCommand((CalculatorCommand)command); + } + } + } + + private static bool IsValidSnapshotReplayCommand(int command) + { + return (command >= (int)CalculatorCommand.Command0 && command <= (int)CalculatorCommand.CommandF) + || (command >= (int)CalculatorCommand.CommandAnd && command <= (int)CalculatorCommand.CommandPWR) + || (command >= (int)CalculatorCommand.CommandLogBaseY && command <= (int)CalculatorCommand.CommandNor) + || command == (int)CalculatorCommand.CommandRSHFL + || command == (int)CalculatorCommand.CommandSIGN + || command == (int)CalculatorCommand.CommandPNT + || (command >= (int)CalculatorCommand.CommandCHOP && command <= (int)CalculatorCommand.CommandPERCENT) + || command == (int)CalculatorCommand.CommandEXP + || command == (int)CalculatorCommand.CommandOPENP + || command == (int)CalculatorCommand.CommandCLOSEP + || (command >= (int)CalculatorCommand.CommandASIN && command <= (int)CalculatorCommand.CommandATANH) + || (command >= (int)CalculatorCommand.CommandSEC && command <= (int)CalculatorCommand.CommandRORC) + || command == (int)NumbersAndOperatorsEnum.Degrees; + } + } } diff --git a/src/CalculatorUITests/StandardModeFunctionalTests.cs b/src/CalculatorUITests/StandardModeFunctionalTests.cs index ae23efa2f4..244a0d4dea 100644 --- a/src/CalculatorUITests/StandardModeFunctionalTests.cs +++ b/src/CalculatorUITests/StandardModeFunctionalTests.cs @@ -8,7 +8,14 @@ using OpenQA.Selenium; using System; +using System.Collections.Generic; +using System.Diagnostics; +using System.IO; +using System.IO.Compression; +using System.Text; +using System.Text.Json; using System.Text.RegularExpressions; +using System.Threading; namespace CalculatorUITests { @@ -786,6 +793,90 @@ public void AoT_ErrorMessage_MessageRetentionUponExitingAoT() #endregion + #region Snapshot Protocol Tests + + [TestMethod] + [DataRow(false)] + [DataRow(true)] + [Priority(0)] + public void SnapshotProtocolRejectsInjectedCommands(bool includeExpression) + { + var standardSnapshot = new Dictionary + { + ["m"] = new Dictionary(), + ["p"] = new Dictionary + { + ["d"] = "12", + ["e"] = false, + }, + ["c"] = new object[] + { + BinaryCommand(322), + OperandCommand(139), + BinaryCommand(121), + OperandCommand(130), + BinaryCommand(102), + BinaryCommand(int.MaxValue), + }, + }; + if (includeExpression) + { + standardSnapshot["e"] = new Dictionary + { + ["t"] = Array.Empty(), + ["c"] = Array.Empty(), + }; + } + + var snapshot = new Dictionary + { + ["m"] = 0, + ["s"] = standardSnapshot, + }; + var json = JsonSerializer.Serialize(snapshot); + var uri = $"ms-calculator:snapshot/{Convert.ToBase64String(Compress(json))}"; + + Process.Start(new ProcessStartInfo(uri) { UseShellExecute = true }); + + Assert.IsTrue( + SpinWait.SpinUntil(() => page.CalculatorResults.GetCalculatorResultText() == "0", TimeSpan.FromSeconds(10)), + $"Snapshot replay produced '{page.CalculatorResults.GetCalculatorResultText()}' instead of rejecting the injected commands."); + } + + private static Dictionary BinaryCommand(int command) + { + return new Dictionary + { + ["$t"] = 1, + ["c"] = command, + }; + } + + private static Dictionary OperandCommand(int command) + { + return new Dictionary + { + ["$t"] = 2, + ["n"] = false, + ["d"] = false, + ["s"] = false, + ["c"] = new[] { command }, + }; + } + + private static byte[] Compress(string value) + { + using var compressed = new MemoryStream(); + using (var stream = new DeflateStream(compressed, CompressionMode.Compress, true)) + { + var bytes = Encoding.UTF8.GetBytes(value); + stream.Write(bytes, 0, bytes.Length); + } + return compressed.ToArray(); + } + + #endregion + #region Error Tests /// /// These automated tests verify errors display gracefully diff --git a/src/CalculatorUnitTests/RationalTest.cpp b/src/CalculatorUnitTests/RationalTest.cpp index 7c8aebb6f0..486248b713 100644 --- a/src/CalculatorUnitTests/RationalTest.cpp +++ b/src/CalculatorUnitTests/RationalTest.cpp @@ -3,6 +3,7 @@ #include "pch.h" #include +#include #include "Header Files/Rational.h" #include "Header Files/RationalMath.h" @@ -40,6 +41,16 @@ TEST_METHOD(TestModuloOperandsNotModified) VERIFY_ARE_EQUAL(rat4, 4); } +TEST_METHOD(TestInt32MinIntegerPower) +{ + PRAT base = Rational(-1).ToPRAT(); + ratpowi32(&base, std::numeric_limits::min(), RATIONAL_PRECISION); + Rational result{ base }; + destroyrat(base); + + VERIFY_ARE_EQUAL(result, 1); +} + TEST_METHOD(TestModuloInteger) { // Check with integers