Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
116 changes: 116 additions & 0 deletions tests/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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()
Expand Down
110 changes: 99 additions & 11 deletions yeti/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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 "
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
Loading