diff --git a/armis_sdk/clients/assets_client.py b/armis_sdk/clients/assets_client.py new file mode 100644 index 0000000..fda06e4 --- /dev/null +++ b/armis_sdk/clients/assets_client.py @@ -0,0 +1,342 @@ +import datetime +from typing import AsyncIterator +from typing import Literal +from typing import Optional +from typing import Type +from typing import Union + +import universalasync + +from armis_sdk.core import response_utils +from armis_sdk.core.armis_error import ArmisError +from armis_sdk.core.armis_error import BulkUpdateError +from armis_sdk.core.armis_error import BulkUpdateItemError +from armis_sdk.core.base_entity_client import BaseEntityClient +from armis_sdk.entities.asset import Asset +from armis_sdk.entities.asset import AssetT +from armis_sdk.entities.device import Device + +AssetIdSource = Literal[ + "ASSET_ID", + "IPV4_ADDRESS", + "IPV6_ADDRESS", + "MAC_ADDRESS", + "SERIAL_NUMBER", +] + + +@universalasync.wrap +class AssetsClient(BaseEntityClient): # pylint: disable=too-few-public-methods + # pylint: disable=line-too-long + """ + A client for interacting with assets. + + The primary entities for this client inherit from [Asset][armis_sdk.entities.asset.Asset]: + + 1. [Device][armis_sdk.entities.device.Device] + """ + + async def list_by_asset_id( + self, + asset_class: Type[AssetT], + asset_ids: Union[list[int], list[str]], + asset_id_source: AssetIdSource = "ASSET_ID", + fields: Optional[list[str]] = None, + ) -> AsyncIterator[AssetT]: + """List assets by asset ID or other identifiers. + + Args: + asset_class: The asset class to list. Must inherit from [Asset][armis_sdk.entities.asset.Asset]. + asset_ids: A list of asset identifiers (int or str depending on asset_id_source). + asset_id_source: The type of identifier provided in asset_ids. + fields: Optional list of fields to retrieve. If None, all non-custom fields are retrieved. + + Yields: + Assets of the specified class matching the provided identifiers. + + Example: + ```python linenums="1" hl_lines="13 17" + import asyncio + + from armis_sdk.clients.assets_client import AssetsClient + from armis_sdk.entities.device import Device + + async def main(): + assets_client = AssetsClient() + + device_ids = [1, 2, 3] + ipv4_addresses = ["1.1.1.1", "2.2.2.2", "3.3.3.3"] + + # List by the default source "ASSET_ID" + async for device in assets_client.list_by_asset_id(Device, device_ids): + print(device) + + # List by explicit source "IPV4_ADDRESS" + async for device in assets_client.list_by_asset_id(Device, ipv4_addresses, asset_id_source="IPV4_ADDRESS"): + print(device) + + asyncio.run(main()) + ``` + """ + filter_ = { + "filter_criteria": "ASSET_ID", + "asset_ids": asset_ids, + "asset_id_source": asset_id_source, + } + async for item in self._list_assets(asset_class, fields, filter_): + yield item + + async def list_by_last_seen( + self, + asset_class: Type[AssetT], + last_seen: Union[datetime.datetime, datetime.timedelta], + fields: Optional[list[str]] = None, + ) -> AsyncIterator[AssetT]: + """List assets by last seen timestamp. + + Args: + asset_class: The asset class to list. Must inherit from [Asset][armis_sdk.entities.asset.Asset]. + last_seen: Either a datetime (assets seen on or after this time) or timedelta (assets seen within this duration). + fields: Optional list of fields to retrieve. If None, all non-custom fields are retrieved. + + Yields: + Assets of the specified class matching the last seen criteria. + + Raises: + ArmisError: If last_seen is neither datetime nor timedelta. + + Example: + ```python linenums="1" hl_lines="11 15" + import asyncio + import datetime + + from armis_sdk.clients.assets_client import AssetsClient + from armis_sdk.entities.device import Device + + async def main(): + assets_client = AssetsClient() + + # List devices seen in the last 24 hours + async for device in assets_client.list_by_last_seen(Device, datetime.timedelta(days=1)): + print(device) + + # List devices seen on or after December 8, 2025 + async for device in assets_client.list_by_last_seen(Device, datetime.datetime(2025, 12, 8)): + print(device) + + asyncio.run(main()) + ``` + """ + filter_: dict[str, Union[str, int]] = {"filter_criteria": "LAST_SEEN"} + + if isinstance(last_seen, datetime.datetime): + filter_["last_seen_ge"] = last_seen.isoformat() + elif isinstance(last_seen, datetime.timedelta): + filter_["last_seen_seconds"] = int(last_seen.total_seconds()) + else: + raise ArmisError(f"Invalid 'last_seen' type {type(last_seen)}") + + async for item in self._list_assets(asset_class, fields, filter_): + yield item + + async def update( + self, + assets: list[AssetT], + fields: list[str], + asset_id_source: AssetIdSource = "ASSET_ID", + ) -> None: + # pylint: disable=line-too-long + """Bulk update assets. + + Args: + assets: A list of assets. Items must inherit from [Asset][armis_sdk.entities.asset.Asset]. + fields: A list of fields to update. Currently only custom properties are supported (i.e. `custom.MyField`). + asset_id_source: From where on the asset to take the unique identifier. + + Raises: + BulkUpdateError: If an error occurs while trying to update any of the assets. + + Example: + ```python linenums="1" hl_lines="13 16" + import asyncio + + from armis_sdk.clients.assets_client import AssetsClient + from armis_sdk.entities.device import Device + + + async def main(): + assets_client = AssetsClient() + + device = Device(device_id=1, ipv4_addresses=["1.2.3.4"], custom={"MyField": "Hello, World"}) + + # Update based on the default source "ASSET_ID" + await assets_client.update([device], ["custom.MyField"]) + + # Update based on the explicit source "IPV4_ADDRESS" + await assets_client.update([device], ["custom.MyField"], asset_id_source="IPV4_ADDRESS") + + asyncio.run(main()) + ``` + """ + if not assets or not fields: + return + + self._validate_asset_class(assets) + + asset_class = type(assets[0]) + self._validate_fields(asset_class, fields, allow_model_members=False) + + items = [] + for index, asset in enumerate(assets): + asset_id = self._get_asset_id(asset, index, asset_id_source) + for field in fields: + items.append(self._create_bulk_update_request(asset, asset_id, field)) + + if not items: + return + + payload = { + "items": items, + "asset_type": asset_class.asset_type, + "asset_id_source": asset_id_source, + } + async with self._armis_client.client() as client: + response = await client.post("/v3/assets/_bulk", json=payload) + data = response_utils.get_data_dict(response) + errors = [ + BulkUpdateItemError(index=index, request=items[index], response=item) + for index, item in enumerate(data["items"]) + if item["status"] != 202 + ] + if errors: + raise BulkUpdateError(errors) + + @classmethod + def _create_bulk_update_request( + cls, + asset: Asset, + asset_id: Union[str, int], + field: str, + ): + request = {"asset_id": asset_id, "key": field} + if cls._is_custom_field(field): + key = field.split(".", 1)[1] + if value := asset.custom.get(key): + request["operation"] = "SET" + request["value"] = value + else: + request["operation"] = "UNSET" + else: + raise ArmisError(f"Updating the field {field!r} is currently not supported") + + return request + + @classmethod + def _get_asset_id( + cls, + asset: Asset, + index: int, + asset_id_source: AssetIdSource, + ) -> Union[str, int]: + if isinstance(asset, Device): + return cls._get_device_asset_id(asset, index, asset_id_source) + + raise ArmisError(f"Can't get {asset_id_source} of asset {asset!r}") + + @classmethod + def _get_device_asset_id( + cls, + device: Device, + index: int, + asset_id_source: AssetIdSource, + ): + if asset_id_source == "ASSET_ID": + if device.device_id is None: + raise ArmisError(f"Device at index {index} doesn't have a device id") + return device.device_id + + if asset_id_source == "MAC_ADDRESS": + if device.mac_addresses is None or len(device.mac_addresses) != 1: + raise ArmisError( + f"Device at index {index} doesn't have exactly one mac address" + ) + return device.mac_addresses[0] + + if asset_id_source == "IPV4_ADDRESS": + if device.ipv4_addresses is None or len(device.ipv4_addresses) != 1: + raise ArmisError( + f"Device at index {index} doesn't have exactly one IPv4 address" + ) + return device.ipv4_addresses[0] + + if asset_id_source == "IPV6_ADDRESS": + if device.ipv6_addresses is None or len(device.ipv6_addresses) != 1: + raise ArmisError( + f"Device at index {index} doesn't have exactly one IPv6 address" + ) + return device.ipv6_addresses[0] + + if asset_id_source == "SERIAL_NUMBER": + if device.serial_numbers is None or len(device.serial_numbers) != 1: + raise ArmisError( + f"Device at index {index} doesn't have exactly one serial number" + ) + return device.serial_numbers[0] + + raise ArmisError(f"Can't get {asset_id_source!r} of device at index {index}") + + @classmethod + def _is_custom_field(cls, field: str) -> bool: + return field.startswith("custom.") + + async def _list_assets( + self, + asset_class: Type[AssetT], + fields: Optional[list[str]], + filter_: dict, + ) -> AsyncIterator[AssetT]: + fields = fields or sorted(asset_class.all_fields()) + + self._validate_fields(asset_class, fields) + + body = { + "asset_type": asset_class.asset_type, + "fields": fields, + "filter": filter_, + } + async for item in self._armis_client.list("/v3/assets/_search", body=body): + yield asset_class.from_search_result(item) + + @classmethod + def _validate_asset_class(cls, assets: list[AssetT]): + asset_types = {type(asset) for asset in assets} + if len(asset_types) > 1: + asset_types_str = ", ".join(sorted(repr(at.__name__) for at in asset_types)) + raise ArmisError( + "All assets must be of the same type, " + f"got {len(asset_types)} types: {asset_types_str}" + ) + + @classmethod + def _validate_fields( + cls, + asset_class: Type[AssetT], + fields: list[str], + allow_model_members=True, + ): + invalid_fields = [] + all_fields = asset_class.all_fields() + for field in fields: + if cls._is_custom_field(field): + continue + + if allow_model_members and field in all_fields: + continue + + invalid_fields.append(field) + + if invalid_fields: + fields_str = ", ".join(map(repr, invalid_fields)) + raise ArmisError( + f"The following fields are not supported with this operation: {fields_str}" + ) diff --git a/armis_sdk/core/armis_client.py b/armis_sdk/core/armis_client.py index 862a37f..1e909c2 100644 --- a/armis_sdk/core/armis_client.py +++ b/armis_sdk/core/armis_client.py @@ -82,11 +82,12 @@ def client(self, retries: Optional[int] = None, backoff: Optional[float] = None) trust_env=True, ) - async def list(self, url: str) -> AsyncIterator[dict]: + async def list(self, url: str, body: Optional[dict] = None) -> AsyncIterator[dict]: """List all items from a paginated endpoint. Args: url (str): The relative endpoint URL. + body (dict): Payload to send as POST request. Returns: An (async) iterator of `dict`s. @@ -113,9 +114,12 @@ async def main(): """ page_size = int(os.getenv(ARMIS_PAGE_SIZE, str(DEFAULT_PAGE_LENGTH))) async with self.client() as client: - params = {"limit": page_size} + params = {"limit": page_size, **(body or {})} while True: - response = await client.get(url, params=params) + if body: + response = await client.post(url, json=params) + else: + response = await client.get(url, params=params) data = response_utils.get_data_dict(response) items = data["items"] for item in items: diff --git a/armis_sdk/core/armis_error.py b/armis_sdk/core/armis_error.py index 530c95b..6c12b9d 100644 --- a/armis_sdk/core/armis_error.py +++ b/armis_sdk/core/armis_error.py @@ -3,6 +3,7 @@ while interacting with the SDK. """ +import json from typing import List from typing import Optional from typing import Union @@ -16,6 +17,13 @@ class DetailItem(BaseModel): msg: str type: str + def __str__(self): + return ( + f"Type: {self.type}\n" + f"Message: {self.msg}\n" + f"Location: {json.dumps(self.loc)}" + ) + class ErrorBody(BaseModel): detail: Union[str, List[DetailItem]] @@ -27,6 +35,24 @@ class ArmisError(Exception): """ +class BulkUpdateItemError(BaseModel): + index: int + request: dict + response: dict + + +class BulkUpdateError(ArmisError): + def __init__(self, items: list[BulkUpdateItemError]): + self.items = items + display = "\n".join( + f"Failed to update item at index {item.index}. " + f"Request: {json.dumps(item.request)}, " + f"Response: {json.dumps(item.response)}" + for item in items + ) + super().__init__(display) + + class ResponseError(ArmisError): # pylint: disable=line-too-long """ @@ -48,7 +74,7 @@ def _get_message(cls, error_body: ErrorBody) -> str: if isinstance(error_body.detail, str): return error_body.detail - return "\n".join(item.msg for item in error_body.detail) + return "\n\n".join(str(item) for item in error_body.detail) class AlreadyExistsError(ResponseError): diff --git a/armis_sdk/core/armis_sdk.py b/armis_sdk/core/armis_sdk.py index ce78abc..7212aaa 100644 --- a/armis_sdk/core/armis_sdk.py +++ b/armis_sdk/core/armis_sdk.py @@ -1,5 +1,6 @@ from typing import Optional +from armis_sdk.clients.assets_client import AssetsClient from armis_sdk.clients.data_export_client import DataExportClient from armis_sdk.clients.device_custom_properties_client import ( DeviceCustomPropertiesClient, @@ -17,6 +18,7 @@ class ArmisSdk: # pylint: disable=too-few-public-methods Attributes: client (ArmisClient): An instance of [ArmisClient][armis_sdk.core.armis_client.ArmisClient] + assets (AssetsClient): An instance of [AssetsClient][armis_sdk.clients.assets_client.AssetsClient] data_export (DataExportClient): An instance of [DataExportClient][armis_sdk.clients.data_export_client.DataExportClient] device_custom_properties (DeviceCustomPropertiesClient): An instance of [DeviceCustomPropertiesClient][armis_sdk.clients.device_custom_properties_client.DeviceCustomPropertiesClient] sites (SitesClient): An instance of [SitesClient][armis_sdk.clients.sites_client.SitesClient] @@ -39,6 +41,7 @@ async def main(): def __init__(self, credentials: Optional[ClientCredentials] = None): self.client: ArmisClient = ArmisClient(credentials=credentials) + self.assets: AssetsClient = AssetsClient(self.client) self.data_export: DataExportClient = DataExportClient(self.client) self.device_custom_properties: DeviceCustomPropertiesClient = ( DeviceCustomPropertiesClient(self.client) diff --git a/armis_sdk/entities/asset.py b/armis_sdk/entities/asset.py new file mode 100644 index 0000000..fa6873d --- /dev/null +++ b/armis_sdk/entities/asset.py @@ -0,0 +1,41 @@ +import collections +from typing import Any +from typing import ClassVar +from typing import DefaultDict +from typing import Literal +from typing import Type +from typing import TypeVar + +from pydantic import Field + +from armis_sdk.core.base_entity import BaseEntity + +AssetT = TypeVar("AssetT", bound="Asset") + + +class Asset(BaseEntity): + """ + A base class for all assets type to inherit from. + """ + + asset_type: ClassVar[Literal["DEVICE"]] + custom: dict[str, Any] = Field(default_factory=dict) + """Custom properties of the asset. Values can by anything.""" + + @classmethod + def from_search_result(cls: Type[AssetT], data: dict) -> AssetT: + fields: DefaultDict[str, Any] = collections.defaultdict(dict) + for key, value in data["fields"].items(): + if len(parts := key.split(".", 1)) > 1: + part1, part2 = parts + fields[part1][part2] = value + else: + fields[key] = value + + return cls(**fields) + + @classmethod + def all_fields(cls) -> set[str]: + # Pylint doesn't recognize that "cls.model_fields" is a dict and not a method + # so it's complaining that the method doesn't have a "keys" attribute. + return set(cls.model_fields.keys()) - {"custom"} # pylint: disable=no-member diff --git a/armis_sdk/entities/boundary.py b/armis_sdk/entities/boundary.py new file mode 100644 index 0000000..52c6818 --- /dev/null +++ b/armis_sdk/entities/boundary.py @@ -0,0 +1,11 @@ +from armis_sdk.core.base_entity import BaseEntity + + +class Boundary(BaseEntity): + """A `Boundary` is a logical segment in the network.""" + + id: int + """The id of the boundary.""" + + name: str + """The name of the boundary.""" diff --git a/armis_sdk/entities/data_export/risk_factor.py b/armis_sdk/entities/data_export/risk_factor.py index 06c4953..b3e07b9 100644 --- a/armis_sdk/entities/data_export/risk_factor.py +++ b/armis_sdk/entities/data_export/risk_factor.py @@ -65,7 +65,7 @@ class RiskFactor(BaseExportedEntity): """ The description of the risk factor - **Example**: `Device is accepting SMBv1 requests.` + **Example**: `Device Supports SMBv1` """ score: int diff --git a/armis_sdk/entities/device.py b/armis_sdk/entities/device.py new file mode 100644 index 0000000..9fba2ad --- /dev/null +++ b/armis_sdk/entities/device.py @@ -0,0 +1,118 @@ +import datetime +from typing import ClassVar +from typing import Literal +from typing import Optional + +from pydantic import Field + +from armis_sdk.entities.asset import Asset +from armis_sdk.entities.boundary import Boundary +from armis_sdk.entities.network_interface import NetworkInterface +from armis_sdk.entities.site import Site + + +class Device(Asset): + # pylint: disable=line-too-long + asset_type: ClassVar[Literal["DEVICE"]] = "DEVICE" + + boundaries: Optional[list[Boundary]] = None + """The list of boundaries the device belongs to.""" + + brand: Optional[str] = None + """ + The device brand. + + Example: `Apple` + """ + + category: Optional[str] = None + """ + The device category. + + Example: `Handheld` + """ + + device_id: Optional[int] = None + """The unique identifier given to the device by thr Armis engine.""" + + display: Optional[str] = None + """ + The display text of the device. + + Example: `My iPhone` + """ + + first_seen: Optional[datetime.datetime] = Field(strict=False, default=None) + """When was the device first seen.""" + + ipv4_addresses: Optional[list[str]] = None + """The list of IPv4 addresses of the device""" + + ipv6_addresses: Optional[list[str]] = None + """The list of IPv6 addresses of the device""" + + last_seen: Optional[datetime.datetime] = Field(strict=False, default=None) + """When was the device last seen.""" + + mac_addresses: Optional[list[str]] = None + """The list of MAC addresses of the device""" + + model: Optional[str] = None + """ + The model of the device. + + Example: `iPhone 17` + """ + + names: Optional[list[str]] = None + """ + List of names of the device + + Example: `["My iPhone 17", "Jane's iPhone"]` + """ + + network_interfaces: Optional[list[NetworkInterface]] = None + """List of network interfaces detected on the device.""" + + os_name: Optional[str] = None + """ + The OS name running on the device. + + Example: `iOS` + """ + + os_version: Optional[str] = None + """ + The OS version running on the device. + + Example: `17` + """ + + purdue_level: Optional[float] = None + """ + The purdue level of the devices. See [Wikipedia](https://en.wikipedia.org/wiki/Purdue_Enterprise_Reference_Architecture) article for more details. + + Example: `4` + """ + + risk_level: Optional[int] = Field(ge=0, le=1000, default=None) + """The risk level given to the device by the Armis engine, between `0` and `100`.""" + + serial_numbers: Optional[list[str]] = None + """The list of serial numbers of the device""" + + site: Optional[Site] = None + """The site in which this device was last seen.""" + + tags: Optional[list[str]] = None + """The tags given to the devices.""" + + type: Optional[str] = None + """ + The type of the device. + + Example: `Mobile Phones` + """ + + visibility: Optional[Literal["Full", "Limited"]] = None + """Whether the device is fully visibly or limited.""" diff --git a/armis_sdk/entities/network_interface.py b/armis_sdk/entities/network_interface.py new file mode 100644 index 0000000..0f3086b --- /dev/null +++ b/armis_sdk/entities/network_interface.py @@ -0,0 +1,49 @@ +from typing import Optional + +from armis_sdk.core.base_entity import BaseEntity + + +class NetworkInterface(BaseEntity): + # pylint: disable=line-too-long + """ + A `NetworkInterface` represents a physical network card of a [Device][armis_sdk.entities.device.Device]. + """ + + alias: Optional[str] + """The alias of the interface.""" + + brand: Optional[str] + """The brand of the interface.""" + + broadcast_ssid: Optional[str] + """The last SSID broadcasted by the interface.""" + + channels: list[int] + """The channels that the interface uses to transmit.""" + + description: Optional[str] + """The description of the interface""" + + hidden_broadcast_ssid: Optional[bool] + """Is the broadcasted SSID hidden.""" + + ipv4_address: Optional[str] + """The last IPv4 address associated with the interface.""" + + ipv6_address: Optional[str] + """The last IPv6 address associated with the interface.""" + + last_connected_ssid: Optional[str] + """The SSID the interface last connected to.""" + + mac_address: Optional[str] + """The MAC address of the interface.""" + + name: Optional[str] + """The name of the interface.""" + + type: Optional[str] + """The type of the interface.""" + + vlan: Optional[int] + """The VLAN of the interface.""" diff --git a/docs/clients/AssetsClient.md b/docs/clients/AssetsClient.md new file mode 100644 index 0000000..ce0c12d --- /dev/null +++ b/docs/clients/AssetsClient.md @@ -0,0 +1 @@ +::: armis_sdk.clients.assets_client.AssetsClient diff --git a/docs/entities/Asset.md b/docs/entities/Asset.md new file mode 100644 index 0000000..169e7d0 --- /dev/null +++ b/docs/entities/Asset.md @@ -0,0 +1 @@ +::: armis_sdk.entities.asset.Asset diff --git a/docs/entities/Boundary.md b/docs/entities/Boundary.md new file mode 100644 index 0000000..8b7a4ff --- /dev/null +++ b/docs/entities/Boundary.md @@ -0,0 +1 @@ +::: armis_sdk.entities.boundary.Boundary diff --git a/docs/entities/Device.md b/docs/entities/Device.md new file mode 100644 index 0000000..28bc576 --- /dev/null +++ b/docs/entities/Device.md @@ -0,0 +1 @@ +::: armis_sdk.entities.device.Device diff --git a/docs/entities/NetworkInterface.md b/docs/entities/NetworkInterface.md new file mode 100644 index 0000000..df1f4a2 --- /dev/null +++ b/docs/entities/NetworkInterface.md @@ -0,0 +1 @@ +::: armis_sdk.entities.network_interface.NetworkInterface diff --git a/mkdocs.yml b/mkdocs.yml index e2133a1..66fa2d7 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -5,15 +5,20 @@ repo_url: https://github.com/ArmisSecurity/armis-sdk-python nav: - Getting started: index.md - Entities: + - Asset: entities/Asset.md - AsqRule: entities/AsqRule.md - DataExport: - Application: entities/data_export/Application.md - DataExport: entities/data_export/DataExport.md - RiskFactor: entities/data_export/RiskFactor.md - Vulnerability: entities/data_export/Vulnerability.md + - Boundary: entities/Boundary.md + - Device: entities/Device.md - DeviceCustomProperty: entities/DeviceCustomProperty.md + - NetworkInterface: entities/NetworkInterface.md - Site: entities/Site.md - Clients: + - AssetsClient: clients/AssetsClient.md - DataExportClient: clients/DataExportClient.md - DeviceCustomPropertiesClient: clients/DeviceCustomPropertiesClient.md - SitesClient: clients/SitesClient.md diff --git a/tests/armis_sdk/clients/assets_client_test.py b/tests/armis_sdk/clients/assets_client_test.py new file mode 100644 index 0000000..51a5bf7 --- /dev/null +++ b/tests/armis_sdk/clients/assets_client_test.py @@ -0,0 +1,405 @@ +import datetime + +import pytest +import pytest_httpx + +from armis_sdk.clients.assets_client import AssetsClient +from armis_sdk.core.armis_error import ArmisError +from armis_sdk.core.armis_error import BulkUpdateError +from armis_sdk.entities.asset import Asset +from armis_sdk.entities.device import Device +from tests.armis_sdk.clients import assets_test_data + +pytest_plugins = ["tests.plugins.auto_setup_plugin"] + + +class NotDevice(Asset): + pass + + +async def test_list_by_last_seen_datetime(httpx_mock: pytest_httpx.HTTPXMock): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_search", + method="POST", + match_json={ + "limit": 100, + "asset_type": "DEVICE", + "fields": assets_test_data.ALL_DEVICE_FIELDS, + "filter": { + "filter_criteria": "LAST_SEEN", + "last_seen_ge": "2025-12-03T00:00:00", + }, + }, + json={ + "items": [ + {"asset_id": 1, "fields": assets_test_data.MOCK_DEVICE_FULL_RAW_DATA} + ] + }, + ) + + assets_client = AssetsClient() + last_seen = datetime.datetime(2025, 12, 3) + devices = [ + device async for device in assets_client.list_by_last_seen(Device, last_seen) + ] + + assert devices == [assets_test_data.MOCK_DEVICE_FULL] + + +async def test_list_by_last_seen_datetime_explicit_fields( + httpx_mock: pytest_httpx.HTTPXMock, +): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_search", + method="POST", + match_json={ + "limit": 100, + "asset_type": "DEVICE", + "fields": ["brand", "custom.MyField1", "custom.MyField2", "purdue_level"], + "filter": { + "filter_criteria": "LAST_SEEN", + "last_seen_ge": "2025-12-03T00:00:00", + }, + }, + json={ + "items": [ + {"asset_id": 1, "fields": assets_test_data.MOCK_DEVICE_PARTIAL_RAW_DATA} + ] + }, + ) + + assets_client = AssetsClient() + last_seen = datetime.datetime(2025, 12, 3) + fields = ["brand", "custom.MyField1", "custom.MyField2", "purdue_level"] + devices = [ + device + async for device in assets_client.list_by_last_seen( + Device, last_seen, fields=fields + ) + ] + + assert devices == [assets_test_data.MOCK_DEVICE_PARTIAL] + + +async def test_list_by_last_seen_timedelta(httpx_mock: pytest_httpx.HTTPXMock): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_search", + method="POST", + match_json={ + "limit": 100, + "asset_type": "DEVICE", + "fields": assets_test_data.ALL_DEVICE_FIELDS, + "filter": {"filter_criteria": "LAST_SEEN", "last_seen_seconds": 3600}, + }, + json={ + "items": [ + {"asset_id": 1, "fields": assets_test_data.MOCK_DEVICE_FULL_RAW_DATA} + ] + }, + ) + + assets_client = AssetsClient() + last_seen = datetime.timedelta(hours=1) + devices = [ + device async for device in assets_client.list_by_last_seen(Device, last_seen) + ] + + assert devices == [assets_test_data.MOCK_DEVICE_FULL] + + +async def test_list_by_last_seen_timedelta_explicit_fields( + httpx_mock: pytest_httpx.HTTPXMock, +): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_search", + method="POST", + match_json={ + "limit": 100, + "asset_type": "DEVICE", + "fields": ["brand", "custom.MyField1", "custom.MyField2", "purdue_level"], + "filter": {"filter_criteria": "LAST_SEEN", "last_seen_seconds": 3600}, + }, + json={ + "items": [ + {"asset_id": 1, "fields": assets_test_data.MOCK_DEVICE_PARTIAL_RAW_DATA} + ] + }, + ) + + assets_client = AssetsClient() + last_seen = datetime.timedelta(hours=1) + fields = ["brand", "custom.MyField1", "custom.MyField2", "purdue_level"] + devices = [ + device + async for device in assets_client.list_by_last_seen( + Device, last_seen, fields=fields + ) + ] + + assert devices == [assets_test_data.MOCK_DEVICE_PARTIAL] + + +async def test_list_by_last_seen_invalid_fields(): + assets_client = AssetsClient() + last_seen = datetime.timedelta(hours=1) + fields = ["device_id", "foo", "bar", "tags"] + + with pytest.raises( + ArmisError, + match="The following fields are not supported with this operation: 'foo', 'bar'", + ): + async for _ in assets_client.list_by_last_seen( + Device, last_seen, fields=fields + ): + pass + + +async def test_list_by_asset_id(httpx_mock: pytest_httpx.HTTPXMock): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_search", + method="POST", + match_json={ + "limit": 100, + "asset_type": "DEVICE", + "fields": assets_test_data.ALL_DEVICE_FIELDS, + "filter": { + "filter_criteria": "ASSET_ID", + "asset_id_source": "IPV4_ADDRESS", + "asset_ids": ["1.1.1.1"], + }, + }, + json={ + "items": [ + {"asset_id": 1, "fields": assets_test_data.MOCK_DEVICE_FULL_RAW_DATA} + ] + }, + ) + + assets_client = AssetsClient() + asset_ids = ["1.1.1.1"] + devices = [ + device + async for device in assets_client.list_by_asset_id( + Device, + asset_ids, + asset_id_source="IPV4_ADDRESS", + ) + ] + + assert devices == [assets_test_data.MOCK_DEVICE_FULL] + + +async def test_list_by_asset_id_explicit_fields(httpx_mock: pytest_httpx.HTTPXMock): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_search", + method="POST", + match_json={ + "limit": 100, + "asset_type": "DEVICE", + "fields": ["brand", "custom.MyField1", "custom.MyField2", "purdue_level"], + "filter": { + "filter_criteria": "ASSET_ID", + "asset_id_source": "IPV4_ADDRESS", + "asset_ids": ["1.1.1.1"], + }, + }, + json={ + "items": [ + {"asset_id": 1, "fields": assets_test_data.MOCK_DEVICE_PARTIAL_RAW_DATA} + ] + }, + ) + + assets_client = AssetsClient() + asset_ids = ["1.1.1.1"] + devices = [ + device + async for device in assets_client.list_by_asset_id( + Device, + asset_ids, + asset_id_source="IPV4_ADDRESS", + fields=["brand", "custom.MyField1", "custom.MyField2", "purdue_level"], + ) + ] + + assert devices == [assets_test_data.MOCK_DEVICE_PARTIAL] + + +async def test_list_by_asset_id_invalid_fields(): + assets_client = AssetsClient() + fields = ["device_id", "foo", "bar", "tags"] + + with pytest.raises( + ArmisError, + match="The following fields are not supported with this operation: 'foo', 'bar'", + ): + async for _ in assets_client.list_by_asset_id(Device, [1, 2, 3], fields=fields): + pass + + +async def test_update(httpx_mock: pytest_httpx.HTTPXMock): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_bulk", + method="POST", + match_json={ + "items": [ + { + "asset_id": 1, + "key": "custom.MyField1", + "operation": "SET", + "value": "value1", + }, + { + "asset_id": 1, + "key": "custom.MyField2", + "operation": "SET", + "value": 2, + }, + { + "asset_id": 2, + "key": "custom.MyField1", + "operation": "SET", + "value": "value3", + }, + {"asset_id": 2, "key": "custom.MyField2", "operation": "UNSET"}, + ], + "asset_type": "DEVICE", + "asset_id_source": "ASSET_ID", + }, + json={"items": [{"status": 202}] * 4}, + ) + + assets_client = AssetsClient() + + assets = [ + Device(device_id=1, custom={"MyField1": "value1", "MyField2": 2}), + Device(device_id=2, custom={"MyField1": "value3"}), + ] + fields = ["custom.MyField1", "custom.MyField2"] + await assets_client.update(assets, fields) + + +async def test_update_with_asset_id_source(httpx_mock: pytest_httpx.HTTPXMock): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_bulk", + method="POST", + match_json={ + "items": [ + { + "asset_id": "1.1.1.1", + "key": "custom.MyField1", + "operation": "SET", + "value": "value1", + }, + { + "asset_id": "1.1.1.1", + "key": "custom.MyField2", + "operation": "SET", + "value": 2, + }, + { + "asset_id": "2.2.2.2", + "key": "custom.MyField1", + "operation": "SET", + "value": "value3", + }, + {"asset_id": "2.2.2.2", "key": "custom.MyField2", "operation": "UNSET"}, + ], + "asset_type": "DEVICE", + "asset_id_source": "IPV4_ADDRESS", + }, + json={"items": [{"status": 202}] * 4}, + ) + + assets_client = AssetsClient() + + assets = [ + Device( + ipv4_addresses=["1.1.1.1"], custom={"MyField1": "value1", "MyField2": 2} + ), + Device(ipv4_addresses=["2.2.2.2"], custom={"MyField1": "value3"}), + ] + fields = ["custom.MyField1", "custom.MyField2"] + await assets_client.update(assets, fields, asset_id_source="IPV4_ADDRESS") + + +async def test_update_with_failed_requests(httpx_mock: pytest_httpx.HTTPXMock): + httpx_mock.add_response( + url="https://api.armis.com/v3/assets/_bulk", + method="POST", + match_json={ + "items": [ + { + "asset_id": 1, + "key": "custom.MyField1", + "operation": "SET", + "value": "value1", + }, + { + "asset_id": 1, + "key": "custom.MyField2", + "operation": "SET", + "value": 2, + }, + { + "asset_id": 2, + "key": "custom.MyField1", + "operation": "SET", + "value": "value3", + }, + {"asset_id": 2, "key": "custom.MyField2", "operation": "UNSET"}, + ], + "asset_type": "DEVICE", + "asset_id_source": "ASSET_ID", + }, + json={ + "items": [ + {"status": 202}, + {"status": 202}, + {"status": 400, "error": "Bad Request"}, + {"status": 202}, + ] + }, + ) + + assets_client = AssetsClient() + + assets = [ + Device(device_id=1, custom={"MyField1": "value1", "MyField2": 2}), + Device(device_id=2, custom={"MyField1": "value3"}), + ] + fields = ["custom.MyField1", "custom.MyField2"] + + with pytest.raises( + BulkUpdateError, + match=( + "Failed to update item at index 2. " + 'Request: {"asset_id": 2, "key": "custom.MyField1", ' + '"operation": "SET", "value": "value3"}, ' + 'Response: {"status": 400, "error": "Bad Request"}' + ), + ): + await assets_client.update(assets, fields) + + +@pytest.mark.parametrize( + ["assets", "fields", "expected_error"], + [ + ( + [Device(), NotDevice()], + ["custom.MyField"], + "All assets must be of the same type, got 2 types: 'Device', 'NotDevice'", + ), + ( + [Device()], + ["custom.MyField", "purdue_level"], + "The following fields are not supported with this operation: 'purdue_level'", + ), + ([Device()], ["custom.MyField"], "Device at index 0 doesn't have a device id"), + ], +) +async def test_update_with_validation_errors(assets, fields, expected_error): + assets_client = AssetsClient() + + with pytest.raises(ArmisError, match=expected_error): + await assets_client.update(assets, fields) diff --git a/tests/armis_sdk/clients/assets_test_data.py b/tests/armis_sdk/clients/assets_test_data.py new file mode 100644 index 0000000..77ad540 --- /dev/null +++ b/tests/armis_sdk/clients/assets_test_data.py @@ -0,0 +1,154 @@ +import datetime + +from armis_sdk.entities.boundary import Boundary +from armis_sdk.entities.device import Device +from armis_sdk.entities.network_interface import NetworkInterface +from armis_sdk.entities.site import Site + +ALL_DEVICE_FIELDS = [ + "boundaries", + "brand", + "category", + "device_id", + "display", + "first_seen", + "ipv4_addresses", + "ipv6_addresses", + "last_seen", + "mac_addresses", + "model", + "names", + "network_interfaces", + "os_name", + "os_version", + "purdue_level", + "risk_level", + "serial_numbers", + "site", + "tags", + "type", + "visibility", +] +MOCK_DEVICE_PARTIAL = Device( + custom={"MyField1": "foo", "MyField2": "bar"}, + brand="VMware", + purdue_level=4.0, +) +MOCK_DEVICE_PARTIAL_RAW_DATA = { + "brand": "VMware", + "custom.MyField1": "foo", + "custom.MyField2": "bar", + "purdue_level": 4.0, +} +MOCK_DEVICE_FULL = Device( + custom={}, + boundaries=[Boundary(id=1, name="Corporate")], + brand="VMware", + category="Computers", + device_id=1, + display="VMware", + first_seen=datetime.datetime(2025, 5, 14, 8, 34, 10), + ipv4_addresses=["10.246.212.12"], + ipv6_addresses=["fe80::4d68:8d3e:d3a5:c930"], + last_seen=datetime.datetime(2025, 12, 3, 13, 52, 45), + mac_addresses=["43:87:a2:05:bc:56"], + model="VMware", + names=["VMware", "044D9BA4_B6E"], + network_interfaces=[ + NetworkInterface( + alias=None, + brand=None, + broadcast_ssid=None, + channels=[], + description=None, + hidden_broadcast_ssid=None, + ipv4_address=None, + ipv6_address=None, + last_connected_ssid=None, + mac_address="43:87:a2:05:bc:56", + name=None, + type=None, + vlan=1, + ) + ], + os_name="Windows", + os_version="Server 2016", + purdue_level=4.0, + risk_level=80, + serial_numbers=None, + site=Site( + id=36481, + name="Geneva Enterprise", + lat=None, + lng=None, + location="Geneva", + parent_id=None, + tier=None, + asq_rule=None, + network_equipment_device_ids=None, + integration_ids=None, + children=[], + ), + tags=[ + "Insecure Traffic and Behavior", + "Critical Vulnerabilities", + "Deprecated SW/HW", + "Misconfigurations", + "Unprotected Sensitive Data", + "Insecure Credentials and Access Control", + ], + type="Virtual Machines", + visibility="Full", +) +MOCK_DEVICE_FULL_RAW_DATA = { + "boundaries": [{"id": 1, "name": "Corporate"}], + "brand": "VMware", + "category": "Computers", + "device_id": 1, + "display": "VMware", + "first_seen": "2025-05-14T08:34:10", + "ipv4_addresses": ["10.246.212.12"], + "ipv6_addresses": ["fe80::4d68:8d3e:d3a5:c930"], + "last_seen": "2025-12-03T13:52:45", + "mac_addresses": ["43:87:a2:05:bc:56"], + "model": "VMware", + "names": ["VMware", "044D9BA4_B6E"], + "network_interfaces": [ + { + "alias": None, + "brand": None, + "broadcast_ssid": None, + "channels": [], + "description": None, + "hidden_broadcast_ssid": None, + "ipv4_address": None, + "ipv6_address": None, + "last_connected_ssid": None, + "mac_address": "43:87:a2:05:bc:56", + "name": None, + "type": None, + "vlan": 1, + } + ], + "os_name": "Windows", + "os_version": "Server 2016", + "purdue_level": 4.0, + "risk_level": 80, + "serial_numbers": None, + "site": { + "id": 36481, + "location": "Geneva", + "name": "Geneva Enterprise", + "tier": None, + }, + "tags": [ + "Insecure Traffic and Behavior", + "Critical Vulnerabilities", + "Deprecated SW/HW", + "Misconfigurations", + "Unprotected Sensitive Data", + "Insecure Credentials and Access Control", + ], + "type": "Virtual Machines", + "visibility": "Full", +}