diff --git a/koboldcpp.py b/koboldcpp.py index d2db55e9c..aa7352a38 100644 --- a/koboldcpp.py +++ b/koboldcpp.py @@ -1152,14 +1152,15 @@ def tryparseint(value,fallback): return 0 try: return int(value) - except ValueError: + except (TypeError, ValueError, OverflowError): return fallback def tryparsefloat(value,fallback): if value is None: return fallback try: - return float(value) - except ValueError: + parsed = float(value) + return parsed if math.isfinite(parsed) else fallback + except (TypeError, ValueError, OverflowError): return fallback def replace_last_in_string(text: str, match: str, replacement: str) -> str: diff --git a/tests/test_numeric_parsing.py b/tests/test_numeric_parsing.py new file mode 100644 index 000000000..877440c83 --- /dev/null +++ b/tests/test_numeric_parsing.py @@ -0,0 +1,39 @@ +import math +import os +import sys +import unittest + + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) + +import koboldcpp + + +class NumericParsingTests(unittest.TestCase): + def test_valid_numeric_values_keep_existing_behavior(self): + self.assertEqual(koboldcpp.tryparseint("12", 17), 12) + self.assertEqual(koboldcpp.tryparseint("true", 17), 1) + self.assertEqual(koboldcpp.tryparseint("false", 17), 0) + self.assertEqual(koboldcpp.tryparsefloat("0.75", 0.25), 0.75) + + def test_tryparseint_returns_fallback_for_incompatible_json_types(self): + for value in ([], {}, [1], {"value": 1}): + with self.subTest(value=value): + self.assertEqual(koboldcpp.tryparseint(value, 17), 17) + + def test_tryparseint_returns_fallback_for_overflow(self): + self.assertEqual(koboldcpp.tryparseint(math.inf, 17), 17) + + def test_tryparsefloat_returns_fallback_for_incompatible_json_types(self): + for value in ([], {}, [1], {"value": 1}): + with self.subTest(value=value): + self.assertEqual(koboldcpp.tryparsefloat(value, 0.25), 0.25) + + def test_tryparsefloat_returns_fallback_for_non_finite_values(self): + for value in (math.inf, -math.inf, math.nan, "inf", "nan", 10**10000): + with self.subTest(value=value): + self.assertEqual(koboldcpp.tryparsefloat(value, 0.25), 0.25) + + +if __name__ == "__main__": + unittest.main()