Просмотр исходного кода

Merge pull request #335 from CloudVE/add-typing

Add comprehensive typing to cloudbridge + mypy tox check
Nuwan Goonasekera 1 месяц назад
Родитель
Сommit
94f6f26db1
40 измененных файлов с 3854 добавлено и 2465 удалено
  1. 3 0
      .github/workflows/integration.yaml
  2. 8 5
      cloudbridge/__init__.py
  3. 33 11
      cloudbridge/base/helpers.py
  4. 7 3
      cloudbridge/base/middleware.py
  5. 36 25
      cloudbridge/base/provider.py
  6. 239 147
      cloudbridge/base/resources.py
  7. 132 75
      cloudbridge/base/services.py
  8. 105 58
      cloudbridge/base/subservices.py
  9. 14 10
      cloudbridge/factory.py
  10. 4 4
      cloudbridge/interfaces/exceptions.py
  11. 34 18
      cloudbridge/interfaces/provider.py
  12. 145 121
      cloudbridge/interfaces/resources.py
  13. 200 103
      cloudbridge/interfaces/services.py
  14. 57 34
      cloudbridge/interfaces/subservices.py
  15. 43 21
      cloudbridge/providers/aws/helpers.py
  16. 20 14
      cloudbridge/providers/aws/provider.py
  17. 226 185
      cloudbridge/providers/aws/resources.py
  18. 294 189
      cloudbridge/providers/aws/services.py
  19. 13 6
      cloudbridge/providers/aws/subservices.py
  20. 168 135
      cloudbridge/providers/azure/azure_client.py
  21. 7 6
      cloudbridge/providers/azure/helpers.py
  22. 23 13
      cloudbridge/providers/azure/provider.py
  23. 263 200
      cloudbridge/providers/azure/resources.py
  24. 369 223
      cloudbridge/providers/azure/services.py
  25. 13 6
      cloudbridge/providers/azure/subservices.py
  26. 35 19
      cloudbridge/providers/gcp/helpers.py
  27. 52 41
      cloudbridge/providers/gcp/provider.py
  28. 300 180
      cloudbridge/providers/gcp/resources.py
  29. 324 200
      cloudbridge/providers/gcp/services.py
  30. 12 6
      cloudbridge/providers/gcp/subservices.py
  31. 5 3
      cloudbridge/providers/mock/provider.py
  32. 12 5
      cloudbridge/providers/openstack/helpers.py
  33. 82 60
      cloudbridge/providers/openstack/provider.py
  34. 227 156
      cloudbridge/providers/openstack/resources.py
  35. 279 167
      cloudbridge/providers/openstack/services.py
  36. 14 6
      cloudbridge/providers/openstack/subservices.py
  37. 0 0
      cloudbridge/py.typed
  38. 40 0
      pyproject.toml
  39. 2 8
      tests/test_compute_service.py
  40. 14 2
      tox.ini

+ 3 - 0
.github/workflows/integration.yaml

@@ -60,6 +60,9 @@ jobs:
       - name: Run tox
         run: tox -e lint
 
+      - name: Run mypy
+        run: tox -e mypy
+
   mock:
     name: Mock-provider tests
     runs-on: ubuntu-latest

+ 8 - 5
cloudbridge/__init__.py

@@ -1,11 +1,12 @@
 """Library setup."""
 import logging
+from typing import Any
 
 # Current version of the library
 __version__ = '4.1.0'
 
 
-def get_version():
+def get_version() -> str:
     """
     Return a string with the current version of the library.
 
@@ -15,7 +16,7 @@ def get_version():
     return __version__
 
 
-def init_logging():
+def init_logging() -> None:
     """
     Initialize logging for testing.
 
@@ -35,7 +36,7 @@ class CBLogger(logging.Logger):
     Add a ``trace`` log level, numeric value 5: ``log.trace("Log message")``
     """
 
-    def trace(self, msg, *args, **kwargs):
+    def trace(self, msg: object, *args: object, **kwargs: Any) -> None:
         """Add ``trace`` log level."""
         self.log(TRACE, msg, *args, **kwargs)
 
@@ -61,7 +62,8 @@ log.addHandler(logging.NullHandler())
 #   cloudbridge.set_file_logger(__name__, '/tmp/log')
 
 
-def set_stream_logger(name, level=TRACE, format_string=None):
+def set_stream_logger(name: str, level: int = TRACE,
+                      format_string: str | None = None) -> None:
     """A convenience method to set the global logger to stream."""
     global log
     if not format_string:
@@ -76,7 +78,8 @@ def set_stream_logger(name, level=TRACE, format_string=None):
     log = logger
 
 
-def set_file_logger(name, filepath, level=logging.INFO, format_string=None):
+def set_file_logger(name: str, filepath: str, level: int = logging.INFO,
+                    format_string: str | None = None) -> None:
     """A convenience method to set the global logger to a file."""
     global log
     if not format_string:

+ 33 - 11
cloudbridge/base/helpers.py

@@ -3,7 +3,13 @@ import functools
 import logging
 import os
 import re
+from collections.abc import Callable
+from collections.abc import Iterator
 from contextlib import contextmanager
+from typing import Any
+from typing import TypeVar
+from typing import cast
+from typing import overload
 
 from cryptography.hazmat.backends import default_backend
 from cryptography.hazmat.primitives import serialization as crypt_serialization
@@ -17,8 +23,11 @@ from ..interfaces.exceptions import InvalidParamException
 
 log = logging.getLogger(__name__)
 
+T = TypeVar("T")
+F = TypeVar("F", bound=Callable[..., Any])
 
-def generate_key_pair():
+
+def generate_key_pair() -> tuple[str, str]:
     """
     This method generates a keypair and returns it as a tuple
     of (public, private) keys.
@@ -38,7 +47,8 @@ def generate_key_pair():
     return public_key, private_key
 
 
-def filter_by(prop_name, kwargs, objs):
+def filter_by(prop_name: str, kwargs: dict[str, Any],
+              objs: list[T]) -> list[T]:
     """
     Utility method for filtering a list of objects by a property.
     If the given property has a non empty value in kwargs, then
@@ -60,7 +70,8 @@ def filter_by(prop_name, kwargs, objs):
         return objs
 
 
-def generic_find(filter_names, kwargs, objs):
+def generic_find(filter_names: list[str], kwargs: dict[str, Any],
+                 objs: list[T]) -> list[T]:
     """
     Utility method for filtering a list of objects by a list of filters.
     """
@@ -78,7 +89,7 @@ def generic_find(filter_names, kwargs, objs):
 
 
 @contextmanager
-def cleanup_action(cleanup_func):
+def cleanup_action(cleanup_func: Callable[[], object]) -> Iterator[None]:
     """
     Context manager to carry out a given
     cleanup action after carrying out a set
@@ -109,7 +120,17 @@ def cleanup_action(cleanup_func):
         log.exception("Error during exception cleanup: ")
 
 
-def get_env(varname, default_value=None):
+@overload
+def get_env(varname: str) -> str | None:
+    ...
+
+
+@overload
+def get_env(varname: str, default_value: T) -> str | T:
+    ...
+
+
+def get_env(varname: str, default_value: object = None) -> object:
     """
     Return the value of the environment variable or default_value.
 
@@ -128,17 +149,18 @@ def get_env(varname, default_value=None):
 # Alias deprecation decorator, following:
 # https://stackoverflow.com/questions/49802412/
 # how-to-implement-deprecation-in-python-with-argument-alias
-def deprecated_alias(**aliases):
-    def deco(f):
+def deprecated_alias(**aliases: str) -> Callable[[F], F]:
+    def deco(f: F) -> F:
         @functools.wraps(f)
-        def wrapper(*args, **kwargs):
+        def wrapper(*args: Any, **kwargs: Any) -> Any:
             rename_kwargs(f.__name__, kwargs, aliases)
             return f(*args, **kwargs)
-        return wrapper
+        return cast(F, wrapper)
     return deco
 
 
-def rename_kwargs(func_name, kwargs, aliases):
+def rename_kwargs(func_name: str, kwargs: dict[str, Any],
+                  aliases: dict[str, str]) -> None:
     for alias, new in aliases.items():
         if alias in kwargs:
             if new in kwargs:
@@ -157,7 +179,7 @@ def rename_kwargs(func_name, kwargs, aliases):
 NON_ALPHA_NUM = re.compile(r"[^A-Za-z0-9]+")
 
 
-def to_resource_name(value, replace_with="-"):
+def to_resource_name(value: str, replace_with: str = "-") -> str:
     """
     Converts a given string to a valid resource name by stripping
     all characters that are not alphanumeric.

+ 7 - 3
cloudbridge/base/middleware.py

@@ -1,4 +1,5 @@
 import logging
+from typing import Any
 
 from pyeventsystem.middleware import dispatch as pyevent_dispatch
 from pyeventsystem.middleware import intercept
@@ -19,12 +20,14 @@ class EventDebugLoggingMiddleware(object):
     access keys.
     """
     @observe(event_pattern="*", priority=100)
-    def pre_log_event(self, event_args, *args, **kwargs):
+    def pre_log_event(self, event_args: dict[str, Any],
+                      *args: Any, **kwargs: Any) -> None:
         log.debug("Event: {0}, args: {1} kwargs: {2}".format(
             event_args.get("event"), args, kwargs))
 
     @observe(event_pattern="*", priority=4900)
-    def post_log_event(self, event_args, *args, **kwargs):
+    def post_log_event(self, event_args: dict[str, Any],
+                       *args: Any, **kwargs: Any) -> None:
         log.debug("Event: {0}, result: {1}".format(
             event_args.get("event"), event_args.get("result")))
 
@@ -34,7 +37,8 @@ class ExceptionWrappingMiddleware(object):
     Wraps all unhandled exceptions in cloudbridge exceptions.
     """
     @intercept(event_pattern="*", priority=1050)
-    def wrap_exception(self, event_args, *args, **kwargs):
+    def wrap_exception(self, event_args: dict[str, Any],
+                       *args: Any, **kwargs: Any) -> Any:
         next_handler = event_args.pop("next_handler")
         if not next_handler:
             return

+ 36 - 25
cloudbridge/base/provider.py

@@ -5,13 +5,17 @@ import logging
 import os
 from configparser import ConfigParser
 from os.path import expanduser
+from typing import Any
+from typing import cast
 
+from pyeventsystem.middleware import MiddlewareManager
 from pyeventsystem.middleware import SimpleMiddlewareManager
 
 from ..base.middleware import ExceptionWrappingMiddleware
 from ..interfaces import CloudProvider
 from ..interfaces.exceptions import ProviderConnectionException
 from ..interfaces.resources import Configuration
+from ..interfaces.resources import PlacementZone
 
 log = logging.getLogger(__name__)
 
@@ -28,11 +32,11 @@ CloudBridgeConfigLocations.append(UserConfigPath)
 
 class BaseConfiguration(Configuration):
 
-    def __init__(self, user_config):
+    def __init__(self, user_config: dict[str, Any]) -> None:
         self.update(user_config)
 
     @property
-    def default_result_limit(self):
+    def default_result_limit(self) -> int:
         """
         Get the maximum number of results to return for a
         list method
@@ -42,28 +46,30 @@ class BaseConfiguration(Configuration):
         """
         log.debug("Maximum number of results for list methods %s",
                   DEFAULT_RESULT_LIMIT)
-        return self.get('default_result_limit', DEFAULT_RESULT_LIMIT)
+        return cast(int, self.get('default_result_limit', DEFAULT_RESULT_LIMIT))
 
     @property
-    def default_wait_timeout(self):
+    def default_wait_timeout(self) -> int:
         """
         Gets the default wait timeout for LifeCycleObjects.
         """
         log.debug("Default wait timeout for LifeCycleObjects %s",
                   DEFAULT_WAIT_TIMEOUT)
-        return self.get('default_wait_timeout', DEFAULT_WAIT_TIMEOUT)
+        return cast(int, self.get('default_wait_timeout',
+                                  DEFAULT_WAIT_TIMEOUT))
 
     @property
-    def default_wait_interval(self):
+    def default_wait_interval(self) -> int:
         """
         Gets the default wait interval for LifeCycleObjects.
         """
         log.debug("Default wait interfal for LifeCycleObjects %s",
                   DEFAULT_WAIT_INTERVAL)
-        return self.get('default_wait_interval', DEFAULT_WAIT_INTERVAL)
+        return cast(int, self.get('default_wait_interval',
+                                  DEFAULT_WAIT_INTERVAL))
 
     @property
-    def debug_mode(self):
+    def debug_mode(self) -> bool:
         """
         A flag indicating whether CloudBridge is in debug mode. Setting
         this to True will cause the underlying provider's debug
@@ -75,59 +81,63 @@ class BaseConfiguration(Configuration):
         :rtype: ``bool``
         :return: Whether debug mode is on.
         """
-        return self.get('cb_debug', os.environ.get('CB_DEBUG', False))
+        return cast(bool, self.get('cb_debug',
+                                   os.environ.get('CB_DEBUG', False)))
 
 
 class BaseCloudProvider(CloudProvider):
-    def __init__(self, config):
+
+    PROVIDER_ID: str
+
+    def __init__(self, config: dict[str, Any]) -> None:
         self._config = BaseConfiguration(config)
         self._config_parser = ConfigParser()
         self._config_parser.read(CloudBridgeConfigLocations)
         self._middleware = SimpleMiddlewareManager()
         self.add_required_middleware()
-        self._region_name = None
-        self._zone_name = None
+        self._region_name: str | None = None
+        self._zone_name: str | None = None
 
     @property
-    def region_name(self):
+    def region_name(self) -> str | None:
         return self._region_name
 
     @property
-    def zone_name(self):
+    def zone_name(self) -> str | None:
         if not self._zone_name:
             region = self.compute.regions.current
-            zone = region.default_zone
+            zone = region.default_zone if region else None
             self._zone_name = zone.name if zone else None
             return self._zone_name
         else:
             try:
                 zone_dict = ast.literal_eval(self._zone_name)
                 if isinstance(zone_dict, dict):
-                    return zone_dict
+                    return cast("str | None", zone_dict)
             except (ValueError, SyntaxError):
                 pass
             return self._zone_name
 
     @property
-    def config(self):
+    def config(self) -> Configuration:
         return self._config
 
     @property
-    def name(self):
+    def name(self) -> str:
         return str(self.__class__.__name__)
 
     @property
-    def middleware(self):
+    def middleware(self) -> MiddlewareManager:
         return self._middleware
 
-    def add_required_middleware(self):
+    def add_required_middleware(self) -> None:
         """
         Adds common middleware that is essential for cloudbridge to function.
         Any other extra middleware can be added through the provider factory.
         """
         self.middleware.add(ExceptionWrappingMiddleware())
 
-    def authenticate(self):
+    def authenticate(self) -> bool:
         """
         A basic implementation which simply runs a low impact command to
         check whether cloud credentials work. Providers should override with
@@ -142,7 +152,7 @@ class BaseCloudProvider(CloudProvider):
             raise ProviderConnectionException(
                 "Authentication with cloud provider failed: %s" % (e,))
 
-    def clone(self, zone=None):
+    def clone(self, zone: PlacementZone | None = None) -> CloudProvider:
         cloned_config = self.config.copy()
         cloned_provider = self.__class__(cloned_config)
         if zone:
@@ -150,11 +160,11 @@ class BaseCloudProvider(CloudProvider):
             cloned_provider._zone_name = zone.name
         return cloned_provider
 
-    def _deepgetattr(self, obj, attr):
+    def _deepgetattr(self, obj: object, attr: str) -> Any:
         """Recurses through an attribute chain to get the ultimate value."""
         return functools.reduce(getattr, attr.split('.'), obj)
 
-    def has_service(self, service_type):
+    def has_service(self, service_type: str) -> bool:
         """
         Checks whether this provider supports a given service.
 
@@ -178,7 +188,8 @@ class BaseCloudProvider(CloudProvider):
                  service_type)
         return False
 
-    def _get_config_value(self, key, default_value=None):
+    def _get_config_value(self, key: str,
+                          default_value: Any = None) -> Any:
         """
         A convenience method to extract a configuration value.
 

Разница между файлами не показана из-за своего большого размера
+ 239 - 147
cloudbridge/base/resources.py


+ 132 - 75
cloudbridge/base/services.py

@@ -2,10 +2,31 @@
 Base implementation for services available through a provider
 """
 import logging
+from abc import abstractmethod
+from typing import Any
+from typing import cast
 
 from cloudbridge.interfaces.exceptions import InvalidParamException
+from cloudbridge.interfaces.provider import CloudProvider
+from cloudbridge.interfaces.resources import Bucket
+from cloudbridge.interfaces.resources import DnsRecord
 from cloudbridge.interfaces.resources import DnsRecordType
+from cloudbridge.interfaces.resources import DnsZone
+from cloudbridge.interfaces.resources import FloatingIP
+from cloudbridge.interfaces.resources import Gateway
+from cloudbridge.interfaces.resources import Instance
+from cloudbridge.interfaces.resources import KeyPair
+from cloudbridge.interfaces.resources import MachineImage
 from cloudbridge.interfaces.resources import Network
+from cloudbridge.interfaces.resources import Region
+from cloudbridge.interfaces.resources import ResultList
+from cloudbridge.interfaces.resources import Router
+from cloudbridge.interfaces.resources import Snapshot
+from cloudbridge.interfaces.resources import Subnet
+from cloudbridge.interfaces.resources import VMFirewall
+from cloudbridge.interfaces.resources import VMFirewallRule
+from cloudbridge.interfaces.resources import VMType
+from cloudbridge.interfaces.resources import Volume
 from cloudbridge.interfaces.services import BucketObjectService
 from cloudbridge.interfaces.services import BucketService
 from cloudbridge.interfaces.services import CloudService
@@ -46,46 +67,47 @@ class BaseCloudService(CloudService):
 
     STANDARD_EVENT_PRIORITY = 2500
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         self._service_event_pattern = "provider"
         self._provider = provider
         # discover and register all middleware
         provider.middleware.add(self)
 
     @property
-    def provider(self):
+    def provider(self) -> CloudProvider:
         return self._provider
 
     @property
-    def events(self):
+    def events(self) -> Any:
         return self._provider.middleware.events
 
 
 class BaseSecurityService(SecurityService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseSecurityService, self).__init__(provider)
 
 
 class BaseKeyPairService(
-        BasePageableObjectMixin, KeyPairService, BaseCloudService):
+        BasePageableObjectMixin[KeyPair], KeyPairService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseKeyPairService, self).__init__(provider)
         self._service_event_pattern += ".security.key_pairs"
 
 
 class BaseVMFirewallService(
-        BasePageableObjectMixin, VMFirewallService, BaseCloudService):
+        BasePageableObjectMixin[VMFirewall], VMFirewallService,
+        BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseVMFirewallService, self).__init__(provider)
         self._service_event_pattern += ".security.vm_firewalls"
 
     @dispatch(event="provider.security.vm_firewalls.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, **kwargs):
-        obj_list = self
+    def find(self, **kwargs: Any) -> ResultList[VMFirewall]:
+        obj_list = list(self)
         filters = ['label']
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
 
@@ -99,21 +121,26 @@ class BaseVMFirewallService(
                                      matches if matches else [])
 
 
-class BaseVMFirewallRuleService(BasePageableObjectMixin,
-                                VMFirewallRuleService,
-                                BaseCloudService):
+# The pageable mixin's list(limit, marker) intentionally differs from this
+# service's list(firewall, limit, marker); the mixin is reused only for its
+# iteration helpers, so the signature clash is expected.
+class BaseVMFirewallRuleService(  # type: ignore[misc]
+        BasePageableObjectMixin[VMFirewallRule],
+        VMFirewallRuleService,
+        BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseVMFirewallRuleService, self).__init__(provider)
         self._provider = provider
 
     @property
-    def provider(self):
+    def provider(self) -> CloudProvider:
         return self._provider
 
     @dispatch(event="provider.security.vm_firewall_rules.get",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def get(self, firewall, rule_id):
+    def get(self, firewall: VMFirewall,
+            rule_id: str) -> VMFirewallRule | None:
         matches = [rule for rule in firewall.rules if rule.id == rule_id]
         if matches:
             return matches[0]
@@ -122,8 +149,9 @@ class BaseVMFirewallRuleService(BasePageableObjectMixin,
 
     @dispatch(event="provider.security.vm_firewall_rules.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, firewall, **kwargs):
-        obj_list = firewall.rules
+    def find(self, firewall: VMFirewall,
+             **kwargs: Any) -> ResultList[VMFirewallRule]:
+        obj_list = list(firewall.rules)
         filters = ['name', 'direction', 'protocol', 'from_port', 'to_port',
                    'cidr', 'src_dest_fw', 'src_dest_fw_id']
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
@@ -132,30 +160,44 @@ class BaseVMFirewallRuleService(BasePageableObjectMixin,
 
 class BaseStorageService(StorageService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseStorageService, self).__init__(provider)
 
+    @property
+    @abstractmethod
+    def _bucket_objects(self) -> BucketObjectService:
+        """
+        Provider-internal service backing bucket-object operations.
+
+        This is the service that ``bucket.objects`` (BucketObjectSubService)
+        and the base multipart-upload code delegate to. It is a base-layer
+        implementation detail, deliberately not part of the public
+        StorageService interface; every provider's storage service implements
+        it.
+        """
+        pass
+
 
 class BaseVolumeService(
-        BasePageableObjectMixin, VolumeService, BaseCloudService):
+        BasePageableObjectMixin[Volume], VolumeService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseVolumeService, self).__init__(provider)
         self._service_event_pattern += ".storage.volumes"
 
 
 class BaseSnapshotService(
-        BasePageableObjectMixin, SnapshotService, BaseCloudService):
+        BasePageableObjectMixin[Snapshot], SnapshotService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseSnapshotService, self).__init__(provider)
         self._service_event_pattern += ".storage.snapshots"
 
 
 class BaseBucketService(
-        BasePageableObjectMixin, BucketService, BaseCloudService):
+        BasePageableObjectMixin[Bucket], BucketService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseBucketService, self).__init__(provider)
         self._service_event_pattern += ".storage.buckets"
 
@@ -163,8 +205,8 @@ class BaseBucketService(
     # provider-specific querying for find method
     @dispatch(event="provider.storage.buckets.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, **kwargs):
-        obj_list = self
+    def find(self, **kwargs: Any) -> ResultList[Bucket]:
+        obj_list = list(self)
         filters = ['name']
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
 
@@ -180,67 +222,67 @@ class BaseBucketService(
 
 class BaseBucketObjectService(BucketObjectService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseBucketObjectService, self).__init__(provider)
         self._service_event_pattern += ".storage._bucket_objects"
-        self._bucket = None
+        self._bucket: Bucket | None = None
 
 
 class BaseComputeService(ComputeService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseComputeService, self).__init__(provider)
 
 
 class BaseImageService(
-        BasePageableObjectMixin, ImageService, BaseCloudService):
+        BasePageableObjectMixin[MachineImage], ImageService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseImageService, self).__init__(provider)
         self._service_event_pattern += ".compute.images"
 
 
 class BaseInstanceService(
-        BasePageableObjectMixin, InstanceService, BaseCloudService):
+        BasePageableObjectMixin[Instance], InstanceService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseInstanceService, self).__init__(provider)
         self._service_event_pattern += ".compute.instances"
 
 
 class BaseVMTypeService(
-        BasePageableObjectMixin, VMTypeService, BaseCloudService):
+        BasePageableObjectMixin[VMType], VMTypeService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseVMTypeService, self).__init__(provider)
         self._service_event_pattern += ".compute.vm_types"
 
     @dispatch(event="provider.compute.vm_types.get",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def get(self, vm_type_id):
+    def get(self, vm_type_id: str) -> VMType | None:
         vm_type = (t for t in self if t.id == vm_type_id)
         return next(vm_type, None)
 
     @dispatch(event="provider.compute.vm_types.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, **kwargs):
-        obj_list = self
+    def find(self, **kwargs: Any) -> ResultList[VMType]:
+        obj_list = list(self)
         filters = ['name']
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
         return ClientPagedResultList(self._provider, list(matches))
 
 
 class BaseRegionService(
-        BasePageableObjectMixin, RegionService, BaseCloudService):
+        BasePageableObjectMixin[Region], RegionService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseRegionService, self).__init__(provider)
         self._service_event_pattern += ".compute.regions"
 
     @dispatch(event="provider.compute.regions.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, **kwargs):
-        obj_list = self
+    def find(self, **kwargs: Any) -> ResultList[Region]:
+        obj_list = list(self)
         filters = ['name']
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
         return ClientPagedResultList(self._provider, list(matches))
@@ -248,23 +290,27 @@ class BaseRegionService(
 
 class BaseNetworkingService(NetworkingService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseNetworkingService, self).__init__(provider)
 
 
 class BaseNetworkService(
-        BasePageableObjectMixin, NetworkService, BaseCloudService):
+        BasePageableObjectMixin[Network], NetworkService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseNetworkService, self).__init__(provider)
         self._service_event_pattern += ".networking.networks"
 
     @property
-    def subnets(self):
-        return [subnet for subnet in self.provider.subnets
-                if subnet.network_id == self.id]
-
-    def get_or_create_default(self):
+    def subnets(self) -> list[Subnet]:  # type: ignore[override]
+        # NOTE: this base implementation is a stub that every provider
+        # overrides; it references attributes that do not exist on the
+        # service, so the accesses are typed through ``Any``.
+        this: Any = self
+        return [subnet for subnet in this.provider.subnets
+                if subnet.network_id == this.id]
+
+    def get_or_create_default(self) -> Network:
         networks = self.provider.networking.networks.find(
             label=BaseNetwork.CB_DEFAULT_NETWORK_LABEL)
 
@@ -278,8 +324,8 @@ class BaseNetworkService(
 
     @dispatch(event="provider.networking.networks.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, **kwargs):
-        obj_list = self
+    def find(self, **kwargs: Any) -> ResultList[Network]:
+        obj_list = list(self)
         filters = ['label']
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
 
@@ -294,15 +340,17 @@ class BaseNetworkService(
 
 
 class BaseSubnetService(
-        BasePageableObjectMixin, SubnetService, BaseCloudService):
+        BasePageableObjectMixin[Subnet], SubnetService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseSubnetService, self).__init__(provider)
         self._service_event_pattern += ".networking.subnets"
 
     @dispatch(event="provider.networking.subnets.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, network=None, **kwargs):
+    def find(self, network: Network | None = None,
+             **kwargs: Any) -> ResultList[Subnet]:
+        obj_list: Any
         if not network:
             obj_list = self
         else:
@@ -311,27 +359,30 @@ class BaseSubnetService(
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
         return ClientPagedResultList(self._provider, list(matches))
 
-    def get_or_create_default(self):
+    def get_or_create_default(self) -> Subnet:
         # Look for a CB-default subnet
-        matches = self.find(label=BaseSubnet.CB_DEFAULT_SUBNET_LABEL)
+        matches: ResultList[Subnet] = self.find(
+            label=BaseSubnet.CB_DEFAULT_SUBNET_LABEL)
         if matches:
             return matches[0]
 
         # No provider-default Subnet exists, try to create it (net + subnets)
-        network = self.provider.networking.networks.get_or_create_default()
+        networks = cast(BaseNetworkService,
+                        self.provider.networking.networks)
+        network = networks.get_or_create_default()
         subnet = self.create(BaseSubnet.CB_DEFAULT_SUBNET_LABEL, network,
                              BaseSubnet.CB_DEFAULT_SUBNET_IPV4RANGE)
         return subnet
 
 
 class BaseRouterService(
-        BasePageableObjectMixin, RouterService, BaseCloudService):
+        BasePageableObjectMixin[Router], RouterService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseRouterService, self).__init__(provider)
         self._service_event_pattern += ".networking.routers"
 
-    def get_or_create_default(self, network):
+    def get_or_create_default(self, network: Network | str) -> Router:
         net_id = network.id if isinstance(network, Network) else network
         routers = self.provider.networking.routers.find(
             label=BaseRouter.CB_DEFAULT_ROUTER_LABEL)
@@ -345,19 +396,20 @@ class BaseRouterService(
 
 class BaseGatewayService(GatewayService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseGatewayService, self).__init__(provider)
 
 
 class BaseFloatingIPService(FloatingIPService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseFloatingIPService, self).__init__(provider)
 
     @dispatch(event="provider.networking.floating_ips.find",
               priority=BaseCloudService.STANDARD_EVENT_PRIORITY)
-    def find(self, gateway, **kwargs):
-        obj_list = gateway.floating_ips
+    def find(self, gateway: Gateway,
+             **kwargs: Any) -> ResultList[FloatingIP]:
+        obj_list = list(gateway.floating_ips)
         filters = ['name', 'public_ip']
         matches = cb_helpers.generic_find(filters, kwargs, obj_list)
         return ClientPagedResultList(self._provider, list(matches))
@@ -365,31 +417,36 @@ class BaseFloatingIPService(FloatingIPService, BaseCloudService):
 
 class BaseDnsService(DnsService, BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseDnsService, self).__init__(provider)
 
 
-class BaseDnsZoneService(BasePageableObjectMixin, DnsZoneService,
+class BaseDnsZoneService(BasePageableObjectMixin[DnsZone], DnsZoneService,
                          BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseDnsZoneService, self).__init__(provider)
 
-    def _get_fully_qualified_dns(self, name):
+    def _get_fully_qualified_dns(self, name: str) -> str:
         # Add a trailing dot to fully qualify
         return name + '.' if not name.endswith('.') else name
 
 
-class BaseDnsRecordService(BasePageableObjectMixin, DnsRecordService,
-                           BaseCloudService):
+# The pageable mixin's list(limit, marker) intentionally differs from this
+# service's list(dns_zone, limit, marker); the mixin is reused only for its
+# iteration helpers, so the signature clash is expected.
+class BaseDnsRecordService(  # type: ignore[misc]
+        BasePageableObjectMixin[DnsRecord],
+        DnsRecordService,
+        BaseCloudService):
 
-    def __init__(self, provider):
+    def __init__(self, provider: CloudProvider) -> None:
         super(BaseDnsRecordService, self).__init__(provider)
 
-    def _get_fully_qualified_dns(self, name):
+    def _get_fully_qualified_dns(self, name: str) -> str:
         # Add a trailing dot to fully qualify
         return name + '.' if not name.endswith('.') else name
 
-    def _standardize_record(self, value, type):
+    def _standardize_record(self, value: str, type: str) -> str:
         return (self._get_fully_qualified_dns(value)
                 if type in (DnsRecordType.CNAME, DnsRecordType.MX) else value)

+ 105 - 58
cloudbridge/base/subservices.py

@@ -1,5 +1,24 @@
+import builtins
 import logging
-
+from typing import Any
+from typing import TYPE_CHECKING
+from typing import cast
+
+from cloudbridge.interfaces.provider import CloudProvider
+from cloudbridge.interfaces.resources import Bucket
+from cloudbridge.interfaces.resources import BucketObject
+from cloudbridge.interfaces.resources import DnsRecord
+from cloudbridge.interfaces.resources import DnsZone
+from cloudbridge.interfaces.resources import FloatingIP
+from cloudbridge.interfaces.resources import Gateway
+from cloudbridge.interfaces.resources import InternetGateway
+from cloudbridge.interfaces.resources import Network
+from cloudbridge.interfaces.resources import ResultList
+from cloudbridge.interfaces.resources import Subnet
+from cloudbridge.interfaces.resources import TrafficDirection
+from cloudbridge.interfaces.resources import VMFirewall
+from cloudbridge.interfaces.resources import VMFirewallRule
+from cloudbridge.interfaces.services import BucketObjectService
 from cloudbridge.interfaces.subservices import BucketObjectSubService
 from cloudbridge.interfaces.subservices import DnsRecordSubService
 from cloudbridge.interfaces.subservices import FloatingIPSubService
@@ -9,194 +28,222 @@ from cloudbridge.interfaces.subservices import VMFirewallRuleSubService
 
 from .resources import BasePageableObjectMixin
 
+if TYPE_CHECKING:
+    from .services import BaseStorageService
+
 log = logging.getLogger(__name__)
 
 
-class BaseBucketObjectSubService(BasePageableObjectMixin,
+class BaseBucketObjectSubService(BasePageableObjectMixin[BucketObject],
                                  BucketObjectSubService):
 
-    def __init__(self, provider, bucket):
+    def __init__(self, provider: CloudProvider, bucket: Bucket) -> None:
         self.__provider = provider
         self.bucket = bucket
 
     @property
-    def _provider(self):
+    def _provider(self) -> CloudProvider:
         return self.__provider
 
-    def get(self, name):
-        return self._provider.storage._bucket_objects.get(self.bucket, name)
+    @property
+    def _bucket_objects(self) -> BucketObjectService:
+        # ``_bucket_objects`` is a base-layer member (BaseStorageService), not
+        # part of the public StorageService interface.
+        storage = cast("BaseStorageService", self._provider.storage)
+        return storage._bucket_objects
+
+    def get(self, name: str) -> BucketObject | None:
+        return self._bucket_objects.get(self.bucket, name)
 
-    def list(self, limit=None, marker=None, prefix=None):
-        return self._provider.storage._bucket_objects.list(self.bucket, limit,
-                                                           marker, prefix)
+    def list(self, limit: int | None = None, marker: str | None = None,
+             prefix: str | None = None) -> ResultList[BucketObject]:
+        return self._bucket_objects.list(self.bucket, limit=limit,
+                                         marker=marker, prefix=prefix)
 
-    def find(self, **kwargs):
-        return self._provider.storage._bucket_objects.find(self.bucket,
-                                                           **kwargs)
+    def find(self, **kwargs: Any) -> ResultList[BucketObject]:
+        return self._bucket_objects.find(self.bucket, **kwargs)
 
-    def create(self, name):
-        return self._provider.storage._bucket_objects.create(self.bucket, name)
+    def create(self, name: str) -> BucketObject:
+        return self._bucket_objects.create(self.bucket, name)
 
 
-class BaseGatewaySubService(GatewaySubService, BasePageableObjectMixin):
+class BaseGatewaySubService(GatewaySubService,
+                            BasePageableObjectMixin[InternetGateway]):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         self._network = network
         self.__provider = provider
 
     @property
-    def _provider(self):
+    def _provider(self) -> CloudProvider:
         return self.__provider
 
-    def get_or_create(self):
+    def get_or_create(self) -> InternetGateway:
         return (self._provider.networking
                               ._gateways
                               .get_or_create(self._network))
 
-    def delete(self, gateway):
+    def delete(self, gateway: Gateway) -> None:
         return (self._provider.networking
                               ._gateways
                               .delete(self._network, gateway))
 
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[InternetGateway]:
         return (self._provider.networking
                               ._gateways
                               .list(self._network, limit, marker))
 
 
-class BaseVMFirewallRuleSubService(BasePageableObjectMixin,
+class BaseVMFirewallRuleSubService(BasePageableObjectMixin[VMFirewallRule],
                                    VMFirewallRuleSubService):
 
-    def __init__(self, provider, firewall):
+    def __init__(self, provider: CloudProvider,
+                 firewall: VMFirewall) -> None:
         self.__provider = provider
         self._firewall = firewall
 
     @property
-    def _provider(self):
+    def _provider(self) -> CloudProvider:
         return self.__provider
 
-    def get(self, rule_id):
+    def get(self, rule_id: str) -> VMFirewallRule | None:
         return self._provider.security._vm_firewall_rules.get(self._firewall,
                                                               rule_id)
 
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[VMFirewallRule]:
         return self._provider.security._vm_firewall_rules.list(self._firewall,
                                                                limit, marker)
 
-    def create(self, direction, protocol=None, from_port=None,
-               to_port=None, cidr=None, src_dest_fw=None):
+    def create(self, direction: TrafficDirection, protocol: str | None = None,
+               from_port: int | None = None,
+               to_port: int | None = None,
+               cidr: str | builtins.list[str] | None = None,
+               src_dest_fw: VMFirewall | None = None) -> VMFirewallRule:
         return (self._provider
                     .security
                     ._vm_firewall_rules
                     .create(self._firewall, direction, protocol, from_port,
                             to_port, cidr, src_dest_fw))
 
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[VMFirewallRule]:
         return self._provider.security._vm_firewall_rules.find(self._firewall,
                                                                **kwargs)
 
-    def delete(self, rule_id):
+    def delete(self, rule_id: str) -> None:
         return (self._provider
                     .security
                     ._vm_firewall_rules
                     .delete(self._firewall, rule_id))
 
 
-class BaseFloatingIPSubService(FloatingIPSubService, BasePageableObjectMixin):
+class BaseFloatingIPSubService(FloatingIPSubService,
+                               BasePageableObjectMixin[FloatingIP]):
 
-    def __init__(self, provider, gateway):
+    def __init__(self, provider: CloudProvider, gateway: Gateway) -> None:
         self.__provider = provider
         self.gateway = gateway
 
     @property
-    def _provider(self):
+    def _provider(self) -> CloudProvider:
         return self.__provider
 
-    def get(self, fip_id):
+    def get(self, fip_id: str) -> FloatingIP | None:
         return self._provider.networking._floating_ips.get(self.gateway,
                                                            fip_id)
 
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[FloatingIP]:
         return self._provider.networking._floating_ips.list(self.gateway,
                                                             limit, marker)
 
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[FloatingIP]:
         return self._provider.networking._floating_ips.find(self.gateway,
                                                             **kwargs)
 
-    def create(self):
+    def create(self) -> FloatingIP:
         return self._provider.networking._floating_ips.create(self.gateway)
 
-    def delete(self, fip):
+    def delete(self, fip: FloatingIP | str) -> None:
         return self._provider.networking._floating_ips.delete(self.gateway,
                                                               fip)
 
 
-class BaseSubnetSubService(SubnetSubService, BasePageableObjectMixin):
+class BaseSubnetSubService(SubnetSubService, BasePageableObjectMixin[Subnet]):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         self.__provider = provider
         self.network = network
 
     @property
-    def _provider(self):
+    def _provider(self) -> CloudProvider:
         return self.__provider
 
-    def get(self, subnet_id):
+    def get(self, subnet_id: str) -> Subnet | None:
         sn = self._provider.networking.subnets.get(subnet_id)
-        if sn.network_id != self.network.id:
+        if sn and sn.network_id != self.network.id:
             log.warning("The SubnetSubService nested in the network '{}' "
                         "returned subnet '{}' which is attached to another "
                         "network '{}'".format(str(self.network), str(sn),
                                               str(sn.network)))
         return sn
 
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[Subnet]:
         return self._provider.networking.subnets.list(network=self.network,
                                                       limit=limit,
                                                       marker=marker)
 
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[Subnet]:
         return self._provider.networking.subnets.find(network=self.network,
                                                       **kwargs)
 
-    def create(self, label, cidr_block):
+    def create(self, label: str, cidr_block: str) -> Subnet:
         return self._provider.networking.subnets.create(label,
                                                         self.network,
                                                         cidr_block)
 
-    def delete(self, subnet):
+    def delete(self, subnet: Subnet | str) -> None:
         return self._provider.networking.subnets.delete(subnet)
 
 
-class BaseDnsRecordSubService(DnsRecordSubService, BasePageableObjectMixin):
+class BaseDnsRecordSubService(DnsRecordSubService,
+                              BasePageableObjectMixin[DnsRecord]):
 
-    def __init__(self, provider, dns_zone):
+    def __init__(self, provider: CloudProvider, dns_zone: DnsZone) -> None:
         self.__provider = provider
         self.dns_zone = dns_zone
 
     @property
-    def _provider(self):
+    def _provider(self) -> CloudProvider:
         return self.__provider
 
-    def get(self, rec_id):
+    def get(self, rec_id: str) -> DnsRecord | None:
         # pylint:disable=protected-access
         return self._provider.dns._records.get(self.dns_zone, rec_id)
 
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[DnsRecord]:
         # pylint:disable=protected-access
         return self._provider.dns._records.list(
             dns_zone=self.dns_zone, limit=limit, marker=marker)
 
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[DnsRecord]:
         # pylint:disable=protected-access
-        return self._provider.dns._records.find(
-            dns_zone=self.dns_zone, **kwargs)
-
-    def create(self, name, type, data, ttl=None):
+        # find/delete are provider-internal extensions not declared on the
+        # DnsRecordService interface; reach them through ``Any``.
+        records: Any = self._provider.dns._records
+        return cast("ResultList[DnsRecord]",
+                    records.find(dns_zone=self.dns_zone, **kwargs))
+
+    def create(self, name: str, type: str, data: str,
+               ttl: int | None = None) -> DnsRecord:
         # pylint:disable=protected-access
         return self._provider.dns._records.create(
             self.dns_zone, name, type, data, ttl)
 
-    def delete(self, rec):
-        return self._provider.dns._records.delete(self.dns_zone, rec)
+    def delete(self, rec: DnsRecord | str) -> None:
+        # pylint:disable=protected-access
+        records: Any = self._provider.dns._records
+        records.delete(self.dns_zone, rec)

+ 14 - 10
cloudbridge/factory.py

@@ -3,6 +3,7 @@ import inspect
 import logging
 import pkgutil
 from collections import defaultdict
+from typing import Any
 
 from cloudbridge import providers
 from cloudbridge.interfaces import CloudProvider
@@ -26,11 +27,12 @@ class CloudProviderFactory(object):
     Get info and handle on the available cloud provider implementations.
     """
 
-    def __init__(self):
-        self.provider_list = defaultdict(dict)
+    def __init__(self) -> None:
+        self.provider_list: defaultdict[str, dict[str, type[CloudProvider]]] \
+            = defaultdict(dict)
         log.debug("Providers List: %s", self.provider_list)
 
-    def register_provider_class(self, cls):
+    def register_provider_class(self, cls: type) -> None:
         """
         Registers a provider class with the factory. The class must
         inherit from cloudbridge.interfaces.CloudProvider
@@ -61,7 +63,7 @@ class CloudProviderFactory(object):
             log.debug("Class: %s does not implement the CloudProvider"
                       "  interface. Ignoring...", cls)
 
-    def discover_providers(self):
+    def discover_providers(self) -> None:
         """
         Discover all available providers within the
         ``cloudbridge.providers`` package.
@@ -74,7 +76,7 @@ class CloudProviderFactory(object):
             except Exception as e:
                 log.debug("Could not import provider: %s", e)
 
-    def _import_provider(self, module_name):
+    def _import_provider(self, module_name: str) -> None:
         """
         Imports and registers providers from the given module name.
         Raises an ImportError if the import does not succeed.
@@ -88,7 +90,7 @@ class CloudProviderFactory(object):
             log.debug("Registering the provider: %s", cls)
             self.register_provider_class(cls)
 
-    def list_providers(self):
+    def list_providers(self) -> dict[str, dict[str, type[CloudProvider]]]:
         """
         Get a list of available providers.
 
@@ -108,7 +110,8 @@ class CloudProviderFactory(object):
         log.debug("List of available providers: %s", self.provider_list)
         return self.provider_list
 
-    def create_provider(self, name, config):
+    def create_provider(self, name: str,
+                        config: dict[str, Any]) -> CloudProvider:
         """
         Searches all available providers for a CloudProvider interface with the
         given name, and instantiates it based on the given config dictionary,
@@ -138,7 +141,7 @@ class CloudProviderFactory(object):
         log.debug("Created '%s' provider", name)
         return provider_class(config)
 
-    def get_provider_class(self, name):
+    def get_provider_class(self, name: str) -> type[CloudProvider] | None:
         """
         Return a class for the requested provider.
 
@@ -155,7 +158,8 @@ class CloudProviderFactory(object):
             log.debug("Provider with the name: %s not found", name)
             return None
 
-    def get_all_provider_classes(self, ignore_mocks=False):
+    def get_all_provider_classes(
+            self, ignore_mocks: bool = False) -> list[type[CloudProvider]]:
         """
         Returns a list of classes for all available provider implementations
 
@@ -167,7 +171,7 @@ class CloudProviderFactory(object):
         :return: A list of all available provider classes or an empty list
         if none found.
         """
-        all_providers = []
+        all_providers: list[type[CloudProvider]] = []
         for impl in self.list_providers().values():
             if ignore_mocks:
                 if not issubclass(impl["class"], TestMockHelperMixin):

+ 4 - 4
cloudbridge/interfaces/exceptions.py

@@ -54,7 +54,7 @@ class InvalidNameException(CloudBridgeBaseException):
     letters, which are not allowed in a resource name.
     """
 
-    def __init__(self, msg):
+    def __init__(self, msg: str) -> None:
         super(InvalidNameException, self).__init__(msg)
 
 
@@ -68,7 +68,7 @@ class InvalidLabelException(InvalidNameException):
     identical.
     """
 
-    def __init__(self, msg):
+    def __init__(self, msg: str) -> None:
         super(InvalidLabelException, self).__init__(msg)
 
 
@@ -79,7 +79,7 @@ class InvalidValueException(CloudBridgeBaseException):
     direction of a firewall rule other than TrafficDirection.INBOUND or
     TrafficDirection.OUTBOUND.
     """
-    def __init__(self, param, value):
+    def __init__(self, param: str, value: object) -> None:
         super(InvalidValueException, self).__init__(
             "Param %s has been given an unrecognised value %s" %
             (param, value))
@@ -100,5 +100,5 @@ class InvalidParamException(InvalidNameException):
     to a service.find() method.
     """
 
-    def __init__(self, msg):
+    def __init__(self, msg: str) -> None:
         super(InvalidParamException, self).__init__(msg)

+ 34 - 18
cloudbridge/interfaces/provider.py

@@ -1,9 +1,25 @@
 """
 Specification for a provider interface
 """
+from __future__ import annotations
+
 from abc import ABCMeta
 from abc import abstractmethod
 from abc import abstractproperty
+from typing import Any
+from typing import TYPE_CHECKING
+
+if TYPE_CHECKING:
+    from pyeventsystem.middleware import MiddlewareManager
+
+    from cloudbridge.interfaces.resources import Configuration
+    from cloudbridge.interfaces.resources import Instance
+    from cloudbridge.interfaces.resources import PlacementZone
+    from cloudbridge.interfaces.services import ComputeService
+    from cloudbridge.interfaces.services import DnsService
+    from cloudbridge.interfaces.services import NetworkingService
+    from cloudbridge.interfaces.services import SecurityService
+    from cloudbridge.interfaces.services import StorageService
 
 
 class CloudProvider(object):
@@ -13,7 +29,7 @@ class CloudProvider(object):
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def __init__(self, config):
+    def __init__(self, config: dict[str, Any]) -> None:
         """
         Create a new provider instance given a dictionary of
         configuration attributes.
@@ -31,7 +47,7 @@ class CloudProvider(object):
         pass
 
     @abstractproperty
-    def config(self):
+    def config(self) -> Configuration:
         """
         Returns the config object associated with this provider. This object
         is a subclass of :class:`dict` and will contain the properties
@@ -58,7 +74,7 @@ class CloudProvider(object):
         pass
 
     @abstractproperty
-    def middleware(self):
+    def middleware(self) -> MiddlewareManager:
         """
         Returns the middleware manager associated with this provider. The
         middleware manager can be used to add or remove middleware from
@@ -72,7 +88,7 @@ class CloudProvider(object):
         pass
 
     @abstractmethod
-    def clone(self, zone=None):
+    def clone(self, zone: PlacementZone | None = None) -> CloudProvider:
         """
         Create a clone of this provider. An optional `zone` parameter can be
         used to clone the provider to use a different zone.
@@ -101,7 +117,7 @@ class CloudProvider(object):
         pass
 
     @abstractmethod
-    def authenticate(self):
+    def authenticate(self) -> bool:
         """
         Checks whether a provider can be successfully authenticated with the
         configured settings. Clients are *not* required to call this method
@@ -127,7 +143,7 @@ class CloudProvider(object):
         pass
 
     @abstractmethod
-    def has_service(self, service_type):
+    def has_service(self, service_type: str) -> bool:
         """
         Checks whether this provider supports a given service.
 
@@ -149,7 +165,7 @@ class CloudProvider(object):
         pass
 
     @abstractproperty
-    def region_name(self):
+    def region_name(self) -> str | None:
         """
         Returns the region that this provider is connected to.
         All provider operations will take place within this region.
@@ -160,7 +176,7 @@ class CloudProvider(object):
         pass
 
     @abstractproperty
-    def zone_name(self):
+    def zone_name(self) -> str | None:
         """
         Returns the placement zone that this provider is connected to.
         All provider operations will take place within this zone. Placement
@@ -183,7 +199,7 @@ class CloudProvider(object):
 #         pass
 
     @abstractproperty
-    def compute(self):
+    def compute(self) -> ComputeService:
         """
         Provides access to all compute related services in this provider.
 
@@ -206,7 +222,7 @@ class CloudProvider(object):
         pass
 
     @abstractproperty
-    def networking(self):
+    def networking(self) -> NetworkingService:
         """
         Provide access to all network related services in this provider.
 
@@ -223,7 +239,7 @@ class CloudProvider(object):
         """
 
     @abstractproperty
-    def security(self):
+    def security(self) -> SecurityService:
         """
         Provides access to key pair management and firewall control
 
@@ -241,7 +257,7 @@ class CloudProvider(object):
         pass
 
     @abstractproperty
-    def storage(self):
+    def storage(self) -> StorageService:
         """
         Provides access to storage related services in this provider.
         This includes the volume, snapshot and bucket services,
@@ -262,7 +278,7 @@ class CloudProvider(object):
         pass
 
     @abstractproperty
-    def dns(self):
+    def dns(self) -> DnsService:
         """
         Provides access to all DNS related services.
 
@@ -288,14 +304,14 @@ class TestMockHelperMixin(object):
     like HTTPretty which take over socket communications.
     """
 
-    def setUpMock(self):
+    def setUpMock(self) -> None:
         """
         Called before a test is started.
         """
         raise NotImplementedError(
             'TestMockHelperMixin.setUpMock not implemented')
 
-    def tearDownMock(self):
+    def tearDownMock(self) -> None:
         """
         Called before test teardown.
         """
@@ -312,11 +328,11 @@ class ContainerProvider(object):
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def create_container(self):
+    def create_container(self) -> None:
         pass
 
     @abstractmethod
-    def delete_container(self):
+    def delete_container(self) -> None:
         pass
 
 
@@ -328,7 +344,7 @@ class DeploymentProvider(object):
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def deploy(self, target):
+    def deploy(self, target: Instance) -> None:
         """
         Deploys on given target, where target is an Instance or Container
         """

Разница между файлами не показана из-за своего большого размера
+ 145 - 121
cloudbridge/interfaces/resources.py


Разница между файлами не показана из-за своего большого размера
+ 200 - 103
cloudbridge/interfaces/services.py


+ 57 - 34
cloudbridge/interfaces/subservices.py

@@ -1,17 +1,31 @@
+from __future__ import annotations
+
+import builtins
 from abc import ABCMeta
 from abc import abstractmethod
+from typing import Any
 
+from cloudbridge.interfaces.resources import BucketObject
+from cloudbridge.interfaces.resources import DnsRecord
+from cloudbridge.interfaces.resources import FloatingIP
+from cloudbridge.interfaces.resources import Gateway
+from cloudbridge.interfaces.resources import InternetGateway
 from cloudbridge.interfaces.resources import PageableObjectMixin
+from cloudbridge.interfaces.resources import ResultList
+from cloudbridge.interfaces.resources import Subnet
+from cloudbridge.interfaces.resources import TrafficDirection
+from cloudbridge.interfaces.resources import VMFirewall
+from cloudbridge.interfaces.resources import VMFirewallRule
 
 
-class BucketObjectSubService(PageableObjectMixin):
+class BucketObjectSubService(PageableObjectMixin[BucketObject]):
     """
     A container service for objects within a bucket.
     """
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def get(self, name):
+    def get(self, name: str) -> BucketObject | None:
         """
         Retrieve a given object from this bucket.
 
@@ -25,7 +39,8 @@ class BucketObjectSubService(PageableObjectMixin):
 
     @abstractmethod
     # pylint:disable=arguments-differ
-    def list(self, limit=None, marker=None, prefix=None):
+    def list(self, limit: int | None = None, marker: str | None = None,
+             prefix: str | None = None) -> ResultList[BucketObject]:
         """
         List objects in this bucket.
 
@@ -44,7 +59,7 @@ class BucketObjectSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[BucketObject]:
         """
         Search for an object by a given list of attributes.
 
@@ -62,7 +77,7 @@ class BucketObjectSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def create(self, name):
+    def create(self, name: str) -> BucketObject:
         """
         Create a new object within this bucket.
 
@@ -72,14 +87,14 @@ class BucketObjectSubService(PageableObjectMixin):
         pass
 
 
-class GatewaySubService(PageableObjectMixin):
+class GatewaySubService(PageableObjectMixin[InternetGateway]):
     """
     Manage internet gateway resources.
     """
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def get_or_create(self):
+    def get_or_create(self) -> InternetGateway:
         """
         Creates new or returns an existing internet gateway for a network.
 
@@ -92,7 +107,7 @@ class GatewaySubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def delete(self, gateway):
+    def delete(self, gateway: Gateway) -> None:
         """
         Delete a gateway.
 
@@ -102,7 +117,8 @@ class GatewaySubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[InternetGateway]:
         """
         List all available internet gateways.
 
@@ -112,14 +128,14 @@ class GatewaySubService(PageableObjectMixin):
         pass
 
 
-class FloatingIPSubService(PageableObjectMixin):
+class FloatingIPSubService(PageableObjectMixin[FloatingIP]):
     """
     Base interface for a FloatingIP Service.
     """
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def get(self, fip_id):
+    def get(self, fip_id: str) -> FloatingIP | None:
         """
         Returns a FloatingIP given its ID or ``None`` if not found.
 
@@ -132,7 +148,8 @@ class FloatingIPSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[FloatingIP]:
         """
         List floating (i.e., static) IP addresses.
 
@@ -142,7 +159,7 @@ class FloatingIPSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[FloatingIP]:
         """
         Searches for a FloatingIP by a given list of attributes.
 
@@ -162,7 +179,7 @@ class FloatingIPSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def create(self):
+    def create(self) -> FloatingIP:
         """
         Allocate a new floating (i.e., static) IP address.
 
@@ -172,7 +189,7 @@ class FloatingIPSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def delete(self, fip_id):
+    def delete(self, fip_id: FloatingIP | str) -> None:
         """
         Delete an existing FloatingIP.
 
@@ -182,14 +199,14 @@ class FloatingIPSubService(PageableObjectMixin):
         pass
 
 
-class VMFirewallRuleSubService(PageableObjectMixin):
+class VMFirewallRuleSubService(PageableObjectMixin[VMFirewallRule]):
     """
     Base interface for Firewall rules.
     """
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def get(self, rule_id):
+    def get(self, rule_id: str) -> VMFirewallRule | None:
         """
         Return a firewall rule given its ID.
 
@@ -212,7 +229,8 @@ class VMFirewallRuleSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[VMFirewallRule]:
         """
         List all firewall rules associated with this firewall.
 
@@ -222,8 +240,10 @@ class VMFirewallRuleSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def create(self, direction, protocol=None, from_port=None,
-               to_port=None, cidr=None, src_dest_fw=None):
+    def create(self, direction: TrafficDirection, protocol: str | None = None,
+               from_port: int | None = None, to_port: int | None = None,
+               cidr: str | builtins.list[str] | None = None,
+               src_dest_fw: VMFirewall | None = None) -> VMFirewallRule:
         """
         Create a VM firewall rule.
 
@@ -274,7 +294,7 @@ class VMFirewallRuleSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[VMFirewallRule]:
         """
         Find a firewall rule filtered by the given parameters.
 
@@ -310,7 +330,7 @@ class VMFirewallRuleSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def delete(self, rule_id):
+    def delete(self, rule_id: str) -> None:
         """
         Delete an existing VMFirewall rule.
 
@@ -320,14 +340,14 @@ class VMFirewallRuleSubService(PageableObjectMixin):
         pass
 
 
-class SubnetSubService(PageableObjectMixin):
+class SubnetSubService(PageableObjectMixin[Subnet]):
     """
     Base interface for a Subnet Service.
     """
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def get(self, subnet_id):
+    def get(self, subnet_id: str) -> Subnet | None:
         """
         Returns a Subnet given its ID or ``None`` if not found.
 
@@ -340,7 +360,8 @@ class SubnetSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[Subnet]:
         """
         List subnets within the network holding this subservice.
 
@@ -350,7 +371,7 @@ class SubnetSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[Subnet]:
         """
         Searches for a Subnet by a given list of attributes.
 
@@ -370,7 +391,7 @@ class SubnetSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def create(self, label, cidr_block):
+    def create(self, label: str, cidr_block: str) -> Subnet:
         """
         Create a new subnet within the network holding this subservice.
 
@@ -387,7 +408,7 @@ class SubnetSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def delete(self, subnet_id):
+    def delete(self, subnet_id: Subnet | str) -> None:
         """
         Delete an existing Subnet.
 
@@ -397,14 +418,14 @@ class SubnetSubService(PageableObjectMixin):
         pass
 
 
-class DnsRecordSubService(PageableObjectMixin):
+class DnsRecordSubService(PageableObjectMixin[DnsRecord]):
     """
     Base interface for a Dns Record Service.
     """
     __metaclass__ = ABCMeta
 
     @abstractmethod
-    def get(self, record_id):
+    def get(self, record_id: str) -> DnsRecord | None:
         """
         Returns a Dns Record given its ID or ``None`` if not found.
 
@@ -417,7 +438,8 @@ class DnsRecordSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def list(self, limit=None, marker=None):
+    def list(self, limit: int | None = None,
+             marker: str | None = None) -> ResultList[DnsRecord]:
         """
         List Dns Records within the Dns Zone holding this subservice.
 
@@ -427,7 +449,7 @@ class DnsRecordSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def find(self, **kwargs):
+    def find(self, **kwargs: Any) -> ResultList[DnsRecord]:
         """
         Searches for a DnsRecord by a given list of attributes.
 
@@ -447,7 +469,8 @@ class DnsRecordSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def create(self, label, type, data, ttl=None):
+    def create(self, label: str, type: str, data: str,
+               ttl: int | None = None) -> DnsRecord:
         """
         Create a new DnsRecord within the Dns Zone holding this subservice.
 
@@ -469,7 +492,7 @@ class DnsRecordSubService(PageableObjectMixin):
         pass
 
     @abstractmethod
-    def delete(self, record_id):
+    def delete(self, record_id: DnsRecord | str) -> None:
         """
         Delete an existing DnsRecord.
 

+ 43 - 21
cloudbridge/providers/aws/helpers.py

@@ -1,5 +1,10 @@
 """A set of AWS-specific helper methods used by the framework."""
+from __future__ import annotations
+
 import logging
+from typing import Any
+from typing import TYPE_CHECKING
+from typing import TypeVar
 
 from boto3.resources.params import create_request_parameters
 
@@ -9,12 +14,19 @@ from botocore.utils import merge_dicts
 
 from cloudbridge.base.resources import ClientPagedResultList
 from cloudbridge.base.resources import ServerPagedResultList
+from cloudbridge.interfaces.resources import CloudResource
+from cloudbridge.interfaces.resources import ResultList
+
+if TYPE_CHECKING:
+    from .provider import AWSCloudProvider
 
 
 log = logging.getLogger(__name__)
 
+T = TypeVar("T")
+
 
-def trim_empty_params(params_dict):
+def trim_empty_params(params_dict: dict[str, Any]) -> dict[str, Any]:
     """
     Given a dict containing potentially null values, trims out
     all the null values. This is to please Boto, which throws
@@ -35,7 +47,7 @@ def trim_empty_params(params_dict):
     return {k: v for k, v in params_dict.items() if v is not None}
 
 
-def find_tag_value(tags, key):
+def find_tag_value(tags: list[dict[str, Any]] | None, key: str) -> Any:
     """
     Finds the value associated with a given key from a list of AWS tags.
 
@@ -59,7 +71,9 @@ class BotoGenericService(object):
     resource, collection and paging support to implement
     basic cloudbridge methods.
     """
-    def __init__(self, provider, cb_resource, boto_conn, boto_collection_name):
+    def __init__(self, provider: AWSCloudProvider,
+                 cb_resource: Any, boto_conn: Any,
+                 boto_collection_name: str) -> None:
         """
         :type provider: :class:`AWSCloudProvider`
         :param provider: CloudBridge AWS provider to use
@@ -86,12 +100,12 @@ class BotoGenericService(object):
         self.boto_resource = self._infer_boto_resource(
             boto_conn, self.boto_collection_model)
 
-    def _infer_collection_model(self, conn, collection_name):
+    def _infer_collection_model(self, conn: Any, collection_name: str) -> Any:
         log.debug("Retrieving boto model for collection: %s", collection_name)
         return next(col for col in conn.meta.resource_model.collections
                     if col.name == collection_name)
 
-    def _infer_boto_resource(self, conn, collection_model):
+    def _infer_boto_resource(self, conn: Any, collection_model: Any) -> Any:
         log.debug("Retrieving resource model for collection: %s",
                   collection_model.name)
         resource_model = next(
@@ -99,7 +113,7 @@ class BotoGenericService(object):
             if sr.resource.model.name == collection_model.resource.model.name)
         return getattr(self.boto_conn, resource_model.name)
 
-    def get_raw(self, resource_id):
+    def get_raw(self, resource_id: str) -> Any:
         """
         Returns a single resource.
 
@@ -124,7 +138,7 @@ class BotoGenericService(object):
             else:
                 raise exc
 
-    def get(self, resource_id):
+    def get(self, resource_id: str) -> Any:
         """
         Returns a single resource.
 
@@ -139,7 +153,7 @@ class BotoGenericService(object):
         else:
             return None
 
-    def _get_list_operation(self):
+    def _get_list_operation(self) -> str:
         """
         This function discovers the list operation for a particular resource
         collection. For example, given the resource collection model for
@@ -147,7 +161,8 @@ class BotoGenericService(object):
         """
         return xform_name(self.boto_collection_model.request.operation)
 
-    def _to_boto_resource(self, collection, params, page):
+    def _to_boto_resource(self, collection: Any, params: Any,
+                          page: Any) -> Any:
         """
         This function duplicates some of the logic of the pages() method in
         boto.resources.collection.ResourceCollection. It will convert a raw
@@ -158,7 +173,8 @@ class BotoGenericService(object):
         # pylint:disable=protected-access
         return collection._handler(collection._parent, params, page)
 
-    def _get_paginated_results(self, limit, marker, collection):
+    def _get_paginated_results(self, limit: int | None, marker: str | None,
+                               collection: Any) -> tuple[Any, Any]:
         """
         If a Boto Paginator is available, use it. The results
         are converted back into BotoResources by directly accessing
@@ -177,7 +193,7 @@ class BotoGenericService(object):
         client = self.boto_conn.meta.client
         list_op = self._get_list_operation()
         paginator = client.get_paginator(list_op)
-        PaginationConfig = {}
+        PaginationConfig: dict[str, Any] = {}
         if limit:
             PaginationConfig = {'MaxItems': limit, 'PageSize': limit}
 
@@ -194,7 +210,8 @@ class BotoGenericService(object):
         resume_token = pages.resume_token
         return (resume_token, boto_objs)
 
-    def _make_query(self, collection, limit, marker):
+    def _make_query(self, collection: Any, limit: int | None,
+                    marker: str | None) -> tuple[str, Any, Any]:
         """
         Decide between server or client pagination,
         depending on the availability of a Boto Paginator.
@@ -213,7 +230,9 @@ class BotoGenericService(object):
                       " limit and page results.")
             return 'client', None, collection
 
-    def list(self, limit=None, marker=None, collection=None, **kwargs):
+    def list(self, limit: int | None = None, marker: str | None = None,
+             collection: Any = None,
+             **kwargs: Any) -> ResultList[CloudResource]:
         """
         List a set of resources.
 
@@ -244,8 +263,9 @@ class BotoGenericService(object):
             return ClientPagedResultList(self.provider, results,
                                          limit=limit, marker=marker)
 
-    def find(self, filters, limit=None, marker=None,
-             **kwargs):
+    def find(self, filters: dict[str, Any], limit: int | None = None,
+             marker: str | None = None,
+             **kwargs: Any) -> ResultList[CloudResource]:
         """
         Return a list of resources by filter.
 
@@ -261,7 +281,7 @@ class BotoGenericService(object):
             collection = collection.filter(**kwargs)
         return self.list(limit=limit, marker=marker, collection=collection)
 
-    def create(self, boto_method, **kwargs):
+    def create(self, boto_method: str, **kwargs: Any) -> Any:
         """
         Creates a resource
 
@@ -281,7 +301,7 @@ class BotoGenericService(object):
         else:
             return self.cb_resource(self.provider, result) if result else None
 
-    def delete(self, resource_id):
+    def delete(self, resource_id: str) -> None:
         """
         Deletes a resource by id
 
@@ -298,8 +318,9 @@ class BotoEC2Service(BotoGenericService):
     """
     Boto EC2 service implementation
     """
-    def __init__(self, provider, cb_resource,
-                 boto_collection_name):
+    def __init__(self, provider: AWSCloudProvider,
+                 cb_resource: Any,
+                 boto_collection_name: str) -> None:
         """
         :type provider: :class:`AWSCloudProvider`
         :param provider: CloudBridge AWS provider to use
@@ -320,8 +341,9 @@ class BotoS3Service(BotoGenericService):
     """
     Boto S3 service implementation.
     """
-    def __init__(self, provider, cb_resource,
-                 boto_collection_name):
+    def __init__(self, provider: AWSCloudProvider,
+                 cb_resource: Any,
+                 boto_collection_name: str) -> None:
         """
         :type provider: :class:`AWSCloudProvider`
         :param provider: CloudBridge AWS provider to use

+ 20 - 14
cloudbridge/providers/aws/provider.py

@@ -1,5 +1,6 @@
 """Provider implementation based on boto library for AWS-compatible clouds."""
 import logging
+from typing import Any
 
 import boto3
 
@@ -7,6 +8,11 @@ from botocore.client import Config
 
 from cloudbridge.base import BaseCloudProvider
 from cloudbridge.base.helpers import get_env
+from cloudbridge.interfaces.services import ComputeService
+from cloudbridge.interfaces.services import DnsService
+from cloudbridge.interfaces.services import NetworkingService
+from cloudbridge.interfaces.services import SecurityService
+from cloudbridge.interfaces.services import StorageService
 
 from .services import AWSComputeService
 from .services import AWSDnsService
@@ -20,9 +26,9 @@ log = logging.getLogger(__name__)
 
 class AWSCloudProvider(BaseCloudProvider):
     '''AWS cloud provider interface'''
-    PROVIDER_ID = 'aws'
+    PROVIDER_ID: str = 'aws'
 
-    def __init__(self, config):
+    def __init__(self, config: dict[str, Any]) -> None:
         super(AWSCloudProvider, self).__init__(config)
 
         # Initialize cloud connection fields
@@ -70,59 +76,59 @@ class AWSCloudProvider(BaseCloudProvider):
         self._dns = AWSDnsService(self)
 
     @property
-    def session(self):
+    def session(self) -> Any:
         '''Get a low-level session object or create one if needed'''
         if not self._session:
             if self.config.debug_mode:
-                boto3.set_stream_logger(level=log.DEBUG)
+                boto3.set_stream_logger(level=logging.DEBUG)
             self._session = boto3.session.Session(
                 region_name=self.region_name, **self.session_cfg)
         return self._session
 
     @property
-    def ec2_conn(self):
+    def ec2_conn(self) -> Any:
         if not self._ec2_conn:
             self._ec2_conn = self._connect_ec2()
         return self._ec2_conn
 
     @property
-    def s3_conn(self):
+    def s3_conn(self) -> Any:
         if not self._s3_conn:
             self._s3_conn = self._connect_s3()
         return self._s3_conn
 
     @property
-    def compute(self):
+    def compute(self) -> ComputeService:
         return self._compute
 
     @property
-    def networking(self):
+    def networking(self) -> NetworkingService:
         return self._networking
 
     @property
-    def security(self):
+    def security(self) -> SecurityService:
         return self._security
 
     @property
-    def storage(self):
+    def storage(self) -> StorageService:
         return self._storage
 
     @property
-    def dns(self):
+    def dns(self) -> DnsService:
         return self._dns
 
-    def _connect_ec2(self):
+    def _connect_ec2(self) -> Any:
         """
         Get a boto ec2 connection object.
         """
         return self._connect_ec2_region(region_name=self.region_name)
 
-    def _connect_ec2_region(self, region_name=None):
+    def _connect_ec2_region(self, region_name: str | None = None) -> Any:
         '''Get an EC2 resource object'''
         return self.session.resource(
             'ec2', region_name=region_name, **self.ec2_cfg)
 
-    def _connect_s3(self):
+    def _connect_s3(self) -> Any:
         '''Get an S3 resource object'''
         return self.session.resource(
             's3', region_name=self.region_name, **self.s3_cfg)

Разница между файлами не показана из-за своего большого размера
+ 226 - 185
cloudbridge/providers/aws/resources.py


Разница между файлами не показана из-за своего большого размера
+ 294 - 189
cloudbridge/providers/aws/services.py


+ 13 - 6
cloudbridge/providers/aws/subservices.py

@@ -6,41 +6,48 @@ from cloudbridge.base.subservices import BaseFloatingIPSubService
 from cloudbridge.base.subservices import BaseGatewaySubService
 from cloudbridge.base.subservices import BaseSubnetSubService
 from cloudbridge.base.subservices import BaseVMFirewallRuleSubService
+from cloudbridge.interfaces.provider import CloudProvider
+from cloudbridge.interfaces.resources import Bucket
+from cloudbridge.interfaces.resources import DnsZone
+from cloudbridge.interfaces.resources import Gateway
+from cloudbridge.interfaces.resources import Network
+from cloudbridge.interfaces.resources import VMFirewall
 
 log = logging.getLogger(__name__)
 
 
 class AWSBucketObjectSubService(BaseBucketObjectSubService):
 
-    def __init__(self, provider, bucket):
+    def __init__(self, provider: CloudProvider, bucket: Bucket) -> None:
         super(AWSBucketObjectSubService, self).__init__(provider, bucket)
 
 
 class AWSGatewaySubService(BaseGatewaySubService):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(AWSGatewaySubService, self).__init__(provider, network)
 
 
 class AWSVMFirewallRuleSubService(BaseVMFirewallRuleSubService):
 
-    def __init__(self, provider, firewall):
+    def __init__(self, provider: CloudProvider,
+                 firewall: VMFirewall) -> None:
         super(AWSVMFirewallRuleSubService, self).__init__(provider, firewall)
 
 
 class AWSFloatingIPSubService(BaseFloatingIPSubService):
 
-    def __init__(self, provider, gateway):
+    def __init__(self, provider: CloudProvider, gateway: Gateway) -> None:
         super(AWSFloatingIPSubService, self).__init__(provider, gateway)
 
 
 class AWSSubnetSubService(BaseSubnetSubService):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(AWSSubnetSubService, self).__init__(provider, network)
 
 
 class AWSDnsRecordSubService(BaseDnsRecordSubService):
 
-    def __init__(self, provider, dns_zone):
+    def __init__(self, provider: CloudProvider, dns_zone: DnsZone) -> None:
         super(AWSDnsRecordSubService, self).__init__(provider, dns_zone)

+ 168 - 135
cloudbridge/providers/azure/azure_client.py

@@ -1,11 +1,8 @@
+from __future__ import annotations
+
 import datetime
 import logging
-
-import tenacity
-from cloudbridge.interfaces.exceptions import (DuplicateResourceException,
-                                               InvalidLabelException,
-                                               ProviderConnectionException,
-                                               WaitStateException)
+from typing import Any
 
 from azure.core.credentials import AzureNamedKeyCredential
 from azure.core.exceptions import (ClientAuthenticationError,
@@ -33,6 +30,13 @@ from azure.mgmt.subscription import SubscriptionClient
 from azure.storage.blob import (BlobBlock, BlobSasPermissions,
                                 BlobServiceClient, generate_blob_sas)
 
+import tenacity
+
+from cloudbridge.interfaces.exceptions import (DuplicateResourceException,
+                                               InvalidLabelException,
+                                               ProviderConnectionException,
+                                               WaitStateException)
+
 from . import helpers as azure_helpers
 
 log = logging.getLogger(__name__)
@@ -167,27 +171,30 @@ class AzureClient(object):
     """
     Azure client is the wrapper on top of azure python sdk
     """
-    def __init__(self, config):
+    def __init__(self, config: dict[str, Any]) -> None:
         self._config = config
         self.subscription_id = str(config.get('azure_subscription_id'))
+        # config.get() yields Any | None; the typed azure.identity SDK wants
+        # str. These auth keys are always supplied by the provider, so the
+        # None branch is not reachable in practice.
         self._credentials = ClientSecretCredential(
-            tenant_id=config.get('azure_tenant'),
-            client_id=config.get('azure_client_id'),
-            client_secret=config.get('azure_secret')
+            tenant_id=config.get('azure_tenant'),  # type: ignore[arg-type]
+            client_id=config.get('azure_client_id'),  # type: ignore[arg-type]
+            client_secret=config.get('azure_secret')  # type: ignore[arg-type]
         )
 
         self._access_token = config.get('azure_access_token')
-        self._resource_client = None
-        self._storage_client = None
-        self._network_management_client = None
-        self._subscription_client = None
-        self._compute_client = None
-        self._dns_client = None
-        self._access_key_result = None
-        self._block_blob_service = None
-        self._table_service_client = None
-        self._public_key_table_client = None
-        self._storage_account = None
+        self._resource_client: Any = None
+        self._storage_client: Any = None
+        self._network_management_client: Any = None
+        self._subscription_client: Any = None
+        self._compute_client: Any = None
+        self._dns_client: Any = None
+        self._access_key_result: Any = None
+        self._block_blob_service: Any = None
+        self._table_service_client: Any = None
+        self._public_key_table_client: Any = None
+        self._storage_account: Any = None
 
         log.debug("azure subscription : %s", self.subscription_id)
 
@@ -204,7 +211,7 @@ class AzureClient(object):
         retry=tenacity.retry_if_exception_type(WaitStateException),
         reraise=True,
     )
-    def access_key_result(self):
+    def access_key_result(self) -> Any:
         if not self._access_key_result:
             storage_account = self.storage_account
 
@@ -225,27 +232,27 @@ class AzureClient(object):
         return self._access_key_result
 
     @property
-    def resource_group(self):
+    def resource_group(self) -> Any:
         return self._config.get('azure_resource_group')
 
     @property
-    def networking_resource_group(self):
+    def networking_resource_group(self) -> Any:
         return self._config.get('azure_networking_resource_group')
 
     @property
-    def storage_account(self):
+    def storage_account(self) -> Any:
         return self._config.get('azure_storage_account')
 
     @property
-    def region_name(self):
+    def region_name(self) -> Any:
         return self._config.get('azure_region_name')
 
     @property
-    def public_key_storage_table_name(self):
+    def public_key_storage_table_name(self) -> Any:
         return self._config.get('azure_public_key_storage_table_name')
 
     @property
-    def storage_client(self):
+    def storage_client(self) -> Any:
         if not self._storage_client:
             self._storage_client = \
                 StorageManagementClient(self._credentials,
@@ -253,13 +260,13 @@ class AzureClient(object):
         return self._storage_client
 
     @property
-    def subscription_client(self):
+    def subscription_client(self) -> Any:
         if not self._subscription_client:
             self._subscription_client = SubscriptionClient(self._credentials)
         return self._subscription_client
 
     @property
-    def resource_client(self):
+    def resource_client(self) -> Any:
         if not self._resource_client:
             self._resource_client = \
                 ResourceManagementClient(self._credentials,
@@ -267,7 +274,7 @@ class AzureClient(object):
         return self._resource_client
 
     @property
-    def compute_client(self):
+    def compute_client(self) -> Any:
         if not self._compute_client:
             self._compute_client = \
                 ComputeManagementClient(self._credentials,
@@ -275,21 +282,21 @@ class AzureClient(object):
         return self._compute_client
 
     @property
-    def network_management_client(self):
+    def network_management_client(self) -> Any:
         if not self._network_management_client:
             self._network_management_client = NetworkManagementClient(
                 self._credentials, self.subscription_id)
         return self._network_management_client
 
     @property
-    def dns_client(self):
+    def dns_client(self) -> Any:
         if not self._dns_client:
             self._dns_client = DnsManagementClient(
                 self._credentials, self.subscription_id)
         return self._dns_client
 
     @property
-    def blob_service(self):
+    def blob_service(self) -> Any:
         self._get_or_create_storage_account()
         if not self._block_blob_service:
             if self._access_token:
@@ -303,7 +310,7 @@ class AzureClient(object):
         return self._block_blob_service
 
     @property
-    def table_service(self):
+    def table_service(self) -> Any:
         self._get_or_create_storage_account()
         if not self._table_service_client:
             credential = AzureNamedKeyCredential(
@@ -318,21 +325,21 @@ class AzureClient(object):
                     table_name=self.public_key_storage_table_name)
         return self._public_key_table_client
 
-    def blob_client(self, container_name, blob_name):
+    def blob_client(self, container_name: str, blob_name: str) -> Any:
         return self.blob_service.get_blob_client(container=container_name, blob=blob_name)
 
-    def get_resource_group(self, name):
+    def get_resource_group(self, name: str) -> Any:
         return self.resource_client.resource_groups.get(name)
 
-    def create_resource_group(self, name, parameters):
+    def create_resource_group(self, name: str, parameters: dict[str, Any]) -> Any:
         return self.resource_client.resource_groups. \
             create_or_update(name, ResourceGroup(**parameters))
 
-    def get_storage_account(self, storage_account):
+    def get_storage_account(self, storage_account: str) -> Any:
         return self.storage_client.storage_accounts. \
             get_properties(self.resource_group, storage_account)
 
-    def create_storage_account(self, name, params):
+    def create_storage_account(self, name: str, params: dict[str, Any]) -> Any:
         return self.storage_client.storage_accounts. \
             begin_create(self.resource_group, name.lower(),
                          StorageAccountCreateParameters(**params)).result()
@@ -342,7 +349,7 @@ class AzureClient(object):
     @tenacity.retry(stop=tenacity.stop.stop_after_attempt(2),
                     retry=tenacity.retry_if_exception_type(HttpResponseError),
                     reraise=True)
-    def _get_or_create_storage_account(self):
+    def _get_or_create_storage_account(self) -> Any:
         if self._storage_account:
             return self._storage_account
         else:
@@ -386,43 +393,44 @@ class AzureClient(object):
                                % exists_err
                     raise InvalidLabelException(mess)
 
-    def list_locations(self):
+    def list_locations(self) -> Any:
         return self.subscription_client.subscriptions. \
             list_locations(self.subscription_id)
 
-    def list_vm_firewall(self):
+    def list_vm_firewall(self) -> Any:
         return self.network_management_client.network_security_groups. \
             list(self.resource_group)
 
-    def create_vm_firewall(self, name, parameters):
+    def create_vm_firewall(self, name: str, parameters: dict[str, Any]) -> Any:
         return self.network_management_client.network_security_groups. \
             begin_create_or_update(
                 self.resource_group, name,
                 NetworkSecurityGroup(**parameters)).result()
 
-    def update_vm_firewall_tags(self, fw_id, tags):
+    def update_vm_firewall_tags(self, fw_id: str, tags: dict[str, str]) -> Any:
         url_params = azure_helpers.parse_url(VM_FIREWALL_RESOURCE_ID,
                                              fw_id)
         name = url_params.get(VM_FIREWALL_NAME, "")
         return self.network_management_client.network_security_groups. \
             update_tags(self.resource_group, name, TagsObject(tags=tags))
 
-    def get_vm_firewall(self, fw_id):
+    def get_vm_firewall(self, fw_id: str) -> Any:
         url_params = azure_helpers.parse_url(VM_FIREWALL_RESOURCE_ID,
                                              fw_id)
         fw_name = url_params.get(VM_FIREWALL_NAME, "")
         return self.network_management_client.network_security_groups. \
             get(self.resource_group, fw_name)
 
-    def delete_vm_firewall(self, fw_id):
+    def delete_vm_firewall(self, fw_id: str) -> None:
         url_params = azure_helpers.parse_url(VM_FIREWALL_RESOURCE_ID,
                                              fw_id)
         name = url_params.get(VM_FIREWALL_NAME, "")
         self.network_management_client \
             .network_security_groups.begin_delete(self.resource_group, name).wait()
 
-    def create_vm_firewall_rule(self, fw_id,
-                                rule_name, parameters):
+    def create_vm_firewall_rule(self, fw_id: str,
+                                rule_name: str,
+                                parameters: dict[str, Any] | SecurityRule) -> Any:
         url_params = azure_helpers.parse_url(VM_FIREWALL_RESOURCE_ID,
                                              fw_id)
         vm_firewall_name = url_params.get(VM_FIREWALL_NAME, "")
@@ -435,20 +443,22 @@ class AzureClient(object):
             begin_create_or_update(self.resource_group, vm_firewall_name,
                                    rule_name, rule).result()
 
-    def delete_vm_firewall_rule(self, fw_rule_id, vm_firewall):
+    def delete_vm_firewall_rule(self, fw_rule_id: str, vm_firewall: str) -> Any:
         url_params = azure_helpers.parse_url(VM_FIREWALL_RULE_RESOURCE_ID,
                                              fw_rule_id)
         name = url_params.get(VM_FIREWALL_RULE_NAME, "")
         return self.network_management_client.security_rules. \
             begin_delete(self.resource_group, vm_firewall, name).result()
 
-    def list_containers(self, prefix=None, limit=None, marker=None):
+    def list_containers(self, prefix: str | None = None,
+                        limit: int | None = None,
+                        marker: str | None = None) -> Any:
         results = self.blob_service.list_containers(name_starts_with=prefix,
                                                     results_per_page=limit,
                                                     marker=marker)
         return results
 
-    def create_container(self, container_name):
+    def create_container(self, container_name: str) -> Any:
         try:
             return self.blob_service.create_container(container_name)
         except ResourceExistsError:
@@ -459,40 +469,46 @@ class AzureClient(object):
                     "in Storage Accounts." % container_name
             raise DuplicateResourceException(msg)
 
-    def get_container(self, container_name):
+    def get_container(self, container_name: str) -> Any:
         return self.blob_service.get_container_client(container_name)
 
-    def delete_container(self, container_name):
+    def delete_container(self, container_name: str) -> None:
         self.blob_service.delete_container(container_name)
 
-    def list_blobs(self, container_name, prefix=None, include=None):
+    def list_blobs(self, container_name: str, prefix: str | None = None,
+                   include: Any = None) -> Any:
         container_client = self.get_container(container_name)
         return container_client.list_blobs(name_starts_with=prefix, include=include)
 
-    def upload_blob(self, container_name, blob_name, data, length=None,
-                    max_concurrency=1):
+    def upload_blob(self, container_name: str, blob_name: str, data: Any,
+                    length: int | None = None,
+                    max_concurrency: int = 1) -> None:
         blob_client = self.blob_client(container_name, blob_name)
         blob_client.upload_blob(data=data, length=length, overwrite=True,
                                 max_concurrency=max_concurrency)
 
-    def stage_block(self, container_name, blob_name, block_id, data):
+    def stage_block(self, container_name: str, blob_name: str,
+                    block_id: str, data: Any) -> None:
         blob_client = self.blob_client(container_name, blob_name)
         blob_client.stage_block(block_id, data)
 
-    def commit_block_list(self, container_name, blob_name, block_ids):
+    def commit_block_list(self, container_name: str, blob_name: str,
+                          block_ids: list[str]) -> None:
         blob_client = self.blob_client(container_name, blob_name)
         block_list = [BlobBlock(block_id=block_id) for block_id in block_ids]
         blob_client.commit_block_list(block_list)
 
-    def get_blob(self, container_name, blob_name):
+    def get_blob(self, container_name: str, blob_name: str) -> Any:
         blob_client = self.blob_client(container_name, blob_name)
         return blob_client.get_blob_properties(container_name, blob_name)
 
-    def delete_blob(self, container_name, blob_name, delete_snapshots="include"):
+    def delete_blob(self, container_name: str, blob_name: str,
+                    delete_snapshots: str = "include") -> None:
         blob_client = self.blob_client(container_name, blob_name)
         blob_client.delete_blob(delete_snapshots)
 
-    def get_blob_url(self, container_name, blob_name, expiry_time, writable):
+    def get_blob_url(self, container_name: Any, blob_name: str,
+                     expiry_time: int, writable: bool) -> str:
         now = datetime.datetime.utcnow()
         expiry = now + datetime.timedelta(
             seconds=expiry_time)
@@ -509,37 +525,37 @@ class AzureClient(object):
         url = f"https://{self.storage_account}.blob.core.windows.net/{container_name}/{blob_name}?{sas}"
         return url
 
-    def create_empty_disk(self, disk_name, params):
+    def create_empty_disk(self, disk_name: str, params: dict[str, Any]) -> Any:
         return self.compute_client.disks.begin_create_or_update(
             self.resource_group,
             disk_name,
             Disk(**params)
         ).result()
 
-    def create_snapshot_disk(self, disk_name, params):
+    def create_snapshot_disk(self, disk_name: str, params: dict[str, Any]) -> Any:
         return self.compute_client.disks.begin_create_or_update(
             self.resource_group,
             disk_name,
             Disk(**params)
         ).result()
 
-    def get_disk(self, disk_id):
+    def get_disk(self, disk_id: str) -> Any:
         url_params = azure_helpers.parse_url(VOLUME_RESOURCE_ID,
                                              disk_id)
         disk_name = url_params.get(VOLUME_NAME, "")
         return self.compute_client.disks.get(self.resource_group, disk_name)
 
-    def list_disks(self):
+    def list_disks(self) -> Any:
         return self.compute_client.disks. \
             list_by_resource_group(self.resource_group)
 
-    def delete_disk(self, disk_id):
+    def delete_disk(self, disk_id: str) -> None:
         url_params = azure_helpers.parse_url(VOLUME_RESOURCE_ID,
                                              disk_id)
         disk_name = url_params.get(VOLUME_NAME, "")
         self.compute_client.disks.begin_delete(self.resource_group, disk_name).wait()
 
-    def update_disk_tags(self, disk_id, tags):
+    def update_disk_tags(self, disk_id: str, tags: dict[str, str]) -> Any:
         url_params = azure_helpers.parse_url(VOLUME_RESOURCE_ID,
                                              disk_id)
         disk_name = url_params.get(VOLUME_NAME, "")
@@ -549,22 +565,25 @@ class AzureClient(object):
             DiskUpdate(tags=tags)
         ).wait()
 
-    def list_snapshots(self):
+    def list_snapshots(self) -> Any:
         return self.compute_client.snapshots. \
             list_by_resource_group(self.resource_group)
 
-    def get_snapshot(self, snapshot_id):
+    def get_snapshot(self, snapshot_id: str) -> Any:
         url_params = azure_helpers.parse_url(SNAPSHOT_RESOURCE_ID,
                                              snapshot_id)
         snapshot_name = url_params.get(SNAPSHOT_NAME, "")
         return self.compute_client.snapshots.get(self.resource_group,
                                                  snapshot_name)
 
-    def create_snapshot(self, snapshot_name, volume, tags):
+    def create_snapshot(self, snapshot_name: str, volume: Any,
+                        tags: dict[str, str] | None) -> Any:
+        # The typed azure.mgmt.compute stub's keyword overload omits the
+        # top-level creation_data kwarg that the runtime model accepts.
         snapshot = self.compute_client.snapshots.begin_create_or_update(
             self.resource_group,
             snapshot_name,
-            Snapshot(
+            Snapshot(  # type: ignore[call-overload]
                 location=volume.location,
                 creation_data=CreationData(
                     create_option='Copy',
@@ -577,14 +596,15 @@ class AzureClient(object):
         self.update_snapshot_tags(snapshot.id, tags)
         return snapshot
 
-    def delete_snapshot(self, snapshot_id):
+    def delete_snapshot(self, snapshot_id: str) -> None:
         url_params = azure_helpers.parse_url(SNAPSHOT_RESOURCE_ID,
                                              snapshot_id)
         snapshot_name = url_params.get(SNAPSHOT_NAME, "")
         self.compute_client.snapshots.begin_delete(self.resource_group,
                                                    snapshot_name).wait()
 
-    def update_snapshot_tags(self, snapshot_id, tags):
+    def update_snapshot_tags(self, snapshot_id: str,
+                             tags: dict[str, str] | None) -> Any:
         url_params = azure_helpers.parse_url(SNAPSHOT_RESOURCE_ID,
                                              snapshot_id)
         snapshot_name = url_params.get(SNAPSHOT_NAME, "")
@@ -594,33 +614,33 @@ class AzureClient(object):
             SnapshotUpdate(tags=tags)
         ).wait()
 
-    def is_gallery_image(self, image_id):
+    def is_gallery_image(self, image_id: str) -> bool:
         url_params = azure_helpers.parse_url(IMAGE_RESOURCE_ID,
                                              image_id)
         # If it is a gallery image, it will always have an offer
         return 'offer' in url_params
 
-    def create_image(self, name, params):
+    def create_image(self, name: str, params: dict[str, Any]) -> Any:
         return self.compute_client.images. \
             begin_create_or_update(
                 self.resource_group, name, Image(**params)).result()
 
-    def delete_image(self, image_id):
+    def delete_image(self, image_id: str) -> None:
         url_params = azure_helpers.parse_url(IMAGE_RESOURCE_ID,
                                              image_id)
         if not self.is_gallery_image(image_id):
             name = url_params.get(IMAGE_NAME, "")
             self.compute_client.images.begin_delete(self.resource_group, name).wait()
 
-    def list_images(self):
+    def list_images(self) -> list[Any]:
         azure_images = list(self.compute_client.images.
                             list_by_resource_group(self.resource_group))
         return azure_images
 
-    def list_gallery_refs(self):
+    def list_gallery_refs(self) -> list[Any]:
         return gallery_image_references
 
-    def get_image(self, image_id):
+    def get_image(self, image_id: str) -> Any:
         url_params = azure_helpers.parse_url(IMAGE_RESOURCE_ID,
                                              image_id)
         if self.is_gallery_image(image_id):
@@ -632,7 +652,7 @@ class AzureClient(object):
             name = url_params.get(IMAGE_NAME, "")
             return self.compute_client.images.get(self.resource_group, name)
 
-    def update_image_tags(self, image_id, tags):
+    def update_image_tags(self, image_id: str, tags: dict[str, str]) -> Any:
         url_params = azure_helpers.parse_url(IMAGE_RESOURCE_ID,
                                              image_id)
         if self.is_gallery_image(image_id):
@@ -643,54 +663,55 @@ class AzureClient(object):
                 self.resource_group, name,
                 ImageUpdate(tags=tags)).result()
 
-    def list_vm_types(self):
+    def list_vm_types(self) -> Any:
         return self.compute_client.virtual_machine_sizes. \
             list(self.region_name)
 
-    def list_networks(self):
+    def list_networks(self) -> Any:
         return self.network_management_client.virtual_networks.list(
             self.networking_resource_group)
 
-    def get_network(self, network_id):
+    def get_network(self, network_id: str) -> Any:
         url_params = azure_helpers.parse_url(NETWORK_RESOURCE_ID,
                                              network_id)
         network_name = url_params.get(NETWORK_NAME, "")
         return self.network_management_client.virtual_networks.get(
             self.networking_resource_group, network_name)
 
-    def create_network(self, name, params):
+    def create_network(self, name: str, params: dict[str, Any]) -> Any:
         return self.network_management_client.virtual_networks. \
             begin_create_or_update(
                 self.networking_resource_group, name,
                 parameters=VirtualNetwork(**params)).result()
 
-    def delete_network(self, network_id):
+    def delete_network(self, network_id: str) -> Any:
         url_params = azure_helpers.parse_url(NETWORK_RESOURCE_ID, network_id)
         network_name = url_params.get(NETWORK_NAME, "")
         return self.network_management_client.virtual_networks. \
             begin_delete(self.networking_resource_group, network_name).wait()
 
-    def update_network_tags(self, network_id, tags):
+    def update_network_tags(self, network_id: str,
+                            tags: dict[str, str]) -> Any:
         url_params = azure_helpers.parse_url(NETWORK_RESOURCE_ID, network_id)
         network_name = url_params.get(NETWORK_NAME, "")
         return self.network_management_client.virtual_networks. \
             update_tags(self.networking_resource_group, network_name,
                         TagsObject(tags=tags))
 
-    def get_network_id_for_subnet(self, subnet_id):
+    def get_network_id_for_subnet(self, subnet_id: str) -> str:
         url_params = azure_helpers.parse_url(SUBNET_RESOURCE_ID, subnet_id)
         network_id = NETWORK_RESOURCE_ID[0]
         for key, val in url_params.items():
             network_id = network_id.replace("{" + key + "}", val)
         return network_id
 
-    def list_subnets(self, network_id):
+    def list_subnets(self, network_id: str) -> Any:
         url_params = azure_helpers.parse_url(NETWORK_RESOURCE_ID, network_id)
         network_name = url_params.get(NETWORK_NAME, "")
         return self.network_management_client.subnets. \
             list(self.networking_resource_group, network_name)
 
-    def get_subnet(self, subnet_id):
+    def get_subnet(self, subnet_id: str) -> Any:
         url_params = azure_helpers.parse_url(SUBNET_RESOURCE_ID,
                                              subnet_id)
         network_name = url_params.get(NETWORK_NAME, "")
@@ -698,7 +719,8 @@ class AzureClient(object):
         return self.network_management_client.subnets. \
             get(self.networking_resource_group, network_name, subnet_name)
 
-    def create_subnet(self, network_id, subnet_name, params):
+    def create_subnet(self, network_id: str, subnet_name: str,
+                      params: dict[str, Any]) -> Any:
         url_params = azure_helpers.parse_url(NETWORK_RESOURCE_ID, network_id)
         network_name = url_params.get(NETWORK_NAME, "")
         result_create = self.network_management_client \
@@ -712,7 +734,8 @@ class AzureClient(object):
 
         return subnet_info
 
-    def __if_subnet_in_use(e):
+    @staticmethod
+    def __if_subnet_in_use(e: BaseException) -> bool:
         # return True if the CloudError exception is due to subnet being in use
         if isinstance(e, HttpResponseError):
             if "InUseSubnetCannotBeDeleted" in e.message:
@@ -723,7 +746,7 @@ class AzureClient(object):
                     retry=tenacity.retry_if_exception(__if_subnet_in_use),
                     wait=tenacity.wait.wait_fixed(5),
                     reraise=True)
-    def delete_subnet(self, subnet_id):
+    def delete_subnet(self, subnet_id: str) -> None:
         url_params = azure_helpers.parse_url(SUBNET_RESOURCE_ID,
                                              subnet_id)
         network_name = url_params.get(NETWORK_NAME, "")
@@ -741,21 +764,22 @@ class AzureClient(object):
             log.exception(cloud_error.message)
             raise cloud_error
 
-    def create_floating_ip(self, public_ip_name, public_ip_parameters):
+    def create_floating_ip(self, public_ip_name: str,
+                           public_ip_parameters: dict[str, Any]) -> Any:
         return self.network_management_client.public_ip_addresses. \
             begin_create_or_update(
                 self.networking_resource_group,
                 public_ip_name,
                 PublicIPAddress(**public_ip_parameters)).result()
 
-    def get_floating_ip(self, public_ip_id):
+    def get_floating_ip(self, public_ip_id: str) -> Any:
         url_params = azure_helpers.parse_url(PUBLIC_IP_RESOURCE_ID,
                                              public_ip_id)
         public_ip_name = url_params.get(PUBLIC_IP_NAME, "")
         return self.network_management_client. \
             public_ip_addresses.get(self.networking_resource_group, public_ip_name)
 
-    def delete_floating_ip(self, public_ip_id):
+    def delete_floating_ip(self, public_ip_id: str) -> None:
         url_params = azure_helpers.parse_url(PUBLIC_IP_RESOURCE_ID,
                                              public_ip_id)
         public_ip_name = url_params.get(PUBLIC_IP_NAME, "")
@@ -763,7 +787,7 @@ class AzureClient(object):
             public_ip_addresses.begin_delete(self.networking_resource_group,
                                              public_ip_name).wait()
 
-    def update_fip_tags(self, fip_id, tags):
+    def update_fip_tags(self, fip_id: str, tags: dict[str, str]) -> None:
         url_params = azure_helpers.parse_url(PUBLIC_IP_RESOURCE_ID,
                                              fip_id)
         fip_name = url_params.get(PUBLIC_IP_NAME, "")
@@ -771,37 +795,37 @@ class AzureClient(object):
             update_tags(self.networking_resource_group, fip_name,
                         TagsObject(tags=tags))
 
-    def list_floating_ips(self):
+    def list_floating_ips(self) -> Any:
         return self.network_management_client.public_ip_addresses.list(
             self.networking_resource_group)
 
-    def list_vm(self):
+    def list_vm(self) -> Any:
         return self.compute_client.virtual_machines.list(
             self.resource_group
         )
 
-    def restart_vm(self, vm_id):
+    def restart_vm(self, vm_id: str) -> Any:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
         return self.compute_client.virtual_machines.begin_restart(
             self.resource_group, vm_name).wait()
 
-    def stop_vm(self, vm_id):
+    def stop_vm(self, vm_id: str) -> None:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
         self.compute_client.virtual_machines. \
             begin_power_off(self.resource_group, vm_name).wait()
 
-    def delete_vm(self, vm_id):
+    def delete_vm(self, vm_id: str) -> Any:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
         return self.compute_client.virtual_machines.begin_delete(
             self.resource_group, vm_name).wait()
 
-    def get_vm(self, vm_id):
+    def get_vm(self, vm_id: str) -> Any:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
@@ -811,13 +835,13 @@ class AzureClient(object):
             expand='instanceView'
         )
 
-    def create_vm(self, vm_name, params):
+    def create_vm(self, vm_name: str, params: dict[str, Any]) -> Any:
         return self.compute_client.virtual_machines. \
             begin_create_or_update(
                 self.resource_group, vm_name,
                 VirtualMachine(**params)).result()
 
-    def update_vm(self, vm_id, params):
+    def update_vm(self, vm_id: str, params: dict[str, Any]) -> Any:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
@@ -826,28 +850,28 @@ class AzureClient(object):
                 self.resource_group, vm_name,
                 VirtualMachine(**params)).wait()
 
-    def deallocate_vm(self, vm_id):
+    def deallocate_vm(self, vm_id: str) -> None:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
         self.compute_client. \
             virtual_machines.begin_deallocate(self.resource_group, vm_name).wait()
 
-    def generalize_vm(self, vm_id):
+    def generalize_vm(self, vm_id: str) -> None:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
         self.compute_client.virtual_machines. \
             generalize(self.resource_group, vm_name)
 
-    def start_vm(self, vm_id):
+    def start_vm(self, vm_id: str) -> None:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
         self.compute_client.virtual_machines. \
             begin_start(self.resource_group, vm_name).wait()
 
-    def update_vm_tags(self, vm_id, tags):
+    def update_vm_tags(self, vm_id: str, tags: dict[str, str]) -> None:
         url_params = azure_helpers.parse_url(VM_RESOURCE_ID,
                                              vm_id)
         vm_name = url_params.get(VM_NAME, "")
@@ -855,21 +879,22 @@ class AzureClient(object):
             self.resource_group, vm_name,
             VirtualMachineUpdate(tags=tags)).result()
 
-    def delete_nic(self, nic_id):
+    def delete_nic(self, nic_id: str) -> None:
         nic_params = azure_helpers.\
             parse_url(NETWORK_INTERFACE_RESOURCE_ID, nic_id)
         nic_name = nic_params.get(NETWORK_INTERFACE_NAME, "")
         self.network_management_client. \
             network_interfaces.begin_delete(self.resource_group, nic_name).wait()
 
-    def get_nic(self, nic_id):
+    def get_nic(self, nic_id: str) -> Any:
         nic_params = azure_helpers.\
             parse_url(NETWORK_INTERFACE_RESOURCE_ID, nic_id)
         nic_name = nic_params.get(NETWORK_INTERFACE_NAME, "")
         return self.network_management_client. \
             network_interfaces.get(self.resource_group, nic_name)
 
-    def update_nic(self, nic_id, params):
+    def update_nic(self, nic_id: str,
+                   params: dict[str, Any] | NetworkInterface) -> Any:
         nic_params = azure_helpers.\
             parse_url(NETWORK_INTERFACE_RESOURCE_ID, nic_id)
         nic_name = nic_params.get(NETWORK_INTERFACE_NAME, "")
@@ -886,7 +911,7 @@ class AzureClient(object):
         nic_info = async_nic_creation.result()
         return nic_info
 
-    def create_nic(self, nic_name, params):
+    def create_nic(self, nic_name: str, params: dict[str, Any]) -> Any:
         return self.network_management_client. \
             network_interfaces.begin_create_or_update(
                 self.resource_group,
@@ -894,21 +919,23 @@ class AzureClient(object):
                 NetworkInterface(**params)
             ).result()
 
-    def create_public_key(self, entity):
+    def create_public_key(self, entity: dict[str, Any]) -> Any:
         return self.table_service.upsert_entity(entity)
 
-    def get_public_key(self, name):
+    def get_public_key(self, name: str) -> Any:
         entities = list(self.table_service.query_entities(
             query_filter="Name eq '{0}'".format(name),
             results_per_page=1))
         return entities[0] if entities else None
 
-    def delete_public_key(self, entity):
+    def delete_public_key(self, entity: dict[str, Any]) -> None:
         self.table_service.delete_entity(
             partition_key=entity['PartitionKey'],
             row_key=entity['RowKey'])
 
-    def list_public_keys(self, partition_key, limit=None, marker=None):
+    def list_public_keys(self, partition_key: str, limit: int | None = None,
+                         marker: str | None = None
+                         ) -> tuple[list[Any], str | None]:
         pager = self.table_service.query_entities(
             query_filter="PartitionKey eq '{0}'".format(partition_key),
             results_per_page=limit).by_page(continuation_token=marker)
@@ -919,12 +946,13 @@ class AzureClient(object):
         items = list(page)
         return (items, pager.continuation_token)
 
-    def delete_route_table(self, route_table_name):
+    def delete_route_table(self, route_table_name: str) -> None:
         self.network_management_client. \
             route_tables.begin_delete(self.resource_group,
                                       route_table_name).wait()
 
-    def attach_subnet_to_route_table(self, subnet_id, route_table_id):
+    def attach_subnet_to_route_table(self, subnet_id: str,
+                                     route_table_id: str) -> Any:
         url_params = azure_helpers.parse_url(SUBNET_RESOURCE_ID,
                                              subnet_id)
         network_name = url_params.get(NETWORK_NAME, "")
@@ -943,12 +971,13 @@ class AzureClient(object):
                  self.resource_group,
                  network_name,
                  subnet_name,
-                 subnet_info)  # type: ignore
+                 subnet_info)
             subnet_info = result_create.result()
 
         return subnet_info
 
-    def detach_subnet_to_route_table(self, subnet_id, route_table_id):
+    def detach_subnet_to_route_table(self, subnet_id: str,
+                                     route_table_id: str) -> Any:
         url_params = azure_helpers.parse_url(SUBNET_RESOURCE_ID,
                                              subnet_id)
         network_name = url_params.get(NETWORK_NAME, "")
@@ -968,64 +997,68 @@ class AzureClient(object):
                  self.resource_group,
                  network_name,
                  subnet_name,
-                 subnet_info)  # type: ignore
+                 subnet_info)
             subnet_info = result_create.result()
 
         return subnet_info
 
-    def list_route_tables(self):
+    def list_route_tables(self) -> Any:
         return self.network_management_client. \
             route_tables.list(self.resource_group)
 
-    def get_route_table(self, router_id):
+    def get_route_table(self, router_id: str) -> Any:
         url_params = azure_helpers.parse_url(ROUTER_RESOURCE_ID,
                                              router_id)
         router_name = url_params.get(ROUTER_NAME, "")
         return self.network_management_client. \
             route_tables.get(self.resource_group, router_name)
 
-    def create_route_table(self, route_table_name, params):
+    def create_route_table(self, route_table_name: str,
+                           params: dict[str, Any]) -> Any:
         return self.network_management_client. \
             route_tables.begin_create_or_update(
              self.resource_group,
              route_table_name, RouteTable(**params)).result()
 
-    def update_route_table_tags(self, route_table_name, tags):
+    def update_route_table_tags(self, route_table_name: str,
+                                tags: dict[str, str]) -> None:
         self.network_management_client.route_tables.update_tags(
             self.resource_group, route_table_name,
             TagsObject(tags=tags))
 
     # DNS operations
-    def get_dns_zone(self, zone_name):
+    def get_dns_zone(self, zone_name: str) -> Any:
         return self.dns_client.zones.get(self.resource_group, zone_name)
 
-    def list_dns_zones(self):
+    def list_dns_zones(self) -> list[Any]:
         return list(self.dns_client.zones.list_by_resource_group(
             self.resource_group))
 
-    def create_dns_zone(self, zone_name, params):
+    def create_dns_zone(self, zone_name: str, params: dict[str, Any]) -> Any:
         return self.dns_client.zones.create_or_update(
             self.resource_group, zone_name, Zone(**params))
 
-    def delete_dns_zone(self, zone_name):
+    def delete_dns_zone(self, zone_name: str) -> None:
         self.dns_client.zones.begin_delete(
             self.resource_group, zone_name).wait()
 
-    def get_dns_record(self, zone_name, relative_record_name, record_type):
+    def get_dns_record(self, zone_name: str, relative_record_name: str,
+                       record_type: str) -> Any:
         return self.dns_client.record_sets.get(
             self.resource_group, zone_name, relative_record_name, record_type)
 
-    def list_dns_records(self, zone_name):
+    def list_dns_records(self, zone_name: str) -> list[Any]:
         return list(self.dns_client.record_sets.list_all_by_dns_zone(
             self.resource_group, zone_name))
 
-    def create_dns_record(self, zone_name, relative_record_name,
-                          record_type, params):
+    def create_dns_record(self, zone_name: str, relative_record_name: str,
+                          record_type: str, params: dict[str, Any]) -> Any:
         from azure.mgmt.dns.models import RecordSet
         return self.dns_client.record_sets.create_or_update(
             self.resource_group, zone_name, relative_record_name,
             record_type, RecordSet(**params))
 
-    def delete_dns_record(self, zone_name, relative_record_name, record_type):
+    def delete_dns_record(self, zone_name: str, relative_record_name: str,
+                          record_type: str) -> None:
         self.dns_client.record_sets.delete(
             self.resource_group, zone_name, relative_record_name, record_type)

+ 7 - 6
cloudbridge/providers/azure/helpers.py

@@ -1,4 +1,5 @@
 import re
+from typing import Any
 
 from cloudbridge.interfaces.exceptions import InvalidValueException
 
@@ -6,7 +7,7 @@ from cloudbridge.interfaces.exceptions import InvalidValueException
 _RG_NAME_RE = re.compile(r'(/resourceGroups/)([^/]+)', re.IGNORECASE)
 
 
-def normalize_rg_case(azure_id):
+def normalize_rg_case(azure_id: str | None) -> str | None:
     # Microsoft.Compute/images list_by_resource_group returns the RG segment
     # in a case that can differ from what create/get echo back (we've seen
     # uppercase from list, lowercase from create/get for the same RG).
@@ -38,7 +39,7 @@ def normalize_rg_case(azure_id):
 #         return list_items
 
 
-def parse_url(template_urls, original_url):
+def parse_url(template_urls: list[str], original_url: str) -> dict[str, str]:
     """
     In Azure all the resource IDs are returned as URIs.
     ex: '/subscriptions/{subscriptionId}/resourceGroups/' \
@@ -52,7 +53,7 @@ def parse_url(template_urls, original_url):
     https://docs.microsoft.com/en-us/azure/virtual-machines/linux/cli-ps-findimage
     """
     if not original_url:
-        raise InvalidValueException(template_urls, original_url)
+        raise InvalidValueException(str(template_urls), original_url)
     original_url_parts = original_url.split('/')
     if len(original_url_parts) == 1:
         original_url_parts = original_url.split(':')
@@ -63,15 +64,15 @@ def parse_url(template_urls, original_url):
         if len(template_url_parts) == len(original_url_parts):
             break
     if len(template_url_parts) != len(original_url_parts):
-        raise InvalidValueException(template_urls, original_url)
-    resource_param = {}
+        raise InvalidValueException(str(template_urls), original_url)
+    resource_param: dict[str, str] = {}
     for key, value in zip(template_url_parts, original_url_parts):
         if key.startswith('{') and key.endswith('}'):
             resource_param.update({key[1:-1]: value})
     return resource_param
 
 
-def generate_urn(gallery_image):
+def generate_urn(gallery_image: Any) -> str:
     """
     This function takes an azure gallery image and outputs a corresponding URN
     :param gallery_image: a GalleryImageReference object

+ 23 - 13
cloudbridge/providers/azure/provider.py

@@ -1,5 +1,6 @@
 import logging
 import uuid
+from typing import Any
 
 from azure.core.exceptions import HttpResponseError
 from azure.core.exceptions import ResourceNotFoundError
@@ -12,6 +13,11 @@ import cloudbridge
 from cloudbridge.base import BaseCloudProvider
 from cloudbridge.base.helpers import get_env
 from cloudbridge.interfaces.exceptions import ProviderConnectionException
+from cloudbridge.interfaces.services import ComputeService
+from cloudbridge.interfaces.services import DnsService
+from cloudbridge.interfaces.services import NetworkingService
+from cloudbridge.interfaces.services import SecurityService
+from cloudbridge.interfaces.services import StorageService
 from cloudbridge.providers.azure.azure_client import AzureClient
 
 from .services import AzureComputeService
@@ -24,9 +30,9 @@ log = logging.getLogger(__name__)
 
 
 class AzureCloudProvider(BaseCloudProvider):
-    PROVIDER_ID = 'azure'
+    PROVIDER_ID: str = 'azure'
 
-    def __init__(self, config):
+    def __init__(self, config: dict[str, Any]) -> None:
         super(AzureCloudProvider, self).__init__(config)
 
         # mandatory config values
@@ -74,7 +80,7 @@ class AzureCloudProvider(BaseCloudProvider):
             'azure_public_key_storage_table_name', get_env(
                 'AZURE_PUBLIC_KEY_STORAGE_TABLE_NAME', 'cbcerts'))
 
-        self._azure_client = None
+        self._azure_client: AzureClient | None = None
 
         self._security = AzureSecurityService(self)
         self._storage = AzureStorageService(self)
@@ -82,7 +88,7 @@ class AzureCloudProvider(BaseCloudProvider):
         self._networking = AzureNetworkingService(self)
         self._dns = AzureDnsService(self)
 
-    def __get_deprecated_username(self, default):
+    def __get_deprecated_username(self, default: str) -> str:
         username = self._get_config_value(
             'azure_vm_default_user_name', get_env(
                 'AZURE_VM_DEFAULT_USER_NAME', None))
@@ -96,31 +102,31 @@ class AzureCloudProvider(BaseCloudProvider):
                 current_version=cloudbridge.__version__,
                 details='AZURE_VM_DEFAULT_USER_NAME was deprecated in favor '
                         'of AZURE_VM_DEFAULT_USERNAME')
-    def __wrap_deprecated_username(self, username):
+    def __wrap_deprecated_username(self, username: str) -> str:
         return username
 
     @property
-    def compute(self):
+    def compute(self) -> ComputeService:
         return self._compute
 
     @property
-    def networking(self):
+    def networking(self) -> NetworkingService:
         return self._networking
 
     @property
-    def security(self):
+    def security(self) -> SecurityService:
         return self._security
 
     @property
-    def storage(self):
+    def storage(self) -> StorageService:
         return self._storage
 
     @property
-    def dns(self):
+    def dns(self) -> DnsService:
         return self._dns
 
     @property
-    def azure_client(self):
+    def azure_client(self) -> Any:
         if not self._azure_client:
 
             # create a dict with both optional and mandatory configuration
@@ -148,12 +154,16 @@ class AzureCloudProvider(BaseCloudProvider):
     @tenacity.retry(stop=tenacity.stop_after_attempt(2),
                     retry=tenacity.retry_if_exception_type(HttpResponseError),
                     reraise=True)
-    def _initialize(self):
+    def _initialize(self) -> None:
         """
         Verifying that resource group and storage account exists
         if not create one with the name provided in the
         configuration
         """
+        # ``_initialize`` is only ever invoked by the ``azure_client`` property
+        # immediately after assigning ``self._azure_client``, so the client is
+        # guaranteed to be set here.
+        assert self._azure_client is not None
         try:
             self._azure_client.get_resource_group(self.resource_group)
 
@@ -164,7 +174,7 @@ class AzureCloudProvider(BaseCloudProvider):
                     create_resource_group(self.resource_group,
                                           resource_group_params)
             except HttpResponseError as cloud_error2:  # pragma: no cover
-                if getattr(cloud_error2, 'error', None) and \
+                if cloud_error2.error and \
                         cloud_error2.error.code == "AuthorizationFailed":
                     mess = 'The following error was returned by Azure:\n' \
                            '%s\n\nThis is likely because the Role' \

Разница между файлами не показана из-за своего большого размера
+ 263 - 200
cloudbridge/providers/azure/resources.py


Разница между файлами не показана из-за своего большого размера
+ 369 - 223
cloudbridge/providers/azure/services.py


+ 13 - 6
cloudbridge/providers/azure/subservices.py

@@ -6,40 +6,47 @@ from cloudbridge.base.subservices import BaseFloatingIPSubService
 from cloudbridge.base.subservices import BaseGatewaySubService
 from cloudbridge.base.subservices import BaseSubnetSubService
 from cloudbridge.base.subservices import BaseVMFirewallRuleSubService
+from cloudbridge.interfaces.provider import CloudProvider
+from cloudbridge.interfaces.resources import Bucket
+from cloudbridge.interfaces.resources import DnsZone
+from cloudbridge.interfaces.resources import Gateway
+from cloudbridge.interfaces.resources import Network
+from cloudbridge.interfaces.resources import VMFirewall
 
 log = logging.getLogger(__name__)
 
 
 class AzureBucketObjectSubService(BaseBucketObjectSubService):
 
-    def __init__(self, provider, bucket):
+    def __init__(self, provider: CloudProvider, bucket: Bucket) -> None:
         super(AzureBucketObjectSubService, self).__init__(provider, bucket)
 
 
 class AzureGatewaySubService(BaseGatewaySubService):
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(AzureGatewaySubService, self).__init__(provider, network)
 
 
 class AzureVMFirewallRuleSubService(BaseVMFirewallRuleSubService):
 
-    def __init__(self, provider, firewall):
+    def __init__(self, provider: CloudProvider,
+                 firewall: VMFirewall) -> None:
         super(AzureVMFirewallRuleSubService, self).__init__(provider, firewall)
 
 
 class AzureFloatingIPSubService(BaseFloatingIPSubService):
 
-    def __init__(self, provider, gateway):
+    def __init__(self, provider: CloudProvider, gateway: Gateway) -> None:
         super(AzureFloatingIPSubService, self).__init__(provider, gateway)
 
 
 class AzureSubnetSubService(BaseSubnetSubService):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(AzureSubnetSubService, self).__init__(provider, network)
 
 
 class AzureDnsRecordSubService(BaseDnsRecordSubService):
 
-    def __init__(self, provider, dns_zone):
+    def __init__(self, provider: CloudProvider, dns_zone: DnsZone) -> None:
         super(AzureDnsRecordSubService, self).__init__(provider, dns_zone)

+ 35 - 19
cloudbridge/providers/gcp/helpers.py

@@ -3,6 +3,10 @@ import collections
 import datetime
 import hashlib
 import re
+from typing import Any
+from typing import Callable
+from typing import Iterator
+from typing import TYPE_CHECKING
 from urllib.parse import quote
 
 from googleapiclient.errors import HttpError
@@ -11,12 +15,15 @@ import tenacity
 
 from cloudbridge.interfaces.exceptions import ProviderInternalException
 
+if TYPE_CHECKING:
+    from .provider import GCPCloudProvider
 
-def gcp_projects(provider):
+
+def gcp_projects(provider: "GCPCloudProvider") -> Any:
     return provider.gcp_compute.projects()
 
 
-def iter_all(resource, **kwargs):
+def iter_all(resource: Any, **kwargs: Any) -> Iterator[Any]:
     token = None
     while True:
         response = resource.list(pageToken=token, **kwargs).execute()
@@ -27,7 +34,7 @@ def iter_all(resource, **kwargs):
         token = response['nextPageToken']
 
 
-def get_common_metadata(provider):
+def get_common_metadata(provider: "GCPCloudProvider") -> Any:
     """
     Get a project's commonInstanceMetadata entry
     """
@@ -36,7 +43,7 @@ def get_common_metadata(provider):
     return metadata["commonInstanceMetadata"]
 
 
-def __if_fingerprint_differs(e):
+def __if_fingerprint_differs(e: BaseException) -> bool:
     # return True if the CloudError exception is due to subnet being in use
     if isinstance(e, HttpError):
         expected_message = 'Supplied fingerprint does not match current ' \
@@ -50,7 +57,8 @@ def __if_fingerprint_differs(e):
                 retry=tenacity.retry_if_exception(__if_fingerprint_differs),
                 wait=tenacity.wait_exponential(max=10),
                 reraise=True)
-def gcp_metadata_save_op(provider, callback):
+def gcp_metadata_save_op(provider: "GCPCloudProvider",
+                         callback: Callable[[Any], Any]) -> None:
     """
     Carries out a metadata save operation. In GCP, a fingerprint based
     locking mechanism is used to prevent lost updates. A new fingerprint
@@ -59,7 +67,7 @@ def gcp_metadata_save_op(provider, callback):
     metadata, and saves the metadata using the original fingerprint
     immediately afterwards, ensuring that update conflicts can be detected.
     """
-    def _save_common_metadata(provider):
+    def _save_common_metadata(provider: "GCPCloudProvider") -> None:
         # get the latest metadata (so we get the latest fingerprint)
         metadata = get_common_metadata(provider)
         # allow callback to do processing on it
@@ -73,8 +81,9 @@ def gcp_metadata_save_op(provider, callback):
     _save_common_metadata(provider)
 
 
-def modify_or_add_metadata_item(provider, key, value):
-    def _update_metadata_key(metadata):
+def modify_or_add_metadata_item(provider: "GCPCloudProvider", key: str,
+                                value: str) -> None:
+    def _update_metadata_key(metadata: Any) -> None:
         entries = [item for item in metadata.get('items', [])
                    if item['key'] == key]
         if entries:
@@ -92,8 +101,9 @@ def modify_or_add_metadata_item(provider, key, value):
 # This function will raise an HttpError with message containing
 # "Metadata has duplicate key" if it's not unique, unlike the previous
 # method which either adds or updates the value corresponding to that key
-def add_metadata_item(provider, key, value):
-    def _add_metadata_key(metadata):
+def add_metadata_item(provider: "GCPCloudProvider", key: str,
+                      value: str) -> None:
+    def _add_metadata_key(metadata: Any) -> None:
         entry = {'key': key, 'value': value}
         entries = metadata.get('items', [])
         entries.append(entry)
@@ -104,7 +114,9 @@ def add_metadata_item(provider, key, value):
     gcp_metadata_save_op(provider, _add_metadata_key)
 
 
-def find_matching_metadata_items(provider, key_regex):
+def find_matching_metadata_items(provider: "GCPCloudProvider",
+                                 key_regex: str | re.Pattern[str]
+                                 ) -> list[Any]:
     metadata = get_common_metadata(provider)
     items = metadata.get('items', [])
     if not items:
@@ -113,7 +125,7 @@ def find_matching_metadata_items(provider, key_regex):
             if re.search(key_regex, item['key'])]
 
 
-def get_metadata_item_value(provider, key):
+def get_metadata_item_value(provider: "GCPCloudProvider", key: str) -> Any:
     metadata = get_common_metadata(provider)
     entries = [item['value'] for item in metadata.get('items', [])
                if item['key'] == key]
@@ -123,8 +135,8 @@ def get_metadata_item_value(provider, key):
         return None
 
 
-def remove_metadata_item(provider, key):
-    def _remove_metadata_by_key(metadata):
+def remove_metadata_item(provider: "GCPCloudProvider", key: str) -> bool:
+    def _remove_metadata_by_key(metadata: Any) -> bool | None:
         items = metadata.get('items', [])
         # No metadata to delete
         if not items:
@@ -144,12 +156,13 @@ def remove_metadata_item(provider, key):
 
             else:
                 metadata['items'] = entries
+                return None
 
     gcp_metadata_save_op(provider, _remove_metadata_by_key)
     return True
 
 
-def __if_label_fingerprint_differs(e):
+def __if_label_fingerprint_differs(e: BaseException) -> bool:
     # return True if the CloudError exception is due to subnet being in use
     if isinstance(e, HttpError):
         expected_message = 'Labels fingerprint either invalid or ' \
@@ -164,7 +177,8 @@ def __if_label_fingerprint_differs(e):
                     __if_label_fingerprint_differs),
                 wait=tenacity.wait_exponential(max=10),
                 reraise=True)
-def change_label(resource, key, value, res_att, request):
+def change_label(resource: Any, key: str, value: str, res_att: str,
+                 request: Any) -> None:
     resource.assert_valid_resource_label(value)
     labels = getattr(resource, res_att).get("labels", {})
     labels[key] = str(value)
@@ -188,9 +202,11 @@ def change_label(resource, key, value, res_att, request):
 
 
 # https://cloud.google.com/storage/docs/access-control/signing-urls-manually#python-sample
-def generate_signed_url(credentials, bucket_name, object_name,
-                        subresource=None, expiration=604800, http_method='GET',
-                        query_parameters=None, headers=None):
+def generate_signed_url(credentials: Any, bucket_name: str, object_name: str,
+                        subresource: str | None = None,
+                        expiration: int = 604800, http_method: str = 'GET',
+                        query_parameters: dict[str, Any] | None = None,
+                        headers: dict[str, str] | None = None) -> str:
 
     if expiration > 604800:
         # max allowed expiration time is 7 days

+ 52 - 41
cloudbridge/providers/gcp/provider.py

@@ -8,8 +8,13 @@ import os
 import re
 import time
 from string import Template
+from typing import Any
+from typing import Callable
 
 import google.auth
+from google.auth.credentials import with_scopes_if_required
+from google.oauth2.service_account import Credentials
+
 import google_auth_httplib2
 
 import googleapiclient
@@ -17,12 +22,13 @@ from googleapiclient import discovery
 
 import httplib2
 
-from google.auth.credentials import with_scopes_if_required
-
-from google.oauth2.service_account import Credentials
-
 from cloudbridge.base import BaseCloudProvider
 from cloudbridge.interfaces.exceptions import ProviderConnectionException
+from cloudbridge.interfaces.services import ComputeService
+from cloudbridge.interfaces.services import DnsService
+from cloudbridge.interfaces.services import NetworkingService
+from cloudbridge.interfaces.services import SecurityService
+from cloudbridge.interfaces.services import StorageService
 
 from .services import GCPComputeService
 from .services import GCPDnsService
@@ -37,12 +43,12 @@ CLOUD_SCOPES = ['https://www.googleapis.com/auth/cloud-platform']
 
 class GCPResourceUrl(object):
 
-    def __init__(self, resource, connection):
+    def __init__(self, resource: str, connection: Any) -> None:
         self._resource = resource
         self._connection = connection
-        self.parameters = {}
+        self.parameters: dict[str, Any] = {}
 
-    def get_resource(self):
+    def get_resource(self) -> Any:
         """
         The format of the returned resource is explained in details in
         https://cloud.google.com/compute/docs/reference/latest/ and
@@ -71,7 +77,7 @@ class GCPResourceUrl(object):
 
 class GCPResources(object):
 
-    def __init__(self, connection, **kwargs):
+    def __init__(self, connection: Any, **kwargs: Any) -> None:
         self._connection = connection
         self._parameter_defaults = kwargs
 
@@ -116,7 +122,7 @@ class GCPResources(object):
         self.RESOURCE_REGEX = re.compile(
             r"(https://.*\.googleapis\.com/{0})(.*)".format(
                 desc['servicePath']))
-        self._resources = {}
+        self._resources: dict[str, dict[str, Any]] = {}
 
         # We will not mutate self._desc; it's OK to use items() in Python 2.x.
         for resource, resource_desc in desc['resources'].items():
@@ -149,7 +155,7 @@ class GCPResources(object):
             self._resources[resource] = {'parameters': parameters,
                                          'pattern': re.compile(pattern)}
 
-    def parse_url(self, url):
+    def parse_url(self, url: str) -> "GCPResourceUrl | None":
         """
         Build a GCPResourceUrl from a resource's URL string. One can then call
         the get() method on the returned object to fetch resource details from
@@ -180,8 +186,11 @@ class GCPResources(object):
             for index, parameter in enumerate(desc['parameters']):
                 out.parameters[parameter] = m.group(index + 1)
             return out
+        return None
 
-    def get_resource_url_with_default(self, resource, url_or_name, **kwargs):
+    def get_resource_url_with_default(
+            self, resource: str, url_or_name: str,
+            **kwargs: Any) -> "GCPResourceUrl | None":
         """
         Build a GCPResourceUrl from a service's name and resource url or name.
         If the url_or_name is a valid GCP resource URL, then we build the
@@ -210,7 +219,7 @@ class GCPResources(object):
 class GCPCloudProvider(BaseCloudProvider):
     PROVIDER_ID = 'gcp'
 
-    def __init__(self, config):
+    def __init__(self, config: dict[str, Any]) -> None:
         super(GCPCloudProvider, self).__init__(config)
 
         # Disable warnings about file_cache not being available when using
@@ -248,12 +257,12 @@ class GCPCloudProvider(BaseCloudProvider):
             self.project_name = os.environ.get('GCP_PROJECT_NAME')
 
         # service connections, lazily initialized
-        self._gcp_compute = None
-        self._gcp_storage = None
-        self._gcp_dns = None
-        self._compute_resources_cache = None
-        self._storage_resources_cache = None
-        self._dns_resources_cache = None
+        self._gcp_compute: Any = None
+        self._gcp_storage: Any = None
+        self._gcp_dns: Any = None
+        self._compute_resources_cache: GCPResources | None = None
+        self._storage_resources_cache: GCPResources | None = None
+        self._dns_resources_cache: GCPResources | None = None
 
         # Initialize provider services
         self._compute = GCPComputeService(self)
@@ -265,49 +274,49 @@ class GCPCloudProvider(BaseCloudProvider):
     # Override base class implementation because it will cause
     # an infinite loop
     @property
-    def zone_name(self):
+    def zone_name(self) -> str | None:
         return self._zone_name
 
     @property
-    def compute(self):
+    def compute(self) -> ComputeService:
         return self._compute
 
     @property
-    def networking(self):
+    def networking(self) -> NetworkingService:
         return self._networking
 
     @property
-    def security(self):
+    def security(self) -> SecurityService:
         return self._security
 
     @property
-    def storage(self):
+    def storage(self) -> StorageService:
         return self._storage
 
     @property
-    def dns(self):
+    def dns(self) -> DnsService:
         return self._dns
 
     @property
-    def gcp_compute(self):
+    def gcp_compute(self) -> Any:
         if not self._gcp_compute:
             self._gcp_compute = self._connect_gcp_compute()
         return self._gcp_compute
 
     @property
-    def gcp_storage(self):
+    def gcp_storage(self) -> Any:
         if not self._gcp_storage:
             self._gcp_storage = self._connect_gcp_storage()
         return self._gcp_storage
 
     @property
-    def gcp_dns(self):
+    def gcp_dns(self) -> Any:
         if not self._gcp_dns:
             self._gcp_dns = self._connect_gcp_dns()
         return self._gcp_dns
 
     @property
-    def _compute_resources(self):
+    def _compute_resources(self) -> GCPResources:
         if not self._compute_resources_cache:
             self._compute_resources_cache = GCPResources(
                 self.gcp_compute,
@@ -317,13 +326,13 @@ class GCPCloudProvider(BaseCloudProvider):
         return self._compute_resources_cache
 
     @property
-    def _storage_resources(self):
+    def _storage_resources(self) -> GCPResources:
         if not self._storage_resources_cache:
             self._storage_resources_cache = GCPResources(self.gcp_storage)
         return self._storage_resources_cache
 
     @property
-    def _dns_resources(self):
+    def _dns_resources(self) -> GCPResources:
         if not self._dns_resources_cache:
             self._dns_resources_cache = GCPResources(
                 self.gcp_dns,
@@ -331,7 +340,7 @@ class GCPCloudProvider(BaseCloudProvider):
         return self._dns_resources_cache
 
     @property
-    def _credentials(self):
+    def _credentials(self) -> Any:
         if not self.credentials_obj:
             if self.credentials_dict:
                 self.credentials_obj = Credentials.from_service_account_info(
@@ -341,42 +350,43 @@ class GCPCloudProvider(BaseCloudProvider):
         return self.credentials_obj
 
     @property
-    def client_id(self):
+    def client_id(self) -> str:
         return self._credentials.service_account_email
 
-    def _get_build_request(self):
+    def _get_build_request(self) -> Callable[..., Any]:
         credentials = with_scopes_if_required(
             self._credentials, list(CLOUD_SCOPES))
 
         # FROM: https://github.com/googleapis/google-api-python-client/blob/
         # master/docs/thread_safety.md
         # Create a new Http() object for every request
-        def build_request(http, *args, **kwargs):
+        def build_request(http: Any, *args: Any, **kwargs: Any) -> Any:
             new_http = google_auth_httplib2.AuthorizedHttp(
                 credentials, http=httplib2.Http())
             return googleapiclient.http.HttpRequest(new_http, *args, **kwargs)
 
         return build_request
 
-    def _connect_gcp_storage(self):
+    def _connect_gcp_storage(self) -> Any:
         return discovery.build('storage', 'v1', credentials=self._credentials,
                                cache_discovery=False,
                                requestBuilder=self._get_build_request()
                                )
 
-    def _connect_gcp_compute(self):
+    def _connect_gcp_compute(self) -> Any:
         return discovery.build('compute', 'v1', credentials=self._credentials,
                                cache_discovery=False,
                                requestBuilder=self._get_build_request()
                                )
 
-    def _connect_gcp_dns(self):
+    def _connect_gcp_dns(self) -> Any:
         return discovery.build('dns', 'v1', credentials=self._credentials,
                                cache_discovery=False,
                                requestBuilder=self._get_build_request()
                                )
 
-    def wait_for_operation(self, operation, region=None, zone=None):
+    def wait_for_operation(self, operation: Any, region: str | None = None,
+                           zone: str | None = None) -> Any:
         args = {'project': self.project_name, 'operation': operation['name']}
         if not region and not zone:
             operations = self.gcp_compute.globalOperations()
@@ -396,11 +406,12 @@ class GCPCloudProvider(BaseCloudProvider):
 
             time.sleep(0.5)
 
-    def parse_url(self, url):
+    def parse_url(self, url: str) -> "GCPResourceUrl | None":
         out = self._compute_resources.parse_url(url)
         return out if out else self._storage_resources.parse_url(url)
 
-    def get_resource(self, resource, url_or_name, **kwargs):
+    def get_resource(self, resource: str, url_or_name: str,
+                     **kwargs: Any) -> Any:
         if not url_or_name:
             return None
         resource_url = (
@@ -424,7 +435,7 @@ class GCPCloudProvider(BaseCloudProvider):
             else:
                 raise
 
-    def authenticate(self):
+    def authenticate(self) -> bool:
         try:
             self.gcp_compute
             return True

Разница между файлами не показана из-за своего большого размера
+ 300 - 180
cloudbridge/providers/gcp/resources.py


Разница между файлами не показана из-за своего большого размера
+ 324 - 200
cloudbridge/providers/gcp/services.py


+ 12 - 6
cloudbridge/providers/gcp/subservices.py

@@ -6,6 +6,12 @@ from cloudbridge.base.subservices import BaseFloatingIPSubService
 from cloudbridge.base.subservices import BaseGatewaySubService
 from cloudbridge.base.subservices import BaseSubnetSubService
 from cloudbridge.base.subservices import BaseVMFirewallRuleSubService
+from cloudbridge.interfaces.provider import CloudProvider
+from cloudbridge.interfaces.resources import Bucket
+from cloudbridge.interfaces.resources import DnsZone
+from cloudbridge.interfaces.resources import Gateway
+from cloudbridge.interfaces.resources import Network
+from cloudbridge.interfaces.resources import VMFirewall
 
 
 log = logging.getLogger(__name__)
@@ -13,34 +19,34 @@ log = logging.getLogger(__name__)
 
 class GCPBucketObjectSubService(BaseBucketObjectSubService):
 
-    def __init__(self, provider, bucket):
+    def __init__(self, provider: CloudProvider, bucket: Bucket) -> None:
         super(GCPBucketObjectSubService, self).__init__(provider, bucket)
 
 
 class GCPGatewaySubService(BaseGatewaySubService):
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(GCPGatewaySubService, self).__init__(provider, network)
 
 
 class GCPVMFirewallRuleSubService(BaseVMFirewallRuleSubService):
 
-    def __init__(self, provider, firewall):
+    def __init__(self, provider: CloudProvider, firewall: VMFirewall) -> None:
         super(GCPVMFirewallRuleSubService, self).__init__(provider, firewall)
 
 
 class GCPFloatingIPSubService(BaseFloatingIPSubService):
 
-    def __init__(self, provider, gateway):
+    def __init__(self, provider: CloudProvider, gateway: Gateway) -> None:
         super(GCPFloatingIPSubService, self).__init__(provider, gateway)
 
 
 class GCPSubnetSubService(BaseSubnetSubService):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(GCPSubnetSubService, self).__init__(provider, network)
 
 
 class GCPDnsRecordSubService(BaseDnsRecordSubService):
 
-    def __init__(self, provider, dns_zone):
+    def __init__(self, provider: CloudProvider, dns_zone: DnsZone) -> None:
         super(GCPDnsRecordSubService, self).__init__(provider, dns_zone)

+ 5 - 3
cloudbridge/providers/mock/provider.py

@@ -6,6 +6,8 @@
     boto being hijacked, which will cause AWS to malfunction.
     See notes below.
 """
+from typing import Any
+
 from moto import mock_aws
 
 from ..aws import AWSCloudProvider
@@ -22,18 +24,18 @@ class MockAWSCloudProvider(AWSCloudProvider, TestMockHelperMixin):
     """
     PROVIDER_ID = 'mock'
 
-    def __init__(self, config):
+    def __init__(self, config: dict[str, Any]) -> None:
         self.setUpMock()
         super(MockAWSCloudProvider, self).__init__(config)
 
-    def setUpMock(self):
+    def setUpMock(self) -> None:
         """
         Let Moto take over all socket communications
         """
         self.mock_aws = mock_aws()
         self.mock_aws.start()
 
-    def tearDownMock(self):
+    def tearDownMock(self) -> None:
         """
         Stop Moto intercepting all socket communications
         """

+ 12 - 5
cloudbridge/providers/openstack/helpers.py

@@ -1,13 +1,19 @@
 """
 Helper functions
 """
+from __future__ import annotations
+
 import itertools
 import logging as log
+from typing import Any
+from typing import Sequence
 
 from cloudbridge.base.resources import ServerPagedResultList
+from cloudbridge.interfaces.provider import CloudProvider
 
 
-def os_result_limit(provider, requested_limit=None):
+def os_result_limit(provider: CloudProvider,
+                    requested_limit: int | None = None) -> int:
     """
     Calculates the limit for OpenStack.
     """
@@ -21,7 +27,8 @@ def os_result_limit(provider, requested_limit=None):
     return limit + 1
 
 
-def to_server_paged_list(provider, objects, limit=None):
+def to_server_paged_list(provider: CloudProvider, objects: Sequence[Any],
+                         limit: int | None = None) -> ServerPagedResultList[Any]:
     """
     A convenience function for wrapping a list of OpenStack native objects in
     a ServerPagedResultList. OpenStack
@@ -32,9 +39,9 @@ def to_server_paged_list(provider, objects, limit=None):
     limit = limit or provider.config.default_result_limit
     is_truncated = len(objects) > limit
     next_token = objects[limit-1].id if is_truncated else None
-    results = ServerPagedResultList(is_truncated,
-                                    next_token,
-                                    False)
+    results: ServerPagedResultList[Any] = ServerPagedResultList(is_truncated,
+                                                                next_token,
+                                                                False)
     for obj in itertools.islice(objects, limit):
         results.append(obj)
     return results

+ 82 - 60
cloudbridge/providers/openstack/provider.py

@@ -1,6 +1,9 @@
 """Provider implementation based on OpenStack Python clients for OpenStack."""
+from __future__ import annotations
 
 import inspect
+from collections.abc import Callable
+from typing import Any
 
 from keystoneauth1 import session
 
@@ -17,8 +20,13 @@ from swiftclient import client as swift_client
 
 from cloudbridge.base import BaseCloudProvider
 from cloudbridge.base.helpers import get_env
-
+from cloudbridge.base.services import BaseCloudService
 from cloudbridge.interfaces.exceptions import ProviderConnectionException
+from cloudbridge.interfaces.services import ComputeService
+from cloudbridge.interfaces.services import DnsService
+from cloudbridge.interfaces.services import NetworkingService
+from cloudbridge.interfaces.services import SecurityService
+from cloudbridge.interfaces.services import StorageService
 
 from .services import OpenStackComputeService
 from .services import OpenStackDnsService
@@ -30,9 +38,9 @@ from .services import OpenStackStorageService
 class OpenStackCloudProvider(BaseCloudProvider):
     """OpenStack provider implementation."""
 
-    PROVIDER_ID = 'openstack'
+    PROVIDER_ID: str = 'openstack'
 
-    def __init__(self, config):
+    def __init__(self, config: dict[str, Any]) -> None:
         super(OpenStackCloudProvider, self).__init__(config)
 
         # Initialize cloud connection fields
@@ -64,14 +72,14 @@ class OpenStackCloudProvider(BaseCloudProvider):
             get_env('OS_USER_DOMAIN_NAME'))
 
         # Service connections, lazily initialized
-        self._nova = None
-        self._keystone = None
-        self._swift = None
-        self._neutron = None
-        self._os_conn = None
+        self._nova: Any = None
+        self._keystone: Any = None
+        self._swift: Any = None
+        self._neutron: Any = None
+        self._os_conn: Any = None
 
         # Additional cached variables
-        self._cached_keystone_session = None
+        self._cached_keystone_session: Any = None
 
         # Initialize provider services
         self._compute = OpenStackComputeService(self)
@@ -81,19 +89,19 @@ class OpenStackCloudProvider(BaseCloudProvider):
         self._dns = OpenStackDnsService(self)
 
     @property
-    def nova(self):
+    def nova(self) -> Any:
         if not self._nova:
             self._nova = self._connect_nova()
         return self._nova
 
     @property
-    def keystone(self):
+    def keystone(self) -> Any:
         if not self._keystone:
             self._keystone = self._connect_keystone()
         return self._keystone
 
     @property
-    def _keystone_version(self):
+    def _keystone_version(self) -> int:
         """
         Return the numeric version of remote Keystone server.
 
@@ -106,7 +114,7 @@ class OpenStackCloudProvider(BaseCloudProvider):
         return 2
 
     @property
-    def _keystone_session(self):
+    def _keystone_session(self) -> Any:
         """
         Connect to Keystone and return a session object.
 
@@ -118,6 +126,7 @@ class OpenStackCloudProvider(BaseCloudProvider):
 
         if self._keystone_version == 3:
             from keystoneauth1.identity import v3
+            auth: Any
             if self.username and self.password:
                 auth = v3.Password(auth_url=self.auth_url,
                                    username=self.username,
@@ -143,7 +152,7 @@ class OpenStackCloudProvider(BaseCloudProvider):
             self._cached_keystone_session = session.Session(auth=auth)
         return self._cached_keystone_session
 
-    def _connect_openstack(self):
+    def _connect_openstack(self) -> Any:
         return connection.Connection(
             region_name=self.region_name,
             user_agent='cloudbridge',
@@ -152,47 +161,47 @@ class OpenStackCloudProvider(BaseCloudProvider):
         )
 
     @property
-    def swift(self):
+    def swift(self) -> Any:
         if not self._swift:
             self._swift = self._connect_swift()
         return self._swift
 
     @property
-    def neutron(self):
+    def neutron(self) -> Any:
         if not self._neutron:
             self._neutron = self._connect_neutron()
         return self._neutron
 
     @property
-    def os_conn(self):
+    def os_conn(self) -> Any:
         if not self._os_conn:
             self._os_conn = self._connect_openstack()
         return self._os_conn
 
     @property
-    def compute(self):
+    def compute(self) -> ComputeService:
         return self._compute
 
     @property
-    def networking(self):
+    def networking(self) -> NetworkingService:
         return self._networking
 
     @property
-    def security(self):
+    def security(self) -> SecurityService:
         return self._security
 
     @property
-    def storage(self):
+    def storage(self) -> StorageService:
         return self._storage
 
     @property
-    def dns(self):
+    def dns(self) -> DnsService:
         return self._dns
 
-    def _connect_nova(self):
+    def _connect_nova(self) -> Any:
         return self._connect_nova_region(self.region_name)
 
-    def _connect_nova_region(self, region_name):
+    def _connect_nova_region(self, region_name: str | None) -> Any:
         """Get an OpenStack Nova (compute) client object."""
         # Force reauthentication with Keystone
         self._cached_keystone_session = None
@@ -215,7 +224,7 @@ class OpenStackCloudProvider(BaseCloudProvider):
                 http_log_debug=True if self.config.debug_mode else False)
         return nova
 
-    def _connect_keystone(self):
+    def _connect_keystone(self) -> Any:
         """Get an OpenStack Keystone (identity) client object."""
         if self._keystone_version == 3:
             return keystone_client.Client(session=self._keystone_session,
@@ -237,7 +246,8 @@ class OpenStackCloudProvider(BaseCloudProvider):
             return keystone
 
     @staticmethod
-    def _clean_options(options, method_to_match):
+    def _clean_options(options: dict[str, Any] | None,
+                       method_to_match: Callable[..., Any]) -> dict[str, Any]:
         """
         Returns a **copy** of the source options with all keys that are not in
         the ``method_to_match`` parameter list removed.
@@ -261,20 +271,17 @@ class OpenStackCloudProvider(BaseCloudProvider):
             then this will be an empty dictionary
         :rtype: ``dict``
         """
-        result = {}
+        result: dict[str, Any] = {}
         if options:
-            try:
-                method_signature = inspect.signature(method_to_match)
-                parameters = set(method_signature.parameters.keys())
-            except AttributeError:
-                parameters = set(inspect.getargspec(method_to_match).args)
+            method_signature = inspect.signature(method_to_match)
+            parameters = set(method_signature.parameters.keys())
             result = {key: val for key, val in options.items() if
                       key in parameters}
             # Don't allow the options to override our authentication
             result.pop('os_options', None)
         return result
 
-    def _connect_swift(self, options=None):
+    def _connect_swift(self, options: dict[str, Any] | None = None) -> Any:
         """
         Get an OpenStack Swift (object store) client connection.
 
@@ -297,42 +304,57 @@ class OpenStackCloudProvider(BaseCloudProvider):
             clean_options['session'] = self._keystone_session
         return swift_client.Connection(**clean_options)
 
-    def _connect_neutron(self):
+    def _connect_neutron(self) -> Any:
         """Get an OpenStack Neutron (networking) client object cloud."""
         return neutron_client.Client(auth_url=self.auth_url,
                                      session=self._keystone_session,
                                      region_name=self.region_name)
 
-    def service_zone_name(self, service):
+    def service_zone_name(self, service: BaseCloudService) -> str | None:
+        # ``service_zone_name`` is an OpenStack-specific attribute set in each
+        # service's __init__; it is not declared on the typed service
+        # interfaces, so reach it through ``Any``.
+        networking: Any = self.networking
+        security: Any = self.security
+        compute: Any = self.compute
+        storage: Any = self.storage
+        # ``zone_name`` is typed ``str | None`` on the interface, but the base
+        # implementation may return a dict at runtime (via ast.literal_eval);
+        # bind it to ``Any`` so the dict branches type-check.
+        zone_name: Any = self.zone_name
         service_name = service._service_event_pattern
         if "networking" in service_name:
-            if self.networking.service_zone_name:
-                return self.networking.service_zone_name
-            elif (isinstance(self.zone_name, dict) and
-                  self.zone_name.get("networking_zone")):
-                return self.zone_name.get("networking_zone")
+            if networking.service_zone_name:
+                return networking.service_zone_name
+            elif (isinstance(zone_name, dict) and
+                  zone_name.get("networking_zone")):
+                return zone_name.get("networking_zone")
         elif "security" in service_name:
-            if self.security.service_zone_name:
-                return self.security.service_zone_name
-            elif (isinstance(self.zone_name, dict) and
-                  self.zone_name.get("security_zone")):
-                return self.zone_name.get("security_zone")
+            if security.service_zone_name:
+                return security.service_zone_name
+            elif (isinstance(zone_name, dict) and
+                  zone_name.get("security_zone")):
+                return zone_name.get("security_zone")
         elif "compute" in service_name:
-            if self.compute.service_zone_name:
-                return self.compute.service_zone_name
-            elif (isinstance(self.zone_name, dict) and
-                  self.zone_name.get("compute_zone")):
-                return self.zone_name.get("compute_zone")
+            if compute.service_zone_name:
+                return compute.service_zone_name
+            elif (isinstance(zone_name, dict) and
+                  zone_name.get("compute_zone")):
+                return zone_name.get("compute_zone")
         elif "storage" in service_name:
-            if self.storage.service_zone_name:
-                return self.storage.service_zone_name
-            elif (isinstance(self.zone_name, dict) and
-                  self.zone_name.get("storage_zone")):
-                return self.zone_name.get("storage_zone")
-        elif (isinstance(self.zone_name, dict) and
-              self.zone_name.get("default_zone")):
-            return self.zone_name.get("default_zone")
-        elif isinstance(self.zone_name, str):
-            return self.zone_name
+            if storage.service_zone_name:
+                return storage.service_zone_name
+            elif (isinstance(zone_name, dict) and
+                  zone_name.get("storage_zone")):
+                return zone_name.get("storage_zone")
+        elif (isinstance(zone_name, dict) and
+              zone_name.get("default_zone")):
+            return zone_name.get("default_zone")
+        elif isinstance(zone_name, str):
+            return zone_name
         else:
             return None
+        # The branches above only return when an inner condition matched;
+        # fall through to ``None`` otherwise (preserving the original
+        # implicit return).
+        return None

Разница между файлами не показана из-за своего большого размера
+ 227 - 156
cloudbridge/providers/openstack/resources.py


Разница между файлами не показана из-за своего большого размера
+ 279 - 167
cloudbridge/providers/openstack/services.py


+ 14 - 6
cloudbridge/providers/openstack/subservices.py

@@ -1,3 +1,5 @@
+from __future__ import annotations
+
 import logging
 
 from cloudbridge.base.subservices import BaseBucketObjectSubService
@@ -6,6 +8,12 @@ from cloudbridge.base.subservices import BaseFloatingIPSubService
 from cloudbridge.base.subservices import BaseGatewaySubService
 from cloudbridge.base.subservices import BaseSubnetSubService
 from cloudbridge.base.subservices import BaseVMFirewallRuleSubService
+from cloudbridge.interfaces.provider import CloudProvider
+from cloudbridge.interfaces.resources import Bucket
+from cloudbridge.interfaces.resources import DnsZone
+from cloudbridge.interfaces.resources import Gateway
+from cloudbridge.interfaces.resources import Network
+from cloudbridge.interfaces.resources import VMFirewall
 
 
 log = logging.getLogger(__name__)
@@ -13,36 +21,36 @@ log = logging.getLogger(__name__)
 
 class OpenStackBucketObjectSubService(BaseBucketObjectSubService):
 
-    def __init__(self, provider, bucket):
+    def __init__(self, provider: CloudProvider, bucket: Bucket) -> None:
         super(OpenStackBucketObjectSubService, self).__init__(provider, bucket)
 
 
 class OpenStackGatewaySubService(BaseGatewaySubService):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(OpenStackGatewaySubService, self).__init__(provider, network)
 
 
 class OpenStackFloatingIPSubService(BaseFloatingIPSubService):
 
-    def __init__(self, provider, gateway):
+    def __init__(self, provider: CloudProvider, gateway: Gateway) -> None:
         super(OpenStackFloatingIPSubService, self).__init__(provider, gateway)
 
 
 class OpenStackVMFirewallRuleSubService(BaseVMFirewallRuleSubService):
 
-    def __init__(self, provider, firewall):
+    def __init__(self, provider: CloudProvider, firewall: VMFirewall) -> None:
         super(OpenStackVMFirewallRuleSubService, self).__init__(
             provider, firewall)
 
 
 class OpenStackSubnetSubService(BaseSubnetSubService):
 
-    def __init__(self, provider, network):
+    def __init__(self, provider: CloudProvider, network: Network) -> None:
         super(OpenStackSubnetSubService, self).__init__(provider, network)
 
 
 class OpenStackDnsRecordSubService(BaseDnsRecordSubService):
 
-    def __init__(self, provider, dns_zone):
+    def __init__(self, provider: CloudProvider, dns_zone: DnsZone) -> None:
         super(OpenStackDnsRecordSubService, self).__init__(provider, dns_zone)

+ 0 - 0
cloudbridge/py.typed


+ 40 - 0
pyproject.toml

@@ -91,6 +91,7 @@ dev = [
     "pydevd",
     "flake8>=3.3.0",
     "flake8-import-order>=0.12",
+    "mypy>=2.1,<3",
 ]
 
 [tool.setuptools.dynamic]
@@ -100,6 +101,10 @@ version = { attr = "cloudbridge.__version__" }
 include = ["cloudbridge*"]
 exclude = ["tests*"]
 
+[tool.setuptools.package-data]
+# Ship the PEP 561 marker so downstream consumers pick up our type hints.
+cloudbridge = ["py.typed"]
+
 [tool.coverage.run]
 branch = true
 source = ["cloudbridge"]
@@ -108,3 +113,38 @@ omit = [
     "cloudbridge/__init__.py",
 ]
 parallel = true
+
+[tool.mypy]
+# CloudBridge is typed gradually. The public API (the interface layer and the
+# factory) is held to a strict baseline so downstream users get a fully-typed
+# API; the base implementations and providers (which wrap untyped cloud SDKs)
+# are temporarily exempted below and ratcheted to strict module-by-module.
+python_version = "3.13"
+files = ["cloudbridge"]
+ignore_missing_imports = true   # untyped cloud SDKs (boto3, azure-*, ...) -> Any
+namespace_packages = true
+warn_unused_configs = true
+# Full strict baseline (applies to every module NOT exempted below), with two
+# documented exceptions that don't fit this codebase:
+strict = true
+# interfaces/__init__.py re-exports the public names without an __all__; keep
+# implicit re-export so `from cloudbridge.interfaces import CloudProvider`
+# keeps working for factory.py and downstream consumers.
+implicit_reexport = true
+# Cross-class property setters (@LabeledCloudResource.label.setter) and the
+# `deprecation` library's decorators are untyped; don't fail the build on them.
+disallow_untyped_decorators = false
+
+# Providers (all typed) wrap untyped cloud SDKs, so they get a pragmatic tier:
+# complete annotations are required (disallow_untyped_defs etc. from the strict
+# baseline), but warn_return_any and disallow_untyped_calls are off -- every
+# getter reads an attribute off an Any-typed SDK object, and forcing those
+# through cast()/ignore adds noise without value.
+[[tool.mypy.overrides]]
+module = ["cloudbridge.providers.*"]
+warn_return_any = false
+disallow_untyped_calls = false
+
+# base/ is held to the FULL strict bar (it orchestrates through the typed
+# interface layer and never touches the cloud SDKs directly), so it falls under
+# the global strict baseline with no override.

+ 2 - 8
tests/test_compute_service.py

@@ -436,7 +436,7 @@ class CloudComputeServiceTestCase(ProviderTestBase):
                                                   subnet=subnet)
 
             # check whether stopping aws instance works
-            resp = test_inst.stop()
+            test_inst.stop()
             test_inst.wait_for([InstanceState.STOPPED])
             test_inst.refresh()
             self.assertTrue(
@@ -445,11 +445,8 @@ class CloudComputeServiceTestCase(ProviderTestBase):
                 "'stop' operation but got %s"
                 % test_inst.state)
 
-            self.assertTrue(resp, "Response from method was suppose to be"
-                            + " True but got False")
-
             # check whether starting aws instance works
-            resp = test_inst.start()
+            test_inst.start()
             test_inst.wait_for([InstanceState.RUNNING])
             test_inst.refresh()
             self.assertTrue(
@@ -457,6 +454,3 @@ class CloudComputeServiceTestCase(ProviderTestBase):
                 "Instance state must be running when refreshing after a "
                 "'start' operation but got %s"
                 % test_inst.state)
-
-            self.assertTrue(resp, "Response from method was suppose to be"
-                            + " True but got False")

+ 14 - 2
tox.ini

@@ -6,7 +6,7 @@
 # running the tests.
 
 [tox]
-envlist = py3.13-{aws,azure,gcp,openstack,mock},lint
+envlist = py3.13-{aws,azure,gcp,openstack,mock},lint,mypy
 
 [testenv]
 commands = # see pyproject.toml for coverage options; setup.cfg for flake8
@@ -88,4 +88,16 @@ deps =
 
 [testenv:lint]
 commands = flake8 cloudbridge tests
-deps = flake8
+deps =
+    flake8
+    flake8-import-order
+
+[testenv:mypy]
+# Type-check the codebase (see [tool.mypy] in pyproject.toml). The provider
+# layer is typed against the cloud SDKs, so the env installs the full deps
+# (via requirements.txt -> .[dev], which includes mypy) rather than running
+# bare: without the SDKs installed, mypy infers SDK values as Any and reports
+# spurious redundant-cast / unused-ignore errors that don't occur in a real
+# dev environment.
+deps = -rrequirements.txt
+commands = mypy

Некоторые файлы не были показаны из-за большого количества измененных файлов