diff --git a/homeassistant/components/onedrive/backup.py b/homeassistant/components/onedrive/backup.py index 3e8b413b9167..a5ff70b09ba0 100644 --- a/homeassistant/components/onedrive/backup.py +++ b/homeassistant/components/onedrive/backup.py @@ -26,6 +26,7 @@ from homeassistant.components.backup import ( from homeassistant.core import HomeAssistant, callback from homeassistant.helpers.aiohttp_client import async_get_clientsession from homeassistant.helpers.json import json_dumps +from homeassistant.util.async_ import gather_with_limited_concurrency from homeassistant.util.json import json_loads_object from .const import CONF_DELETE_PERMANENTLY, DATA_BACKUP_AGENT_LISTENERS, DOMAIN @@ -36,6 +37,7 @@ MAX_CHUNK_SIZE = 60 * 1024 * 1024 # largest chunk possible, must be <= 60 MiB TARGET_CHUNKS = 20 TIMEOUT = ClientTimeout(connect=10, total=43200) # 12 hours CACHE_TTL = 300 +METADATA_DOWNLOAD_CONCURRENCY = 10 async def async_get_backup_agents( @@ -276,7 +278,7 @@ class OneDriveBackupAgent(BackupAgent): item.name for item in items if item.name and item.name.endswith(".tar") } - metadata_files: dict[str, AgentBackup] = {} + metadata_item_ids: list[str] = [] for item in items: if item.name and item.name.endswith(".metadata.json"): # Check if corresponding backup file exists @@ -288,10 +290,15 @@ class OneDriveBackupAgent(BackupAgent): item.name, ) continue - if metadata := await _download_metadata(item.id): - metadata_files[metadata.backup_id] = metadata + metadata_item_ids.append(item.id) - self._cache_backup_metadata = metadata_files + metadata_contents = await gather_with_limited_concurrency( + METADATA_DOWNLOAD_CONCURRENCY, + *(_download_metadata(item_id) for item_id in metadata_item_ids), + ) + self._cache_backup_metadata = { + metadata.backup_id: metadata for metadata in metadata_contents if metadata + } self._cache_expiration = time() + CACHE_TTL return self._cache_backup_metadata diff --git a/homeassistant/components/onedrive_for_business/backup.py b/homeassistant/components/onedrive_for_business/backup.py index f6c7d8cd1656..e43c2949a9a0 100644 --- a/homeassistant/components/onedrive_for_business/backup.py +++ b/homeassistant/components/onedrive_for_business/backup.py @@ -26,6 +26,7 @@ from homeassistant.components.backup import ( from homeassistant.core import HomeAssistant, callback from homeassistant.helpers.aiohttp_client import async_get_clientsession from homeassistant.helpers.json import json_dumps +from homeassistant.util.async_ import gather_with_limited_concurrency from homeassistant.util.json import json_loads_object from . import OneDriveConfigEntry @@ -36,6 +37,7 @@ MAX_CHUNK_SIZE = 60 * 1024 * 1024 # largest chunk possible, must be <= 60 MiB TARGET_CHUNKS = 20 TIMEOUT = ClientTimeout(connect=10, total=43200) # 12 hours CACHE_TTL = 300 +METADATA_DOWNLOAD_CONCURRENCY = 10 async def async_get_backup_agents( @@ -264,7 +266,7 @@ class OneDriveBackupAgent(BackupAgent): item.name for item in items if item.name and item.name.endswith(".tar") } - metadata_files: dict[str, AgentBackup] = {} + metadata_item_ids: list[str] = [] for item in items: if item.name and item.name.endswith(".metadata.json"): # Check if corresponding backup file exists @@ -276,10 +278,15 @@ class OneDriveBackupAgent(BackupAgent): item.name, ) continue - if metadata := await _download_metadata(item.id): - metadata_files[metadata.backup_id] = metadata + metadata_item_ids.append(item.id) - self._cache_backup_metadata = metadata_files + metadata_contents = await gather_with_limited_concurrency( + METADATA_DOWNLOAD_CONCURRENCY, + *(_download_metadata(item_id) for item_id in metadata_item_ids), + ) + self._cache_backup_metadata = { + metadata.backup_id: metadata for metadata in metadata_contents if metadata + } self._cache_expiration = time() + CACHE_TTL return self._cache_backup_metadata diff --git a/tests/components/onedrive/test_backup.py b/tests/components/onedrive/test_backup.py index 1382fec938f4..581a9c4d8981 100644 --- a/tests/components/onedrive/test_backup.py +++ b/tests/components/onedrive/test_backup.py @@ -1,7 +1,10 @@ """Test the backups for OneDrive.""" +import asyncio from collections.abc import AsyncGenerator +from dataclasses import replace from io import StringIO +from json import dumps from unittest.mock import Mock, patch from onedrive_personal_sdk.exceptions import ( @@ -14,6 +17,7 @@ import pytest from homeassistant.components.backup import DOMAIN as BACKUP_DOMAIN, AgentBackup from homeassistant.components.onedrive.backup import ( + METADATA_DOWNLOAD_CONCURRENCY, async_register_backup_agents_listener, ) from homeassistant.components.onedrive.const import DATA_BACKUP_AGENT_LISTENERS, DOMAIN @@ -118,6 +122,49 @@ async def test_agents_list_backups_with_download_failure( assert response["result"]["backups"] == [] +async def test_agents_list_backups_downloads_metadata_concurrently( + hass: HomeAssistant, + hass_ws_client: WebSocketGenerator, + mock_onedrive_client: MagicMock, + mock_backup_file: File, + mock_metadata_file: File, +) -> None: + """Test metadata files are downloaded concurrently, up to the limit.""" + backup_ids = [f"backup_{i}" for i in range(METADATA_DOWNLOAD_CONCURRENCY + 2)] + mock_onedrive_client.list_drive_items.return_value = [ + file + for backup_id in backup_ids + for file in ( + replace(mock_backup_file, id=f"{backup_id}_tar", name=f"{backup_id}.tar"), + replace( + mock_metadata_file, id=backup_id, name=f"{backup_id}.metadata.json" + ), + ) + ] + in_flight = max_in_flight = 0 + + async def download_drive_item(item_id: str) -> Mock: + nonlocal in_flight, max_in_flight + in_flight += 1 + max_in_flight = max(max_in_flight, in_flight) + await asyncio.sleep(0) + in_flight -= 1 + metadata = dumps({**BACKUP_METADATA, "backup_id": item_id}) + return Mock(read=AsyncMock(return_value=metadata.encode())) + + mock_onedrive_client.download_drive_item.side_effect = download_drive_item + client = await hass_ws_client(hass) + await client.send_json_auto_id({"type": "backup/info"}) + response = await client.receive_json() + + assert response["success"] + assert response["result"]["agent_errors"] == {} + assert {backup["backup_id"] for backup in response["result"]["backups"]} == set( + backup_ids + ) + assert max_in_flight == METADATA_DOWNLOAD_CONCURRENCY + + async def test_agents_get_backup( hass: HomeAssistant, hass_ws_client: WebSocketGenerator, diff --git a/tests/components/onedrive_for_business/test_backup.py b/tests/components/onedrive_for_business/test_backup.py index a81f286f1992..72c6a5142f45 100644 --- a/tests/components/onedrive_for_business/test_backup.py +++ b/tests/components/onedrive_for_business/test_backup.py @@ -1,7 +1,10 @@ """Test the backups for OneDrive for Business.""" +import asyncio from collections.abc import AsyncGenerator +from dataclasses import replace from io import StringIO +from json import dumps from unittest.mock import Mock, patch from onedrive_personal_sdk.exceptions import ( @@ -14,6 +17,7 @@ import pytest from homeassistant.components.backup import DOMAIN as BACKUP_DOMAIN, AgentBackup from homeassistant.components.onedrive_for_business.backup import ( + METADATA_DOWNLOAD_CONCURRENCY, async_register_backup_agents_listener, ) from homeassistant.components.onedrive_for_business.const import ( @@ -124,6 +128,49 @@ async def test_agents_list_backups_with_download_failure( assert response["result"]["backups"] == [] +async def test_agents_list_backups_downloads_metadata_concurrently( + hass: HomeAssistant, + hass_ws_client: WebSocketGenerator, + mock_onedrive_client: MagicMock, + mock_backup_file: File, + mock_metadata_file: File, +) -> None: + """Test metadata files are downloaded concurrently, up to the limit.""" + backup_ids = [f"backup_{i}" for i in range(METADATA_DOWNLOAD_CONCURRENCY + 2)] + mock_onedrive_client.list_drive_items.return_value = [ + file + for backup_id in backup_ids + for file in ( + replace(mock_backup_file, id=f"{backup_id}_tar", name=f"{backup_id}.tar"), + replace( + mock_metadata_file, id=backup_id, name=f"{backup_id}.metadata.json" + ), + ) + ] + in_flight = max_in_flight = 0 + + async def download_drive_item(item_id: str) -> Mock: + nonlocal in_flight, max_in_flight + in_flight += 1 + max_in_flight = max(max_in_flight, in_flight) + await asyncio.sleep(0) + in_flight -= 1 + metadata = dumps({**BACKUP_METADATA, "backup_id": item_id}) + return Mock(read=AsyncMock(return_value=metadata.encode())) + + mock_onedrive_client.download_drive_item.side_effect = download_drive_item + client = await hass_ws_client(hass) + await client.send_json_auto_id({"type": "backup/info"}) + response = await client.receive_json() + + assert response["success"] + assert response["result"]["agent_errors"] == {} + assert {backup["backup_id"] for backup in response["result"]["backups"]} == set( + backup_ids + ) + assert max_in_flight == METADATA_DOWNLOAD_CONCURRENCY + + async def test_agents_get_backup( hass: HomeAssistant, hass_ws_client: WebSocketGenerator,