diff --git a/litellm/litellm_core_utils/duration_parser.py b/litellm/litellm_core_utils/duration_parser.py index 7ac2299b7d8..06eb0aa2920 100644 --- a/litellm/litellm_core_utils/duration_parser.py +++ b/litellm/litellm_core_utils/duration_parser.py @@ -21,13 +21,19 @@ _BUDGET_DURATION_WORD_ALIASES: Final[dict[str, str]] = { "monthly": "30d", } - -def _normalize_duration(duration: str) -> str: - return _BUDGET_DURATION_WORD_ALIASES.get(duration.strip().lower(), duration) +_DURATION_RE: Final[re.Pattern[str]] = re.compile(r"(\d+)(mo|[smhdw])") -def _extract_from_regex(duration: str) -> tuple[int, str]: - match: Final = re.match(r"(\d+)(mo|[smhdw]?)", duration) +def _normalize_duration(duration: object) -> object: + if not isinstance(duration, str): + return duration + return _BUDGET_DURATION_WORD_ALIASES.get(duration.strip().lower(), duration.strip()) + + +def _extract_from_regex(duration: object) -> tuple[int, str]: + if not isinstance(duration, str): + raise ValueError("Invalid duration format") + match: Final = _DURATION_RE.fullmatch(duration.strip()) if not match: raise ValueError("Invalid duration format") @@ -189,9 +195,11 @@ def _setup_timezone(current_time: datetime, timezone_str: str = "UTC") -> tuple[ return current_time, tz -def _parse_duration(duration: str) -> tuple[int | None, str | None]: +def _parse_duration(duration: object) -> tuple[int | None, str | None]: """Parse the duration string into value and unit.""" - match: Final = re.match(r"(\d+)([a-z]+)", duration) + if not isinstance(duration, str): + return None, None + match: Final = _DURATION_RE.fullmatch(duration.strip()) if not match: return None, None diff --git a/tests/unit/litellm_core_utils/test_duration_parser.py b/tests/unit/litellm_core_utils/test_duration_parser.py index cb9f273a0a7..1904eb0dee4 100644 --- a/tests/unit/litellm_core_utils/test_duration_parser.py +++ b/tests/unit/litellm_core_utils/test_duration_parser.py @@ -385,5 +385,40 @@ class TestWordFormBudgetDurations(unittest.TestCase): self.assertIn("garbage", mock_warning.call_args.args) +class TestDurationPrefixRejection(unittest.TestCase): + def test_rejects_prefix_matched_garbage(self): + for bad in ("30dabc", "30days", "30s; DROP", "1moabc", "30dabc ", "30s;DROP TABLE"): + with self.subTest(bad=bad): + with self.assertRaises(ValueError): + duration_in_seconds(bad) + base_time = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) + with patch.object(duration_parser.verbose_logger, "warning") as mock_warning: + result = get_next_standardized_reset_time(bad, base_time, "UTC") + self.assertEqual(result, datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc)) + mock_warning.assert_called_once() + self.assertIn("Unrecognized budget_duration", mock_warning.call_args.args[0]) + self.assertEqual(mock_warning.call_args.args[1], bad) + + def test_accepts_valid_with_optional_strip(self): + self.assertEqual(duration_in_seconds("30d"), 30 * 86400) + self.assertEqual(duration_in_seconds("30d "), 30 * 86400) + self.assertEqual(duration_in_seconds(" 30d"), 30 * 86400) + with self.assertRaises(ValueError): + duration_in_seconds("30") + with self.assertRaises(ValueError): + duration_in_seconds("30D") + + def test_non_string_duration_is_rejected(self): + for bad in (None, 123, 30, 1.5, [], {}): # type: ignore[arg-type] + with self.subTest(bad=bad): + with self.assertRaises(ValueError): + duration_in_seconds(bad) # type: ignore[arg-type] + base_time = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) + with patch.object(duration_parser.verbose_logger, "warning") as mock_warning: + result = get_next_standardized_reset_time(bad, base_time, "UTC") # type: ignore[arg-type] + self.assertEqual(result, datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc)) + mock_warning.assert_called_once() + + if __name__ == "__main__": unittest.main()