diff --git a/tests/api.py b/tests/api.py index 2708ad8..0f048f6 100644 --- a/tests/api.py +++ b/tests/api.py @@ -6,6 +6,20 @@ import requests +def _json_response(content: bytes) -> MagicMock: + response = MagicMock() + response.content = content + return response + + +def _error_response(status_code: int) -> MagicMock: + error = requests.exceptions.HTTPError(f"{status_code} Client Error") + error.response = MagicMock(status_code=status_code, text="error_message") + response = MagicMock() + response.raise_for_status.side_effect = error + return response + + class TestYetiApi(unittest.TestCase): def setUp(self): self.api = YetiApi("http://fake-url") @@ -23,6 +37,108 @@ def test_auth_api_key(self, mock_post): headers={"x-yeti-apikey": "fake_apikey"}, ) + @patch("yeti.api.requests.Session.post") + def test_auth_api_key_revoked(self, mock_post): + """A key revoked mid-session raises instead of retrying without end.""" + mock_post.return_value = _json_response(b'{"access_token": "fake_token"}') + self.api.auth_api_key("fake_apikey") + mock_post.reset_mock() + mock_post.return_value = _error_response(401) + + with self.assertRaises(errors.YetiAuthError): + self.api.search_indicators(name="test") + # The search, then one attempt to renew the session. + self.assertEqual(mock_post.call_count, 2) + + @patch("yeti.api.requests.Session.post") + def test_auth_id_token(self, mock_post): + mock_post.return_value = _json_response(b'{"access_token": "fake_token"}') + + self.api.auth_id_token(lambda: "fake_id_token") + self.assertEqual(self.api.client.headers["authorization"], "Bearer fake_token") + mock_post.assert_called_with( + "http://fake-url/api/v2/auth/oidc-callback-token", + json={"id_token": "fake_id_token"}, + ) + + @patch("yeti.api.requests.Session.post") + def test_auth_google_access_token(self, mock_post): + mock_post.return_value = _json_response(b'{"access_token": "fake_token"}') + + self.api.auth_google_access_token(lambda: "fake_google_token") + self.assertEqual(self.api.client.headers["authorization"], "Bearer fake_token") + mock_post.assert_called_with( + "http://fake-url/api/v2/auth/google-access-token", + json={"access_token": "fake_google_token"}, + ) + + @patch("yeti.api.requests.Session.post") + def test_auth_google_access_token_refresh(self, mock_post): + """An expired session is renewed with a new token from the provider.""" + google_tokens = iter(["google_token_1", "google_token_2"]) + mock_post.side_effect = [ + _json_response(b'{"access_token": "session_token_1"}'), + _error_response(401), + _json_response(b'{"access_token": "session_token_2"}'), + _json_response(b'{"indicators": [{"name": "test"}]}'), + ] + + self.api.auth_google_access_token(lambda: next(google_tokens)) + result = self.api.search_indicators(name="test") + + self.assertEqual(result, [{"name": "test"}]) + exchanges = [ + call.kwargs["json"] + for call in mock_post.call_args_list + if call.args[0] == "http://fake-url/api/v2/auth/google-access-token" + ] + self.assertEqual( + exchanges, + [{"access_token": "google_token_1"}, {"access_token": "google_token_2"}], + ) + self.assertEqual( + self.api.client.headers["authorization"], "Bearer session_token_2" + ) + + @patch("yeti.api.requests.Session.post") + def test_auth_google_access_token_rejected(self, mock_post): + mock_post.return_value = _error_response(401) + token_provider = MagicMock(return_value="fake_google_token") + + with self.assertRaises(errors.YetiAuthError): + self.api.auth_google_access_token(token_provider) + token_provider.assert_called_once() + mock_post.assert_called_once() + + @patch("yeti.api.requests.Session.post") + def test_auth_google_access_token_not_enabled(self, mock_post): + mock_post.return_value = _error_response(404) + + with self.assertRaises(errors.YetiApiError) as raised: + self.api.auth_google_access_token(lambda: "fake_google_token") + self.assertEqual(raised.exception.status_code, 404) + + def test_auth_token_methods_require_provider(self): + with self.assertRaises(ValueError): + self.api.auth_id_token() + with self.assertRaises(ValueError): + self.api.auth_google_access_token() + + @patch("yeti.api.requests.Session.post") + def test_set_session_token_override(self, mock_post): + """Subclasses with their own transport receive the session token.""" + + class CustomTransportApi(YetiApi): + def _set_session_token(self, access_token): + self.session_token = access_token + + api = CustomTransportApi("http://fake-url") + mock_post.return_value = _json_response(b'{"access_token": "fake_token"}') + + api.auth_google_access_token(lambda: "fake_google_token") + self.assertEqual(api.session_token, "fake_token") + self.assertNotIn("authorization", api.client.headers) + @patch("yeti.api.requests.Session.post") def test_search_indicators(self, mock_post): mock_response = MagicMock() diff --git a/yeti/api.py b/yeti/api.py index 12b5974..af62211 100644 --- a/yeti/api.py +++ b/yeti/api.py @@ -4,7 +4,7 @@ import logging import urllib.parse import warnings -from typing import Any, Sequence +from typing import Any, Callable, Sequence import requests import requests_toolbelt.multipart.encoder as encoder @@ -20,6 +20,7 @@ OIDC_CALLBACK_ENDPOINT = "/api/v2/auth/oidc-callback-token" API_TOKEN_ENDPOINT = "/api/v2/auth/api-token" +GOOGLE_ACCESS_TOKEN_ENDPOINT = "/api/v2/auth/google-access-token" DFIQ_TYPE_DEPRECATION = ( "dfiq_type is ignored: Yeti infers the DFIQ type from the payload. The " @@ -63,6 +64,10 @@ YetiObject = dict[str, Any] YetiLinkObject = dict[str, Any] +# Returns a credential to exchange for a Yeti session token. Auth methods call +# it again for a fresh credential each time the session token expires. +TokenProvider = Callable[[], str] + logger = logging.getLogger(__name__) handler = logging.StreamHandler() @@ -91,11 +96,17 @@ def __init__(self, url_root: str, tls_cert: str | None = None): self._url_root = url_root self._auth_function = "" - self._auth_function_map = { + # refresh_auth calls these without arguments, so each auth method must + # remember the credential it was given. + self._auth_function_map: dict[str, Callable[[], None]] = { "auth_api_key": self.auth_api_key, + "auth_id_token": self.auth_id_token, + "auth_google_access_token": self.auth_google_access_token, } self._apikey = None + self._id_token_provider: TokenProvider | None = None + self._google_access_token_provider: TokenProvider | None = None def do_request( self, @@ -170,24 +181,101 @@ def auth_api_key(self, apikey: str | None = None) -> None: if not self._apikey: raise ValueError("No API key provided.") + self._start_session(API_TOKEN_ENDPOINT, headers={"x-yeti-apikey": self._apikey}) + self._auth_function = "auth_api_key" + + def auth_id_token(self, token_provider: TokenProvider | None = None) -> None: + """Authenticates a session using an OpenID Connect ID token. + + The server must use OIDC authentication. It verifies the token as a + Google ID token, checks that its audience is the server's OIDC client ID + or one of `auth.oidc_extra_client_audiences`, and maps its email address + to an existing, enabled user. + + Args: + token_provider: Returns an ID token. Called again for a fresh token + each time the session token expires. Defaults to the provider + from the previous call. + + Raises: + ValueError: If no provider was given in this or a previous call. + YetiAuthError: If the server rejects the token. + """ + if token_provider is not None: + self._id_token_provider = token_provider + if self._id_token_provider is None: + raise ValueError("No ID token provider set.") + + self._start_session( + OIDC_CALLBACK_ENDPOINT, json_data={"id_token": self._id_token_provider()} + ) + self._auth_function = "auth_id_token" + + def auth_google_access_token( + self, token_provider: TokenProvider | None = None + ) -> None: + """Authenticates a session using a Google OAuth access token. + + This suits clients that can get Google access tokens without a browser, + but not ID tokens. The server must use OIDC authentication and set + `auth.google_access_token_client_ids`, or the endpoint answers 404. The + token must be issued to one of those OAuth clients, carry only identity + scopes (openid, email, profile), and belong to the verified email + address of an existing, enabled user. + + Args: + token_provider: Returns a Google access token, for example by + refreshing google-auth credentials and returning their token. + Called again for a fresh token each time the session token + expires. Defaults to the provider from the previous call. + + Raises: + ValueError: If no provider was given in this or a previous call. + YetiAuthError: If the server rejects the token. + YetiApiError: If the server doesn't offer the exchange (404), or + can't validate tokens at the moment (503). + """ + if token_provider is not None: + self._google_access_token_provider = token_provider + if self._google_access_token_provider is None: + raise ValueError("No Google access token provider set.") + + self._start_session( + GOOGLE_ACCESS_TOKEN_ENDPOINT, + json_data={"access_token": self._google_access_token_provider()}, + ) + self._auth_function = "auth_google_access_token" + + def _start_session( + self, + endpoint: str, + json_data: dict[str, Any] | None = None, + headers: dict[str, Any] | None = None, + ) -> None: + """Exchanges a credential at an auth endpoint for a session token.""" + # No retries: a 401 here means the credential was rejected, and a retry + # would call refresh_auth, which calls back into this method. response = self.do_request( "POST", - f"{self._url_root}{API_TOKEN_ENDPOINT}", - headers={"x-yeti-apikey": self._apikey}, + f"{self._url_root}{endpoint}", + json_data=json_data, + headers=headers, + retries=0, ) - access_token = json.loads(response).get("access_token") if not access_token: raise RuntimeError( f"Failed to find access token in the response: {response}" ) - authd_session = requests.Session() - if self._tls_cert: - authd_session.verify = self._tls_cert # type: ignore - authd_session.headers.update({"authorization": f"Bearer {access_token}"}) - self.client = authd_session + self._set_session_token(access_token) - self._auth_function = "auth_api_key" + def _set_session_token(self, access_token: str) -> None: + """Sends a session token with subsequent requests. + + Subclasses that send requests through another transport override this + to attach the token there. + """ + self.client.headers["authorization"] = f"Bearer {access_token}" def refresh_auth(self): if self._auth_function: