This commit is contained in:
Arnab Mandal 2026-10-04 05:30:12 -04:00 • committed by GitHub
commit bf2ac44510
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 50 additions and 7 deletions

View file

@ -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

View file

@ -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()