Skip to content
69 changes: 27 additions & 42 deletions src/firebase_functions/options.py
Original file line number Diff line number Diff line change
Expand Up @@ -482,7 +482,6 @@ def _required_apis(self) -> list[_manifest.ManifestRequiredApi]:
]


# TODO refactor Storage & Database options to use this base class.
@_dataclasses.dataclass(frozen=True, kw_only=True)
class EventHandlerOptions(RuntimeOptions):
"""
Expand All @@ -502,11 +501,21 @@ def _endpoint(
assert kwargs["event_filters"] is not None
assert kwargs["event_type"] is not None

event_trigger = _manifest.EventTrigger(
eventType=kwargs["event_type"],
retry=self.retry if self.retry is not None else False,
eventFilters=kwargs["event_filters"],
)
event_filters_path_patterns = kwargs.get("event_filters_path_patterns") or None
retry = self.retry if self.retry is not None else False
if event_filters_path_patterns is not None:
event_trigger = _manifest.EventTrigger(
eventType=kwargs["event_type"],
retry=retry,
eventFilters=kwargs["event_filters"],
eventFilterPathPatterns=event_filters_path_patterns,
)
else:
event_trigger = _manifest.EventTrigger(
eventType=kwargs["event_type"],
retry=retry,
eventFilters=kwargs["event_filters"],
)

kwargs_merged = {
**_dataclasses.asdict(super()._endpoint(**kwargs)),
Expand Down Expand Up @@ -906,7 +915,7 @@ def _required_apis(self) -> list[_manifest.ManifestRequiredApi]:


@_dataclasses.dataclass(frozen=True, kw_only=True)
class StorageOptions(RuntimeOptions):
class StorageOptions(EventHandlerOptions):
"""
Options specific to Cloud Storage function types.
Internal use only.
Expand Down Expand Up @@ -936,21 +945,11 @@ def _endpoint(
event_filters: _typing.Any = {
"bucket": bucket,
}
event_trigger = _manifest.EventTrigger(
eventType=kwargs["event_type"],
retry=False,
eventFilters=event_filters,
)

kwargs_merged = {
**_dataclasses.asdict(super()._endpoint(**kwargs)),
"eventTrigger": event_trigger,
}
return _manifest.ManifestEndpoint(**_typing.cast(dict, kwargs_merged))
return super()._endpoint(**kwargs, event_filters=event_filters)


@_dataclasses.dataclass(frozen=True, kw_only=True)
class DatabaseOptions(RuntimeOptions):
class DatabaseOptions(EventHandlerOptions):
"""
Options specific to Realtime Database function types.
Internal use only.
Expand Down Expand Up @@ -989,19 +988,12 @@ def _endpoint(
else:
event_filters["instance"] = event_filter_instance

event_trigger = _manifest.EventTrigger(
eventType=kwargs["event_type"],
retry=False,
eventFilters=event_filters,
eventFilterPathPatterns=event_filters_path_patterns,
return super()._endpoint(
**kwargs,
event_filters=event_filters,
event_filters_path_patterns=event_filters_path_patterns,
)

kwargs_merged = {
**_dataclasses.asdict(super()._endpoint(**kwargs)),
"eventTrigger": event_trigger,
}
return _manifest.ManifestEndpoint(**_typing.cast(dict, kwargs_merged))


@_dataclasses.dataclass(frozen=True, kw_only=True)
class BlockingOptions(RuntimeOptions):
Expand Down Expand Up @@ -1056,7 +1048,7 @@ def _required_apis(self) -> list[_manifest.ManifestRequiredApi]:


@_dataclasses.dataclass(frozen=True, kw_only=True)
class FirestoreOptions(RuntimeOptions):
class FirestoreOptions(EventHandlerOptions):
"""
Options specific to Firestore function types.
Internal use only.
Expand Down Expand Up @@ -1096,19 +1088,12 @@ def _endpoint(
event_filters_path_patterns["document"] = event_filter_document
else:
event_filters["document"] = event_filter_document
event_trigger = _manifest.EventTrigger(
eventType=kwargs["event_type"],
retry=False,
eventFilters=event_filters,
eventFilterPathPatterns=event_filters_path_patterns,
return super()._endpoint(
**kwargs,
event_filters=event_filters,
event_filters_path_patterns=event_filters_path_patterns,
)

kwargs_merged = {
**_dataclasses.asdict(super()._endpoint(**kwargs)),
"eventTrigger": event_trigger,
}
return _manifest.ManifestEndpoint(**_typing.cast(dict, kwargs_merged))


@_dataclasses.dataclass(frozen=True, kw_only=True)
class HttpsOptions(RuntimeOptions):
Expand Down
18 changes: 18 additions & 0 deletions tests/test_db.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,24 @@ class TestDb(unittest.TestCase):
Tests for the db module.
"""

def test_database_decorator_retry_option(self):
func = mock.Mock(__name__="example_func")
decorated_func = db_fn.on_value_written(reference="/items/{itemId}", retry=True)(func)

endpoint = decorated_func.__firebase_endpoint__

self.assertIsNotNone(endpoint.eventTrigger)
self.assertTrue(endpoint.eventTrigger["retry"])

def test_database_decorator_retry_defaults_false(self):
func = mock.Mock(__name__="example_func")
decorated_func = db_fn.on_value_written(reference="/items/{itemId}")(func)

endpoint = decorated_func.__firebase_endpoint__

self.assertIsNotNone(endpoint.eventTrigger)
self.assertFalse(endpoint.eventTrigger["retry"])

def test_calls_init_function(self):
hello = None

Expand Down
39 changes: 39 additions & 0 deletions tests/test_firestore_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,45 @@ class TestFirestore(TestCase):
firestore_fn tests.
"""

def test_firestore_decorator_retry_option(self):
with patch.dict("sys.modules", mocked_modules):
from firebase_functions import firestore_fn

func = Mock(__name__="example_func")
decorated_func = firestore_fn.on_document_created(
document="/foo/{bar}",
retry=True,
)(func)

endpoint = decorated_func.__firebase_endpoint__

self.assertIsNotNone(endpoint.eventTrigger)
self.assertTrue(endpoint.eventTrigger["retry"])

def test_firestore_decorator_retry_defaults_false(self):
with patch.dict("sys.modules", mocked_modules):
from firebase_functions import firestore_fn

func = Mock(__name__="example_func")
decorated_func = firestore_fn.on_document_created(document="/foo/{bar}")(func)

endpoint = decorated_func.__firebase_endpoint__

self.assertIsNotNone(endpoint.eventTrigger)
self.assertFalse(endpoint.eventTrigger["retry"])

def test_firestore_decorator_omits_empty_path_patterns(self):
with patch.dict("sys.modules", mocked_modules):
from firebase_functions import firestore_fn

func = Mock(__name__="example_func")
decorated_func = firestore_fn.on_document_created(document="/foo/bar")(func)

endpoint = decorated_func.__firebase_endpoint__

self.assertIsNotNone(endpoint.eventTrigger)
self.assertNotIn("eventFilterPathPatterns", endpoint.eventTrigger)

def test_firestore_endpoint_handler_calls_function_with_correct_args(self):
with patch.dict("sys.modules", mocked_modules):
from cloudevents.http import CloudEvent
Expand Down
18 changes: 18 additions & 0 deletions tests/test_storage_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,24 @@ class TestStorage(unittest.TestCase):
Storage function tests.
"""

def test_storage_decorator_retry_option(self):
func = Mock(__name__="example_func")
decorated_func = storage_fn.on_object_finalized(bucket="bucket", retry=True)(func)

endpoint = decorated_func.__firebase_endpoint__

self.assertIsNotNone(endpoint.eventTrigger)
self.assertTrue(endpoint.eventTrigger["retry"])

def test_storage_decorator_retry_defaults_false(self):
func = Mock(__name__="example_func")
decorated_func = storage_fn.on_object_finalized(bucket="bucket")(func)

endpoint = decorated_func.__firebase_endpoint__

self.assertIsNotNone(endpoint.eventTrigger)
self.assertFalse(endpoint.eventTrigger["retry"])

def test_calls_init(self):
hello = None

Expand Down
Loading