mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge a86760cd1d into bb4f7211d7
This commit is contained in:
commit
bf2ac44510
2 changed files with 50 additions and 7 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue