diff --git a/README.md b/README.md index 9c2bc80..12ab185 100644 --- a/README.md +++ b/README.md @@ -32,6 +32,43 @@ cio.track(customer_id="5", name="purchased") cio.track(customer_id="5", name="purchased", data={"price": 23.45}) ``` +### Quick start + +Create a client with your secret key (`ak_…`). It sends tracking calls to the Track API and transactional messages to the App API, and picks the US or EU region from the key. + +```python +import os + +from customerio import Client + +cio = Client(api_key=os.environ["CIO_API_KEY"]) +cio.identify(id="user_123", email="a@example.com") +cio.track(customer_id="user_123", name="order_completed") +cio.send_email( + { + "to": "a@example.com", + "identifiers": {"id": "user_123"}, + "transactional_message_id": 5, + } +) +``` + +Pass `region` to override the region in the key. Legacy App API keys also work here; they default to `Regions.US` unless you pass `region`. + +`identify`, `track`, `track_anonymous`, `pageview` and the `send_*` methods are on the client directly. The rest of the Track and App API are on `cio.track_client` and `cio.api_client`. + +If the key can't do what you asked, the call raises a `KeyCapabilityError` (a subclass of `CustomerIOException`). A public key (`wk_…`) raises it before any request when you send a transactional message. A 401 or 403 from the server is re-raised as a `KeyCapabilityError` too, chained from the original error. + +```python +from customerio import KeyCapabilityError + +try: + cio.send_email(request) +except KeyCapabilityError as e: + # e.capability == "transactional", e.status_code == 401 + print(e) # This key can't send transactional messages (401): ... +``` + ### Instantiating customer.io object Create an instance of the client with your [Customer.io credentials](https://fly.customer.io/settings/api_credentials). diff --git a/customerio/__init__.py b/customerio/__init__.py index 2955934..9c8b397 100644 --- a/customerio/__init__.py +++ b/customerio/__init__.py @@ -7,14 +7,17 @@ SendSMSRequest, SendWhatsAppRequest, ) -from customerio.client_base import CustomerIOException +from customerio.client import Client +from customerio.client_base import CustomerIOException, KeyCapabilityError from customerio.regions import Regions from customerio.track import CustomerIO __all__ = [ "APIClient", + "Client", "CustomerIO", "CustomerIOException", + "KeyCapabilityError", "Regions", "SendEmailRequest", "SendInAppRequest", diff --git a/customerio/client.py b/customerio/client.py new file mode 100644 index 0000000..36b4d7b --- /dev/null +++ b/customerio/client.py @@ -0,0 +1,203 @@ +""" +Implements a single client for the Track and App APIs, authenticated with one key. +""" + +import re +from functools import wraps +from urllib.parse import urlsplit + +from .api import APIClient +from .client_base import CustomerIOException, KeyCapabilityError +from .regions import Region, Regions +from .track import CustomerIO + +# Keys in the shared format look like `{ak|wk}_{us|eu}_{random}_{checksum}`. +# The server checks the checksum; the client only reads the prefix to pick a host. +KEY_PREFIX = re.compile(r"^(ak|wk)_(us|eu)_") + +CAPABILITY_ACTIONS = { + "track": "send tracking data", + "transactional": "send transactional messages", +} + + +class _BearerTrackClient(CustomerIO): + """Track API client that sends the key as a Bearer token instead of Basic auth.""" + + def _build_session(self): + session = super()._build_session() + session.auth = None + session.headers["Authorization"] = f"Bearer {self.api_key}" + return session + + +def _guard(capability): + """Re-raises a 401 or 403 as a KeyCapabilityError so a key mismatch is never a bare + auth failure.""" + + def decorator(method): + @wraps(method) + def wrapper(*args, **kwargs): + try: + return method(*args, **kwargs) + except KeyCapabilityError: + raise + except CustomerIOException as e: + status = getattr(e, "status_code", None) + if status not in (401, 403): + raise + message = ( + f"This key can't {CAPABILITY_ACTIONS[capability]} " + f"({status}): {_server_message(e.response)}" + ) + raise KeyCapabilityError(capability, message, status_code=status) from e + + return wrapper + + return decorator + + +def _server_message(response): + try: + return response.json()["meta"]["error"] + except Exception: + return response.text + + +class Client: + """One client for tracking and transactional messages, authenticated with a single key. + + Tracking methods (``identify``, ``track``, ``track_anonymous``, ``pageview``) go to + the Track API; ``send_*`` methods go to the App API. Both send the key as a Bearer + token. For ``ak_us_…`` / ``ak_eu_…`` keys the region comes from the key; an explicit + ``region`` wins. Legacy keys default to ``Regions.US``. + + When the key can't do what a method asks, the method raises a + :class:`KeyCapabilityError`: before the request when the key's type says so (a + public ``wk_`` key can't send transactional messages), otherwise when the server + answers 401 or 403. + + Other Track and App API methods are on ``track_client`` and ``api_client``. + """ + + def __init__( + self, + api_key, + region=None, + track_url=None, + api_url=None, + retries=3, + timeout=10, + backoff_factor=0.02, + use_connection_pooling=True, + ): + if not api_key: + raise CustomerIOException("api_key is required") + if region is not None and not isinstance(region, Region): + raise CustomerIOException("invalid region provided") + + match = KEY_PREFIX.match(api_key) + self._key_kind = match.group(1) if match else "legacy" + key_region = (Regions.EU if match.group(2) == "eu" else Regions.US) if match else None + self.region = region or key_region or Regions.US + + options = dict( + retries=retries, + timeout=timeout, + backoff_factor=backoff_factor, + use_connection_pooling=use_connection_pooling, + ) + + track_host = track_port = track_prefix = None + if track_url: + parts = urlsplit(track_url) + track_host, track_port, track_prefix = parts.hostname, parts.port, parts.path or None + + #: Track API client, authenticated with the key as a Bearer token. + self.track_client = _BearerTrackClient( + api_key=api_key, + host=track_host, + port=track_port, + url_prefix=track_prefix, + region=self.region, + **options, + ) + self._api_client = APIClient(api_key, url=api_url, region=self.region, **options) + + def __enter__(self): + return self + + def __exit__(self, *args): + self.close() + + def close(self): + try: + self.track_client.close() + finally: + self._api_client.close() + + @property + def api_client(self): + """App API client, authenticated with the key. + + Raises :class:`KeyCapabilityError` if the key is a public ``wk_`` key. + """ + if self._key_kind == "wk": + raise KeyCapabilityError( + "transactional", + "This key can't use the App API: public keys (wk_) are for tracking only. " + "Use a secret key (ak_).", + ) + return self._api_client + + @_guard("track") + def identify(self, id, **kwargs): + """Create or update a person. See :meth:`CustomerIO.identify`.""" + return self.track_client.identify(id, **kwargs) + + @_guard("track") + def track(self, customer_id, name, data=None, id=None, timestamp=None): + """Track an event for a person. See :meth:`CustomerIO.track`.""" + return self.track_client.track(customer_id, name, data=data, id=id, timestamp=timestamp) + + @_guard("track") + def track_anonymous(self, anonymous_id, name, data=None, id=None, timestamp=None): + """Track an event for an anonymous visitor. See :meth:`CustomerIO.track_anonymous`.""" + return self.track_client.track_anonymous( + anonymous_id, name, data=data, id=id, timestamp=timestamp + ) + + @_guard("track") + def pageview(self, customer_id, page, **data): + """Track a page view for a person. See :meth:`CustomerIO.pageview`.""" + return self.track_client.pageview(customer_id, page, **data) + + @_guard("transactional") + def send_email(self, request): + """Send a transactional email. See :meth:`APIClient.send_email`.""" + return self.api_client.send_email(request) + + @_guard("transactional") + def send_push(self, request): + """Send a transactional push. See :meth:`APIClient.send_push`.""" + return self.api_client.send_push(request) + + @_guard("transactional") + def send_sms(self, request): + """Send a transactional SMS. See :meth:`APIClient.send_sms`.""" + return self.api_client.send_sms(request) + + @_guard("transactional") + def send_whatsapp(self, request): + """Send a transactional WhatsApp message. See :meth:`APIClient.send_whatsapp`.""" + return self.api_client.send_whatsapp(request) + + @_guard("transactional") + def send_inbox_message(self, request): + """Send a transactional inbox message. See :meth:`APIClient.send_inbox_message`.""" + return self.api_client.send_inbox_message(request) + + @_guard("transactional") + def send_in_app(self, request): + """Send a transactional in-app message. See :meth:`APIClient.send_in_app`.""" + return self.api_client.send_in_app(request) diff --git a/customerio/client_base.py b/customerio/client_base.py index 4e63dd7..6c7bbe8 100644 --- a/customerio/client_base.py +++ b/customerio/client_base.py @@ -50,6 +50,23 @@ class CustomerIOException(Exception): pass +class KeyCapabilityError(CustomerIOException): + """Raised by :class:`customerio.Client` when its key can't do what a method asks. + + Raised before the request when the key's type rules it out (a public ``wk_`` + key can't send transactional messages). Otherwise the server's 401 or 403 is + re-raised as this error, chained from the original :class:`CustomerIOException`. + + ``capability`` is ``"track"`` or ``"transactional"``. ``status_code`` is the + HTTP status when the server refused the key, else ``None``. + """ + + def __init__(self, capability, message, status_code=None): + super().__init__(message) + self.capability = capability + self.status_code = status_code + + class ClientBase: def __init__(self, retries=3, timeout=10, backoff_factor=0.02, use_connection_pooling=True): self.timeout = timeout @@ -100,7 +117,10 @@ def send_request(self, method, url, data): result_status = response.status_code if result_status < 200 or result_status >= 300: - raise CustomerIOException(f"{result_status}: {url} {data} {response.text}") + error = CustomerIOException(f"{result_status}: {url} {data} {response.text}") + error.status_code = result_status + error.response = response + raise error return response except CustomerIOException: diff --git a/tests/test_client.py b/tests/test_client.py new file mode 100644 index 0000000..84feea1 --- /dev/null +++ b/tests/test_client.py @@ -0,0 +1,247 @@ +import json +import unittest +from functools import partial + +import urllib3 + +from customerio import ( + Client, + CustomerIOException, + KeyCapabilityError, + Regions, + SendEmailRequest, + SendInAppRequest, + SendInboxMessageRequest, + SendPushRequest, + SendSMSRequest, + SendWhatsAppRequest, +) +from tests.server import HTTPSTestCase + +# test uses a self signed certificate so disable the warning messages +urllib3.disable_warnings() + +US_KEY = "ak_us_0123456789ABCDEFGHIJKLMNOPQRSTUV_000000" +EU_KEY = "ak_eu_0123456789ABCDEFGHIJKLMNOPQRSTUV_000000" +PUBLIC_KEY = "wk_us_0123456789ABCDEFGHIJKLMNOPQRSTUV_000000" +LEGACY_KEY = "abc123" + +EMAIL = {"to": "a@example.com", "identifiers": {"id": "user_123"}, "transactional_message_id": 5} + + +class FakeResponse: + def __init__(self, status_code, body): + self.status_code = status_code + self.text = json.dumps(body) + + def json(self): + return json.loads(self.text) + + +class FakeSession: + def __init__(self, response): + self.response = response + self.request_count = 0 + + def request(self, *args, **kwargs): + self.request_count += 1 + return self.response + + def close(self): + pass + + +class TestClientSetup(unittest.TestCase): + def test_requires_api_key(self): + with self.assertRaises(CustomerIOException): + Client(api_key="") + + def test_rejects_invalid_region(self): + with self.assertRaises(CustomerIOException): + Client(api_key=US_KEY, region="eu") + + def test_region_comes_from_the_key(self): + self.assertEqual(Client(api_key=US_KEY).region, Regions.US) + self.assertEqual(Client(api_key=EU_KEY).region, Regions.EU) + self.assertEqual(Client(api_key=PUBLIC_KEY).region, Regions.US) + + cio = Client(api_key=EU_KEY) + self.assertEqual(cio.track_client.base_url, f"https://{Regions.EU.track_host}/api/v1") + self.assertEqual(cio.api_client.url, f"https://{Regions.EU.api_host}") + + def test_explicit_region_wins_over_the_key(self): + cio = Client(api_key=EU_KEY, region=Regions.US) + self.assertEqual(cio.region, Regions.US) + self.assertEqual(cio.track_client.base_url, f"https://{Regions.US.track_host}/api/v1") + self.assertEqual(cio.api_client.url, f"https://{Regions.US.api_host}") + + def test_legacy_keys_keep_the_region_option_and_us_default(self): + self.assertEqual(Client(api_key=LEGACY_KEY).region, Regions.US) + self.assertEqual(Client(api_key=LEGACY_KEY, region=Regions.EU).region, Regions.EU) + + def test_both_clients_use_bearer_auth(self): + cio = Client(api_key=US_KEY) + self.assertIsNone(cio.track_client.http.auth) + self.assertEqual(cio.track_client.http.headers["Authorization"], f"Bearer {US_KEY}") + self.assertEqual(cio.api_client.http.headers["Authorization"], f"Bearer {US_KEY}") + + +class TestClientRequests(HTTPSTestCase): + """Starts server which the client connects to in the following tests""" + + def setUp(self): + url = f"https://{self.server.server_address[0]}:{self.server.server_port}" + self.cio = self._client(US_KEY, url) + + def _client(self, key, url): + cio = Client(api_key=key, track_url=url, api_url=url, retries=5, backoff_factor=0) + + # do not verify the ssl certificate as it is self signed + # should only be done for tests + cio.track_client.http.verify = False + cio._api_client.http.verify = False + return cio + + def _check_request(self, resp, rq, *args, **kwargs): + request = resp.request + self.assertEqual(request.method, rq["method"]) + self.assertEqual(request.headers["Authorization"], rq["authorization"]) + self.assertTrue( + request.url.endswith(rq["url_suffix"]), + "url: {} expected suffix: {}".format(request.url, rq["url_suffix"]), + ) + + def _expect(self, http, method, url_suffix, key=US_KEY): + http.hooks = dict( + response=partial( + self._check_request, + rq={"method": method, "authorization": f"Bearer {key}", "url_suffix": url_suffix}, + ) + ) + + def test_tracking_goes_to_the_track_api_with_bearer(self): + http = self.cio.track_client.http + + self._expect(http, "PUT", "/api/v1/customers/user_123") + self.cio.identify("user_123", email="a@example.com") + + self._expect(http, "POST", "/api/v1/customers/user_123/events") + self.cio.track("user_123", "order_completed", data={"price": 1}) + + self._expect(http, "POST", "/api/v1/events") + self.cio.track_anonymous("anon_1", "viewed") + + self._expect(http, "POST", "/api/v1/customers/user_123/events") + self.cio.pageview("user_123", "/pricing") + + def test_send_methods_go_to_the_app_api_with_bearer(self): + http = self.cio.api_client.http + identifiers = {"id": "user_123"} + sends = [ + (self.cio.send_email, SendEmailRequest(**EMAIL), "email"), + (self.cio.send_push, SendPushRequest(identifiers=identifiers), "push"), + (self.cio.send_sms, SendSMSRequest(identifiers=identifiers), "sms"), + (self.cio.send_whatsapp, SendWhatsAppRequest(identifiers=identifiers), "whatsapp"), + ( + self.cio.send_inbox_message, + SendInboxMessageRequest(identifiers=identifiers), + "inbox_message", + ), + (self.cio.send_in_app, SendInAppRequest(identifiers=identifiers), "in_app"), + ] + + for send, request, path in sends: + self._expect(http, "POST", f"/v1/send/{path}") + self.assertEqual(send(request), {}) + + self._expect(http, "POST", "/v1/send/email") + self.cio.send_email(dict(EMAIL)) + + def test_public_key_can_still_track(self): + url = f"https://{self.server.server_address[0]}:{self.server.server_port}" + cio = self._client(PUBLIC_KEY, url) + + self._expect(cio.track_client.http, "POST", "/api/v1/customers/user_123/events", PUBLIC_KEY) + cio.track("user_123", "viewed") + + +class TestClientKeyCapability(unittest.TestCase): + def _fail_with(self, client, status, body): + session = FakeSession(FakeResponse(status, body)) + client._build_session = lambda: session + client._current_session = None + return session + + def test_public_key_raises_on_transactional_calls_without_a_request(self): + cio = Client(api_key=PUBLIC_KEY) + session = self._fail_with(cio._api_client, 200, {}) + + for send in [ + cio.send_email, + cio.send_push, + cio.send_sms, + cio.send_whatsapp, + cio.send_inbox_message, + cio.send_in_app, + ]: + with self.assertRaises(KeyCapabilityError) as ctx: + send(dict(EMAIL)) + self.assertEqual(ctx.exception.capability, "transactional") + self.assertIsNone(ctx.exception.status_code) + self.assertIn("public keys (wk_) are for tracking only", str(ctx.exception)) + + with self.assertRaises(KeyCapabilityError): + cio.api_client # noqa: B018 + + self.assertEqual(session.request_count, 0) + + def test_401_from_the_app_api_becomes_key_capability_error(self): + cio = Client(api_key=US_KEY) + self._fail_with(cio._api_client, 401, {"meta": {"error": "unauthorized"}}) + + with self.assertRaises(KeyCapabilityError) as ctx: + cio.send_email(dict(EMAIL)) + + self.assertEqual(ctx.exception.capability, "transactional") + self.assertEqual(ctx.exception.status_code, 401) + self.assertEqual( + str(ctx.exception), "This key can't send transactional messages (401): unauthorized" + ) + self.assertIsInstance(ctx.exception.__cause__, CustomerIOException) + self.assertNotIsInstance(ctx.exception.__cause__, KeyCapabilityError) + + def test_403_from_the_track_api_becomes_key_capability_error(self): + cio = Client(api_key=US_KEY) + self._fail_with(cio.track_client, 403, {"meta": {"error": "forbidden"}}) + + with self.assertRaises(KeyCapabilityError) as ctx: + cio.identify("1") + + self.assertEqual(ctx.exception.capability, "track") + self.assertEqual(ctx.exception.status_code, 403) + self.assertEqual(str(ctx.exception), "This key can't send tracking data (403): forbidden") + self.assertIsInstance(ctx.exception.__cause__, CustomerIOException) + + def test_other_request_errors_pass_through_unchanged(self): + cio = Client(api_key=US_KEY) + self._fail_with(cio.track_client, 400, {"meta": {"error": "bad request"}}) + + with self.assertRaises(CustomerIOException) as ctx: + cio.identify("1") + + self.assertNotIsInstance(ctx.exception, KeyCapabilityError) + self.assertIn("400", str(ctx.exception)) + + def test_missing_params_still_raise_before_a_request(self): + cio = Client(api_key=US_KEY) + session = self._fail_with(cio.track_client, 200, {}) + + with self.assertRaises(CustomerIOException) as ctx: + cio.identify("") + + self.assertNotIsInstance(ctx.exception, KeyCapabilityError) + self.assertEqual(session.request_count, 0) + + +if __name__ == "__main__": + unittest.main()