Skip to content

Commit 411c34d

Browse files
fix(fcm): encode durations without floating-point rounding errors
encode_ttl and encode_milliseconds computed the nanoseconds from the float returned by timedelta.total_seconds(), so some durations were sent with rounding errors: an AndroidConfig.ttl of 86400.9 seconds (or timedelta(days=1, microseconds=900000)) was encoded as "86400.899999999s", and 1.123457 as "1.123456999s". The Duration string is now built from the exact days/seconds/microseconds of the timedelta.
1 parent de45607 commit 411c34d

2 files changed

Lines changed: 15 additions & 12 deletions

File tree

‎firebase_admin/_messaging_encoder.py‎

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,6 @@
1616

1717
import datetime
1818
import json
19-
import math
2019
import numbers
2120
import re
2221
import warnings
@@ -275,14 +274,9 @@ def encode_ttl(cls, ttl):
275274
if not isinstance(ttl, datetime.timedelta):
276275
raise ValueError('AndroidConfig.ttl must be a duration in seconds or an instance of '
277276
'datetime.timedelta.')
278-
total_seconds = ttl.total_seconds()
279-
if total_seconds < 0:
277+
if ttl < datetime.timedelta(0):
280278
raise ValueError('AndroidConfig.ttl must not be negative.')
281-
seconds = int(math.floor(total_seconds))
282-
nanos = int((total_seconds - seconds) * 1e9)
283-
if nanos:
284-
return f'{seconds}.{str(nanos).zfill(9)}s'
285-
return f'{seconds}s'
279+
return cls.encode_duration(ttl)
286280

287281
@classmethod
288282
def encode_milliseconds(cls, label, msec):
@@ -294,11 +288,17 @@ def encode_milliseconds(cls, label, msec):
294288
if not isinstance(msec, datetime.timedelta):
295289
raise ValueError(
296290
f'{label} must be a duration in milliseconds or an instance of datetime.timedelta.')
297-
total_seconds = msec.total_seconds()
298-
if total_seconds < 0:
291+
if msec < datetime.timedelta(0):
299292
raise ValueError(f'{label} must not be negative.')
300-
seconds = int(math.floor(total_seconds))
301-
nanos = int((total_seconds - seconds) * 1e9)
293+
return cls.encode_duration(msec)
294+
295+
@classmethod
296+
def encode_duration(cls, duration):
297+
"""Encodes a non-negative ``datetime.timedelta`` into a protobuf Duration string."""
298+
# Use the exact integer fields of the timedelta instead of total_seconds(), which
299+
# is a float and can turn e.g. 86400.9 seconds into 86400.899999999s.
300+
seconds = duration.days * 86400 + duration.seconds
301+
nanos = duration.microseconds * 1000
302302
if nanos:
303303
return f'{seconds}.{str(nanos).zfill(9)}s'
304304
return f'{seconds}s'

‎tests/test_messaging.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -466,6 +466,9 @@ def test_android_config(self):
466466
(123, '123s'),
467467
(123.45, '123.450000000s'),
468468
(datetime.timedelta(days=1, seconds=100), '86500s'),
469+
(1.123457, '1.123457000s'),
470+
(86400.9, '86400.900000000s'),
471+
(datetime.timedelta(days=1, microseconds=900000), '86400.900000000s'),
469472
])
470473
def test_android_ttl(self, ttl):
471474
msg = messaging.Message(

0 commit comments

Comments
 (0)