"""Base class for all clients.""" from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Collection, Type from ddtrace import tracer from stream.client.base import BaseStreamClient from stream.utils import validate_feed_slug, validate_user_id from ..integrations import flask from ..throttling import BaseThrottler, RateLimitThrottler from ..utils import datadog, injector, ratelimit if TYPE_CHECKING: from stream.feed.base import BaseFeed __all__ = ["BaseExtendedClient", "Activity"] Activity = dict[str, Any] class BaseExtendedClient(BaseStreamClient, ABC): """Base class for both sync and async clients.""" __trace_methods: tuple[str, ...] = ( "follow_many", "unfollow_many", "update_activities", "get_activities", "activities_partial_update", ) def __init__( self, *args, throttler: BaseThrottler = RateLimitThrottler(), inject_trace_context: bool = True, inject_correlation_id: bool = True, **kwargs, ): """Patch StreamClient methods to add datadog traces. :param control_rpm: Whether to proactively reduce request rate. :param throttler: Delay calculator used to control request rate limiting. :param inject_trace_context: Whether to inject trace context to activity data. :param inject_correlation_id: Whether to inject correlation id to activity data. """ super().__init__(*args, **kwargs) self._inject_trace_context = inject_trace_context self._inject_correlation_id = inject_correlation_id self._throttler = throttler self._ratelimit_info: ratelimit.RatelimitInfo = ratelimit.RatelimitInfo() self._correlation_id: str | None = None datadog.wrap_methods(self, self.__trace_methods) @abstractmethod def _get_feed_cls(self) -> Type["BaseFeed"]: """Return the feed class (sync or async). Must be overridden by subclasses. """ ... def feed(self, feed_slug: str, user_id: str) -> "BaseFeed": """Override to use custom Feed class.""" feed_slug = validate_feed_slug(feed_slug) user_id = validate_user_id(user_id) token = self.create_jwt_token("feed", "*", feed_id="*") feed_cls = self._get_feed_cls() return feed_cls(self, feed_slug, user_id, token) @property def ratelimit_info(self) -> ratelimit.RatelimitInfo: """Return last API call ratelimit info.""" return self._ratelimit_info def set_correlation_id(self, correlation_id: str | None) -> None: """Set correlation id for future API calls.""" self._correlation_id = correlation_id def get_correlation_id(self) -> str | None: """Get correlation id from different sources by priority.""" return self._correlation_id or flask.get_correlation_id() def inject_trace_data(self, dest: Activity | Collection[Activity]) -> None: """Inject trace data into activity data. :param dest: activity or collection of activities to inject trace data. """ if isinstance(dest, dict): dest = [dest] context = tracer.current_trace_context() correlation_id = self.get_correlation_id() for activity in dest: if self._inject_trace_context: injector.inject_trace_context(activity, context) if self._inject_correlation_id: injector.inject_correlation_id(activity, correlation_id)