diff --git a/.gitignore b/.gitignore index 0e32d89ea..ccb892270 100644 --- a/.gitignore +++ b/.gitignore @@ -151,3 +151,4 @@ debug_* htmlcov htmlcov /CLAUDE.md +.env.example diff --git a/CHANGELOG.rst b/CHANGELOG.rst index a1904b884..1a858204e 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -6,6 +6,23 @@ History All release highlights of this project will be documented in this file. +4.6.1 - Sep 13, 2026 +____________________ + + +**Added** + + - ``SAORGClient`` New client for using an Organization API Key to perform organization-level operations across teams. + + - ``SAORGClient.list_teams()`` Returns the teams in an organization. + + - ``SAORGClient.get_team_client(team_id)`` Returns a SAClient instance for the specified team. + +### Updated + + - ``SAClient()`` Added support for authentication with Organization API Keys by providing a `team_id`. + + 4.6.0 - Aug 16, 2026 ____________________ diff --git a/docs/source/api_reference/api_org_client.rst b/docs/source/api_reference/api_org_client.rst new file mode 100644 index 000000000..85b4e4c65 --- /dev/null +++ b/docs/source/api_reference/api_org_client.rst @@ -0,0 +1,8 @@ +===================== +SAORGClient interface +===================== + +.. _ref_org_client: + +.. automethod:: superannotate.SAORGClient.get_team_client +.. automethod:: superannotate.SAORGClient.list_teams diff --git a/docs/source/api_reference/index.rst b/docs/source/api_reference/index.rst index 7ce246c14..31c31d7a0 100644 --- a/docs/source/api_reference/index.rst +++ b/docs/source/api_reference/index.rst @@ -10,5 +10,6 @@ Contents :maxdepth: 2 api_client + api_org_client api_metadata helpers diff --git a/docs/source/userguide/quickstart.rst b/docs/source/userguide/quickstart.rst index 5a72270eb..888a351ac 100644 --- a/docs/source/userguide/quickstart.rst +++ b/docs/source/userguide/quickstart.rst @@ -38,8 +38,14 @@ on the team setup page, for more details please visit our documentation at https - **Team API key** — scoped to one team. Works with ``SAClient``. - **Personal (team-user) API key** — scoped to one team, tied to your user. Works with ``SAClient``. -- **Organization API key** — not scoped to a team. Not supported by the SDK; - ``SAClient`` will reject it. +- **Organization API key** — not scoped to a team, so the team to operate in must be supplied alongside it: + + .. code-block:: python + + SAClient(token="", team_id=) + + Instead of passing ``team_id``, you can set ``SA_TEAM_ID`` as an environment variable or in your config file. An explicit ``team_id`` argument takes precedence. Omitting the team entirely raises ``AppException``. + To work across several teams, or when the team isn't known in advance, use ``SAORGClient`` instead — see below. SAClient can be used with or without arguments @@ -76,6 +82,13 @@ ______________________________________________ sa_client = SAClient(token="") +An Organization API key carries no team, so it is passed together with the team to +operate in: + +.. code-block:: python + + sa_client = SAClient(token="", team_id=) + *Method 2:* Create a custom config file: @@ -93,9 +106,12 @@ Custom config.ini example: [DEFAULT] SA_TOKEN = + ; Only an Organization API key needs it; other keys carry their own team. + SA_TEAM_ID = LOGGING_LEVEL = INFO LOGGING_PATH = /Users/username/data/superannotate_logs + ---------- @@ -177,3 +193,38 @@ A team contributor can be invited to the team with: .. code-block:: python sa_client.invite_contributors_to_team(emails=["admin@superannotate.com"], admin=False) + + +---------- + + +SAORGClient: organization-scoped access +======================================== + +``SAORGClient`` authorizes the SDK at the organization level rather than within a single team. Use it when a script operates across several teams, or when the team isn't known in advance. + +.. code-block:: python + + from superannotate import SAORGClient + + + org_client = SAORGClient(token="") + # List the teams in the organization + org_client.list_teams() + + # Get a team-scoped client for one of them + sa_client = org_client.get_team_client(team_id=12345) + sa_client.list_projects() + +``get_team_client(team_id)`` returns a standard ``SAClient`` bound to that team. It supports the full team-level SDK surface, and no Team API key is created. +As with ``SAClient``, if no ``token`` argument is given, ``SAORGClient`` reads ``SA_TOKEN`` from the environment or from your config file. + +**Important notes:** + +- ``SAORGClient`` requires an Organization API key. Passing a Team or Personal API + key raises ``AppException``. +- The team is fixed when ``get_team_client()`` is called. To work in a different + team, call ``get_team_client()`` again with the other ID. + +For the full list of organization-level methods, see the +:ref:`full method reference `. diff --git a/pytest.ini b/pytest.ini index c0f66b58e..29a68f0a7 100644 --- a/pytest.ini +++ b/pytest.ini @@ -3,4 +3,4 @@ minversion = 3.7 log_cli=true python_files = test_*.py ;pytest_plugins = ['pytest_profiling'] -addopts = -n 6 --dist loadscope +;addopts = -n 8 --dist loadscope diff --git a/src/superannotate/__init__.py b/src/superannotate/__init__.py index 7af435220..5796d01eb 100644 --- a/src/superannotate/__init__.py +++ b/src/superannotate/__init__.py @@ -2,7 +2,7 @@ import os import sys -__version__ = "4.6.0" +__version__ = "4.6.1" os.environ.update({"sa_version": __version__}) @@ -15,10 +15,12 @@ from lib.core import PACKAGE_VERSION_INFO_MESSAGE from lib.core import PACKAGE_VERSION_MAJOR_UPGRADE from lib.core.exceptions import AppException +from lib.core.exceptions import SAAuthError from lib.core.exceptions import FileChangedError from superannotate.lib.app.input_converters import export_annotation from superannotate.lib.app.input_converters import import_annotation from superannotate.lib.app.interface.sdk_interface import SAClient +from superannotate.lib.app.interface.sdk_interface import SAORGClient from superannotate.lib.app.interface.sdk_interface import ItemContext SESSIONS = {} @@ -27,10 +29,12 @@ __all__ = [ "__version__", "SAClient", + "SAORGClient", "ItemContext", # Utils "enums", "AppException", + "SAAuthError", "FileChangedError", "import_annotation", "export_annotation", diff --git a/src/superannotate/lib/app/interface/base_interface.py b/src/superannotate/lib/app/interface/base_interface.py index 4b0db4080..07a442a43 100644 --- a/src/superannotate/lib/app/interface/base_interface.py +++ b/src/superannotate/lib/app/interface/base_interface.py @@ -1,11 +1,13 @@ +from __future__ import annotations + import functools import json +import logging import os import platform import sys from collections.abc import Iterable from collections.abc import Sized -from functools import lru_cache from inspect import signature from pathlib import Path from types import FunctionType @@ -13,73 +15,156 @@ import lib.core as constants from lib.app.interface.types import validate_arguments from lib.core import CONFIG +from lib.core import CREDENTIALS_NOT_FOUND_ERROR +from lib.core import INVALID_CREDENTIALS_ERROR +from lib.core import INVALID_TOKEN_ERROR from lib.core import setup_logging from lib.core.entities.base import ConfigEntity -from lib.core.entities.base import TokenStr from lib.core.exceptions import AppException -from lib.infrastructure.controller import Controller +from lib.core.exceptions import SAAuthError +from lib.infrastructure.controller import BaseController from lib.infrastructure.utils import extract_project_folder_inputs from lib.infrastructure.validators import wrap_error from mixpanel import Mixpanel from pydantic import ValidationError +logger = logging.getLogger("sa") + class BaseInterfaceFacade: - REGISTRY = [] + """Credential resolution, shared by every facade. + + What a token has to resolve to is the facade's own business - named by its + CONTROLLER_CLASS rather than configured through this __init__: SAClient needs a + team, SAORGClient needs an organization key. + """ + + #: The controller class this facade drives. Named for the class, not the + #: instance: `self.controller` is the built one. + CONTROLLER_CLASS: type[BaseController] @validate_arguments - def __init__(self, token: TokenStr | None = None, config_path: str | None = None): + def __init__( + self, + # Plain str, not TokenStr: the token's shape is checked by ConfigEntity in + # _resolve_config, so a malformed one is reported as the credential failure it + # is (SAAuthError) rather than as a generic bad argument. + token: str | None = None, + config_path: str | None = None, + team_id: int | None = None, + *, + config: dict | None = None, + ): + resolved = self._resolve_config(token, config_path, config) + if team_id is not None: + # An explicit team_id wins over whatever the config source provided. + resolved.TEAM_ID = team_id + setup_logging(resolved.LOGGING_LEVEL, resolved.LOGGING_PATH) + self.controller = self.CONTROLLER_CLASS(resolved) + + @classmethod + def _resolve_config( + cls, + token: str | None, + config_path: str | None, + settings: dict | None = None, + ) -> ConfigEntity: + """Resolve credentials and settings. + + Credentials come from the first source that carries them: the ``token`` + argument, a config file, inline ``settings``, then the environment and the + default config files. Inline settings are applied over whatever that produced, + so a caller can configure a client - another backend, bigger chunks - without + writing a file for it. + + Shared by every facade's ``__init__`` (``SAClient``, ``SAORGClient``). + """ + settings = cls._validated_settings(settings) try: if token: - config = ConfigEntity(SA_TOKEN=token) + config = ConfigEntity(**{**settings, "SA_TOKEN": token}) elif config_path: - config_path = Path(config_path).expanduser() - if not Path(config_path).is_file() or not os.access( - config_path, os.R_OK - ): - raise AppException( - f"SuperAnnotate config file {str(config_path)} not found." - ) - if config_path.suffix == ".json": - config = self._retrieve_configs_from_json(config_path) - else: - config = self._retrieve_configs_from_ini(config_path) - + config = cls._merge_settings( + cls._resolve_config_from_path(config_path), settings + ) + elif "SA_TOKEN" in settings: + config = ConfigEntity(**settings) else: - config = self._retrieve_configs_from_env() - if not config: - if Path(constants.CONFIG_INI_FILE_LOCATION).exists(): - config = self._retrieve_configs_from_ini( - constants.CONFIG_INI_FILE_LOCATION - ) - elif Path(constants.CONFIG_JSON_FILE_LOCATION).exists(): - config = self._retrieve_configs_from_json( - constants.CONFIG_JSON_FILE_LOCATION - ) - else: - raise AppException( - "Credentials not found: SA_TOKEN environment variable is not set and " - f"config file '{constants.CONFIG_INI_FILE_LOCATION}' was not found." - ) from None + config = cls._merge_settings( + cls._resolve_config_from_env_or_files(), settings + ) except ValidationError as e: - raise AppException(wrap_error(e)) + raise SAAuthError(wrap_error(e)) if not config: - raise AppException("Credentials not provided.") - setup_logging(config.LOGGING_LEVEL, config.LOGGING_PATH) - self.controller = Controller(config) - BaseInterfaceFacade.REGISTRY.append(self) + raise SAAuthError(INVALID_CREDENTIALS_ERROR) + return config + + @staticmethod + def _validated_settings(settings: dict | None) -> dict: + """Inline settings, with anything the SDK has no such setting for rejected. + + ConfigEntity ignores what it does not recognise, so a mistyped key would + otherwise be dropped without a word - and a mistyped SA_TOKEN would send the + client off to authenticate as whatever the environment happens to hold. + """ + if not settings: + return {} + known = set(ConfigEntity.model_fields) | { + field.alias + for field in ConfigEntity.model_fields.values() + if field.alias is not None + } + unknown = sorted(set(settings) - known) + if unknown: + raise AppException( + f"Unknown configuration: {', '.join(unknown)}. " + f"Available: {', '.join(sorted(known))}." + ) + return dict(settings) @staticmethod - def _retrieve_configs_from_json(path: Path) -> ConfigEntity: + def _merge_settings(config: ConfigEntity, settings: dict) -> ConfigEntity: + """``config`` with inline settings applied over it.""" + if not settings: + return config + return ConfigEntity(**{**config.model_dump(by_alias=True), **settings}) + + @classmethod + def _resolve_config_from_path(cls, config_path: str) -> ConfigEntity: + """A config file the caller named explicitly (``.ini`` or ``.json``).""" + path = Path(config_path).expanduser() + if not path.is_file() or not os.access(path, os.R_OK): + raise AppException(f"SuperAnnotate config file {config_path} not found.") + if path.suffix == ".json": + return cls._retrieve_configs_from_json(path) + return cls._retrieve_configs_from_ini(path) + + @classmethod + def _resolve_config_from_env_or_files(cls) -> ConfigEntity: + """No token or config path given: ``SA_TOKEN``, then the default config files.""" + config = cls._retrieve_configs_from_env() + if config: + return config + if Path(constants.CONFIG_INI_FILE_LOCATION).exists(): + return cls._retrieve_configs_from_ini(constants.CONFIG_INI_FILE_LOCATION) + if Path(constants.CONFIG_JSON_FILE_LOCATION).exists(): + return cls._retrieve_configs_from_json(constants.CONFIG_JSON_FILE_LOCATION) + raise SAAuthError(CREDENTIALS_NOT_FOUND_ERROR) + + @staticmethod + def _retrieve_configs_from_json(path: Path | str) -> ConfigEntity: with open(path) as json_file: json_data = json.load(json_file) token = json_data["token"] try: config = ConfigEntity(SA_TOKEN=token) except ValidationError: - raise AppException("Invalid token.") + raise SAAuthError(INVALID_TOKEN_ERROR) host = json_data.get("main_endpoint") verify_ssl = json_data.get("ssl_verify") + team_id = json_data.get("team_id") + if team_id is not None: + config.TEAM_ID = int(team_id) if host: config.API_URL = host if verify_ssl: @@ -87,7 +172,7 @@ def _retrieve_configs_from_json(path: Path) -> ConfigEntity: return config @staticmethod - def _retrieve_configs_from_ini(path: Path) -> ConfigEntity: + def _retrieve_configs_from_ini(path: Path | str) -> ConfigEntity: import configparser config_parser = configparser.ConfigParser() @@ -114,23 +199,30 @@ def _retrieve_configs_from_env() -> ConfigEntity | None: class Tracker: - def get_mp_instance(self) -> Mixpanel: - client = self.get_client() - if client.controller._config.API_URL == constants.BACKEND_URL: # noqa + def get_mp_instance(self, client, explicit_credentials: bool = False) -> Mixpanel: + # client may have no .controller yet (e.g. __init__ failed before setting one). + controller = getattr(client, "controller", None) + if controller is not None: + api_url = controller.config["SA_URL"] + elif explicit_credentials: + api_url = constants.BACKEND_URL + else: + api_url = os.environ.get("SA_URL", constants.BACKEND_URL) + if api_url == constants.BACKEND_URL: mp_token = "ca95ed96f80e8ec3be791e2d3097cf51" else: mp_token = "e741d4863e7e05b1a45833d01865ef0d" return Mixpanel(mp_token) @staticmethod - @lru_cache - def get_default_payload(team_name, user_email, auth_type): + def get_default_payload(team_name, user_email, auth_type) -> dict: + """Built fresh per event: a cached dict would be shared between callers, and + would freeze sa_version and SA_ENV as they were on the first tracked call. + """ return { "SDK": True, "Team": team_name, "User Email": user_email, - # How the client authenticated: "sdk" for a legacy team token, - # "api_key" for a scoped API key. "Auth Type": auth_type, "Version": os.environ["sa_version"], "Python version": platform.python_version(), @@ -140,27 +232,16 @@ def get_default_payload(team_name, user_email, auth_type): def __init__(self, function): self.function = function - self._client = None - self.skip_flag = os.environ.get("SA_SKIP_METRICS", "False").lower() in ( - "true", - "1", - "t", - ) functools.update_wrapper(self, function) - def get_client(self): - if not self._client: - if BaseInterfaceFacade.REGISTRY: - return BaseInterfaceFacade.REGISTRY[-1] - else: - from lib.app.interface.sdk_interface import SAClient + @staticmethod + def _metrics_disabled() -> bool: + """Whether the caller turned telemetry off. - try: - return SAClient() - except Exception: - pass - elif hasattr(self._client, "controller"): - return self._client + Read per call: a Tracker is built while its class is being created, so a value + captured in __init__ would ignore anything set after ``import superannotate``. + """ + return os.environ.get("SA_SKIP_METRICS", "False").lower() in ("true", "1", "t") @staticmethod def extract_arguments(function, *args, **kwargs) -> dict: @@ -194,38 +275,86 @@ def default_parser(function_name: str, kwargs: dict) -> tuple: properties[key] = str(value) return function_name, properties - def _track(self, user_id: str, event_name: str, data: dict): - if "pytest" in sys.modules or self.skip_flag: + def _track( + self, + user_id: str, + event_name: str, + data: dict, + *, + client, + explicit_credentials: bool = False, + ): + if "pytest" in sys.modules: return - self.get_mp_instance().track(user_id, event_name, data) + self.get_mp_instance(client, explicit_credentials).track( + user_id, event_name, data + ) + + @classmethod + def _failure_reason( + cls, function_name: str, success: bool, error: BaseException | None + ): + """The original error message, for the "Auth Failure" event property. - def _track_method(self, args, kwargs, success: bool): + Scoped to __init__ auth/credential failures only. + """ + if success or function_name != "__init__" or error is None: + return None + return str(error) if isinstance(error, SAAuthError) else None + + def _track_method( + self, + instance, + args, + kwargs, + success: bool, + error: BaseException | None = None, + ): + # Before anything is gathered: building the payload reads controller.team_name, + # which fetches the team from the backend. A caller who turned metrics off + # should not pay for a request that is never sent. + if self._metrics_disabled(): + return try: - client = self.get_client() - if not client: - return function_name = self.function.__name__ if self.function else "" arguments = self.extract_arguments(self.function, *args, **kwargs) event_name, properties = self.default_parser(function_name, arguments) - user_email = client.controller.current_user.email - team_name = client.controller.team_name - auth_type = client.controller.token_context.auth_type + + # instance is args[0] - the actual object the call was made on, captured + # locally per call (see __call__), not shared/mutable state. It has no + # .controller yet when __init__ fails before setting one. + controller = getattr(instance, "controller", None) + user_email = team_name = auth_type = None + if controller is not None: + user_email = controller.current_user.email + auth_type = controller.token_context.scope.label + team_name = controller.team_name + elif instance is None: + return properties["Success"] = success + properties["Class"] = instance.__class__.__name__ + if error: + properties["Failure Reason"] = self._failure_reason( + function_name, success, error + ) default = self.get_default_payload( team_name=team_name, user_email=user_email, auth_type=auth_type ) self._track( - user_email, + user_email or "", event_name, {**default, **properties, **CONFIG.get_current_session().data}, + client=instance, + explicit_credentials=bool( + arguments.get("token") or arguments.get("config_path") + ), ) except BaseException: - pass + logger.debug("Skipped telemetry for this call.", exc_info=True) def __get__(self, obj, owner=None): if obj is not None: - self._client = obj tmp = functools.partial(self.__call__, obj) functools.update_wrapper(tmp, self.function) return tmp @@ -233,15 +362,24 @@ def __get__(self, obj, owner=None): def __call__(self, *args, **kwargs): success = True + error = None + # The instance the call is bound to (set by __get__ via functools.partial) - + # captured locally here, per call, rather than read back from shared state. + instance = args[0] if args else None try: result = self.function(*args, **kwargs) - except Exception as e: + except BaseException as e: + # BaseException, not Exception: a KeyboardInterrupt used to skip this and + # leave the call reported as a success. success = False - raise e + error = e + raise else: return result finally: - self._track_method(args=args, kwargs=kwargs, success=success) + self._track_method( + instance, args=args, kwargs=kwargs, success=success, error=error + ) class TrackableMeta(type): @@ -251,6 +389,7 @@ def __new__(mcs, name, bases, attrs): attr_value, FunctionType ) and not attr_value.__name__.startswith("_"): attrs[attr_name] = Tracker(validate_arguments(attr_value)) - attrs["__init__"] = Tracker(validate_arguments(attrs["__init__"])) + if "__init__" in attrs: + attrs["__init__"] = Tracker(validate_arguments(attrs["__init__"])) tmp = super().__new__(mcs, name, bases, attrs) return tmp diff --git a/src/superannotate/lib/app/interface/sdk_interface.py b/src/superannotate/lib/app/interface/sdk_interface.py index 587992093..1e8968871 100644 --- a/src/superannotate/lib/app/interface/sdk_interface.py +++ b/src/superannotate/lib/app/interface/sdk_interface.py @@ -32,7 +32,8 @@ from pydantic import TypeAdapter import lib.core as constants -from lib.infrastructure.controller import Controller +from lib.infrastructure.controller import OrgController +from lib.infrastructure.controller import TeamController from lib.app.helpers import get_annotation_paths from lib.app.helpers import get_name_url_duplicated_from_csv from lib.app.helpers import wrap_error as wrap_validation_errors @@ -67,6 +68,8 @@ from lib.core.enums import ProjectType from lib.core.enums import ClassTypeEnum from lib.core.exceptions import AppException +from lib.core.exceptions import SAAuthError +from lib.core import INVALID_TEAM_ID_ERROR from lib.core.types import PriorityScoreEntity from lib.infrastructure.annotation_adapter import BaseMultimodalAnnotationAdapter from lib.infrastructure.annotation_adapter import MultimodalSmallAnnotationAdapter @@ -141,7 +144,7 @@ class ItemContext: def __init__( self, - controller: Controller, + controller: TeamController, project: ProjectEntity, folder: FolderEntity, item: BaseItemEntity, @@ -301,10 +304,41 @@ class SAClient(BaseInterfaceFacade, metaclass=TrackableMeta): :param config_path: path to config file :type config_path: path-like (str or Path) + :param team_id: the team to operate in. Required for an Organization API key, which + is not bound to a team; for any other key it is optional and, when given, must + match the team the key grants access to. + :type team_id: int + + :param config: configuration applied on creation, instead of or on top of a config + file. Keys are a config file's keys - ``SA_TOKEN``, ``SA_URL``, ``SA_TEAM_ID``, + ``VERIFY_SSL``, ``LOGGING_LEVEL``, ``LOGGING_PATH``, ``ANNOTATION_CHUNK_SIZE``, + ``ITEM_CHUNK_SIZE``, ``MAX_THREAD_COUNT``, ``MAX_COROUTINE_COUNT`` - and an + unrecognised one is an error rather than silently ignored. Explicit ``token`` + and ``team_id`` arguments win over the same keys here. + :type config: dict + + Request Example: + :: + + sa = SAClient(config={"SA_TOKEN": "", "SA_URL": ""}) """ - def __init__(self, token: str | None = None, config_path: str | None = None): - super().__init__(token, config_path) + CONTROLLER_CLASS = TeamController + + def __init__( + self, + token: str | None = None, + config_path: str | None = None, + team_id: int | None = None, + *, + config: dict | None = None, + ): + super().__init__(token, config_path, team_id=team_id, config=config) + + @property + def team_id(self) -> int: + """The team this client operates in.""" + return self.controller.team_id def get_project_by_id(self, project_id: int): """Returns the project metadata @@ -6135,3 +6169,78 @@ def remove_users_from_project( logger.info( f"Successfully removed {success} users(s) out of the {len(users)} provided from the project {project.name}." ) + + +class SAORGClient(BaseInterfaceFacade, metaclass=TrackableMeta): + """Create SAORGClient instance to authorize SDK in an organization scope. + In case of no argument has been provided, SA_TOKEN environmental variable + will be checked or $HOME/.superannotate/config.ini will be used. + + Requires an Organization API key. + + :param token: Organization API key + :type token: str + + :param config_path: path to config file + :type config_path: str + + :param config: configuration applied on creation, instead of or on top of a config + file. Keys are a config file's keys - ``SA_TOKEN``, ``SA_URL``, ``SA_TEAM_ID``, + ``VERIFY_SSL``, ``LOGGING_LEVEL``, ``LOGGING_PATH``, ``ANNOTATION_CHUNK_SIZE``, + ``ITEM_CHUNK_SIZE``, ``MAX_THREAD_COUNT``, ``MAX_COROUTINE_COUNT`` - and an + unrecognised one is an error rather than silently ignored. Explicit ``token`` + and ``team_id`` arguments win over the same keys here. + :type config: dict + """ + + CONTROLLER_CLASS = OrgController + + def __init__( + self, + token: str | None = None, + config_path: str | None = None, + *, + config: dict | None = None, + ): + super().__init__(token, config_path, config=config) + + def get_team_client(self, team_id: int) -> SAClient: + """Returns a normal SAClient backed by the organization token, with team context returned. + + :param team_id: ID of the team to operate in. + :type team_id: int + + :return: a team-scoped client exposing the full team-level SDK surface. The + team is fixed at construction; team_id is available as a read-only + property. + :rtype: SAClient + + Request Example: + :: + + team_client = org_client.get_team_client(team_id=12345) + team_client.list_projects(name__contains="My Project") + """ + # The whole configuration this client was built with, re-pointed at one team. + # Passing the token alone would rebuild the config from defaults and lose + # API_URL with it, which means silently moving to production. + try: + return SAClient(config={**self.controller.config, "SA_TEAM_ID": team_id}) + except SAAuthError: + raise AppException(INVALID_TEAM_ID_ERROR) + + def list_teams(self) -> list[dict]: + """Returns the teams in the given organization. Must use Organization API Key for this function. + + :return: the organization's teams; an empty list if it has none. + :rtype: list of dicts + + Request Example: + :: + + org_client.list_teams() + """ + response = self.controller.list_teams() + return BaseSerializer.serialize_iterable( + response.data or [], by_alias=True, exclude_unset=True + ) diff --git a/src/superannotate/lib/core/__init__.py b/src/superannotate/lib/core/__init__.py index c03d4852e..325878ef1 100644 --- a/src/superannotate/lib/core/__init__.py +++ b/src/superannotate/lib/core/__init__.py @@ -18,16 +18,29 @@ CONFIG = Config() BACKEND_URL = "https://api.superannotate.com" -HOME_PATH = expanduser("~/.superannotate") +#: Where the SDK keeps its config. The display form keeps the "~": a message shown to +#: a user should name the file, not spell out their home directory. +HOME_DISPLAY_PATH = "~/.superannotate" +HOME_PATH = expanduser(HOME_DISPLAY_PATH) -CONFIG_JSON_PATH = f"{HOME_PATH}/config.json" -CONFIG_INI_PATH = f"{HOME_PATH}/config.ini" -CONFIG_JSON_FILE_LOCATION = CONFIG_JSON_PATH -CONFIG_INI_FILE_LOCATION = CONFIG_INI_PATH +CONFIG_JSON_FILE_LOCATION = f"{HOME_PATH}/config.json" +CONFIG_INI_FILE_LOCATION = f"{HOME_PATH}/config.ini" +CONFIG_INI_DISPLAY_PATH = f"{HOME_DISPLAY_PATH}/config.ini" LOG_FILE_LOCATION = f"{HOME_PATH}/logs" DEFAULT_LOGGING_LEVEL = "INFO" +INVALID_TOKEN_ERROR = "Invalid token." +INVALID_TEAM_ID_ERROR = "Invalid team id provided." +INVALID_CREDENTIALS_ERROR = "Invalid credentials provided." +INVALID_TEAM_CONTEXT = ( + 'Team context not provided. An Organization API key requires a "team_id".' +) +CREDENTIALS_NOT_FOUND_ERROR = ( + "Credentials not found: SA_TOKEN environment variable is not set and " + f"config file '{CONFIG_INI_DISPLAY_PATH}' was not found." +) + def setup_logging(level=DEFAULT_LOGGING_LEVEL, file_path=LOG_FILE_LOCATION): logger = logging.getLogger("sa") diff --git a/src/superannotate/lib/core/entities/__init__.py b/src/superannotate/lib/core/entities/__init__.py index a344186ee..4876e4987 100644 --- a/src/superannotate/lib/core/entities/__init__.py +++ b/src/superannotate/lib/core/entities/__init__.py @@ -2,6 +2,8 @@ from lib.core.entities.base import ConfigEntity from lib.core.entities.base import SubSetEntity from lib.core.entities.classes import AnnotationClassEntity +from lib.core.entities.context import TokenContext +from lib.core.entities.context import TokenScope from lib.core.entities.folder import FolderEntity from lib.core.entities.integrations import IntegrationEntity from lib.core.entities.items import CategoryEntity @@ -16,6 +18,7 @@ from lib.core.entities.multimodal_form import generate_classes_from_form from lib.core.entities.project import AttachmentEntity from lib.core.entities.project import CustomFieldEntity +from lib.core.entities.project import OrgTeamEntity from lib.core.entities.project import ProjectEntity from lib.core.entities.project import SettingEntity from lib.core.entities.project import StepEntity @@ -30,6 +33,8 @@ __all__ = [ # base "ConfigEntity", + "TokenContext", + "TokenScope", "SettingEntity", "SubSetEntity", "CustomFieldEntity", @@ -49,12 +54,12 @@ "WorkflowEntity", "CategoryEntity", "WMProjectUserEntity", - "ConfigEntity", "StepEntity", "FolderEntity", "S3FileEntity", "AnnotationClassEntity", "TeamEntity", + "OrgTeamEntity", "UserEntity", "IntegrationEntity", "PROJECT_ITEM_ENTITY_MAP", diff --git a/src/superannotate/lib/core/entities/base.py b/src/superannotate/lib/core/entities/base.py index 2208b60d3..246ebe9ba 100644 --- a/src/superannotate/lib/core/entities/base.py +++ b/src/superannotate/lib/core/entities/base.py @@ -4,6 +4,7 @@ from typing import Literal from lib.core import BACKEND_URL +from lib.core import INVALID_TOKEN_ERROR from lib.core import LOG_FILE_LOCATION from pydantic import AfterValidator from pydantic import BaseModel @@ -119,7 +120,7 @@ def is_legacy_token(value: str) -> bool: def _validate_token(value: str) -> str: """Validate token format.""" if not is_legacy_token(value) and not API_KEY_PATTERN.match(value): - raise ValueError("Invalid token.") + raise ValueError(INVALID_TOKEN_ERROR) return value @@ -128,9 +129,12 @@ def _validate_token(value: str) -> str: class ConfigEntity(BaseModel): - model_config = ConfigDict(extra="ignore") + model_config = ConfigDict(extra="ignore", populate_by_name=True) API_TOKEN: TokenStr = Field(alias="SA_TOKEN") + #: The team to operate in. Only an organization API key needs it — its scope carries + #: no team; every other token resolves its own team. + TEAM_ID: int | None = Field(alias="SA_TEAM_ID", default=None) API_URL: str = Field(alias="SA_URL", default=BACKEND_URL) LOGGING_LEVEL: Literal[ "NOTSET", "DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL" diff --git a/src/superannotate/lib/core/entities/context.py b/src/superannotate/lib/core/entities/context.py new file mode 100644 index 000000000..74183feae --- /dev/null +++ b/src/superannotate/lib/core/entities/context.py @@ -0,0 +1,106 @@ +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum + +from lib.core.entities.project import UserEntity + +#: Wire values of the ``authtype`` header, which is how a token authenticates rather +#: than what it was issued for. Derived from the scope; see TokenScope.auth_type. +SDK_AUTH_TYPE = "sdk" +API_KEY_AUTH_TYPE = "api_key" + + +class TokenScope(str, Enum): + """What a token was issued for. + + This is the one fact that decides how a token authenticates, whether it carries + its own team, and what it is allowed to do - so it is held once, as a scope, + rather than spread across an auth type and a nullable scope string. + """ + + #: Issued for the team itself, with no user behind it. It acts on the team's + #: behalf, so the backend denies operations only a user can perform (changing a + #: team admin's permissions, for one). + TEAM = "team" + #: Issued for one user of a team; it acts as that user. + TEAM_USER = "teamuser" + #: Issued for an organization. It carries no team, so the team to operate in has + #: to be named explicitly. + ORGANIZATION = "organization" + #: A legacy team-owner token, which carries its team id in the token itself and + #: resolves offline. The backend reports no scope for one, so this is the SDK's + #: own name for it rather than a value it receives. + LEGACY = "legacy" + + def __str__(self) -> str: + # Formats as the value it stands for, so log lines and test skip messages read + # "team" rather than "TokenScope.TEAM". + return self.value + + @classmethod + def of_api_key(cls, scope_type: str | None) -> TokenScope | None: + """The scope an API key reports, or None if it is not one the SDK knows. + + LEGACY is never a valid answer: a legacy token is recognised from its own + shape, never from a scope the backend reports. + """ + try: + scope = cls(scope_type) + except ValueError: + return None + return None if scope is cls.LEGACY else scope + + @property + def carries_team(self) -> bool: + """Whether the scope names its own team, so no explicit team_id is needed.""" + return self is not TokenScope.ORGANIZATION + + @property + def auth_type(self) -> str: + """The ``authtype`` a token of this scope authenticates with.""" + return SDK_AUTH_TYPE if self is TokenScope.LEGACY else API_KEY_AUTH_TYPE + + @property + def label(self) -> str: + """Human-readable name, for telemetry.""" + return _SCOPE_LABELS[self] + + +_SCOPE_LABELS = { + TokenScope.TEAM: "Team API Key", + TokenScope.TEAM_USER: "Personal API Key", + TokenScope.ORGANIZATION: "Org API Key", + TokenScope.LEGACY: "SDK Token", +} + + +@dataclass(frozen=True) +class TokenContext: + """Everything a client needs to act on the caller's behalf: the resolved form of + the credentials in ``ConfigEntity``. + + Clients are built from one of these rather than from loose token / team_id / + auth_type arguments, so "which team am I in, and as whom" is answered in a single + place. A team-less, organization-scoped client is expressed as ``team_id=None`` + here instead of being re-derived at every call site. + + Frozen on purpose: a client caches its HTTP session with the headers of the context + it was built from, so re-scoping a live context would leave session and context + disagreeing. Resolve a new context (and a new client) for a different team. + """ + + #: The API key or legacy token sent as the Authorization header. + token: str + #: The team every request is scoped to. None for an organization-scoped client with + #: no team bound (SAORGClient), which scopes its requests to no team at all. + team_id: int | None + #: What the token was issued for; compare against TokenScope directly. + scope: TokenScope + #: The user acting behind the token; None until it has been resolved. + user: UserEntity | None = None + + @property + def auth_type(self) -> str: + """The ``authtype`` header this context authenticates with.""" + return self.scope.auth_type diff --git a/src/superannotate/lib/core/entities/project.py b/src/superannotate/lib/core/entities/project.py index 7bd5f322a..8f52548b8 100644 --- a/src/superannotate/lib/core/entities/project.py +++ b/src/superannotate/lib/core/entities/project.py @@ -115,7 +115,7 @@ def __eq__(self, other): class UserEntity(BaseModel): - model_config = ConfigDict(extra="ignore") + model_config = ConfigDict(extra="allow") id: str | None = None first_name: str | None = None @@ -140,5 +140,18 @@ class TeamEntity(BaseModel): scores: list[str] | None = None +class OrgTeamEntity(BaseModel): + """A team as returned by SAORGClient.list_teams() - only what that endpoint sends.""" + + model_config = ConfigDict(extra="ignore", coerce_numbers_to_str=True) + + id: int | None = None + name: str | None = None + description: str | None = None + creator_id: str | None = None + owner_id: str | None = None + owner_type: str | None = None + + class CustomFieldEntity(BaseModel): model_config = ConfigDict(extra="allow") diff --git a/src/superannotate/lib/core/entities/work_managament.py b/src/superannotate/lib/core/entities/work_managament.py index eff5363c6..f279d790c 100644 --- a/src/superannotate/lib/core/entities/work_managament.py +++ b/src/superannotate/lib/core/entities/work_managament.py @@ -86,7 +86,7 @@ class WMUserEntity(TimedBaseModel): id: int | None = None team_id: int | None = None - role: WMUserTypeEnum + role: WMUserTypeEnum | None = None email: str | None = None state: WMUserStateEnum | None = None custom_fields: dict | None = Field(default_factory=dict, alias="customField") diff --git a/src/superannotate/lib/core/exceptions.py b/src/superannotate/lib/core/exceptions.py index 6a945338a..2291566ac 100644 --- a/src/superannotate/lib/core/exceptions.py +++ b/src/superannotate/lib/core/exceptions.py @@ -13,6 +13,18 @@ def __str__(self): return self.message +class SAAuthError(AppException): + """Credentials the SDK cannot act on. + + Raised when a token is missing, malformed, or does not grant what was asked for - + the team or the organization. Distinguishing these from every other AppException + lets a caller retry authentication rather than the operation, and lets telemetry + recognise an auth failure by type instead of by matching on the message. + + Subclasses AppException, so existing ``except AppException`` still catches it. + """ + + class BackendError(AppException): """ Backend Error diff --git a/src/superannotate/lib/core/service_types.py b/src/superannotate/lib/core/service_types.py index d625411af..0740fdf09 100644 --- a/src/superannotate/lib/core/service_types.py +++ b/src/superannotate/lib/core/service_types.py @@ -162,6 +162,10 @@ class TeamResponse(ServiceResponse): res_data: entities.TeamEntity = None +class OrgTeamsResponse(ServiceResponse): + res_data: list[entities.OrgTeamEntity] = None + + class UserResponse(ServiceResponse): res_data: entities.UserEntity = None diff --git a/src/superannotate/lib/core/serviceproviders.py b/src/superannotate/lib/core/serviceproviders.py index 64fb30579..92061b8e7 100644 --- a/src/superannotate/lib/core/serviceproviders.py +++ b/src/superannotate/lib/core/serviceproviders.py @@ -11,6 +11,7 @@ from lib.core import entities from lib.core.conditions import Condition from lib.core.entities import CategoryEntity +from lib.core.entities import TokenContext from lib.core.entities import WMAnnotationClassEntity from lib.core.entities.project_entities import BaseEntity from lib.core.enums import CustomFieldEntityEnum @@ -22,6 +23,7 @@ from lib.core.service_types import IntegrationListResponse from lib.core.service_types import ListCategoryResponse from lib.core.service_types import ListProjectCategoryResponse +from lib.core.service_types import OrgTeamsResponse from lib.core.service_types import ProjectListResponse from lib.core.service_types import ProjectResponse from lib.core.service_types import ServiceResponse @@ -43,20 +45,14 @@ class BaseClient(ABC): - DEFAULT_AUTH_TYPE = "sdk" - - def __init__( - self, - api_url: str, - token: str, - team_id: int, - auth_type: str = DEFAULT_AUTH_TYPE, - ): - self.team_id = team_id - + def __init__(self, api_url: str, context: TokenContext): self._api_url = api_url - self._token = token - self._auth_type = auth_type + self._context = context + + @property + def context(self) -> TokenContext: + """The session the client speaks for: its token, team and acting user.""" + return self._context @property def api_url(self): @@ -64,15 +60,27 @@ def api_url(self): @property def token(self): - return self._token + return self._context.token @property def auth_type(self): - return self._auth_type + return self._context.auth_type + + @property + def team_id(self) -> int | None: + """The team requests are scoped to; None for a team-less organization client.""" + return self._context.team_id + + @property + @abstractmethod + def default_headers(self) -> dict: + """Headers sent with every request.""" + raise NotImplementedError @property @abstractmethod - def default_headers(self): + def default_query_params(self) -> dict: + """Query params sent with every request.""" raise NotImplementedError @abstractmethod @@ -87,6 +95,7 @@ def paginate( chunk_size: int = 2000, query_params: dict[str, Any] | None = None, headers: dict | None = None, + method: Literal["get", "post"] = "get", ) -> ServiceResponse: raise NotImplementedError @@ -233,6 +242,10 @@ def update_user_activity( def list_scores(self) -> WMScoreListResponse: raise NotImplementedError + @abstractmethod + def list_project_scores(self, project_id: int) -> WMScoreListResponse: + raise NotImplementedError + @abstractmethod def create_score( self, @@ -850,6 +863,10 @@ def get_annotation_status_name( def get_team(self, team_id: int) -> TeamResponse: raise NotImplementedError + @abstractmethod + def list_teams(self) -> OrgTeamsResponse: + raise NotImplementedError + @abstractmethod def get_user(self, team_id: int) -> UserResponse: raise NotImplementedError diff --git a/src/superannotate/lib/core/usecases/annotations.py b/src/superannotate/lib/core/usecases/annotations.py index 9946174fb..288896c3a 100644 --- a/src/superannotate/lib/core/usecases/annotations.py +++ b/src/superannotate/lib/core/usecases/annotations.py @@ -62,6 +62,8 @@ ANNOTATION_CHUNK_SIZE_MB = 10 * 1024 * 1024 URI_THRESHOLD = 4 * 1024 - 120 +STATUS_CHANGE_ERROR_MSG = "Failed to change status." + @dataclass class Report: @@ -418,20 +420,25 @@ def execute(self): {i.item.name for i in items_to_upload} - set(self._report.failed_annotations).union(set(skipped)) ) - workflow = self._service_provider.work_management.get_workflow( - self._project.workflow_id - ) - if workflow.is_system(): - if uploaded_annotations and not self._keep_status: - statuses_changed = set_annotation_statuses_in_progress( - service_provider=self._service_provider, - project=self._project, - folder=self._folder, - item_names=uploaded_annotations, - ) - if not statuses_changed: - self._response.errors = AppException("Failed to change status.") - + try: + workflow = self._service_provider.work_management.get_workflow( + self._project.workflow_id + ) + if workflow.is_system(): + if uploaded_annotations and not self._keep_status: + statuses_changed = set_annotation_statuses_in_progress( + service_provider=self._service_provider, + project=self._project, + folder=self._folder, + item_names=uploaded_annotations, + ) + if not statuses_changed: + self._response.errors = AppException( + STATUS_CHANGE_ERROR_MSG + ) + except AppException as e: + if e.message != "Forbidden": + raise e self._response.data = { "succeeded": uploaded_annotations, "failed": failed, @@ -741,19 +748,22 @@ def execute(self): name_path_mappings.keys() - set(self._report.failed_annotations).union(set(missing_annotations)) ) - workflow = self._service_provider.work_management.get_workflow( - self._project.workflow_id - ) - if workflow.is_system() and uploaded_annotations and not self._keep_status: - statuses_changed = set_annotation_statuses_in_progress( - service_provider=self._service_provider, - project=self._project, - folder=self._folder, - item_names=uploaded_annotations, + try: + workflow = self._service_provider.work_management.get_workflow( + self._project.workflow_id ) - if not statuses_changed: - self._response.errors = AppException("Failed to change status.") - + if workflow.is_system() and uploaded_annotations and not self._keep_status: + statuses_changed = set_annotation_statuses_in_progress( + service_provider=self._service_provider, + project=self._project, + folder=self._folder, + item_names=uploaded_annotations, + ) + if not statuses_changed: + self._response.errors = AppException(STATUS_CHANGE_ERROR_MSG) + except AppException as e: + if e.message != "Forbidden": + raise e if missing_annotations: logger.warning( f"Couldn't find {len(missing_annotations)}/{len(name_path_mappings.keys())} " @@ -950,20 +960,25 @@ def execute(self): self.reporter.log_warning( f"Couldn't find attribute {attr}." ) - workflow = self._service_provider.work_management.get_workflow( - self._project.workflow_id - ) - if workflow.is_system() and not self._keep_status: - statuses_changed = set_annotation_statuses_in_progress( - service_provider=self._service_provider, - project=self._project, - folder=self._folder, - item_names=[self._image.name], + try: + workflow = self._service_provider.work_management.get_workflow( + self._project.workflow_id ) - if not statuses_changed: - self._response.errors = AppException( - "Failed to change status." + if workflow.is_system() and not self._keep_status: + statuses_changed = set_annotation_statuses_in_progress( + service_provider=self._service_provider, + project=self._project, + folder=self._folder, + item_names=[self._image.name], ) + if not statuses_changed: + self._response.errors = AppException( + STATUS_CHANGE_ERROR_MSG + ) + except AppException as e: + if e.message != "Forbidden": + raise e + if self._verbose: self.reporter.log_info( f"Uploading annotations for image {str(self._image.name)} in project {self._project.name}." @@ -1719,7 +1734,7 @@ def execute(self): ) ) except Exception as e: - logger.error(e) + logger.exception(e) self._response.errors = AppException("Can't get annotations.") return self._response self.reporter.stop_spinner() @@ -2021,22 +2036,26 @@ def execute(self): folder_id=folder.id, item_id_category_map=item_id_category_map, ) - workflow = self._service_provider.work_management.get_workflow( - self._project.workflow_id - ) uploaded.extend(uploaded_annotations) - if workflow.is_system(): - if uploaded_annotations and not self._keep_status: - statuses_changed = set_annotation_statuses_in_progress( - service_provider=self._service_provider, - project=self._project, - folder=folder, - item_names=uploaded_annotations, - ) - if not statuses_changed: - self._response.errors = AppException( - "Failed to change status." + try: + workflow = self._service_provider.work_management.get_workflow( + self._project.workflow_id + ) + if workflow.is_system(): + if uploaded_annotations and not self._keep_status: + statuses_changed = set_annotation_statuses_in_progress( + service_provider=self._service_provider, + project=self._project, + folder=folder, + item_names=uploaded_annotations, ) + if not statuses_changed: + self._response.errors = AppException( + STATUS_CHANGE_ERROR_MSG + ) + except AppException as e: + if e.message != "Forbidden": + raise e self.reporter.finish_progress() self._report.failed_annotations = [] diff --git a/src/superannotate/lib/core/usecases/images.py b/src/superannotate/lib/core/usecases/images.py index 8542c5acb..f3d5d34a1 100644 --- a/src/superannotate/lib/core/usecases/images.py +++ b/src/superannotate/lib/core/usecases/images.py @@ -712,15 +712,19 @@ def __init__( self._s3_repo = s3_repo self._service_provider = service_provider if annotation_status_value is None: - workflow = self._service_provider.work_management.get_workflow( - self._project.workflow_id - ) - if workflow.is_system(): - annotation_status_value = ( - self._service_provider.get_annotation_status_value( - self._project, "NotStarted" - ) + try: + workflow = self._service_provider.work_management.get_workflow( + self._project.workflow_id ) + if workflow.is_system(): + annotation_status_value = ( + self._service_provider.get_annotation_status_value( + self._project, "NotStarted" + ) + ) + except AppException as e: + if e.message != "Forbidden": + raise e self._annotation_status_value = annotation_status_value self._auth_data = None diff --git a/src/superannotate/lib/core/usecases/items.py b/src/superannotate/lib/core/usecases/items.py index 744867c8d..9e6348628 100644 --- a/src/superannotate/lib/core/usecases/items.py +++ b/src/superannotate/lib/core/usecases/items.py @@ -703,7 +703,7 @@ def execute(self): status_changed = self._service_provider.items.set_statuses( project=self._project, folder=self._folder, - item_names=self._item_names[i : i + self.CHUNK_SIZE], # noqa: E203, + item_names=self._item_names[i : i + self.CHUNK_SIZE], # noqa: E203 annotation_status=self._annotation_status_code, ) if not status_changed.ok: @@ -772,7 +772,7 @@ def execute(self): response = self._service_provider.items.set_approval_statuses( project=self._project, folder=self._folder, - item_names=self._item_names[i : i + self.CHUNK_SIZE], # noqa: E203, + item_names=self._item_names[i : i + self.CHUNK_SIZE], # noqa: E203 approval_status=self._approval_status_code, ) if not response.ok: @@ -826,10 +826,13 @@ def execute(self): item_ids = [item.id for item in items] for i in range(0, len(item_ids), self.CHUNK_SIZE): - self._service_provider.items.delete_multiple( + response = self._service_provider.items.delete_multiple( project=self._project, item_ids=item_ids[i : i + self.CHUNK_SIZE], # noqa: E203 ) + if not response.ok: + self._response.errors = response.error + return self._response logger.info( f"Items deleted in project {self._project.name}{'/' + self._folder.name if not self._folder.is_root else ''}" ) @@ -864,7 +867,7 @@ def __init__( def __filter_duplicates( self, ): - def uniqueQ(item, seen): + def _unique(item, seen): result = True if "id" in item: if item["id"] in seen: @@ -880,13 +883,13 @@ def uniqueQ(item, seen): return result seen = set() - uniques = [x for x in self.items if uniqueQ(x, seen)] + uniques = [x for x in self.items if _unique(x, seen)] return uniques def __filter_invalid_items( self, ): - def validQ(item): + def _valid(item): if "id" in item: return True if "name" in item and "path" in item: @@ -894,7 +897,7 @@ def validQ(item): self.results["skipped"].append(item) return False - filtered_items = [x for x in self.items if validQ(x)] + filtered_items = [x for x in self.items if _valid(x)] return filtered_items diff --git a/src/superannotate/lib/core/usecases/projects.py b/src/superannotate/lib/core/usecases/projects.py index 302c5de75..fa3c134de 100644 --- a/src/superannotate/lib/core/usecases/projects.py +++ b/src/superannotate/lib/core/usecases/projects.py @@ -28,6 +28,7 @@ from lib.core.usecases.base import BaseUseCase from lib.core.usecases.base import BaseUserBasedUseCase from pydantic import ValidationError +from superannotate import SAAuthError logger = logging.getLogger("sa") @@ -780,8 +781,12 @@ def execute(self): try: response = self._service_provider.get_team(self._team_id) if not response.ok: + if response.status_code == 403: + raise SAAuthError(constants.INVALID_TEAM_ID_ERROR) raise AppException(response.error) self._response.data = response.data + except SAAuthError: + raise except Exception as e: raise AppException( "Unable to retrieve team data. Please verify your credentials." @@ -789,6 +794,24 @@ def execute(self): return self._response +class ListTeamsUseCase(BaseUseCase): + def __init__(self, service_provider: BaseServiceProvider): + super().__init__() + self._service_provider = service_provider + + def execute(self): + try: + response = self._service_provider.list_teams() + if not response.ok: + raise AppException(response.error) + self._response.data = response.data + except Exception: + raise AppException( + "Unable to retrieve team data. Please verify your credentials." + ) from None + return self._response + + class GetCurrentUserUseCase(BaseUseCase): def __init__(self, service_provider: BaseServiceProvider, team_id: int): super().__init__() diff --git a/src/superannotate/lib/infrastructure/annotation_adapter.py b/src/superannotate/lib/infrastructure/annotation_adapter.py index 2ef38ce79..f123304a1 100644 --- a/src/superannotate/lib/infrastructure/annotation_adapter.py +++ b/src/superannotate/lib/infrastructure/annotation_adapter.py @@ -9,7 +9,7 @@ from lib.core.entities import ProjectEntity from lib.core.utils import run_async from lib.core.utils import set_last_action -from lib.infrastructure.controller import Controller +from lib.infrastructure.controller import TeamController class BaseMultimodalAnnotationAdapter(ABC): @@ -18,7 +18,7 @@ def __init__( project: ProjectEntity, folder: FolderEntity, item: BaseItemEntity, - controller: Controller, + controller: TeamController, annotation: dict = None, ): self._project = project @@ -93,7 +93,7 @@ def __init__( project: ProjectEntity, folder: FolderEntity, item: BaseItemEntity, - controller: Controller, + controller: TeamController, overwrite: bool = True, annotation: dict = None, ): diff --git a/src/superannotate/lib/infrastructure/controller.py b/src/superannotate/lib/infrastructure/controller.py index 2da47bf7a..348279943 100644 --- a/src/superannotate/lib/infrastructure/controller.py +++ b/src/superannotate/lib/infrastructure/controller.py @@ -3,8 +3,8 @@ import copy import io import logging -import os from abc import ABCMeta +from abc import abstractmethod from collections.abc import Callable from pathlib import Path from typing import Any @@ -12,6 +12,7 @@ import lib.core as constants from lib.core import ApprovalStatus +from lib.core import INVALID_CREDENTIALS_ERROR from lib.core import usecases from lib.core.conditions import Condition from lib.core.conditions import CONDITION_EQ as EQ @@ -25,6 +26,7 @@ from lib.core.entities import ProjectEntity from lib.core.entities import SettingEntity from lib.core.entities import TeamEntity +from lib.core.entities import TokenContext from lib.core.entities import UserEntity from lib.core.entities import WMAnnotationClassEntity from lib.core.entities import WMProjectUserEntity @@ -42,6 +44,7 @@ from lib.core.enums import ProjectType from lib.core.exceptions import AppException from lib.core.exceptions import FileChangedError +from lib.core.exceptions import SAAuthError from lib.core.jsx_conditions import EmptyQuery from lib.core.jsx_conditions import Filter from lib.core.jsx_conditions import Join @@ -60,8 +63,8 @@ from lib.infrastructure.query_builder import TeamUserFilterHandler from lib.infrastructure.repositories import S3Repository from lib.infrastructure.serviceprovider import ServiceProvider -from lib.infrastructure.services.auth import resolve_token_context -from lib.infrastructure.services.auth import TokenContext +from lib.infrastructure.services.auth import resolve_organization_context +from lib.infrastructure.services.auth import resolve_team_context from lib.infrastructure.services.http_client import HttpClient from lib.infrastructure.utils import divide_to_chunks from lib.infrastructure.utils import extract_project_folder @@ -320,7 +323,9 @@ def get_user_scores( scored_user: str, provided_score_names: list[str] | None = None, ): - score_fields_res = self.service_provider.work_management.list_scores() + score_fields_res = self.service_provider.work_management.list_project_scores( + project_id=project.id + ) # validate provided score names all_score_names = [s.name for s in score_fields_res.data] @@ -995,7 +1000,7 @@ def delete_multiple(self, project: ProjectEntity, folders: list[FolderEntity]): return use_case.execute() def get_by_name(self, project: ProjectEntity, name: str = None): - name = Controller.get_folder_name(name) + name = TeamController.get_folder_name(name) use_case = usecases.GetFolderUseCase( project=project, folder_name=name, @@ -1661,84 +1666,104 @@ def add_items(self, project: ProjectEntity, subset: str, items: list[dict]): class BaseController(metaclass=ABCMeta): - SESSIONS = {} + """What every client needs regardless of what its token was issued for: a resolved + session, an HTTP client built from it, and the user it acts as. - def __init__(self, config: ConfigEntity): - self._config = config - self._logger = logging.getLogger("sa") - self._testing = os.getenv("SA_TESTING", "False").lower() in ("true", "1", "t") - self._token = config.API_TOKEN - self._s3_upload_auth_data = None - self._projects = None - self._folders = None - self._teams = None - self._images = None - self._items = None - self._integrations = None - self._user_id = None - self._reporter = None + Each subclass resolves the session it needs and hands it here - what a token must + resolve to is where the two part company, so it is named at the point each one + constructs itself rather than configured through a flag. Anything team-scoped lives + on TeamController, not here: an organization-scoped client has no team to apply it + to. + """ - self._token_context = resolve_token_context( - api_url=config.API_URL, - token=config.API_TOKEN, - verify_ssl=config.VERIFY_SSL, - ) - self._team_id = self._token_context.team_id - - http_client = HttpClient( - api_url=config.API_URL, - token=config.API_TOKEN, - team_id=self._team_id, - auth_type=self._token_context.auth_type, - verify_ssl=config.VERIFY_SSL, + def __init__(self, config: ConfigEntity, context: TokenContext): + self._config = config + self._token_context = context + self.service_provider = ServiceProvider( + HttpClient( + api_url=config.API_URL, + context=context, + verify_ssl=config.VERIFY_SSL, + ) ) - - self.service_provider = ServiceProvider(http_client) self._user = self.get_current_user() - # An API key already resolved its team, so the team data is only fetched once - # something actually needs it (the organization id, mostly). - self._team = self.get_team().data if self._token_context.is_legacy else None - self.annotation_classes = AnnotationClassManager(self.service_provider) - self.projects = ProjectManager(self.service_provider, team=lambda: self.team) - self.work_management = WorkManagementManager(self.service_provider) - self.folders = FolderManager(self.service_provider) - self.items = ItemManager(self.service_provider) - self.annotations = AnnotationManager(self.service_provider, config) - self.custom_fields = CustomFieldManager(self.service_provider) - self.subsets = SubsetManager(self.service_provider) - self.integrations = IntegrationManager(self.service_provider) - @property - def reporter(self): - return self._reporter + @abstractmethod + def get_current_user(self) -> UserEntity: + """The user the token acts as.""" + raise NotImplementedError @property - def org_id(self): - return self.team.owner_id + def config(self) -> dict: + """The configuration this client was built from. + + Keyed the way configuration is written - ``SA_TOKEN``, ``SA_URL``, + ``SA_TEAM_ID`` and the rest - rather than by ConfigEntity's field names, so it + is the form ``SAClient(config=...)`` takes and a client can be rebuilt from + another one. A fresh dict each time: mutating it changes nothing. + """ + return self._config.model_dump(by_alias=True) @property - def current_user(self): + def current_user(self) -> UserEntity: return self._user @property def token_context(self) -> TokenContext: - """The scope the client authenticated with (team / team-user / legacy).""" + """The session the client authenticated with: its token, team and user.""" return self._token_context @property - def team(self) -> TeamEntity: - if self._team is None: - self._team = self.get_team().data - return self._team + def team_name(self) -> str | None: + """The team the client operates in, for telemetry; None when it has no team.""" + return None - def get_team(self): - return usecases.GetTeamUseCase( - service_provider=self.service_provider, team_id=self.team_id + +class OrgController(BaseController): + """Operates on an organization rather than inside one of its teams. + + It is not team-bound, so it offers none of the team-level surface: listing the + organization's teams is the whole of it. A team-scoped client for one of those + teams is a TeamController, built separately. + """ + + def __init__(self, config: ConfigEntity): + super().__init__(config, resolve_organization_context(config)) + + def get_current_user(self) -> UserEntity: + # Resolved with the scope: an organization key has no team to look a user up + # in, so the creator behind the key is the only user there is. + user = self._token_context.user + if user is None: + raise SAAuthError(INVALID_CREDENTIALS_ERROR) + return user + + def list_teams(self): + return usecases.ListTeamsUseCase( + service_provider=self.service_provider ).execute() + +class TeamController(BaseController): + """Operates inside one team, which every request is scoped to.""" + + def __init__(self, config: ConfigEntity): + super().__init__(config, resolve_team_context(config)) + self._reporter = None + # An API key already resolved its team, so the team data itself is fetched only + # once something needs it (the organization id, and telemetry's team name). + self._team = self.get_team().data + self.annotation_classes = AnnotationClassManager(self.service_provider) + self.projects = ProjectManager(self.service_provider, team=lambda: self.team) + self.work_management = WorkManagementManager(self.service_provider) + self.folders = FolderManager(self.service_provider) + self.items = ItemManager(self.service_provider) + self.annotations = AnnotationManager(self.service_provider, config) + self.custom_fields = CustomFieldManager(self.service_provider) + self.subsets = SubsetManager(self.service_provider) + self.integrations = IntegrationManager(self.service_provider) + def get_current_user(self) -> UserEntity: - # An API key resolves its own user (or its creator, for team-scoped keys) while the - # team is being resolved, so there is nothing left to look up. if self._token_context.user: return self._token_context.user response = usecases.GetCurrentUserUseCase( @@ -1749,19 +1774,35 @@ def get_current_user(self) -> UserEntity: return response.data @property - def team_name(self) -> str: - """The name of the team the client operates in. + def reporter(self): + return self._reporter - Used by telemetry. An API key resolves only the team id on init, so the first - call fetches the team; the result is cached on the controller from then on. + @property + def team_id(self) -> int | None: + """The team every request is scoped to.""" + return self._token_context.team_id + + @property + def team(self) -> TeamEntity | None: + if self._team is None: + self._team = self.get_team().data + return self._team + + @property + def team_name(self) -> str | None: + """Used by telemetry. An API key resolves only the team id on init, so the + first call fetches the team; every later reader is served from the cache. """ return self.team.name @property - def team_id(self) -> int: - if not self._token or not self._team_id: - raise AppException("Invalid credentials provided.") - return self._team_id + def org_id(self): + return self.team.owner_id + + def get_team(self): + return usecases.GetTeamUseCase( + service_provider=self.service_provider, team_id=self.team_id + ).execute() @staticmethod def get_default_reporter( @@ -1776,10 +1817,6 @@ def get_default_reporter( def s3_repo(self): return S3Repository - -class Controller(BaseController): - DEFAULT = None - def get_folder_by_id(self, folder_id: int, project_id: int): response = self.folders.get_by_id( folder_id=folder_id, project_id=project_id, team_id=self.team_id diff --git a/src/superannotate/lib/infrastructure/serviceprovider.py b/src/superannotate/lib/infrastructure/serviceprovider.py index 5cfb60f27..1883b9a00 100644 --- a/src/superannotate/lib/infrastructure/serviceprovider.py +++ b/src/superannotate/lib/infrastructure/serviceprovider.py @@ -1,12 +1,12 @@ from __future__ import annotations -import base64 import datetime import lib.core as constants from lib.core import entities from lib.core.enums import ApprovalStatus from lib.core.enums import CustomFieldEntityEnum +from lib.core.service_types import OrgTeamsResponse from lib.core.service_types import TeamResponse from lib.core.service_types import UploadAnnotationAuthDataResponse from lib.core.service_types import UserLimitsResponse @@ -16,6 +16,7 @@ from lib.infrastructure.services.annotation_class import AnnotationClassService from lib.infrastructure.services.explore import ExploreService from lib.infrastructure.services.folder import FolderService +from lib.infrastructure.services.http_client import encode_entity_context from lib.infrastructure.services.http_client import HttpClient from lib.infrastructure.services.integration import IntegrationService from lib.infrastructure.services.item import ItemService @@ -29,6 +30,7 @@ class ServiceProvider(BaseServiceProvider): URL_TEAM = "api/v1/team" + URL_TEAMS = "api/v1/teams" URL_GET_LIMITS = "project/{project_id}/limitationDetails" URL_GET_TEMPLATES = "templates" URL_PREPARE_EXPORT = "export" @@ -55,21 +57,19 @@ def __init__(self, client: HttpClient): self.integrations = IntegrationService(client) self.explore = ExploreService(client) self.telemetry_scoring = TelemetryScoringService(client) + # The sibling services talk to their own hosts on the same session as the + # main client, so they share its context rather than re-deriving one. self.work_management = WorkManagementService( HttpClient( api_url=self._get_work_management_url(client), - token=client.token, - team_id=client.team_id, - auth_type=client.auth_type, + context=client.context, verify_ssl=client.verify_ssl, ) ) self.item_service = SeparateItemService( HttpClient( api_url=self._get_item_service_url(client), - token=client.token, - team_id=client.team_id, - auth_type=client.auth_type, + context=client.context, verify_ssl=client.verify_ssl, ) ) @@ -204,6 +204,15 @@ def get_team(self, team_id: int) -> TeamResponse: params={"include_users": False}, ) + def list_teams(self) -> OrgTeamsResponse: + # The backend wraps the array as {"count": N, "data": [...]}. + return self.client.request( + self.URL_TEAMS, + "get", + content_type=OrgTeamsResponse, + dispatcher="data", + ) + def get_user(self, team_id: int) -> UserResponse: return self.client.request( self.URL_USER, "get", params={"team_id": team_id}, content_type=UserResponse @@ -366,9 +375,9 @@ def create_custom_workflow(self, org_id: str, data: dict): url=self.URL_CREATE_WORKFLOW, method="post", headers={ - "x-sa-entity-context": base64.b64encode( - f'{{"team_id":{self.client.team_id},"organization_id":"{org_id}"}}'.encode() - ).decode() + "x-sa-entity-context": encode_entity_context( + team_id=self.client.team_id, organization_id=org_id + ) }, data=data, ) diff --git a/src/superannotate/lib/infrastructure/services/auth.py b/src/superannotate/lib/infrastructure/services/auth.py index 81bf61355..61c3cd383 100644 --- a/src/superannotate/lib/infrastructure/services/auth.py +++ b/src/superannotate/lib/infrastructure/services/auth.py @@ -1,101 +1,109 @@ from __future__ import annotations import logging -from dataclasses import dataclass +from typing import TYPE_CHECKING import lib.core as constants import requests +from lib.core import INVALID_CREDENTIALS_ERROR +from lib.core import INVALID_TEAM_CONTEXT +from lib.core import INVALID_TEAM_ID_ERROR from lib.core.entities.base import is_legacy_token +from lib.core.entities.context import API_KEY_AUTH_TYPE +from lib.core.entities.context import TokenContext +from lib.core.entities.context import TokenScope from lib.core.entities.project import UserEntity -from lib.core.exceptions import AppException +from lib.core.exceptions import SAAuthError -logger = logging.getLogger("sa") +if TYPE_CHECKING: + from lib.core.entities.base import ConfigEntity -SDK_AUTH_TYPE = "sdk" -API_KEY_AUTH_TYPE = "api_key" +logger = logging.getLogger("sa") URL_TOKEN_CONTEXT = "users/me" -#: A key issued for the team itself, with no user behind it. It acts on the team's -#: behalf, so the backend denies operations that only a user can perform (changing a -#: team admin's permissions, for one). -TEAM_SCOPE_TYPE = "team" -#: A key issued for one user of a team; it acts as that user. -TEAM_USER_SCOPE_TYPE = "teamuser" -#: Token scopes that carry a team, and therefore need no explicit team_id. -TEAM_SCOPED_TYPES = (TEAM_SCOPE_TYPE, TEAM_USER_SCOPE_TYPE) - -ORGANIZATION_API_KEY_ERROR = ( - "SAClient does not accept an Organization API key — it requires a Team or " - "Personal API key." -) -AUTHENTICATION_ERROR = ( - "Unable to authenticate the provided token. Please verify your credentials." -) - - -@dataclass -class TokenContext: - """The team the client operates in, plus the user acting behind the token.""" - - team_id: int - auth_type: str - user: UserEntity | None = None - #: Scope the key was issued for ("team", "teamuser", "organization"); None for a - #: legacy token, whose scope is not reported by the backend. - scope_type: str | None = None - - @property - def is_legacy(self) -> bool: - return self.auth_type == SDK_AUTH_TYPE - - @property - def is_team_key(self) -> bool: - """Whether the token acts as the team rather than as a user.""" - return self.scope_type == TEAM_SCOPE_TYPE - - @property - def is_personal_key(self) -> bool: - """Whether the token acts as one specific user of the team.""" - return self.scope_type == TEAM_USER_SCOPE_TYPE - - -def resolve_token_context( - api_url: str, - token: str, - verify_ssl: bool = True, -) -> TokenContext: - """Resolve the team (and acting user) a token grants access to. - - Legacy team-owner tokens carry the team id, so they are resolved offline. New-style - API keys are resolved against the work-management service, which reports the scope - the key was issued for. The SDK operates within a single team, so a key that is not - scoped to one is rejected. + +def resolve_team_context(config: ConfigEntity) -> TokenContext: + """The team a token grants access to, plus the user acting behind it. + + A legacy team-owner token carries its team id, so it resolves offline. An API key + resolves against the work-management service, which reports the scope it was issued + for: a team or team-user key names its own team, while an organization key names + none, so the team has to come from the config (``SAClient(team_id=...)``, + ``SA_TEAM_ID``). """ + token = config.API_TOKEN if is_legacy_token(token): - return TokenContext(team_id=int(token.split("=")[-1]), auth_type=SDK_AUTH_TYPE) + team_id = int(token.split("=")[-1]) + _validate_requested_team(config.TEAM_ID, team_id) + return TokenContext(token=token, team_id=team_id, scope=TokenScope.LEGACY) - data = _fetch_token_context(api_url, token, verify_ssl) + scope, scope_team_id, user = _resolve_api_key(config) + if scope is None: + raise SAAuthError(INVALID_TEAM_ID_ERROR) + team_id = _team_for_scope(scope, config.TEAM_ID, scope_team_id) + logger.debug(f"Token resolved to {scope} scope, team {team_id}.") + return TokenContext(token=token, team_id=int(team_id), scope=scope, user=user) + + +def resolve_organization_context(config: ConfigEntity) -> TokenContext: + """An organization-scoped session, bound to no team, for ``SAORGClient``. + + Only an Organization API key is accepted: every other token acts within one team + and so cannot act for the organization. Any ``team_id`` in the config is ignored - + an organization client operates outside any single team by definition. + """ + if is_legacy_token(config.API_TOKEN): + raise SAAuthError(INVALID_CREDENTIALS_ERROR) + scope, _, user = _resolve_api_key(config) + if scope is not TokenScope.ORGANIZATION: + raise SAAuthError(INVALID_CREDENTIALS_ERROR) + logger.debug("Token resolved to organization scope, with no team.") + return TokenContext(token=config.API_TOKEN, team_id=None, scope=scope, user=user) + + +def _resolve_api_key( + config: ConfigEntity, +) -> tuple[TokenScope | None, int | None, UserEntity | None]: + """What an API key reports: the scope it was issued for (None when the SDK does not + know it), the team that scope names, and the user behind the key. + """ + data = _fetch_token_context(config.API_URL, config.API_TOKEN, config.VERIFY_SSL) token_data = data.get("token") or {} - scope = token_data.get("scope") or {} - scope_type = token_data.get("scope_type") - token_team_id = scope.get("team_id") - - # Anything outside the allowlist (an organization key, today) has no team to operate - # in; the team_id check keeps a malformed response from resolving to no team at all. - if scope_type not in TEAM_SCOPED_TYPES or token_team_id is None: - logger.debug(f"Rejected a token of {scope_type} scope.") - raise AppException(ORGANIZATION_API_KEY_ERROR) - - logger.debug(f"Token resolved to {scope_type} scope, team {token_team_id}.") - return TokenContext( - team_id=int(token_team_id), - auth_type=API_KEY_AUTH_TYPE, - user=_build_user(data.get("user"), token_data.get("created_by")), - scope_type=scope_type, + reported_scope = token_data.get("scope_type") + scope = TokenScope.of_api_key(reported_scope) + if scope is None: + logger.debug(f"Got a token of unknown {reported_scope} scope.") + return ( + scope, + (token_data.get("scope") or {}).get("team_id"), + _build_user(data.get("user"), token_data.get("created_by")), ) +def _team_for_scope(scope: TokenScope, requested_team_id, scope_team_id) -> int: + """The team a team-scoped client operates in.""" + if not scope.carries_team: + # An organization key has no team of its own, so the caller has to name one. + if requested_team_id is None: + raise SAAuthError(INVALID_TEAM_CONTEXT) + return requested_team_id + # The team_id check keeps a malformed response from resolving to no team at all. + if scope_team_id is None: + logger.debug(f"Got a {scope} scoped token with no team.") + raise SAAuthError(INVALID_TEAM_ID_ERROR) + _validate_requested_team(requested_team_id, scope_team_id) + return scope_team_id + + +def _validate_requested_team(requested_team_id, token_team_id) -> None: + """A team_id passed alongside a team-carrying token must agree with it.""" + if requested_team_id is None: + return + if int(requested_team_id) != int(token_team_id): + raise SAAuthError(INVALID_TEAM_ID_ERROR) + + def _get_work_management_url(api_url: str) -> str: # The token scope has to be resolved before there is a client to ask, so the # work-management host is derived here as well as in the service provider. @@ -117,17 +125,14 @@ def _fetch_token_context(api_url: str, token: str, verify_ssl: bool) -> dict: }, verify=verify_ssl, ) - except (requests.RequestException, ConnectionError) as e: - raise AppException(f"Unable to authenticate the provided token: {e}.") - if not response.ok: - logger.debug( - f"Got {response.status_code} response from backend: {response.text}" - ) - raise AppException(AUTHENTICATION_ERROR) - try: + if not response.ok: + logger.debug( + f"Got {response.status_code} response from backend {url}: {response.text}" + ) + raise ValueError("non-ok response") return response.json() - except ValueError: - raise AppException(AUTHENTICATION_ERROR) + except (requests.RequestException, ConnectionError, ValueError): + raise SAAuthError(INVALID_CREDENTIALS_ERROR) from None def _build_user(user: dict | None, created_by: str | None) -> UserEntity | None: diff --git a/src/superannotate/lib/infrastructure/services/http_client.py b/src/superannotate/lib/infrastructure/services/http_client.py index 8e627cda5..57c2ad9fe 100644 --- a/src/superannotate/lib/infrastructure/services/http_client.py +++ b/src/superannotate/lib/infrastructure/services/http_client.py @@ -18,6 +18,7 @@ import aiohttp import requests +from lib.core.entities import TokenContext from lib.core.exceptions import AppException from lib.core.jsx_conditions import EmptyQuery from lib.core.jsx_conditions import Limit @@ -42,16 +43,23 @@ def default(self, obj): return json.JSONEncoder.default(self, obj) +def encode_entity_context(**context) -> str: + """Build the ``x-sa-entity-context`` header value: base64-encoded JSON. + + The backend reads the entity a request applies to (team, project, organization) + from this header rather than from the path. + """ + return base64.b64encode(json.dumps(context).encode("utf-8")).decode("utf-8") + + class HttpClient(BaseClient): def __init__( self, api_url: str, - token: str, - team_id: int, - auth_type: str = BaseClient.DEFAULT_AUTH_TYPE, + context: TokenContext, verify_ssl: bool = True, ): - super().__init__(api_url, token, team_id, auth_type) + super().__init__(api_url, context) self._verify_ssl = verify_ssl self._version = os.environ.get("sa_version") self._env = os.environ.get("SA_ENV") @@ -78,22 +86,28 @@ def get_session(self): ) @property - def default_headers(self): - return { - "Authorization": self._token, - "authtype": self._auth_type, + def default_headers(self) -> dict: + headers = { + "Authorization": self.token, + "authtype": self.auth_type, "Content-Type": "application/json", - "x-sa-entity-context": base64.b64encode( - json.dumps( - { - "team_id": self.team_id, - } - ).encode("utf-8") - ).decode("utf-8"), "User-Agent": f"Python-SDK-Version: {self._version}; Python: {platform.python_version()};" - f"OS: {platform.system()}; Team: {self.team_id}" + f"OS: {platform.system()}" + f"{f'; Team: {self.team_id}' if self.team_id is not None else ''}" f"{'; Env: ' + self._env if self._env else ''}", } + # A team-less, organization-scoped client (SAORGClient) has nothing to scope the + # request to, so the header is left out rather than sent as null. + if self.team_id is not None: + headers["x-sa-entity-context"] = encode_entity_context(team_id=self.team_id) + return headers + + @property + def default_query_params(self) -> dict: + """Query params every request carries; empty for a team-less client.""" + if self.team_id is None: + return {} + return {"team_id": self.team_id} @property def safe_api(self): @@ -130,7 +144,7 @@ def _request(self, url, method, session, retried=0, **kwargs): ) if response.status_code > 299: logger.debug( - f"Got {response.status_code} from {url} response from backend" + f"Got {method} {response.status_code} from {url} response from backend {response.text}" ) return response @@ -152,7 +166,7 @@ def request( dispatcher: str = None, ) -> ServiceResponse: _url = self._get_url(url) - kwargs = {"params": {"team_id": self.team_id}} + kwargs = {"params": dict(self.default_query_params)} if data: kwargs["data"] = json.dumps(data, cls=PydanticEncoder) if params: @@ -173,6 +187,7 @@ def paginate( chunk_size: int = 2000, query_params: dict[str, Any] = None, headers: dict = None, + method: Literal["get", "post"] = "get", ) -> ServiceResponse: offset = 0 total = [] @@ -182,7 +197,7 @@ def paginate( _url = f"{url}{splitter}offset={offset}" _response = self.request( _url, - method="get", + method=method, params=query_params, dispatcher="data", headers=headers, diff --git a/src/superannotate/lib/infrastructure/services/work_management.py b/src/superannotate/lib/infrastructure/services/work_management.py index 1eb571e43..326b10c71 100644 --- a/src/superannotate/lib/infrastructure/services/work_management.py +++ b/src/superannotate/lib/infrastructure/services/work_management.py @@ -1,7 +1,5 @@ from __future__ import annotations -import base64 -import json from typing import Literal from lib.core.entities import CategoryEntity @@ -30,6 +28,7 @@ from lib.core.service_types import WMScoreListResponse from lib.core.service_types import WMUserListResponse from lib.core.serviceproviders import BaseWorkManagementService +from lib.infrastructure.services.http_client import encode_entity_context def prepare_validation_error(func): @@ -69,6 +68,7 @@ class WorkManagementService(BaseWorkManagementService): URL_CREATE_CATEGORIES = "categories/bulk" URL_CUSTOM_FIELD_TEMPLATES = "customfieldtemplates" URL_SCORES = "scores" + URL_PROJECT_SCORES = "scores/getProjectScores" URL_DELETE_SCORE = "scores/{score_id}" URL_CUSTOM_FIELD_TEMPLATE_DELETE = "customfieldtemplates/{template_id}" URL_SET_CUSTOM_ENTITIES = "customentities/{pk}" @@ -82,11 +82,6 @@ class WorkManagementService(BaseWorkManagementService): URL_PERMISSION_GROUPS = "permissiongroups" URL_UPDATE_ANNOTATION_CLASS = "classes/{class_id}" - @staticmethod - def _generate_context(**kwargs): - encoded_context = base64.b64encode(json.dumps(kwargs).encode("utf-8")) - return encoded_context.decode("utf-8") - def list_folders(self, project_id: int, query: Query) -> ServiceResponse: result = self.client.jsx_paginate( self.URL_LIST_FOLDERS, @@ -94,7 +89,7 @@ def list_folders(self, project_id: int, query: Query) -> ServiceResponse: item_type=FolderEntity, method="post", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ), }, @@ -109,7 +104,7 @@ def list_project_categories( item_type=entity, query_params={"project_id": project_id}, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -126,7 +121,7 @@ def create_project_categories( params={"project_id": project_id}, data={"bulk": [{"name": i} for i in categories]}, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ), }, @@ -144,7 +139,7 @@ def remove_project_categories( method="delete", url=f"{self.URL_CREATE_CATEGORIES}?{query.build_query()}", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ), }, @@ -167,7 +162,7 @@ def list_workflows(self, query: Query): f"{self.URL_LIST}?{query.build_query()}", item_type=WorkflowEntity, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -179,7 +174,7 @@ def list_workflow_statuses(self, project_id: int, workflow_id: int): url=self.URL_LIST_STATUSES.format(workflow_id=workflow_id), method="get", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ) }, @@ -193,7 +188,7 @@ def list_workflow_roles(self, project_id: int, workflow_id: int): url=self.URL_LIST_ROLES.format(workflow_id=workflow_id), method="get", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ) }, @@ -207,7 +202,7 @@ def create_custom_role(self, org_id: str, data: dict): url=self.URL_CREATE_ROLE, method="post", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, organization_id=org_id ) }, @@ -219,7 +214,7 @@ def create_custom_status(self, org_id: str, data: dict): url=self.URL_CREATE_STATUS, method="post", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, organization_id=org_id ) }, @@ -238,7 +233,7 @@ def list_custom_field_templates( url=self.URL_CUSTOM_FIELD_TEMPLATES, method="get", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, **context ), }, @@ -254,7 +249,7 @@ def create_project_custom_field_template(self, data: dict): method="post", data=data, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -269,7 +264,7 @@ def list_project_custom_entities(self, project_id: int): url=self.URL_SET_CUSTOM_ENTITIES.format(pk=project_id), method="get", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ), }, @@ -290,7 +285,7 @@ def list_projects(self, body_query: Query, chunk_size=100) -> WMProjectListRespo "parentEntity": "Team", }, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -307,7 +302,7 @@ def search_projects( method="post", body_query=body_query, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -332,10 +327,10 @@ def list_users( url = self.URL_SEARCH_PROJECT_USERS if project_id is None: user_entity = WMUserEntity - entity_context = self._generate_context(team_id=self.client.team_id) + entity_context = encode_entity_context(team_id=self.client.team_id) else: user_entity = WMProjectUserEntity - entity_context = self._generate_context( + entity_context = encode_entity_context( team_id=self.client.team_id, project_id=project_id, ) @@ -382,7 +377,7 @@ def create_custom_field_template( "access": access if access is not None else {}, }, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, **entity_context ), }, @@ -405,7 +400,7 @@ def delete_custom_field_template( "parentEntity": parent_entity.value, }, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, **entity_context ), }, @@ -426,7 +421,7 @@ def set_custom_field_value( url=self.URL_SET_CUSTOM_ENTITIES.format(pk=entity_id), method="patch", headers={ - "x-sa-entity-context": self._generate_context(**context), + "x-sa-entity-context": encode_entity_context(**context), }, data={"customField": {"custom_field_values": {template_id: data}}}, params={ @@ -448,7 +443,7 @@ def update_user_activity( method="post", data=body, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -458,13 +453,25 @@ def list_scores(self) -> WMScoreListResponse: return self.client.paginate( url=self.URL_SCORES, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, item_type=WMScoreEntity, ) + def list_project_scores(self, project_id: int) -> WMScoreListResponse: + return self.client.paginate( + url=self.URL_PROJECT_SCORES, + headers={ + "x-sa-entity-context": encode_entity_context( + team_id=self.client.team_id, project_id=project_id + ), + }, + item_type=WMScoreEntity, + method="post", + ) + def create_score( self, name: str, @@ -482,7 +489,7 @@ def create_score( url=self.URL_SCORES, method="post", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=int(self.client.team_id) ), }, @@ -494,7 +501,7 @@ def delete_score(self, score_id: int) -> ServiceResponse: url=self.URL_DELETE_SCORE.format(score_id=score_id), method="delete", headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -533,7 +540,7 @@ def set_remove_contributor_categories( "body": {"categories": [{"id": i} for i in category_ids]}, }, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ), }, @@ -547,7 +554,7 @@ def list_permission_groups(self) -> WMPermissionGroupListResponse: return self.client.paginate( url=self.URL_PERMISSION_GROUPS, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id ), }, @@ -587,7 +594,7 @@ def edit_project_user_permissions( }, }, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ), }, @@ -619,7 +626,7 @@ def set_team_user_permissions( "body": {"userPermissions": [{"id": i} for i in permission_ids]}, }, headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, ), }, @@ -647,7 +654,7 @@ def update_annotation_class( exclude_unset=True, by_alias=True, mode="json", exclude_none=True ), headers={ - "x-sa-entity-context": self._generate_context( + "x-sa-entity-context": encode_entity_context( team_id=self.client.team_id, project_id=project_id ), }, diff --git a/tests/README.md b/tests/README.md new file mode 100644 index 000000000..28e32c5b5 --- /dev/null +++ b/tests/README.md @@ -0,0 +1,76 @@ +# Running the tests + +Unit tests need no credentials: + +```bash +pytest tests/unit +``` + +The integration tests talk to a real team, and take their credentials from a `.env` file +in the repository root (copy `.env.example`): + +```ini +SA_TOKEN= +SA_URL=https://api.devsuperannotate.com +# Required for an organization API key, which carries no team of its own. +SA_TEAM_ID=6085 +``` + +```bash +pytest tests/integration +``` + +The file is read before any test module is imported, so the modules that build their +client at import time with a bare `SAClient()` pick it up. Variables already exported in +the environment win over the file, so CI can provide them without a `.env`; with neither, +the SDK falls back to `~/.superannotate/config.ini`. + +Unit tests are hidden from these variables (`tests/unit/conftest.py`) - they assert how +the SDK itself resolves credentials. + +## Running as a different token type + +What the backend allows depends on the token in the `.env`: + +| Token | Acts as | Notes | +| --- | --- | --- | +| Organization API key | the organization | carries no team — `SA_TEAM_ID` is required | +| Team API key | the team, with no user behind it | user-level operations are denied | +| Personal (team-user) API key | the user it was issued for | owner or team admin, per key | +| Legacy team-owner token | the team owner | carries its team in the token | + +To run the suite as another type, put that token in the `.env` and run it again. The +`sa_client` fixture is the client the run authenticates as. + +## Tests that bring their own token + +A test that describes one specific kind of key does not depend on which key the run +uses: it builds its own client from a key named by its own `.env` variable, and is +skipped while that variable is unset (see `tests/env.py`): + +```python +from tests import env + +@env.requires_env_vars(env.SA_ORGANIZATION_TOKEN_ENV) +class TestOrganizationToken(TestCase): + ... +``` + +The variables the suites use: + +```ini +# What a project-admin contributor may do (tests/integration/client). +SA_OWNER_PERSONAL_TOKEN= +SA_PROJECT_ADMIN_TOKEN= +``` + +`test_project_admin_token.py` runs its setup as the owner - it creates two projects and +makes the contributor a ProjectAdmin of one of them - and then does everything else as +the contributor, so the role's reach is measured against a project it was never given. +Two of its tests are `xfail`: a project-admin key cannot list team users, and so cannot +add contributors either. Both break inside the SDK, and the reasons on the tests say +where. + +`test_token_scopes.py` and `test_org_client.py` work the same way: what an +organization key grants comes from `SA_ORGANIZATION_TOKEN`, and what it is refused - +along with what a team-bound key grants - from `SA_OWNER_PERSONAL_TOKEN`. diff --git a/tests/conftest.py b/tests/conftest.py index d0c1b94ac..ad271ecd7 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,8 +1,13 @@ -import os - import pytest +from tests import env + +# Read before any test module is imported, so a module that builds its client at import +# time authenticates with the project's .env rather than falling back to the SDK's own +# ~/.superannotate/config.ini. +env.load_dotenv() -@pytest.fixture(autouse=True) -def tests_setup(): - os.environ.update({"SA_TESTING": "True", "SA_VERSION_CHECK": "False"}) +@pytest.fixture(scope="session") +def sa_client(): + """The client the suite runs as, built from the .env credentials.""" + return env.get_client() diff --git a/tests/env.py b/tests/env.py new file mode 100644 index 000000000..4e0c5f078 --- /dev/null +++ b/tests/env.py @@ -0,0 +1,224 @@ +"""Credentials the test suite runs with, taken from a ``.env`` file. + +The suite talks to a real team, and which token it uses changes what the backend allows: +a team-scoped API key acts as the team (there is no user behind it), a personal key acts +as the user it was issued for, and an organization key is not bound to a team at all, so +it only works together with a team id. + +Put the credentials in a ``.env`` file at the repository root (override the path with +``SA_TEST_ENV_FILE``):: + + SA_TOKEN= + SA_URL=https://api.devsuperannotate.com + # Only an organization key needs it; any other key carries its own team. + SA_TEAM_ID=6085 + # For SAORGClient's own tests - a second, independent key. + SA_ORGANIZATION_TOKEN= + SA_ORGANIZATION_TEAM_ID= + +The file is read before the integration modules build their clients, so a plain +``SAClient()`` picks it up. Values already set in the environment win over the file, +which is how CI provides them. With no ``.env`` and no environment the suite falls back +to the SDK's own ``~/.superannotate/config.ini``, as it always did. + +A test that only makes sense for one kind of token builds a client from a token of that +kind, declared under its own variable, and is skipped while that variable is unset:: + + @env.requires_env_vars(env.SA_CONTRIBUTOR_TOKEN_ENV) + class TestSomething(TestCase): + @classmethod + def setUpClass(cls): + cls.client = env.build_client(env.var(env.SA_CONTRIBUTOR_TOKEN_ENV)) +""" + +import configparser +import contextlib +import os +import tempfile +import unittest +from functools import lru_cache +from pathlib import Path + +#: Overrides the location of the .env file. +ENV_FILE_ENV = "SA_TEST_ENV_FILE" +DEFAULT_ENV_FILE = Path(__file__).parent.parent / ".env" + +#: Tokens the suite can build a client with, beyond the ``SA_TOKEN`` it runs as. +#: A suite that needs one declares it (``requires_env_vars``) and is skipped without it. +OWNER_PERSONAL_TOKEN_ENV = "SA_OWNER_PERSONAL_TOKEN" +SA_CONTRIBUTOR_TOKEN_ENV = "SA_CONTRIBUTOR_TOKEN" +#: A key for SAORGClient's own tests, independent of SA_TOKEN's scope, plus a team it +#: can reach. +SA_ORGANIZATION_TOKEN_ENV = "SA_ORGANIZATION_TOKEN" +SA_URL = "SA_URL" +SA_ORGANIZATION_TEAM_ID_ENV = "SA_ORGANIZATION_TEAM_ID" + + +def env_file() -> Path: + return Path(os.environ.get(ENV_FILE_ENV) or DEFAULT_ENV_FILE).expanduser() + + +def dotenv_values(path=None) -> dict: + """What a ``.env`` file holds, read straight from the file. + + Unlike ``load_dotenv`` this neither consults nor touches the environment, so it + still answers when the environment has been scrubbed. + """ + path = Path(path) if path else env_file() + if not path.is_file(): + return {} + values = {} + for line in path.read_text().splitlines(): + line = line.strip() + if not line or line.startswith("#") or "=" not in line: + continue + key, _, value = line.partition("=") + key = key.strip().removeprefix("export ").strip() + if key: + values[key] = value.strip().strip("\"'") + return values + + +def load_dotenv(path=None) -> dict: + """Read a ``.env`` file into the environment and return what it set. + + Only keys that are not already in the environment are set, so an explicitly exported + variable (CI, or a one-off run) always wins over the file. + """ + loaded = {} + for key, value in dotenv_values(path).items(): + if key not in os.environ: + os.environ[key] = value + loaded[key] = value + return loaded + + +@lru_cache(maxsize=None) +def get_client(): + """The client the suite runs as, built from the environment (``.env``). + + Cached: building one costs a token-scope round trip to the backend. + """ + from src.superannotate import SAClient + + load_dotenv() + return SAClient() + + +@contextlib.contextmanager +def environ(**values): + """Temporarily set environment variables; a ``None`` value unsets one.""" + saved = {key: os.environ.get(key) for key in values} + try: + for key, value in values.items(): + if value is None: + os.environ.pop(key, None) + else: + os.environ[key] = value + yield + finally: + for key, value in saved.items(): + if value is None: + os.environ.pop(key, None) + else: + os.environ[key] = value + + +def _sdk_environ(token: str, team_id: int | None = None) -> dict: + """The whole environment a client is built in: every variable the SDK reads. + + Every one is set here rather than inherited, so building a client does not depend + on what the ambient environment happens to hold. Two ways that used to bite: + a suite that scrubs the environment (``tests/unit/conftest.py``) would silently get + a client pointed at production, and a stray ``SA_TEAM_ID`` would attach itself to a + token that never asked for one. The backend comes from the environment first, then + the ``.env`` file, so an exported override still wins. + """ + from_file = dotenv_values() + values = { + "SA_TOKEN": token, + "SA_TEAM_ID": str(team_id) if team_id is not None else None, + } + # SA_URL only: the SDK reads SA_SSL but can never act on it, since + # _retrieve_configs_from_env assigns VERIFY_SSL only when it is already True. + values["SA_URL"] = os.environ.get("SA_URL") or from_file.get("SA_URL") + return values + + +def build_client(token: str, team_id: int | None = None, team_id_via_env: bool = False): + """An ``SAClient`` for an ad-hoc token, on the backend the ``.env`` names. + + The token reaches the SDK the way the suite's own credentials do - through the + environment - so ``SA_URL`` still applies. Passing it as ``SAClient(token=...)`` + would not: only the no-argument path reads ``SA_URL``. + + The team is passed as the ``team_id`` argument, or as ``SA_TEAM_ID`` when + ``team_id_via_env`` is set; both are paths a caller has. It is never inherited from + the ``.env``, so a client can be built with no team at all. + """ + from src.superannotate import SAClient + + with environ(**_sdk_environ(token, team_id if team_id_via_env else None)): + return SAClient(team_id=None if team_id_via_env else team_id) + + +@contextlib.contextmanager +def _config_file(token: str): + """A throwaway ini config naming the credentials and the backend. + + Given ``config_path``, the SDK reads the file and consults the environment for + nothing, so a client built from one can neither be perturbed by the ambient + environment nor leak into it. The backend still comes from the environment first + and the ``.env`` file second, so an exported override wins as it always did. + """ + settings = {"SA_TOKEN": token} + url = os.environ.get("SA_URL") or dotenv_values().get("SA_URL") + if url: + settings["SA_URL"] = url + with tempfile.TemporaryDirectory() as directory: + path = Path(directory) / "config.ini" + parser = configparser.ConfigParser() + parser.optionxform = str + parser["DEFAULT"] = settings + with path.open("w") as handle: + parser.write(handle) + yield str(path) + + +def build_org_client(token: str): + """An ``SAORGClient`` for an ad-hoc token, on the backend the ``.env`` names. + + Built from a config file rather than the environment: an organization client is + bound to no team, so a stray ``SA_TEAM_ID`` must not reach it, and a scrubbed + ``SA_URL`` must not move it to another backend. Nothing about it is inherited. + """ + from src.superannotate import SAORGClient + + with _config_file(token) as config_path: + return SAORGClient(config_path=config_path) + + +def var(name: str) -> str: + """A token the ``.env`` provides under ``name`` (one of the ``*_TOKEN_ENV``).""" + value = os.environ.get(name) or dotenv_values().get(name) + if not value: + raise KeyError(name) + return value + + +def missing_env_vars(*names: str) -> list[str]: + """Which of these ``.env`` variables (tokens, team ids, ...) are not provided.""" + from_file = dotenv_values() + return [name for name in names if not (os.environ.get(name) or from_file.get(name))] + + +def requires_env_vars(*names: str): + """Run only when the ``.env`` provides every one of these variables. + + This asks nothing of the backend: a variable is either in the environment or it is + not, so a test or a whole ``TestCase`` is skipped on the spot. + """ + missing = missing_env_vars(*names) + return unittest.skipIf( + bool(missing), f"needs {', '.join(missing)} in the .env (see tests/env.py)" + ) diff --git a/tests/integration/__init__.py b/tests/integration/__init__.py index e69de29bb..7a0d196d9 100644 --- a/tests/integration/__init__.py +++ b/tests/integration/__init__.py @@ -0,0 +1,6 @@ +"""Integration tests, run against a real team. + +The credentials come from the project's ``.env`` file (see ``tests/env.py``), which +``tests/conftest.py`` reads into the environment before any of these modules is imported - +they build their client at import time with a bare ``SAClient()``. +""" diff --git a/tests/integration/annotations/validations/test_gen_ai_annotation_validation.py b/tests/integration/annotations/validations/test_gen_ai_annotation_validation.py deleted file mode 100644 index 30075bb2f..000000000 --- a/tests/integration/annotations/validations/test_gen_ai_annotation_validation.py +++ /dev/null @@ -1,17 +0,0 @@ -from unittest import TestCase -from unittest.mock import patch - -from src.superannotate import SAClient - -sa = SAClient() - - -class TestVectorValidators(TestCase): - PROJECT_TYPE = "Multimodal" - - @patch("builtins.print") - def test_validate_annotation_without_metadata(self, mock_print): - # Failed because the BED does not have a validation schema for Multimodal projects. - is_valid = sa.validate_annotations(self.PROJECT_TYPE, {"instances": []}) - assert not is_valid - mock_print.assert_any_call("'metadata' is a required property") diff --git a/tests/integration/client/__init__.py b/tests/integration/client/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/integration/client/test_org_client.py b/tests/integration/client/test_org_client.py new file mode 100644 index 000000000..f501b9a9b --- /dev/null +++ b/tests/integration/client/test_org_client.py @@ -0,0 +1,90 @@ +"""What SAORGClient can do: list an organization's teams, and mint a team-scoped +SAClient on demand. + +Every test here builds its client from a key named in the .env - SA_ORGANIZATION_TOKEN +for what an org key can do, SA_OWNER_PERSONAL_TOKEN for what it rejects - and is skipped +while that key is unset (see tests/env.py). +""" + +import os +from unittest import TestCase + +import pytest +from src.superannotate import AppException +from src.superannotate import SAClient +from src.superannotate.lib.core.entities.context import TokenScope +from superannotate import SAAuthError +from tests import env + + +@env.requires_env_vars(env.SA_ORGANIZATION_TOKEN_ENV, env.SA_ORGANIZATION_TEAM_ID_ENV) +class TestOrgClient(TestCase): + @classmethod + def setUpClass(cls): + cls.org_client = env.build_org_client(env.var(env.SA_ORGANIZATION_TOKEN_ENV)) + cls.team_id = int(os.environ[env.SA_ORGANIZATION_TEAM_ID_ENV]) + + def test_authenticates_with_no_team(self): + context = self.org_client.controller.token_context + assert context.scope == TokenScope.ORGANIZATION + assert context.team_id is None + + def test_list_teams_contains_the_configured_team(self): + teams = self.org_client.list_teams() + + assert teams + matching = next((team for team in teams if team["id"] == self.team_id), None) + assert matching is not None + # Only the documented fields - the backend also sends type/user_role/is_default. + assert matching.keys() == { + "id", + "name", + "description", + "creator_id", + "owner_id", + "owner_type", + } + + def test_get_team_client_returns_a_working_team_scoped_client(self): + team_client = self.org_client.get_team_client(self.team_id) + + assert isinstance(team_client, SAClient) + assert team_client.team_id == self.team_id + # A real, team-scoped call: proves the returned client actually authenticates + assert team_client.controller.current_user.email + team_all_projects = team_client.list_projects() + assert team_all_projects + for p in team_all_projects: + assert p["team_id"] == self.team_id + + self.org_client.list_teams() + + def test_get_team_client_rejects_a_non_integer_team_id(self): + with self.assertRaisesRegex(AppException, r"Input should be a valid integer"): + self.org_client.get_team_client("not-an-id") + + def test_invalid_team_id(self): + with self.assertRaisesRegex(AppException, r"Invalid team id provided."): + self.org_client.get_team_client(1) + + +@env.requires_env_vars(env.SA_ORGANIZATION_TOKEN_ENV) +def test_sa_client_via_org_token(): + with pytest.raises( + SAAuthError, + match='Team context not provided. An Organization API key requires a "team_id".', + ): + sa = SAClient( + config={ + "SA_TOKEN": env.var(env.SA_ORGANIZATION_TOKEN_ENV), + "SA_URL": env.var(env.SA_URL), + } + ) + + +@env.requires_env_vars(env.OWNER_PERSONAL_TOKEN_ENV) +def test_a_team_bound_token_is_rejected(): + # SAORGClient takes only an organization key: a key bound to one team - personal + # here, but a team key or a le File "/Users/vaghinak.basentsyan/env/superannotate-python-sdk3-14/lib/python3.14/site-packages/pydantic/_internal/_validate_call.py", line 137, in __call__gacy token the same way - cannot act for the org. + with pytest.raises(AppException, match=r"Invalid credentials provided\."): + env.build_org_client(env.var(env.OWNER_PERSONAL_TOKEN_ENV)) diff --git a/tests/integration/client/test_project_admin_token.py b/tests/integration/client/test_project_admin_token.py new file mode 100644 index 000000000..90b4be47d --- /dev/null +++ b/tests/integration/client/test_project_admin_token.py @@ -0,0 +1,442 @@ +"""What a project-admin contributor's API key can do. + +The suite runs only when the ``.env`` holds a personal API key of a team contributor +(``SA_PROJECT_ADMIN_TOKEN``); it is skipped otherwise. The setup runs as the team owner +(``SA_OWNER_PERSONAL_TOKEN``) and creates two projects, making that contributor a +ProjectAdmin of one of them - so what the role grants can be told apart from what it +does not. Neither client comes from the suite's own ``SA_TOKEN``: these tests describe +the project-admin key itself, whichever token the rest of the run uses. + +Two of them are ``xfail``: a project-admin key cannot list team users, and therefore +cannot add contributors either. Both fail inside the SDK rather than at the backend, see +the reasons on the tests. +""" + +import contextlib +import json +import os +import time +import uuid +from pathlib import Path +from unittest import TestCase + +from src.superannotate import AppException +from tests import env +from tests.integration.work_management.data_set import SCORE_TEMPLATES + + +class BaseProjectAdminTest(TestCase): + #: The project the contributor administers. + PROJECT_NAME = "TestProjectAdminToken" + #: A project they are never added to, so it has to stay out of their reach. + FOREIGN_PROJECT_NAME = "TestProjectAdminTokenForeign" + PROJECT_DESCRIPTION = "project-admin token suite" + PROJECT_TYPE = "Multimodal" + FOLDER_NAME = "created-by-project-admin" + SETTINGS = [ + {"attribute": "TemplateState", "value": 1}, + {"attribute": "CategorizeItems", "value": 2}, + {"attribute": "UploadImages", "value": 1}, + {"attribute": "DeleteImages", "value": 1}, + ] + MULTIMODAL_FORM = { + "components": [ + { + "id": "r_qx07c6", + "type": "audio", + "permissions": [], + "hasTooltip": False, + "exclude": False, + "label": "", + "value": "", + } + ], + "readme": "", + } + + def setUp(self) -> None: + #: The team owner, who sets the projects up and cleans them up. + self.owner = env.build_client(env.var(env.OWNER_PERSONAL_TOKEN_ENV)) + #: The client under test: a contributor's key, made project admin below. + self.project_admin = env.build_client(env.var(env.SA_CONTRIBUTOR_TOKEN_ENV)) + #: The user that key acts as - the one the owner promotes. + self.project_admin_email = self.project_admin.controller.current_user.email + + self._delete_projects() + self._project = self.owner.create_project( + self.PROJECT_NAME, + self.PROJECT_DESCRIPTION, + self.PROJECT_TYPE, + settings=self.SETTINGS, + form=self.MULTIMODAL_FORM, + ) + self.owner.create_project( + self.FOREIGN_PROJECT_NAME, + self.PROJECT_DESCRIPTION, + self.PROJECT_TYPE, + settings=[ + {"attribute": "TemplateState", "value": 1}, + {"attribute": "CategorizeItems", "value": 2}, + ], + form=self.MULTIMODAL_FORM, + ) + added, skipped = self.owner.add_contributors_to_project( + self.PROJECT_NAME, [self.project_admin_email], "ProjectAdmin" + ) + assert self.project_admin_email in added + skipped, ( + f"{self.project_admin_email} is out of the team scope, so it cannot be made " + f"a project admin - {env.SA_CONTRIBUTOR_TOKEN_ENV} has to belong to a member " + f"of the team {env.OWNER_PERSONAL_TOKEN_ENV} owns" + ) + + def tearDown(self) -> None: + self._delete_projects() + + def _delete_projects(self) -> None: + for name in (self.PROJECT_NAME, self.FOREIGN_PROJECT_NAME): + for project in self.owner.list_projects(name=name): + with contextlib.suppress(Exception): + self.owner.delete_project(project["id"]) + + +@env.requires_env_vars(env.OWNER_PERSONAL_TOKEN_ENV, env.SA_CONTRIBUTOR_TOKEN_ENV) +class TestProjectAdminTokenFullAccess(BaseProjectAdminTest): + #: The project the contributor administers. + PROJECT_NAME = "TestProjectAdminToken" + #: A project they are never added to, so it has to stay out of their reach. + FOREIGN_PROJECT_NAME = "TestProjectAdminTokenForeign" + PROJECT_DESCRIPTION = "project-admin token suite" + PROJECT_TYPE = "Multimodal" + FOLDER_NAME = "created-by-project-admin" + + MULTIMODAL_FORM = { + "components": [ + { + "id": "r_qx07c6", + "type": "audio", + "permissions": [], + "hasTooltip": False, + "exclude": False, + "label": "", + "value": "", + } + ], + "readme": "", + } + + def _team_contributor(self): + """A team contributor for the project admin to add, found as the owner. + + The lookup runs as the owner on purpose: a project-admin key cannot list team + users (see ``test_lists_team_users``), so it cannot pick its own candidate. + """ + for user in self.owner.list_users(): + if ( + user["role"] == "Contributor" + and user["email"] != self.project_admin_email + ): + return user + self.skipTest("the team has no other contributor to add to a project") + + def test_lists_only_the_projects_it_has_access_to(self): + visible = {p["name"] for p in self.project_admin.list_projects()} + + assert self.PROJECT_NAME in visible + # The second project was never shared, so the role must not surface it. + assert self.FOREIGN_PROJECT_NAME not in visible + assert self.FOREIGN_PROJECT_NAME in { + p["name"] for p in self.owner.list_projects() + } + + def test_add_remove_a_contributor_to_its_project(self): + # TODO should raise error on ProjectAdmin deletion + scapegoat = self._team_contributor() + + self.project_admin.add_contributors_to_project( + self.PROJECT_NAME, [scapegoat["email"]], "ProjectAdmin" + ) + + project_roles = { + user["email"]: user["role"] + for user in self.project_admin.list_users(project=self.PROJECT_NAME) + } + assert project_roles.get(scapegoat["email"]) == "ProjectAdmin" + self.project_admin.remove_users_from_project( + self.PROJECT_NAME, [scapegoat["email"]] + ) + project_roles = { + user["email"]: user["role"] + for user in self.project_admin.list_users(project=self.PROJECT_NAME) + } + assert scapegoat["email"] not in project_roles + + def test_lists_team_users(self): + team_users = self.project_admin.list_users() + + assert self.project_admin_email in {user["email"] for user in team_users} + + def test_lists_the_users_of_its_project_with_categories(self): + project_users = self.project_admin.list_users(project=self.PROJECT_NAME) + + project_roles = {user["email"]: user["role"] for user in project_users} + assert project_roles[self.project_admin_email] == "ProjectAdmin" + scapegoat = self._team_contributor() + + self.project_admin.add_contributors_to_project( + self.PROJECT_NAME, [scapegoat["email"]], "Annotator" + ) + self.project_admin.create_categories(self.PROJECT_NAME, ["test"]) + categories = self.project_admin.list_categories(self.PROJECT_NAME) + assert len(categories) == 1 + self.project_admin.set_contributors_categories( + self.PROJECT_NAME, [scapegoat["email"]], categories=["test"] + ) + users = self.project_admin.list_users( + project=self.PROJECT_NAME, email=scapegoat["email"], include=["categories"] + ) + assert users[0]["categories"][0]["name"] == "test" + + def test_creates_a_folder_in_its_project(self): + folder = self.project_admin.create_folder(self.PROJECT_NAME, self.FOLDER_NAME) + + assert folder["name"] == self.FOLDER_NAME + assert self.FOLDER_NAME in { + f["name"] for f in self.project_admin.list_folders(self.PROJECT_NAME) + } + + def test_get_list_delete_items(self): + self.project_admin.generate_items(self.PROJECT_NAME, count=5, name="test") + + items = self.project_admin.list_items( + self.PROJECT_NAME, include=["categories", "custom_metadata"] + ) + item = self.project_admin.get_item_metadata(self.PROJECT_NAME, items[0]["name"]) + assert len(items) == 5 + assert item is not None + self.project_admin.delete_items(self.PROJECT_NAME) + items = self.project_admin.list_items( + self.PROJECT_NAME, include=["categories", "custom_metadata"] + ) + assert len(items) == 0 + + def test_set_item_status(self): + self.project_admin.generate_items(self.PROJECT_NAME, count=1, name="test") + + items = self.project_admin.list_items( + self.PROJECT_NAME, include=["categories", "custom_metadata"] + ) + self.project_admin.set_annotation_statuses( + self.PROJECT_NAME, "Completed", [items[0]["name"]] + ) + items = self.project_admin.list_items( + self.PROJECT_NAME, include=["categories", "custom_metadata"] + ) + assert items[0]["annotation_status"] == "Completed" + + def test_get_set_annotation(self): + self.project_admin.generate_items(self.PROJECT_NAME, count=5, name="test") + annotations = self.project_admin.get_annotations( + self.PROJECT_NAME, data_spec="multimodal" + ) + assert len(annotations) == 5 + response = self.project_admin.upload_annotations( + self.PROJECT_NAME, annotations, data_spec="multimodal" + ) + assert len(response["succeeded"]) == 5 + + def test_get_project_metadata(self): + self.project_admin.get_project_metadata( + project=self.PROJECT_NAME, + include_annotation_classes=True, + include_settings=True, + # include_workflow=True, + include_contributors=True, + include_complete_item_count=True, + ) + + def test_delete_project(self): + # TODO fix project admin should not be able to delete project + self.project_admin.delete_project(self.PROJECT_NAME) + projects = self.project_admin.list_projects(name=self.PROJECT_NAME) + assert not projects + + +@env.requires_env_vars(env.OWNER_PERSONAL_TOKEN_ENV, env.SA_CONTRIBUTOR_TOKEN_ENV) +class TestProjectAdminSemiAccess(BaseProjectAdminTest): + PROJECT_NAME = "TestProjectAdminSemiAccess" + FOREIGN_PROJECT_NAME = "TestProjectAdminSemiAccessFOREIGN" + SETTINGS = [ + {"attribute": "TemplateState", "value": 1}, + {"attribute": "CategorizeItems", "value": 2}, + {"attribute": "UploadImages", "value": 0}, + {"attribute": "DeleteImages", "value": 0}, + ] + + def test_item_deletion(self): + self.owner.generate_items(self.PROJECT_NAME, count=5, name="test") + with self.assertRaisesRegex( + AppException, "You do not have sufficient access to delete this items." + ): + self.project_admin.delete_items(self.PROJECT_NAME) + + def test_create_items(self): + # todo update error message + with self.assertRaisesRegex( + AppException, "You do not have sufficient access export." + ): + self.project_admin.generate_items(self.PROJECT_NAME, count=5, name="test") + + +@env.requires_env_vars(env.OWNER_PERSONAL_TOKEN_ENV, env.SA_CONTRIBUTOR_TOKEN_ENV) +class TestProjectAdminUserScoring(TestCase): + """ + Test using mock Multimodal form template with dynamically generated scores created during setup. + """ + + PROJECT_NAME = "TestProjectAdminUserScoring" + PROJECT_TYPE = "Multimodal" + PROJECT_DESCRIPTION = "DESCRIPTION" + EDITOR_TEMPLATE_PATH = os.path.join( + Path(__file__).parent.parent.parent, + "data_set/editor_templates/form_with_scores.json", + ) + CLASSES_TEMPLATE_PATH = os.path.join( + Path(__file__).parent.parent.parent, + "data_set/editor_templates/form1_classes.json", + ) + MULTIMODAL_FORM = { + "components": [ + { + "id": "r_qx07c6", + "type": "audio", + "permissions": [], + "hasTooltip": False, + "exclude": False, + "label": "", + "value": "", + } + ], + "readme": "", + } + + def setUp(self, *args, **kwargs) -> None: + # setup user scores for test + self.owner = env.build_client(env.var(env.OWNER_PERSONAL_TOKEN_ENV)) + self.tearDown() + #: The client under test: a contributor's key, made project admin below. + self.project_admin = env.build_client(env.var(env.SA_CONTRIBUTOR_TOKEN_ENV)) + self.project_admin_email = self.project_admin.controller.current_user.email + self._project = self.owner.create_project( + self.PROJECT_NAME, + self.PROJECT_DESCRIPTION, + self.PROJECT_TYPE, + settings=[{"attribute": "TemplateState", "value": 1}], + form=self.MULTIMODAL_FORM, + ) + team = self.owner.controller.team + project = self.owner.controller.get_project(self.PROJECT_NAME) + self.owner.add_contributors_to_project( + self.PROJECT_NAME, [self.project_admin_email], "ProjectAdmin" + ) + time.sleep(5) + + # setup form template from crated scores + with open(self.EDITOR_TEMPLATE_PATH) as f: + template_data = json.load(f) + for data in SCORE_TEMPLATES: + req = ( + self.owner.controller.service_provider.work_management.create_score( + **data + ) + ) + assert req.status_code == 201 + for component in template_data["components"]: + if "scoring" in component and component["type"] == req.data["type"]: + component["scoring"]["id"] = req.data["id"] + + res = ( + self.owner.controller.service_provider.projects.attach_editor_template( + team, project, template=template_data + ) + ) + assert res.ok + self.owner.create_annotation_classes_from_classes_json( + self.PROJECT_NAME, self.CLASSES_TEMPLATE_PATH + ) + + users = self.owner.list_users() + scapegoat = [ + u for u in users if u["role"] == "Contributor" and u["state"] == "Confirmed" + ][0] + self.scapegoat = scapegoat + self.owner.add_contributors_to_project( + self.PROJECT_NAME, [scapegoat["email"]], "Annotator" + ) + + def tearDown(self) -> None: + # cleanup test scores and project + projects = self.owner.list_projects(name__in=[self.PROJECT_NAME]) + for project in projects: + try: + self.owner.delete_project(project=project["id"]) + except Exception as _: + pass + + score_templates_name_id_map = { + s.name: s.id + for s in self.owner.controller.service_provider.work_management.list_scores().data + } + for data in SCORE_TEMPLATES: + score_id = score_templates_name_id_map.get(data["name"]) + if score_id: + self.owner.controller.service_provider.work_management.delete_score( + score_id + ) + + def _attach_item(self, path, name): + self.owner.attach_items(path, [{"name": name, "url": "url"}]) + + def test_set_get_scores(self): + scores_name_payload_map = { + "SDK-my-score-1": { + "component_id": "r_34k7k7", # rating type score + "value": 5, + "weight": 0.5, + }, + "SDK-my-score-2": { + "component_id": "r_ioc7wd", # number type score + "value": 45, + "weight": 1.5, + }, + "SDK-my-score-3": { + "component_id": "r_tcof7o", # radio type score + "value": None, + "weight": None, + }, + } + item_name = f"test_item_{uuid.uuid4()}" + self._attach_item(self.PROJECT_NAME, item_name) + + with self.assertLogs("sa", level="INFO") as cm: + self.project_admin.set_user_scores( + project=self.PROJECT_NAME, + item=item_name, + scored_user=self.scapegoat["email"], + scores=list(scores_name_payload_map.values()), + ) + assert cm.output[0] == "INFO:sa:Scores successfully set." + + created_scores = self.project_admin.get_user_scores( + project=self.PROJECT_NAME, + item=item_name, + scored_user=self.scapegoat["email"], + score_names=[s["name"] for s in SCORE_TEMPLATES], + ) + assert len(created_scores) == len(SCORE_TEMPLATES) + for score in created_scores: + score_pyload = scores_name_payload_map[score["name"]] + assert score["value"] == score_pyload["value"] + assert score["weight"] == score_pyload["weight"] + assert score["id"] + assert score["createdAt"] + assert score["updatedAt"] diff --git a/tests/integration/client/test_token_scopes.py b/tests/integration/client/test_token_scopes.py new file mode 100644 index 000000000..4425cd58a --- /dev/null +++ b/tests/integration/client/test_token_scopes.py @@ -0,0 +1,70 @@ +"""What each kind of token grants, checked against the backend. + +A test that only makes sense for one kind of key builds a client from a key of that +kind, named by its own .env variable, and is skipped while that variable is unset (see +tests/env.py). Only the first test uses the ambient SA_TOKEN, whatever kind it is. +""" + +import os +from unittest import TestCase + +from src.superannotate import AppException +from src.superannotate.lib.core.entities.context import TokenScope +from tests import env + + +def test_the_suite_token_authenticates(sa_client): + assert sa_client.controller.team_id + # Every token resolves the user it acts as, or the creator behind a team key. + assert sa_client.controller.current_user.email + + +@env.requires_env_vars(env.SA_ORGANIZATION_TOKEN_ENV, env.SA_ORGANIZATION_TEAM_ID_ENV) +class TestOrganizationToken(TestCase): + """An organization key is not bound to a team, so the caller names one.""" + + @classmethod + def setUpClass(cls): + cls.token = env.var(env.SA_ORGANIZATION_TOKEN_ENV) + cls.team_id = int(os.environ[env.SA_ORGANIZATION_TEAM_ID_ENV]) + + def test_team_id_as_an_argument(self): + client = env.build_client(self.token, team_id=self.team_id) + + assert client.controller.token_context.scope == TokenScope.ORGANIZATION + assert client.controller.team_id == self.team_id + + def test_team_id_from_the_environment(self): + client = env.build_client( + self.token, team_id=self.team_id, team_id_via_env=True + ) + + assert client.controller.config["SA_TEAM_ID"] == self.team_id + assert client.controller.team_id == self.team_id + + def test_without_a_team_id_is_rejected(self): + # The team is not part of the key, so there is nothing to fall back on. + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): + env.build_client(self.token) + + +@env.requires_env_vars(env.OWNER_PERSONAL_TOKEN_ENV) +class TestPersonalToken(TestCase): + """A personal key acts as the user it was issued for, and names its own team.""" + + @classmethod + def setUpClass(cls): + cls.token = env.var(env.OWNER_PERSONAL_TOKEN_ENV) + cls.client = env.build_client(cls.token) + + def test_acts_as_a_user_of_its_own_team(self): + context = self.client.controller.token_context + + assert context.scope == TokenScope.TEAM_USER + assert context.team_id + assert self.client.controller.current_user.email + + def test_rejects_a_conflicting_team_id(self): + # The key names its own team, so a team_id that disagrees is a caller mistake. + with self.assertRaisesRegex(AppException, r"Invalid team id provided\."): + env.build_client(self.token, team_id=self.client.team_id + 1) diff --git a/tests/integration/mixpanel/test_mixpanel_decorator.py b/tests/integration/mixpanel/test_mixpanel_decorator.py index 3212035c3..a53c571fe 100644 --- a/tests/integration/mixpanel/test_mixpanel_decorator.py +++ b/tests/integration/mixpanel/test_mixpanel_decorator.py @@ -1,12 +1,10 @@ import copy import platform import tempfile -import threading from configparser import ConfigParser from unittest import TestCase from unittest.mock import patch -import pytest from src.superannotate import __version__ from src.superannotate import AppException from src.superannotate import SAClient @@ -20,11 +18,12 @@ class TestMixpanel(TestCase): "SDK": True, "Team": sa.get_team_metadata()["name"], "User Email": sa.controller.current_user.email, - "Auth Type": sa.controller.token_context.auth_type, + "Auth Type": sa.controller.token_context.scope.label, "Version": __version__, "Success": True, "Python version": platform.python_version(), "Python interpreter type": platform.python_implementation(), + "Class": "SAClient", } PROJECT_NAME = "TEST_MIX" PROJECT_DESCRIPTION = "Desc" @@ -61,9 +60,16 @@ def test_init(self, track_method): SAClient() result = list(track_method.call_args)[0] payload = self.default_payload - # team_id is part of the SAClient signature, so it is tracked like every - # other argument (None unless the token needs an explicit team). - payload.update({"sa_token": "False", "config_path": "False", "team_id": None}) + # Every argument of the signature is tracked: team_id and config are None + # unless the caller gives them, and a token is reduced to whether it was given. + payload.update( + { + "sa_token": "False", + "config_path": "False", + "team_id": None, + "config": None, + } + ) assert result[1] == "__init__" assert payload == result[2] @@ -71,6 +77,8 @@ def test_init(self, track_method): @patch("lib.core.usecases.GetTeamUseCase") @patch("lib.infrastructure.serviceprovider.ServiceProvider.get_user") def test_init_via_token(self, get_user, get_team_use_case, track_method): + get_team_use_case().execute().data.name = "Mocked Team" + get_user().data.email = "mocked@example.com" SAClient(token="test=3232") result = list(track_method.call_args)[0] payload = self.default_payload @@ -79,10 +87,11 @@ def test_init_via_token(self, get_user, get_team_use_case, track_method): "sa_token": "True", "config_path": "False", "team_id": None, + "config": None, # A legacy "=" token, whatever the ambient one is. - "Auth Type": "sdk", - "Team": get_team_use_case().execute().data.name, - "User Email": get_user().data.email, + "Auth Type": "SDK Token", + "Team": "Mocked Team", + "User Email": "mocked@example.com", } ) assert result[1] == "__init__" @@ -92,6 +101,8 @@ def test_init_via_token(self, get_user, get_team_use_case, track_method): @patch("lib.core.usecases.GetTeamUseCase") @patch("lib.infrastructure.serviceprovider.ServiceProvider.get_user") def test_init_via_config_file(self, get_user, get_team_use_case, track_method): + get_team_use_case().execute().data.name = "Mocked Team" + get_user().data.email = "mocked@example.com" with tempfile.TemporaryDirectory() as config_dir: config_ini_path = f"{config_dir}/config.ini" with patch("lib.core.CONFIG_INI_FILE_LOCATION", config_ini_path): @@ -108,9 +119,10 @@ def test_init_via_config_file(self, get_user, get_team_use_case, track_method): "sa_token": "False", "config_path": "True", "team_id": None, - "Auth Type": "sdk", - "Team": get_team_use_case().execute().data.name, - "User Email": get_user().data.email, + "config": None, + "Auth Type": "SDK Token", + "Team": "Mocked Team", + "User Email": "mocked@example.com", } ) assert result[1] == "__init__" @@ -161,39 +173,9 @@ def test_create_project(self, track_method): result = list(track_method.call_args)[0] payload = self.default_payload payload["Success"] = False + # Only a failed __init__ carries a reason; this one just marks the failure. + payload["Failure Reason"] = None payload.update(kwargs) payload["settings"] = list(kwargs["settings"].keys()) assert result[1] == "create_project" assert payload == result[2] - - @pytest.mark.skip("Need to adjust") - @patch("lib.app.interface.base_interface.Tracker._track") - def test_create_project_multi_thread(self, track_method): - project_1 = self.PROJECT_NAME + "_1" - project_2 = self.PROJECT_NAME + "_2" - try: - kwargs_1 = { - "project_name": project_1, - "project_description": self.PROJECT_DESCRIPTION, - "project_type": self.PROJECT_TYPE, - } - kwargs_2 = { - "project_name": project_2, - "project_description": self.PROJECT_DESCRIPTION, - "project_type": self.PROJECT_TYPE, - } - thread_1 = threading.Thread(target=sa.create_project, kwargs=kwargs_1) - thread_2 = threading.Thread(target=sa.create_project, kwargs=kwargs_2) - thread_1.start() - thread_2.start() - thread_1.join() - thread_2.join() - r1, r2 = track_method.call_args_list - r1_pr_name = r1[0][2].pop("project_name") - r2_pr_name = r2[0][2].pop("project_name") - assert r1_pr_name == project_1 - assert r2_pr_name == project_2 - assert r1[0][2] == r2[0][2] - finally: - self._safe_delete_project(project_1) - self._safe_delete_project(project_2) diff --git a/tests/integration/test_cli.py b/tests/integration/test_cli.py index 4d862767d..38ff48ba0 100644 --- a/tests/integration/test_cli.py +++ b/tests/integration/test_cli.py @@ -1,18 +1,18 @@ import os import tempfile from configparser import ConfigParser +from importlib.metadata import version from os.path import dirname from pathlib import Path from unittest import TestCase from unittest.mock import patch -import pkg_resources import src.superannotate.lib.core as constants from src.superannotate import SAClient from src.superannotate.lib.app.interface.cli_interface import CLIFacade try: - CLI_VERSION = pkg_resources.get_distribution("superannotate").version + CLI_VERSION = version("superannotate") except Exception: CLI_VERSION = None diff --git a/tests/integration/work_management/test_team_admin_user_permissions.py b/tests/integration/work_management/test_team_admin_user_permissions.py index 6ee69a7ec..9ac37dcc6 100644 --- a/tests/integration/work_management/test_team_admin_user_permissions.py +++ b/tests/integration/work_management/test_team_admin_user_permissions.py @@ -4,6 +4,7 @@ from unittest import TestCase from lib.core import TEAM_USER_PERMISSION_DEPRECATED_IDS +from lib.core.entities.context import TokenScope from lib.core.exceptions import AppException from src.superannotate import SAClient @@ -16,7 +17,7 @@ #: a user and may update team admin permissions, which is what the bulk of this #: module asserts. The suite therefore picks its expectations from the token the #: client was built with. -IS_TEAM_KEY = sa.controller.token_context.is_team_key +IS_TEAM_KEY = sa.controller.token_context.scope == TokenScope.TEAM TEAM_KEY_ONLY = "requires a team-scoped API key" USER_KEY_ONLY = "requires a personal (team-user) API key or a legacy token" #: Reason line the SDK adds to every permission-update failure, and the only one @@ -64,8 +65,7 @@ def _confirmed_admins(cls): return [ u for u in sa.list_users() - if u.get("state") == "Confirmed" - and u.get("role") in ("TeamAdmin", "TeamOwner") + if u.get("state") == "Confirmed" and u.get("role") in ("TeamAdmin") ] @classmethod diff --git a/tests/integration/work_management/test_user_scoring.py b/tests/integration/work_management/test_user_scoring.py index 3363d44e6..20888bab1 100644 --- a/tests/integration/work_management/test_user_scoring.py +++ b/tests/integration/work_management/test_user_scoring.py @@ -2,15 +2,17 @@ import os import time import uuid -from pathlib import Path from unittest import TestCase from lib.core.exceptions import AppException from src.superannotate import SAClient +from tests import DATA_SET_PATH from tests.integration.work_management.data_set import SCORE_TEMPLATES sa = SAClient() +print(sa.controller.token_context.token) + class TestUserScoring(TestCase): """ @@ -21,12 +23,10 @@ class TestUserScoring(TestCase): PROJECT_TYPE = "Multimodal" PROJECT_DESCRIPTION = "DESCRIPTION" EDITOR_TEMPLATE_PATH = os.path.join( - Path(__file__).parent.parent.parent, - "data_set/editor_templates/form_with_scores.json", + DATA_SET_PATH / "editor_templates" / "form_with_scores.json" ) CLASSES_TEMPLATE_PATH = os.path.join( - Path(__file__).parent.parent.parent, - "data_set/editor_templates/form1_classes.json", + DATA_SET_PATH / "editor_templates" / "form1_classes.json" ) MULTIMODAL_FORM = { "components": [ diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py new file mode 100644 index 000000000..a4098f6f4 --- /dev/null +++ b/tests/unit/conftest.py @@ -0,0 +1,30 @@ +"""Unit tests read no credentials. + +The suite loads the project's ``.env`` (``tests/conftest.py``) for the sake of modules +that build a client at import time. The tests here assert how the SDK picks up +credentials, so they must not see whatever happens to be in the developer's ``.env`` - +neither through the environment nor through a ``load_dotenv()`` of the code under test. +""" + +import pytest +from tests import env + +CREDENTIAL_VARS = ( + "SA_TOKEN", + "SA_URL", + "SA_TEAM_ID", + "SA_SSL", + env.OWNER_PERSONAL_TOKEN_ENV, + env.SA_CONTRIBUTOR_TOKEN_ENV, + env.SA_ORGANIZATION_TOKEN_ENV, + env.SA_ORGANIZATION_TEAM_ID_ENV, +) + + +@pytest.fixture(autouse=True) +def hide_credentials(tmp_path_factory): + overrides = {var: None for var in CREDENTIAL_VARS} + # A .env path that does not exist, so re-reading the file cannot bring them back. + overrides[env.ENV_FILE_ENV] = str(tmp_path_factory.mktemp("no-dotenv") / ".env") + with env.environ(**overrides): + yield diff --git a/tests/unit/test_env.py b/tests/unit/test_env.py new file mode 100644 index 000000000..a24118a0d --- /dev/null +++ b/tests/unit/test_env.py @@ -0,0 +1,262 @@ +"""The test suite's own credential plumbing (tests/env.py).""" + +import json +import os +import tempfile +from pathlib import Path +from unittest import TestCase +from unittest.mock import MagicMock +from unittest.mock import patch + +import superannotate # noqa: F401 +from tests import env + +# Imported for its side effect as much as its contents: importing the package puts the +# SDK's internal `lib` package on sys.path, which the patch targets below address. + +TOKEN = "sa_SOZVLlnbheUITTGb_PXlk2ON5QtqNPWY9bHZJctzlx4EPTkImzncQgRmybgh" + + +class LoadDotenvTestCase(TestCase): + def setUp(self): + self._dir = tempfile.TemporaryDirectory() + self.addCleanup(self._dir.cleanup) + self.env_path = Path(self._dir.name) / ".env" + # The developer's own credentials must not leak into the assertions. + patcher = patch.dict(os.environ, {}, clear=False) + patcher.start() + self.addCleanup(patcher.stop) + for key in ("SA_TOKEN", "SA_URL", "SA_TEAM_ID"): + os.environ.pop(key, None) + + def test_values_are_read_into_the_environment(self): + self.env_path.write_text( + "# credentials\n" + f"SA_TOKEN={TOKEN}\n" + "\n" + 'SA_URL="https://api.devsuperannotate.com"\n' + "export SA_TEAM_ID = 6085\n" + "not a pair\n" + ) + loaded = env.load_dotenv(self.env_path) + + assert loaded == { + "SA_TOKEN": TOKEN, + "SA_URL": "https://api.devsuperannotate.com", + "SA_TEAM_ID": "6085", + } + assert os.environ["SA_TOKEN"] == TOKEN + # Quotes are stripped, comments and malformed lines are ignored. + assert os.environ["SA_URL"] == "https://api.devsuperannotate.com" + assert os.environ["SA_TEAM_ID"] == "6085" + + def test_the_environment_wins_over_the_file(self): + # CI exports the credentials; a leftover .env must not override them. + self.env_path.write_text(f"SA_TOKEN={TOKEN}\nSA_URL=from-file\n") + with patch.dict(os.environ, {"SA_URL": "from-environment"}): + loaded = env.load_dotenv(self.env_path) + assert os.environ["SA_URL"] == "from-environment" + assert "SA_URL" not in loaded + + def test_missing_file_is_not_an_error(self): + # Without a .env the SDK falls back to its own config, as it always did. + assert env.load_dotenv(Path(self._dir.name) / "absent") == {} + + def test_path_is_overridable(self): + self.env_path.write_text(f"SA_TOKEN={TOKEN}\n") + with patch.dict(os.environ, {env.ENV_FILE_ENV: str(self.env_path)}): + assert env.env_file() == self.env_path + assert env.load_dotenv() == {"SA_TOKEN": TOKEN} + + def test_dotenv_credentials_reach_the_client(self): + from superannotate import SAClient + + self.env_path.write_text("SA_TOKEN=token=6085\nSA_URL=https://sa.test\n") + env.load_dotenv(self.env_path) + with patch("lib.infrastructure.controller.TeamController.get_team"), patch( + "lib.infrastructure.controller.TeamController.get_current_user" + ): + client = SAClient() + assert client.controller.team_id == 6085 + assert client.controller.config["SA_URL"] == "https://sa.test" + + +class RequiresEnvVarsTestCase(TestCase): + """The gate in front of the suites that need extra variables from the .env.""" + + def _decorate(self, *names): + @env.requires_env_vars(*names) + class Suite(TestCase): + pass + + return Suite + + def test_runs_when_every_token_is_there(self): + with env.environ(**{env.SA_CONTRIBUTOR_TOKEN_ENV: TOKEN}): + assert env.missing_env_vars(env.SA_CONTRIBUTOR_TOKEN_ENV) == [] + suite = self._decorate(env.SA_CONTRIBUTOR_TOKEN_ENV) + assert getattr(suite, "__unittest_skip__", False) is False + + def test_skips_naming_only_the_missing_ones(self): + with env.environ( + **{env.SA_CONTRIBUTOR_TOKEN_ENV: TOKEN, env.OWNER_PERSONAL_TOKEN_ENV: None} + ): + names = (env.OWNER_PERSONAL_TOKEN_ENV, env.SA_CONTRIBUTOR_TOKEN_ENV) + assert env.missing_env_vars(*names) == [env.OWNER_PERSONAL_TOKEN_ENV] + suite = self._decorate(*names) + assert suite.__unittest_skip__ is True + assert env.OWNER_PERSONAL_TOKEN_ENV in suite.__unittest_skip_why__ + assert env.SA_CONTRIBUTOR_TOKEN_ENV not in suite.__unittest_skip_why__ + + +class _BuilderFixture(TestCase): + """A .env naming a backend, plus canned token-scope responses. + + Shared by the two builders' cases. The unit suite scrubs the credential variables + (tests/unit/conftest.py), so a builder that inherited any of them would behave + differently here than in an integration run - and would quietly point a client at + production. + """ + + TEAM_TOKEN = { + "user": None, + "token": { + "scope": {"team_id": 6085}, + "scope_type": "team", + "created_by": "a@b.com", + "status": "ACTIVE", + }, + } + ORG_TOKEN = { + "user": None, + "token": { + "scope": {"organization_id": "org-1"}, + "scope_type": "organization", + "created_by": "a@b.com", + "status": "ACTIVE", + }, + } + API_KEY = "sa_SOZVLlnbheUITTGb_PXlk2ON5QtqNPWY9bHZJctzlx4EPTkImzncQgRmybgh" + + def setUp(self): + self.env_path = Path(tempfile.mkdtemp()) / ".env" + self.env_path.write_text("SA_URL=https://sa.test\n") + patcher = patch.dict(os.environ, {env.ENV_FILE_ENV: str(self.env_path)}) + patcher.start() + self.addCleanup(patcher.stop) + + @staticmethod + def _response(payload): + response = MagicMock() + response.ok = True + response.status_code = 200 + response.text = json.dumps(payload) + response.json.return_value = payload + return response + + def _respond_with(self, payload): + return patch( + "lib.infrastructure.services.auth.requests.post", + return_value=self._response(payload), + ) + + +class BuildClientTestCase(_BuilderFixture): + """env.build_client sets every variable the SDK reads, inheriting none.""" + + def test_the_backend_comes_from_the_dotenv_even_when_the_environment_is_scrubbed( + self, + ): + with self._respond_with(self.TEAM_TOKEN): + client = env.build_client(self.API_KEY) + + assert client.controller.config["SA_URL"] == "https://sa.test" + + def test_an_exported_url_still_wins_over_the_file(self): + with patch.dict(os.environ, {"SA_URL": "https://exported.test"}): + with self._respond_with(self.TEAM_TOKEN): + client = env.build_client(self.API_KEY) + + assert client.controller.config["SA_URL"] == "https://exported.test" + + def test_an_ambient_team_id_does_not_attach_itself_to_the_token(self): + # A team key names its own team; an inherited SA_TEAM_ID that disagreed would + # be rejected as a conflicting team id. + with patch.dict(os.environ, {"SA_TEAM_ID": "42"}): + with self._respond_with(self.TEAM_TOKEN): + client = env.build_client(self.API_KEY) + + assert client.controller.team_id == 6085 + + def test_nothing_is_left_behind_in_the_environment(self): + # It used to call load_dotenv, which writes to os.environ for good. + with self._respond_with(self.TEAM_TOKEN): + env.build_client(self.API_KEY) + + assert "SA_TOKEN" not in os.environ + assert "SA_URL" not in os.environ + + +class BuildOrgClientTestCase(_BuilderFixture): + """env.build_org_client goes through a config file, not the environment. + + An organization client is bound to no team, so a stray SA_TEAM_ID must not reach + it; and config_path makes the SDK ignore the environment altogether, so there is + nothing to perturb it and nothing to put back. + """ + + def test_it_is_built_with_no_team_whatever_the_environment_holds(self): + with patch.dict(os.environ, {"SA_TEAM_ID": "6085"}): + with self._respond_with(self.ORG_TOKEN): + client = env.build_org_client(self.API_KEY) + + assert client.controller.token_context.team_id is None + assert client.controller.config["SA_URL"] == "https://sa.test" + + def test_the_environment_is_untouched_while_it_is_built(self): + # The point of the config file: nothing is written to the environment, not even + # for the duration of the call. Observed from inside the auth request, which + # happens mid-construction - checking afterwards proves nothing, because a + # save-and-restore approach also leaves the environment as it found it. + seen = {} + + def record(*args, **kwargs): + seen["SA_TOKEN"] = os.environ.get("SA_TOKEN") + seen["SA_TEAM_ID"] = os.environ.get("SA_TEAM_ID") + return self._response(self.ORG_TOKEN) + + with patch.dict(os.environ, {"SA_TEAM_ID": "6085"}), patch( + "lib.infrastructure.services.auth.requests.post", side_effect=record + ): + client = env.build_org_client(self.API_KEY) + + assert client.controller.token_context.token == self.API_KEY + # The token never entered the environment, and the ambient team id was left + # exactly as it was rather than being cleared and put back. + assert seen["SA_TOKEN"] is None + assert seen["SA_TEAM_ID"] == "6085" + + +class MissingEnvVarsTestCase(TestCase): + """The gate reads the file directly, so it answers with the environment scrubbed.""" + + def setUp(self): + self.env_path = Path(tempfile.mkdtemp()) / ".env" + patcher = patch.dict(os.environ, {env.ENV_FILE_ENV: str(self.env_path)}) + patcher.start() + self.addCleanup(patcher.stop) + + def test_a_variable_only_the_file_provides_counts_as_present(self): + self.env_path.write_text(f"{env.SA_CONTRIBUTOR_TOKEN_ENV}={TOKEN}\n") + + assert env.missing_env_vars(env.SA_CONTRIBUTOR_TOKEN_ENV) == [] + assert env.var(env.SA_CONTRIBUTOR_TOKEN_ENV) == TOKEN + # ... and reading it did not put it in the environment. + assert env.SA_CONTRIBUTOR_TOKEN_ENV not in os.environ + + def test_a_variable_nobody_provides_is_reported_missing(self): + self.env_path.write_text("") + + assert env.missing_env_vars(env.SA_CONTRIBUTOR_TOKEN_ENV) == [ + env.SA_CONTRIBUTOR_TOKEN_ENV + ] diff --git a/tests/unit/test_http_client.py b/tests/unit/test_http_client.py index d6edb3337..2cbff7b6a 100644 --- a/tests/unit/test_http_client.py +++ b/tests/unit/test_http_client.py @@ -3,6 +3,8 @@ from unittest import TestCase from unittest.mock import patch +from src.superannotate.lib.core.entities.context import TokenContext +from src.superannotate.lib.core.entities.context import TokenScope from src.superannotate.lib.infrastructure.services.http_client import HttpClient @@ -11,10 +13,13 @@ def setUp(self): self.api_url = "https://api.example.com" self.team_id = 123 self.token = f"test_token={self.team_id}" + self.context = TokenContext( + token=self.token, team_id=self.team_id, scope=TokenScope.LEGACY + ) @patch.dict(os.environ, {"sa_version": "1.0.0", "SA_ENV": "test"}) def test_default_headers_with_env(self): - client = HttpClient(self.api_url, self.token, self.team_id) + client = HttpClient(self.api_url, self.context) headers = client.default_headers expected_user_agent = ( @@ -30,7 +35,12 @@ def test_default_headers_with_env(self): @patch.dict(os.environ, {"sa_version": "1.0.0"}) def test_default_headers_auth_type(self): client = HttpClient( - self.api_url, "sa_public_id_secret", self.team_id, auth_type="api_key" + self.api_url, + TokenContext( + token="sa_public_id_secret", + team_id=self.team_id, + scope=TokenScope.TEAM, + ), ) headers = client.default_headers @@ -40,7 +50,7 @@ def test_default_headers_auth_type(self): @patch.dict(os.environ, {"sa_version": "2.0.0"}, clear=True) def test_default_headers_without_env(self): - client = HttpClient(self.api_url, self.token, self.team_id) + client = HttpClient(self.api_url, self.context) headers = client.default_headers expected_user_agent = ( @@ -53,7 +63,7 @@ def test_default_headers_without_env(self): def test_default_headers_no_version(self): with patch.dict(os.environ, {}, clear=True): - client = HttpClient(self.api_url, self.token, self.team_id) + client = HttpClient(self.api_url, self.context) headers = client.default_headers expected_user_agent = ( @@ -61,3 +71,47 @@ def test_default_headers_no_version(self): f"OS: {platform.system()}; Team: {self.team_id}" ) assert headers["User-Agent"] == expected_user_agent + + +class TestTeamScoping(TestCase): + """The context is the only place a team is named: a client with one scopes every + request to it, a team-less (organization) client scopes to nothing.""" + + API_URL = "https://api.example.com" + + def _client(self, team_id, scope=TokenScope.TEAM): + return HttpClient( + self.API_URL, + TokenContext( + token="sa_public_id_secret", + team_id=team_id, + scope=scope, + ), + ) + + def test_a_team_context_is_sent_as_a_header_and_a_query_param(self): + client = self._client(123) + + assert client.team_id == 123 + assert client.default_query_params == {"team_id": 123} + # base64 of {"team_id": 123} + assert client.default_headers["x-sa-entity-context"] == ( + "eyJ0ZWFtX2lkIjogMTIzfQ==" + ) + assert "Team: 123" in client.default_headers["User-Agent"] + + def test_a_team_less_context_scopes_requests_to_no_team(self): + client = self._client(None, scope=TokenScope.ORGANIZATION) + + assert client.team_id is None + assert client.default_query_params == {} + assert "x-sa-entity-context" not in client.default_headers + assert "Team:" not in client.default_headers["User-Agent"] + + def test_default_query_params_cannot_be_mutated_through_a_request(self): + # request() copies them per call, so one request cannot leak params into the next. + client = self._client(123) + params = client.default_query_params + params["project_id"] = 7 + + assert client.default_query_params == {"team_id": 123} diff --git a/tests/unit/test_init.py b/tests/unit/test_init.py index c3f8b01b8..ff3d2b776 100644 --- a/tests/unit/test_init.py +++ b/tests/unit/test_init.py @@ -10,7 +10,12 @@ import superannotate.lib.core as constants from superannotate import AppException +from superannotate import SAAuthError from superannotate import SAClient +from superannotate import SAORGClient +from superannotate.lib.app.interface.base_interface import BaseInterfaceFacade +from superannotate.lib.core.entities import OrgTeamEntity +from superannotate.lib.core.entities.context import TokenScope class ClientInitTestCase(TestCase): @@ -21,7 +26,7 @@ def test_init_via_invalid_token(self): with self.assertRaisesRegex(AppException, r"Invalid token\."): SAClient(token=_token) - @patch("lib.infrastructure.controller.Controller.get_current_user") + @patch("lib.infrastructure.controller.TeamController.get_current_user") @patch("lib.core.usecases.GetTeamUseCase") def test_init_via_token(self, get_team_use_case, get_current_user): sa = SAClient(token=self._token) @@ -29,10 +34,10 @@ def test_init_via_token(self, get_team_use_case, get_current_user): self._token.split("=")[-1] ) assert get_current_user.call_count == 1 - assert sa.controller._config.API_TOKEN == self._token - assert sa.controller._config.API_URL == constants.BACKEND_URL + assert sa.controller.config["SA_TOKEN"] == self._token + assert sa.controller.config["SA_URL"] == constants.BACKEND_URL - @patch("lib.infrastructure.controller.Controller.get_current_user") + @patch("lib.infrastructure.controller.TeamController.get_current_user") @patch("lib.core.usecases.GetTeamUseCase") def test_init_via_config_json(self, get_team_use_case, get_current_user): with tempfile.TemporaryDirectory() as config_dir: @@ -46,8 +51,8 @@ def test_init_via_config_json(self, get_team_use_case, get_current_user): for kwargs in ({}, {"config_path": f"{config_dir}/config.json"}): sa = SAClient(**kwargs) - assert sa.controller._config.API_TOKEN == self._token - assert sa.controller._config.API_URL == constants.BACKEND_URL + assert sa.controller.config["SA_TOKEN"] == self._token + assert sa.controller.config["SA_URL"] == constants.BACKEND_URL assert get_team_use_case.call_args_list[0].kwargs["team_id"] == int( self._token.split("=")[-1] ) @@ -66,7 +71,7 @@ def test_init_via_config_json_invalid_json(self): with self.assertRaisesRegex(AppException, r"Invalid token\."): SAClient(**kwargs) - @patch("lib.infrastructure.controller.Controller.get_current_user") + @patch("lib.infrastructure.controller.TeamController.get_current_user") @patch("lib.core.usecases.GetTeamUseCase") def test_init_via_config_ini(self, get_team_use_case, get_current_user): with tempfile.TemporaryDirectory() as config_dir: @@ -85,15 +90,15 @@ def test_init_via_config_ini(self, get_team_use_case, get_current_user): config_parser.write(config_ini) for kwargs in ({}, {"config_path": f"{config_dir}/config.ini"}): sa = SAClient(**kwargs) - assert sa.controller._config.API_TOKEN == self._token - assert sa.controller._config.LOGGING_LEVEL == "DEBUG" - assert sa.controller._config.API_URL == constants.BACKEND_URL + assert sa.controller.config["SA_TOKEN"] == self._token + assert sa.controller.config["LOGGING_LEVEL"] == "DEBUG" + assert sa.controller.config["SA_URL"] == constants.BACKEND_URL assert get_team_use_case.call_args_list[0].kwargs["team_id"] == int( self._token.split("=")[-1] ) assert get_current_user.call_count == 2 - @patch("lib.infrastructure.controller.Controller.get_current_user") + @patch("lib.infrastructure.controller.TeamController.get_current_user") @patch("lib.core.usecases.GetTeamUseCase") def test_init_via_config_relative_filepath( self, get_team_use_case, get_current_user @@ -117,21 +122,21 @@ def test_init_via_config_relative_filepath( {"config_path": f"~/{Path(config_dir).name}/config.ini"}, ): sa = SAClient(**kwargs) - assert sa.controller._config.API_TOKEN == self._token - assert sa.controller._config.LOGGING_LEVEL == "DEBUG" - assert sa.controller._config.API_URL == constants.BACKEND_URL + assert sa.controller.config["SA_TOKEN"] == self._token + assert sa.controller.config["LOGGING_LEVEL"] == "DEBUG" + assert sa.controller.config["SA_URL"] == constants.BACKEND_URL assert get_team_use_case.call_args_list[0].kwargs["team_id"] == int( self._token.split("=")[-1] ) assert get_current_user.call_count == 2 - @patch("lib.infrastructure.controller.Controller.get_current_user") - @patch("lib.infrastructure.controller.Controller.get_team") + @patch("lib.infrastructure.controller.TeamController.get_current_user") + @patch("lib.infrastructure.controller.TeamController.get_team") @patch.dict(os.environ, {"SA_URL": "SOME_URL", "SA_TOKEN": "SOME_TOKEN=123"}) def test_init_env(self, get_team, get_current_user): sa = SAClient() - assert sa.controller._config.API_TOKEN == "SOME_TOKEN=123" - assert sa.controller._config.API_URL == "SOME_URL" + assert sa.controller.config["SA_TOKEN"] == "SOME_TOKEN=123" + assert sa.controller.config["SA_URL"] == "SOME_URL" assert get_team.call_count == 1 assert get_current_user.call_count == 1 @@ -222,7 +227,7 @@ def _mock_response(payload: dict, ok: bool = True, status_code: int = 200): } -@patch("lib.infrastructure.controller.Controller.get_team") +@patch("lib.infrastructure.controller.TeamController.get_team") @patch("lib.infrastructure.services.auth.requests.post") class ApiKeyInitTestCase(TestCase): _token = "sa_SOZVLlnbheUITTGb_PXlk2ON5QtqNPWY9bHZJctzlx4EPTkImzncQgRmybgh" @@ -247,10 +252,7 @@ def test_init_via_team_token(self, post, get_team): # The scope is kept: a team key acts as the team, so it is not allowed to # perform user-level operations (updating a team admin's permissions). context = sa.controller.token_context - assert context.scope_type == "team" - assert context.is_team_key - assert not context.is_personal_key - assert not context.is_legacy + assert context.scope == TokenScope.TEAM def test_token_context_request(self, post, get_team): post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) @@ -287,9 +289,7 @@ def test_init_via_team_user_token(self, post, get_team): # A personal key acts as its user, so it may do what that user may do. context = sa.controller.token_context - assert context.scope_type == "teamuser" - assert context.is_personal_key - assert not context.is_team_key + assert context.scope == TokenScope.TEAM_USER def test_nested_service_clients_share_team_context(self, post, get_team): post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) @@ -302,20 +302,43 @@ def test_nested_service_clients_share_team_context(self, post, get_team): assert service.client.team_id == 6085 assert service.client.auth_type == "api_key" - def test_organization_api_key_rejected(self, post, get_team): + def test_organization_api_key_without_team_id_rejected(self, post, get_team): + # An organization key carries no team, so it cannot resolve one on its own. post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) - with self.assertRaisesRegex( - AppException, r"does not accept an Organization API key" - ): + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): SAClient(token=self._token) + def test_organization_api_key_with_team_id(self, post, get_team): + # The team is not part of the key, so the caller names the team to operate in. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + sa = SAClient(token=self._token, team_id=6085) + + assert sa.controller.team_id == 6085 + context = sa.controller.token_context + assert context.scope == TokenScope.ORGANIZATION + # A team-less key has no user behind it either, so it falls back to its creator. + assert sa.controller.current_user.email == "vaghinak@superannotate.com" + + client = sa.controller.service_provider.client + assert client.team_id == 6085 + assert client.auth_type == "api_key" + + def test_team_id_matching_the_token_accepted(self, post, get_team): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + sa = SAClient(token=self._token, team_id=6085) + assert sa.controller.team_id == 6085 + + def test_team_id_mismatching_the_token_rejected(self, post, get_team): + # A team key names its own team; a conflicting team_id is a caller mistake. + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + with self.assertRaisesRegex(AppException, r"Invalid team id provided\."): + SAClient(token=self._token, team_id=42) + def test_unknown_scope_type_rejected(self, post, get_team): response = deepcopy(TEAM_TOKEN_RESPONSE) response["token"]["scope_type"] = "something-new" post.return_value = _mock_response(response) - with self.assertRaisesRegex( - AppException, r"does not accept an Organization API key" - ): + with self.assertRaisesRegex(AppException, r"Invalid team id provided\."): SAClient(token=self._token) def test_team_scope_without_team_id_rejected(self, post, get_team): @@ -323,20 +346,85 @@ def test_team_scope_without_team_id_rejected(self, post, get_team): response = deepcopy(TEAM_TOKEN_RESPONSE) response["token"]["scope"] = {} post.return_value = _mock_response(response) - with self.assertRaisesRegex( - AppException, r"does not accept an Organization API key" - ): + with self.assertRaisesRegex(AppException, r"Invalid team id provided\."): SAClient(token=self._token) def test_authentication_failure(self, post, get_team): post.return_value = _mock_response({}, ok=False, status_code=401) - with self.assertRaisesRegex(AppException, r"Unable to authenticate"): + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): SAClient(token=self._token) +@patch("lib.infrastructure.controller.TeamController.get_team") +@patch("lib.infrastructure.services.auth.requests.post") +class TeamIdFromConfigTestCase(TestCase): + """The team an organization key operates in may come from any config source.""" + + _token = "sa_SOZVLlnbheUITTGb_PXlk2ON5QtqNPWY9bHZJctzlx4EPTkImzncQgRmybgh" + + def setUp(self): + self._config_dir = tempfile.TemporaryDirectory() + config_dir = self._config_dir.name + self._ini_path = f"{config_dir}/config.ini" + self._json_path = f"{config_dir}/config.json" + patches = ( + patch("lib.core.CONFIG_INI_FILE_LOCATION", self._ini_path), + patch("lib.core.CONFIG_JSON_FILE_LOCATION", self._json_path), + ) + for p in patches: + p.start() + self.addCleanup(p.stop) + self.addCleanup(self._config_dir.cleanup) + + def _write_ini(self, **values): + config_parser = ConfigParser() + config_parser.optionxform = str + config_parser["DEFAULT"] = {k: str(v) for k, v in values.items()} + with open(self._ini_path, "w") as config_ini: + config_parser.write(config_ini) + + def test_team_id_from_config_ini(self, post, get_team): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + self._write_ini(SA_TOKEN=self._token, SA_TEAM_ID=6085) + # Both the default location and an explicit path read the same file. + for kwargs in ({}, {"config_path": self._ini_path}): + sa = SAClient(**kwargs) + assert sa.controller.config["SA_TEAM_ID"] == 6085 + assert sa.controller.team_id == 6085 + + def test_team_id_from_config_ini_by_field_name(self, post, get_team): + # The ini keys are read as-is, so the internal field name works as well. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + self._write_ini(SA_TOKEN=self._token, TEAM_ID=6085) + assert SAClient().controller.team_id == 6085 + + def test_org_token_in_config_ini_without_team_id_rejected(self, post, get_team): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + self._write_ini(SA_TOKEN=self._token) + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): + SAClient() + + def test_team_id_from_config_json(self, post, get_team): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + with open(self._json_path, "w") as config_json: + json.dump({"token": self._token, "team_id": 6085}, config_json) + for kwargs in ({}, {"config_path": self._json_path}): + assert SAClient(**kwargs).controller.team_id == 6085 + + def test_team_id_from_env(self, post, get_team): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + with patch.dict(os.environ, {"SA_TOKEN": self._token, "SA_TEAM_ID": "6085"}): + assert SAClient().controller.team_id == 6085 + + def test_explicit_team_id_overrides_config_ini(self, post, get_team): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + self._write_ini(SA_TOKEN=self._token, SA_TEAM_ID=6085) + assert SAClient(team_id=1).controller.team_id == 1 + + class LegacyTokenTestCase(TestCase): - @patch("lib.infrastructure.controller.Controller.get_current_user") - @patch("lib.infrastructure.controller.Controller.get_team") + @patch("lib.infrastructure.controller.TeamController.get_current_user") + @patch("lib.infrastructure.controller.TeamController.get_team") @patch("lib.infrastructure.services.auth.requests.post") def test_legacy_token_resolves_offline(self, post, get_team, get_current_user): sa = SAClient(token="token=123") @@ -344,15 +432,337 @@ def test_legacy_token_resolves_offline(self, post, get_team, get_current_user): assert post.call_count == 0 assert sa.controller.team_id == 123 assert sa.controller.service_provider.client.auth_type == "sdk" - # No scope is reported for a legacy token, and it is not a team key: it - # acts as the team owner, so it may update team admin permissions. + # The backend reports no scope for a legacy token; the SDK names it LEGACY. + # It acts as the team owner, so it may update team admin permissions. context = sa.controller.token_context - assert context.is_legacy - assert context.scope_type is None - assert not context.is_team_key + assert context.scope == TokenScope.LEGACY - @patch("lib.infrastructure.controller.Controller.get_current_user") - @patch("lib.infrastructure.controller.Controller.get_team") + @patch("lib.infrastructure.controller.TeamController.get_current_user") + @patch("lib.infrastructure.controller.TeamController.get_team") def test_legacy_token_team_id_mismatch_raises(self, get_team, get_current_user): - with self.assertRaisesRegex(AppException, r"does not match the team"): + with self.assertRaisesRegex(AppException, r"Invalid team id provided\."): SAClient(token="token=123", team_id=42) + + +@patch("lib.infrastructure.services.auth.requests.post") +class OrgClientInitTestCase(TestCase): + """SAORGClient authenticates outside any team, so it accepts only an org key.""" + + _token = "sa_SOZVLlnbheUITTGb_PXlk2ON5QtqNPWY9bHZJctzlx4EPTkImzncQgRmybgh" + + def test_init_via_organization_token(self, post): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + org = SAORGClient(token=self._token) + + context = org.controller.token_context + assert context.scope == TokenScope.ORGANIZATION + # No team to scope to, so no team header and no team query param either. + assert context.team_id is None + client = org.controller.service_provider.client + assert client.team_id is None + assert client.default_query_params == {} + assert "x-sa-entity-context" not in client.default_headers + # A team-less key has no user behind it, so it falls back to its creator. + assert org.controller.current_user.email == "vaghinak@superannotate.com" + + def test_organization_client_is_not_team_bound(self, post): + # SA_TEAM_ID in the environment does not make an org client team-scoped. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + with patch.dict(os.environ, {"SA_TEAM_ID": "6085"}): + org = SAORGClient(token=self._token) + + assert org.controller.token_context.team_id is None + # Nothing team-level is reachable through an organization client. + assert not hasattr(org.controller, "team_id") + assert not hasattr(org.controller, "projects") + # ... and telemetry gets no team name to report, rather than an error. + assert org.controller.team_name is None + + def test_team_token_rejected(self, post): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): + SAORGClient(token=self._token) + + def test_personal_token_rejected(self, post): + post.return_value = _mock_response(TEAM_USER_TOKEN_RESPONSE) + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): + SAORGClient(token=self._token) + + def test_legacy_token_rejected(self, post): + # A legacy token resolves offline and is bound to its own team. + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): + SAORGClient(token="token=123") + assert post.call_count == 0 + + def test_unknown_scope_rejected(self, post): + response = deepcopy(ORGANIZATION_TOKEN_RESPONSE) + response["token"]["scope_type"] = "something-new" + post.return_value = _mock_response(response) + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): + SAORGClient(token=self._token) + + def test_organization_token_with_no_creator_is_rejected(self, post): + # There is no team to look a user up in, so a key with no creator behind it + # has no user at all - and no team-scoped fallback to find one. + response = deepcopy(ORGANIZATION_TOKEN_RESPONSE) + response["token"]["created_by"] = None + post.return_value = _mock_response(response) + with self.assertRaisesRegex(AppException, r"Invalid credentials provided\."): + SAORGClient(token=self._token) + + @patch("lib.infrastructure.serviceprovider.ServiceProvider.list_teams") + def test_list_teams_returns_only_the_documented_fields(self, list_teams, post): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + list_teams.return_value = MagicMock( + ok=True, + data=[ + OrgTeamEntity( + id=6085, + name="Team A", + description="", + creator_id="a@b.com", + owner_id="org-1", + owner_type="organization", + ) + ], + ) + + teams = SAORGClient(token=self._token).list_teams() + + assert [team["id"] for team in teams] == [6085] + assert teams[0].keys() == { + "id", + "name", + "description", + "creator_id", + "owner_id", + "owner_type", + } + + def _org_config_ini(self, directory, **settings): + path = f"{directory}/config.ini" + parser = ConfigParser() + parser.optionxform = str + parser["DEFAULT"] = {"SA_TOKEN": self._token, **settings} + with open(path, "w") as handle: + parser.write(handle) + return path + + @patch("lib.infrastructure.controller.TeamController.get_team") + def test_get_team_client_keeps_the_configuration_it_was_built_with( + self, get_team, post + ): + # An organization client is configured like any other - a config file, or + # SA_URL. The team client it hands back has to run on that same configuration: + # rebuilding one from the token alone reverts to the defaults, which means + # production. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + with tempfile.TemporaryDirectory() as directory: + config_path = self._org_config_ini( + directory, SA_URL="https://sa.test", ANNOTATION_CHUNK_SIZE="100" + ) + org = SAORGClient(config_path=config_path) + assert org.controller.config["SA_URL"] == "https://sa.test" + + team_client = org.get_team_client(6085) + + assert team_client.team_id == 6085 + assert team_client.controller.config["SA_URL"] == "https://sa.test" + assert team_client.controller.config["ANNOTATION_CHUNK_SIZE"] == 100 + # ... and every request it makes goes to that backend, not just the config. + assert ( + team_client.controller.service_provider.client.api_url == "https://sa.test" + ) + + def test_config_round_trips_through_the_public_constructor(self, post): + # controller.config is exactly the dict SAClient(config=...) takes, so a client + # can be rebuilt from another one - which is all get_team_client does. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + with tempfile.TemporaryDirectory() as directory: + config_path = self._org_config_ini( + directory, SA_URL="https://sa.test", ITEM_CHUNK_SIZE="250" + ) + org = SAORGClient(config_path=config_path) + + config = org.controller.config + rebuilt = SAORGClient(config=config) + + assert config["SA_URL"] == "https://sa.test" + assert config["ITEM_CHUNK_SIZE"] == 250 + assert config["SA_TOKEN"] == self._token + # Every key it hands out is one the constructor accepts, and nothing is lost. + assert rebuilt.controller.config == config + + def test_config_is_a_copy(self, post): + # It is dumped per access, so a caller cannot reconfigure a live client through + # the dict it was handed. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + org = SAORGClient(token=self._token) + + org.controller.config["SA_URL"] = "https://mutated.test" + + assert org.controller.config["SA_URL"] == constants.BACKEND_URL + + @patch("lib.infrastructure.controller.TeamController.get_team") + def test_get_team_client_is_scoped_to_the_requested_team(self, get_team, post): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + org = SAORGClient(token=self._token) + + team_client = org.get_team_client(6085) + + assert isinstance(team_client, SAClient) + assert team_client.team_id == 6085 + # The same organization key and user, now scoping every request to one team. + context = team_client.controller.token_context + assert context.token == org.controller.token_context.token + assert context.scope == TokenScope.ORGANIZATION + assert context.user == org.controller.token_context.user + assert context.team_id == 6085 + client = team_client.controller.service_provider.client + assert client.default_query_params == {"team_id": 6085} + # The organization client it came from keeps its own, team-less session. + assert org.controller.token_context.team_id is None + + @patch("lib.infrastructure.controller.TeamController.get_team") + def test_get_team_client_resolves_the_token_a_second_time(self, get_team, post): + # The price of going through the public constructor: it is handed a config, so + # it resolves the token from scratch and asks the backend about it again. The + # alternative was constructing the client around __init__, which a user-facing + # class should not have to expose. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + org = SAORGClient(token=self._token) + resolved_once = post.call_count + + org.get_team_client(6085) + + assert post.call_count == resolved_once + 1 + + @patch("lib.infrastructure.serviceprovider.ServiceProvider.get_team") + def test_get_team_client_does_not_check_the_team(self, get_team, post): + # An organization key is authorised across the organization, so construction + # accepts whatever team it is given, even one the backend refuses. The failure + # is reported on the first team-scoped call instead - with GetTeamUseCase's + # one message, which is what the integration test asserts against a real + # backend. + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + get_team.return_value = _mock_response( + {"error": "team not found"}, ok=False, status_code=404 + ) + org = SAORGClient(token=self._token) + + client = org.get_team_client(999_999_999) + + assert client.team_id == 999_999_999 + with self.assertRaisesRegex(AppException, r"Unable to retrieve team data"): + client.get_team_metadata() + + +@patch("lib.infrastructure.services.auth.requests.post") +class AuthErrorTestCase(TestCase): + """Credential failures are SAAuthError; everything else stays an AppException. + + A caller can act on the difference - reauthenticate rather than retry - and + telemetry recognises an auth failure by type instead of by matching the message. + """ + + _token = "sa_SOZVLlnbheUITTGb_PXlk2ON5QtqNPWY9bHZJctzlx4EPTkImzncQgRmybgh" + + def test_a_malformed_token_is_an_auth_error(self, post): + with self.assertRaisesRegex(SAAuthError, r"Invalid token\."): + SAClient(token="nope") + # Rejected on shape alone, without asking the backend. + assert post.call_count == 0 + + def test_a_team_id_the_token_disagrees_with_is_an_auth_error(self, post): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + with self.assertRaisesRegex(SAAuthError, r"Invalid team id provided\."): + SAClient(token=self._token, team_id=42) + + def test_the_wrong_kind_of_key_is_an_auth_error(self, post): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + with self.assertRaisesRegex(SAAuthError, r"Invalid credentials provided\."): + SAORGClient(token=self._token) + + def test_missing_credentials_are_an_auth_error(self, post): + with patch.object( + BaseInterfaceFacade, "_retrieve_configs_from_env", return_value=None + ), patch("lib.core.CONFIG_INI_FILE_LOCATION", "/nonexistent.ini"), patch( + "lib.core.CONFIG_JSON_FILE_LOCATION", "/nonexistent.json" + ): + with self.assertRaisesRegex(SAAuthError, r"Credentials not found"): + SAClient() + + def test_an_auth_error_is_still_an_app_exception(self, post): + # Callers catching AppException today keep working. + with self.assertRaises(AppException): + SAClient(token="nope") + + def test_a_bad_argument_is_not_an_auth_error(self, post): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + with self.assertRaises(AppException) as caught: + SAClient(token=self._token, team_id="not-an-int") + assert not isinstance(caught.exception, SAAuthError) + + +@patch("lib.infrastructure.controller.TeamController.get_team") +@patch("lib.infrastructure.services.auth.requests.post") +class InlineConfigTestCase(TestCase): + """``config=`` configures a client on creation, instead of through a file.""" + + _token = "sa_SOZVLlnbheUITTGb_PXlk2ON5QtqNPWY9bHZJctzlx4EPTkImzncQgRmybgh" + + def test_config_alone_is_enough_to_build_a_client(self, post, get_team): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + + sa = SAClient( + config={ + "SA_TOKEN": self._token, + "SA_URL": "https://sa.test", + "MAX_THREAD_COUNT": 8, + } + ) + + assert sa.controller.config["SA_TOKEN"] == self._token + assert sa.controller.config["SA_URL"] == "https://sa.test" + assert sa.controller.config["MAX_THREAD_COUNT"] == 8 + + def test_config_applies_on_top_of_a_token_argument(self, post, get_team): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + + sa = SAClient(token=self._token, config={"SA_URL": "https://sa.test"}) + + assert sa.controller.config["SA_TOKEN"] == self._token + assert sa.controller.config["SA_URL"] == "https://sa.test" + + def test_config_applies_on_top_of_the_environment(self, post, get_team): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + with patch.dict( + os.environ, {"SA_TOKEN": self._token, "SA_URL": "https://from-env.test"} + ): + sa = SAClient(config={"SA_URL": "https://sa.test"}) + + assert sa.controller.config["SA_URL"] == "https://sa.test" + + def test_an_explicit_argument_wins_over_the_same_key_in_config( + self, post, get_team + ): + post.return_value = _mock_response(ORGANIZATION_TOKEN_RESPONSE) + + sa = SAClient(team_id=6085, config={"SA_TOKEN": self._token, "SA_TEAM_ID": 42}) + + assert sa.controller.team_id == 6085 + + def test_an_unrecognised_key_is_rejected(self, post, get_team): + # ConfigEntity ignores what it does not know, so an unchecked typo would be + # dropped in silence - and a mistyped SA_TOKEN would send the client off to + # authenticate as whatever the environment happens to hold. + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + + with self.assertRaisesRegex(AppException, r"Unknown configuration: SA_URLL"): + SAClient(token=self._token, config={"SA_URLL": "https://typo.test"}) + + def test_config_is_keyword_only(self, post, get_team): + post.return_value = _mock_response(TEAM_TOKEN_RESPONSE) + + with self.assertRaises(AppException): + SAClient(self._token, None, None, {"SA_URL": "https://sa.test"}) diff --git a/tests/unit/test_token_scope.py b/tests/unit/test_token_scope.py new file mode 100644 index 000000000..1b861aef7 --- /dev/null +++ b/tests/unit/test_token_scope.py @@ -0,0 +1,50 @@ +"""The scope a token was issued for: what it implies, and how it is read off a key.""" + +from unittest import TestCase + +from superannotate.lib.core.entities.context import API_KEY_AUTH_TYPE +from superannotate.lib.core.entities.context import SDK_AUTH_TYPE +from superannotate.lib.core.entities.context import TokenScope + + +class TokenScopeTestCase(TestCase): + def test_only_an_organization_key_carries_no_team(self): + assert TokenScope.TEAM.carries_team + assert TokenScope.TEAM_USER.carries_team + assert TokenScope.LEGACY.carries_team + assert not TokenScope.ORGANIZATION.carries_team + + def test_only_a_legacy_token_authenticates_as_sdk(self): + assert TokenScope.LEGACY.auth_type == SDK_AUTH_TYPE + for scope in (TokenScope.TEAM, TokenScope.TEAM_USER, TokenScope.ORGANIZATION): + assert scope.auth_type == API_KEY_AUTH_TYPE + + def test_every_scope_has_the_label_telemetry_reports(self): + assert TokenScope.TEAM.label == "Team API Key" + assert TokenScope.TEAM_USER.label == "Personal API Key" + assert TokenScope.ORGANIZATION.label == "Org API Key" + assert TokenScope.LEGACY.label == "SDK Token" + + def test_formats_as_the_value_it_stands_for(self): + # Log lines interpolate a scope directly; Enum's own __str__ would print + # "TokenScope.TEAM" instead of the value the backend uses. + assert f"{TokenScope.TEAM}" == "team" + assert TokenScope.TEAM == "team" + + +class OfApiKeyTestCase(TestCase): + """Reading the scope off what an API key reports.""" + + def test_reads_the_scopes_the_backend_reports(self): + assert TokenScope.of_api_key("team") is TokenScope.TEAM + assert TokenScope.of_api_key("teamuser") is TokenScope.TEAM_USER + assert TokenScope.of_api_key("organization") is TokenScope.ORGANIZATION + + def test_an_unknown_scope_is_not_resolved(self): + assert TokenScope.of_api_key("something-new") is None + assert TokenScope.of_api_key(None) is None + + def test_legacy_is_never_read_off_a_key(self): + # LEGACY is the SDK's own name for a token that resolves offline. An API key + # reporting it would otherwise authenticate as "sdk" and fail. + assert TokenScope.of_api_key("legacy") is None diff --git a/tests/unit/test_tracker.py b/tests/unit/test_tracker.py new file mode 100644 index 000000000..84973e699 --- /dev/null +++ b/tests/unit/test_tracker.py @@ -0,0 +1,224 @@ +"""The telemetry decorator every public SDK method is wrapped in. + +Exercised against throwaway Trackable classes rather than SAClient, so what is under +test is the decorator itself: which event it reports, what lands in the payload, and +what it must not do. Nothing here touches the network - Tracker._track is patched, and +it would refuse to send from a test run anyway. +""" + +import os +from unittest import TestCase +from unittest.mock import patch + +from superannotate.lib.app.interface.base_interface import TrackableMeta +from superannotate.lib.app.interface.base_interface import Tracker + + +class _Recorder: + """Collects what Tracker._track was handed, in place of sending it.""" + + def __init__(self): + self.events = [] + + def as_track(self): + """A stand-in for Tracker._track. + + A plain function, so it binds as a method - an instance with __call__ would + not, and the resulting TypeError would be swallowed along with everything + else _track_method guards against. + """ + events = self.events + + def _track( + tracker, user_id, event_name, data, *, client, explicit_credentials=False + ): + events.append({"user_id": user_id, "event": event_name, "data": data}) + + return _track + + @property + def events_named(self): + return [event["event"] for event in self.events] + + @property + def last(self): + return self.events[-1] + + +class _Subject(metaclass=TrackableMeta): + def __init__(self, controller=None): + self.controller = controller + + def work(self, project, item_name=None, count=None, options=None): + return "ok" + + def blows_up(self): + raise ValueError("nope") + + def interrupted(self): + raise KeyboardInterrupt("ctrl-c") + + def _private(self): # not wrapped: TrackableMeta skips underscored names + return "quiet" + + +@patch.dict(os.environ, {"sa_version": "4.6.0"}) +class TrackerTestCase(TestCase): + def setUp(self): + self.recorder = _Recorder() + patcher = patch.object(Tracker, "_track", self.recorder.as_track()) + patcher.start() + self.addCleanup(patcher.stop) + + def test_reports_the_method_name_and_the_calling_class(self): + _Subject().work(project="p") + + assert self.recorder.last["event"] == "work" + assert self.recorder.last["data"]["Class"] == "_Subject" + + def test_underscored_methods_are_not_tracked(self): + _Subject()._private() + + assert self.recorder.events_named == ["__init__"] + + def test_a_failing_call_is_reported_as_a_failure_and_still_raises(self): + with self.assertRaises(ValueError): + _Subject().blows_up() + + assert self.recorder.last["event"] == "blows_up" + assert self.recorder.last["data"]["Success"] is False + + def test_an_interrupted_call_is_not_reported_as_a_success(self): + # __call__ catches BaseException for this: with `except Exception` a + # KeyboardInterrupt skipped the failure flag and the call was logged as done. + with self.assertRaises(KeyboardInterrupt): + _Subject().interrupted() + + assert self.recorder.last["data"]["Success"] is False + + def test_the_original_traceback_is_not_reframed(self): + # A bare `raise` keeps this decorator out of the traceback the caller sees. + try: + _Subject().blows_up() + except ValueError as e: + frames = [] + tb = e.__traceback__ + while tb is not None: + frames.append(tb.tb_frame.f_code.co_name) + tb = tb.tb_next + + assert frames[-1] == "blows_up" + + +class SkipMetricsTestCase(TestCase): + def setUp(self): + self.recorder = _Recorder() + patcher = patch.object(Tracker, "_track", self.recorder.as_track()) + patcher.start() + self.addCleanup(patcher.stop) + + def test_nothing_is_gathered_once_metrics_are_turned_off(self): + # Read per call, not at import: a Tracker is built while its class is being + # created, so a value captured then would ignore anything set afterwards. + with patch.dict(os.environ, {"SA_SKIP_METRICS": "true"}): + _Subject().work(project="p") + + assert self.recorder.events == [] + + def test_metrics_resume_when_it_is_turned_back_on(self): + with patch.dict(os.environ, {"SA_SKIP_METRICS": "false"}): + _Subject().work(project="p") + + assert self.recorder.events_named == ["__init__", "work"] + + +class PayloadTestCase(TestCase): + """What reaches the payload, and what is held back.""" + + def test_credentials_are_reduced_to_whether_they_were_given(self): + _, properties = Tracker.default_parser( + "__init__", {"self": object(), "token": "sa_secret", "config_path": None} + ) + + assert properties["sa_token"] == "True" + assert "token" not in properties + assert properties["config_path"] == "False" + + def test_a_project_path_is_split_into_project_and_folder(self): + _, properties = Tracker.default_parser("work", {"project": "Proj/batch1"}) + + assert properties["project_name"] == "Proj" + assert properties["folder_name"] == "batch1" + + def test_containers_are_reduced_to_their_shape(self): + _, properties = Tracker.default_parser( + "work", + { + "options": {"a": 1, "b": 2}, + "count": 3, + "names": ["x", "y"], + "flag": True, + }, + ) + + # A dict contributes its keys, a sized value its length - not the contents. + assert properties["options"] == ["a", "b"] + assert properties["names"] == 2 + assert properties["count"] == 3 + assert properties["flag"] is True + + def test_defaults_are_filled_in_for_arguments_the_caller_omitted(self): + arguments = Tracker.extract_arguments( + _Subject.work.__wrapped__, _Subject(), project="p" + ) + + assert arguments["item_name"] is None + assert arguments["count"] is None + + +class DefaultPayloadTestCase(TestCase): + """The per-event envelope, which used to be cached.""" + + def _payload(self): + return Tracker.get_default_payload("Team A", "a@b.com", "Team API Key") + + def test_carries_the_identity_of_the_client(self): + with patch.dict(os.environ, {"sa_version": "4.6.0"}): + payload = self._payload() + + assert payload["Team"] == "Team A" + assert payload["User Email"] == "a@b.com" + assert payload["Auth Type"] == "Team API Key" + assert payload["Version"] == "4.6.0" + assert payload["SDK"] is True + + def test_each_event_gets_its_own_dict(self): + # It was lru_cached, so every event shared one dict: a caller mutating the + # payload it was handed corrupted every later event. + with patch.dict(os.environ, {"sa_version": "4.6.0"}): + first = self._payload() + first["Team"] = "MUTATED" + + assert self._payload()["Team"] == "Team A" + + def test_the_environment_is_read_per_event(self): + # Caching also froze SA_ENV and sa_version as they were on the first call. + with patch.dict(os.environ, {"sa_version": "4.6.0"}): + assert self._payload()["Env"] == "N/A" + with patch.dict(os.environ, {"SA_ENV": "staging"}): + assert self._payload()["Env"] == "staging" + + +class TrackableMetaTestCase(TestCase): + def test_a_subclass_need_not_define_its_own_init(self): + # attrs["__init__"] used to be read unconditionally, so a Trackable subclass + # without one raised KeyError while the class was being created. + class Inheriting(_Subject): + def extra(self): + return "ok" + + assert Inheriting().extra() == "ok" + + def test_the_wrapped_method_keeps_its_identity(self): + assert _Subject.work.__name__ == "work" + assert _Subject().work.__name__ == "work"