From c94279bfacfbc69a86dcf93b8a00a6a1860c5bd5 Mon Sep 17 00:00:00 2001 From: Yue Zhang Date: Mon, 2 Nov 2015 16:11:14 +0800 Subject: [PATCH] Refactor protocol detection code --- azurelinuxagent/distro/coreos/deprovision.py | 3 + .../distro/coreos/handlerFactory.py | 2 +- azurelinuxagent/distro/default/deprovision.py | 2 + azurelinuxagent/distro/default/env.py | 10 +- azurelinuxagent/distro/default/extension.py | 9 +- .../distro/default/handlerFactory.py | 30 ++--- azurelinuxagent/distro/default/init.py | 3 + azurelinuxagent/distro/default/provision.py | 93 ++++++------- .../distro/default/resourceDisk.py | 2 + azurelinuxagent/distro/default/run.py | 17 ++- azurelinuxagent/distro/default/scvmm.py | 2 + azurelinuxagent/distro/ubuntu/deprovision.py | 3 + .../distro/ubuntu/handlerFactory.py | 4 +- azurelinuxagent/distro/ubuntu/provision.py | 56 ++++---- azurelinuxagent/event.py | 7 +- azurelinuxagent/protocol/__init__.py | 3 - .../{distro/default => protocol}/dhcp.py | 126 ++++++++++-------- azurelinuxagent/protocol/factory.py | 124 +++++++++++++++++ azurelinuxagent/protocol/ovfenv.py | 18 ++- azurelinuxagent/protocol/protocolFactory.py | 113 ---------------- azurelinuxagent/protocol/v1.py | 13 +- azurelinuxagent/protocol/v2.py | 14 +- tests/test_deprovision.py | 4 +- tests/test_dhcp.py | 32 +++-- tests/test_distroLoader.py | 1 - tests/test_envmon.py | 12 +- tests/test_event.py | 4 +- tests/test_ext.py | 2 +- tests/test_protocolFactory.py | 37 ----- tests/test_resourcedisk.py | 6 +- 30 files changed, 388 insertions(+), 364 deletions(-) rename azurelinuxagent/{distro/default => protocol}/dhcp.py (80%) create mode 100644 azurelinuxagent/protocol/factory.py delete mode 100644 azurelinuxagent/protocol/protocolFactory.py delete mode 100644 tests/test_protocolFactory.py diff --git a/azurelinuxagent/distro/coreos/deprovision.py b/azurelinuxagent/distro/coreos/deprovision.py index 99d3a40..1dc5218 100644 --- a/azurelinuxagent/distro/coreos/deprovision.py +++ b/azurelinuxagent/distro/coreos/deprovision.py @@ -21,6 +21,9 @@ import azurelinuxagent.utils.fileutil as fileutil from azurelinuxagent.distro.default.deprovision import DeprovisionHandler, DeprovisionAction class CoreOSDeprovisionHandler(DeprovisionHandler): + def __init__(self, handlers): + self.handlers = handlers + def setup(self, deluser): warnings, actions = super(CoreOSDeprovisionHandler, self).setup(deluser) warnings.append("WARNING! /etc/machine-id will be removed.") diff --git a/azurelinuxagent/distro/coreos/handlerFactory.py b/azurelinuxagent/distro/coreos/handlerFactory.py index 58f476c..01a2afb 100644 --- a/azurelinuxagent/distro/coreos/handlerFactory.py +++ b/azurelinuxagent/distro/coreos/handlerFactory.py @@ -23,5 +23,5 @@ from azurelinuxagent.distro.default.handlerFactory import DefaultHandlerFactory class CoreOSHandlerFactory(DefaultHandlerFactory): def __init__(self): super(CoreOSHandlerFactory, self).__init__() - self.deprovision_handler = CoreOSDeprovisionHandler() + self.deprovision_handler = CoreOSDeprovisionHandler(self) diff --git a/azurelinuxagent/distro/default/deprovision.py b/azurelinuxagent/distro/default/deprovision.py index 854e754..8bbb52d 100644 --- a/azurelinuxagent/distro/default/deprovision.py +++ b/azurelinuxagent/distro/default/deprovision.py @@ -35,6 +35,8 @@ class DeprovisionAction(object): self.func(*self.args, **self.kwargs) class DeprovisionHandler(object): + def __init__(self, handlers): + self.handlers = handlers def del_root_password(self, warnings, actions): warnings.append("WARNING! root password will be disabled. " diff --git a/azurelinuxagent/distro/default/env.py b/azurelinuxagent/distro/default/env.py index 28bf718..fe7a227 100644 --- a/azurelinuxagent/distro/default/env.py +++ b/azurelinuxagent/distro/default/env.py @@ -23,6 +23,7 @@ import threading import time import azurelinuxagent.logger as logger import azurelinuxagent.conf as conf +import azurelinuxagent.protocol.dhcp as dhcp from azurelinuxagent.utils.osutil import OSUTIL class EnvHandler(object): @@ -34,7 +35,8 @@ class EnvHandler(object): If new scsi disk found, set """ def __init__(self, handlers): - self.monitor = EnvMonitor(handlers.dhcp_handler) + self.handlers = handlers + self.monitor = EnvMonitor() def start(self): self.monitor.start() @@ -44,14 +46,14 @@ class EnvHandler(object): class EnvMonitor(object): - def __init__(self, dhcp_handler): - self.dhcp_handler = dhcp_handler + def __init__(self): self.stopped = True self.hostname = None self.dhcpid = None self.server_thread=None def start(self): + self.dhcp_resp = dhcp.DHCPCLIENT.get_dhcp_resp() if not self.stopped: logger.info("Stop existing env monitor service.") self.stop() @@ -102,7 +104,7 @@ class EnvMonitor(object): if newpid is not None and newpid != self.dhcpid: logger.info("EnvMonitor: Detected dhcp client restart. " "Restoring routing table.") - self.dhcp_handler.conf_routes() + self.dhcp_resp.conf_routes() self.dhcpid = newpid def stop(self): diff --git a/azurelinuxagent/distro/default/extension.py b/azurelinuxagent/distro/default/extension.py index 0ddf00f..f29bb92 100644 --- a/azurelinuxagent/distro/default/extension.py +++ b/azurelinuxagent/distro/default/extension.py @@ -25,7 +25,8 @@ import shutil import azurelinuxagent.logger as logger from azurelinuxagent.future import text from azurelinuxagent.utils.osutil import OSUTIL -import azurelinuxagent.protocol as prot +from azurelinuxagent.protocol.factory import PROT_FACTORY +import azurelinuxagent.protocol.common as prot from azurelinuxagent.metadata import AGENT_VERSION from azurelinuxagent.event import add_event, WALAEventOperation from azurelinuxagent.exception import ExtensionError @@ -146,10 +147,12 @@ class ExtHandlerState(object): class ExtHandlersHandler(object): + def __init__(self, handlers): + self.handlers = handlers def process(self): try: - protocol = prot.FACTORY.get_default_protocol() + protocol = PROT_FACTORY.get_protocol() ext_handlers = protocol.get_ext_handlers() except prot.ProtocolError as e: add_event(name="WALA", is_success=False, message = text(e)) @@ -195,7 +198,7 @@ class ExtHandlersHandler(object): if handler.ext_status is not None: try: - protocol = prot.FACTORY.get_default_protocol() + protocol = PROT_FACTORY.get_protocol() protocol.report_ext_status(handler.name, handler.ext.name, handler.ext_status) except prot.ProtocolError as e: diff --git a/azurelinuxagent/distro/default/handlerFactory.py b/azurelinuxagent/distro/default/handlerFactory.py index dceb2a3..55a48f1 100644 --- a/azurelinuxagent/distro/default/handlerFactory.py +++ b/azurelinuxagent/distro/default/handlerFactory.py @@ -16,25 +16,23 @@ # # Requires Python 2.4+ and Openssl 1.0+ # -from .init import InitHandler -from .run import MainHandler -from .scvmm import ScvmmHandler -from .dhcp import DhcpHandler -from .env import EnvHandler -from .provision import ProvisionHandler -from .resourceDisk import ResourceDiskHandler -from .extension import ExtHandlersHandler -from .deprovision import DeprovisionHandler +from azurelinuxagent.distro.default.init import InitHandler +from azurelinuxagent.distro.default.run import MainHandler +from azurelinuxagent.distro.default.scvmm import ScvmmHandler +from azurelinuxagent.distro.default.env import EnvHandler +from azurelinuxagent.distro.default.provision import ProvisionHandler +from azurelinuxagent.distro.default.resourceDisk import ResourceDiskHandler +from azurelinuxagent.distro.default.extension import ExtHandlersHandler +from azurelinuxagent.distro.default.deprovision import DeprovisionHandler class DefaultHandlerFactory(object): def __init__(self): - self.init_handler = InitHandler() + self.init_handler = InitHandler(self) self.main_handler = MainHandler(self) - self.scvmm_handler = ScvmmHandler() - self.dhcp_handler = DhcpHandler() + self.scvmm_handler = ScvmmHandler(self) self.env_handler = EnvHandler(self) - self.provision_handler = ProvisionHandler() - self.resource_disk_handler = ResourceDiskHandler() - self.ext_handlers_handler = ExtHandlersHandler() - self.deprovision_handler = DeprovisionHandler() + self.provision_handler = ProvisionHandler(self) + self.resource_disk_handler = ResourceDiskHandler(self) + self.ext_handlers_handler = ExtHandlersHandler(self) + self.deprovision_handler = DeprovisionHandler(self) diff --git a/azurelinuxagent/distro/default/init.py b/azurelinuxagent/distro/default/init.py index db74fef..3c6671c 100644 --- a/azurelinuxagent/distro/default/init.py +++ b/azurelinuxagent/distro/default/init.py @@ -25,6 +25,9 @@ import azurelinuxagent.utils.fileutil as fileutil class InitHandler(object): + def __init__(self, handlers): + self.handlers = handlers + def init(self, verbose): #Init stdout log level = logger.LogLevel.VERBOSE if verbose else logger.LogLevel.INFO diff --git a/azurelinuxagent/distro/default/provision.py b/azurelinuxagent/distro/default/provision.py index 424f083..67fede1 100644 --- a/azurelinuxagent/distro/default/provision.py +++ b/azurelinuxagent/distro/default/provision.py @@ -24,9 +24,10 @@ import azurelinuxagent.logger as logger from azurelinuxagent.future import text import azurelinuxagent.conf as conf from azurelinuxagent.event import add_event, WALAEventOperation -from azurelinuxagent.exception import * +from azurelinuxagent.exception import ProvisionError from azurelinuxagent.utils.osutil import OSUTIL, OSUtilError -import azurelinuxagent.protocol as prot +from azurelinuxagent.protocol.factory import PROT_FACTORY, ProvisionStatus, \ + ProtocolError import azurelinuxagent.protocol.ovfenv as ovf import azurelinuxagent.utils.shellutil as shellutil import azurelinuxagent.utils.fileutil as fileutil @@ -35,58 +36,65 @@ CUSTOM_DATA_FILE="CustomData" class ProvisionHandler(object): + def __init__(self, handlers): + self.handlers = handlers + + def report_event(self, message, is_success=False): + add_event(name="WALA", message=message, is_success=is_success, + op=WALAEventOperation.Provision) + + def report_not_ready(self, protocol, sub_status, description): + status = ProvisionStatus(status="NotReady", subStatus=sub_status, + description=description) + try: + protocol.report_provision_status(status) + except ProtocolError as e: + self.report_event(text(e)) + + def report_ready(self, protocol, thumbprint=None): + status = ProvisionStatus(status="Ready") + status.properties.certificateThumbprint = thumbprint + try: + protocol.report_provision_status(status) + except ProtocolError as e: + self.report_event(text(e)) + def process(self): #If provision is not enabled, return if not conf.get_switch("Provisioning.Enabled", True): logger.info("Provisioning is disabled. Skip.") - return + return provisioned = os.path.join(OSUTIL.get_lib_dir(), "provisioned") if os.path.isfile(provisioned): return - logger.info("run provision handler.") - protocol = prot.FACTORY.get_default_protocol() + logger.info("Run provision handler.") + logger.info("Copy ovf-env.xml.") try: - status = prot.ProvisionStatus(status="NotReady", - subStatus="Provisioning", - description="Starting") - try: - protocol.report_provision_status(status) - except prot.ProtocolError as e: - add_event(name="WALA", is_success=False, message=text(e), - op=WALAEventOperation.Provision) - - self.provision() + ovfenv = ovf.copy_ovf_env() + except ProtocolError as e: + self.report_event("Failed to copy ovf-env.xml: {0}".format(e)) + return + + protocol = PROT_FACTORY.detect_protocol_by_file() + self.report_not_ready(protocol, "Provisioning", "Starting") + + try: + logger.info("Start provisioning") + self.provision(ovfenv) fileutil.write_file(provisioned, "") thumbprint = self.reg_ssh_host_key() - logger.info("Finished provisioning") - status = prot.ProvisionStatus(status="Ready") - status.properties.certificateThumbprint = thumbprint - - try: - protocol.report_provision_status(status) - except prot.ProtocolError as pe: - add_event(name="WALA", is_success=False, message=text(pe), - op=WALAEventOperation.Provision) - - add_event(name="WALA", is_success=True, message="", - op=WALAEventOperation.Provision) except ProvisionError as e: logger.error("Provision failed: {0}", e) - status = prot.ProvisionStatus(status="NotReady", - subStatus="ProvisioningFailed", - description= text(e)) - try: - protocol.report_provision_status(status) - except prot.ProtocolError as pe: - add_event(name="WALA", is_success=False, message=text(pe), - op=WALAEventOperation.Provision) - - add_event(name="WALA", is_success=False, message=text(e), - op=WALAEventOperation.Provision) + self.report_not_ready(protocol, "ProvisioningFailed", text(e)) + self.report_event(text(e)) + return + self.report_ready(protocol, thumbprint) + self.report_event("Provision succeed", is_success=True) + def reg_ssh_host_key(self): keypair_type = conf.get("Provisioning.SshHostKeyPairType", "rsa") if conf.get_switch("Provisioning.RegenerateSshHostKeyPair"): @@ -105,14 +113,7 @@ class ProvisionHandler(object): raise ProvisionError(("Failed to generate ssh host key: " "ret={0}, out= {1}").format(ret[0], ret[1])) - - def provision(self): - logger.info("Copy ovf-env.xml.") - try: - ovfenv = ovf.copy_ovf_env() - except prot.ProtocolError as e: - raise ProvisionError("Failed to copy ovf-env.xml: {0}".format(e)) - + def provision(self, ovfenv): logger.info("Handle ovf-env.xml.") try: logger.info("Set host name.") diff --git a/azurelinuxagent/distro/default/resourceDisk.py b/azurelinuxagent/distro/default/resourceDisk.py index 734863c..ff548d3 100644 --- a/azurelinuxagent/distro/default/resourceDisk.py +++ b/azurelinuxagent/distro/default/resourceDisk.py @@ -41,6 +41,8 @@ For additional details to please refer to the MSDN documentation at : http://msd """ class ResourceDiskHandler(object): + def __init__(self, handlers): + self.handlers = handlers def start_activate_resource_disk(self): disk_thread = threading.Thread(target = self.run) diff --git a/azurelinuxagent/distro/default/run.py b/azurelinuxagent/distro/default/run.py index dfd3b03..cbb47c6 100644 --- a/azurelinuxagent/distro/default/run.py +++ b/azurelinuxagent/distro/default/run.py @@ -28,7 +28,8 @@ from azurelinuxagent.metadata import AGENT_LONG_NAME, AGENT_VERSION, \ DISTRO_FULL_NAME, PY_VERSION_MAJOR, \ PY_VERSION_MINOR, PY_VERSION_MICRO import azurelinuxagent.event as event -import azurelinuxagent.protocol as prot +import azurelinuxagent.protocol.dhcp as dhcp +from azurelinuxagent.protocol.factory import PROT_FACTORY from azurelinuxagent.utils.osutil import OSUTIL import azurelinuxagent.utils.fileutil as fileutil @@ -43,27 +44,25 @@ class MainHandler(object): logger.info("Python: {0}.{1}.{2}", PY_VERSION_MAJOR, PY_VERSION_MINOR, PY_VERSION_MICRO) - event.enable_unhandled_err_dump("Azure Linux Agent") + event.enable_unhandled_err_dump(AGENT_LONG_NAME) fileutil.write_file(OSUTIL.get_agent_pid_file_path(), text(os.getpid())) if conf.get_switch("DetectScvmmEnv", False): if self.handlers.scvmm_handler.detect_scvmm_env(): return - - self.handlers.dhcp_handler.probe() - - prot.detect_default_protocol() - - event.EventMonitor().start() + + PROT_FACTORY.wait_for_network() self.handlers.provision_handler.process() + + PROT_FACTORY.detect_protocol() if conf.get_switch("ResourceDisk.Format", False): self.handlers.resource_disk_handler.start_activate_resource_disk() + event.EventMonitor().start() self.handlers.env_handler.start() - protocol = prot.FACTORY.get_default_protocol() while True: #Handle extensions self.handlers.ext_handlers_handler.process() diff --git a/azurelinuxagent/distro/default/scvmm.py b/azurelinuxagent/distro/default/scvmm.py index 680c04b..21e2906 100644 --- a/azurelinuxagent/distro/default/scvmm.py +++ b/azurelinuxagent/distro/default/scvmm.py @@ -26,6 +26,8 @@ VMM_CONF_FILE_NAME = "linuxosconfiguration.xml" VMM_STARTUP_SCRIPT_NAME= "install" class ScvmmHandler(object): + def __init__(self, handlers): + self.handlers = handlers def detect_scvmm_env(self): logger.info("Detecting Microsoft System Center VMM Environment") diff --git a/azurelinuxagent/distro/ubuntu/deprovision.py b/azurelinuxagent/distro/ubuntu/deprovision.py index 0c3c4e5..4d790db 100644 --- a/azurelinuxagent/distro/ubuntu/deprovision.py +++ b/azurelinuxagent/distro/ubuntu/deprovision.py @@ -33,6 +33,9 @@ def del_resolv(): class UbuntuDeprovisionHandler(DeprovisionHandler): + def __init__(self, handlers): + self.handlers = handlers + def setup(self, deluser): warnings, actions = super(UbuntuDeprovisionHandler, self).setup(deluser) warnings.append("WARNING! Nameserver configuration in " diff --git a/azurelinuxagent/distro/ubuntu/handlerFactory.py b/azurelinuxagent/distro/ubuntu/handlerFactory.py index 11f7f04..9446594 100644 --- a/azurelinuxagent/distro/ubuntu/handlerFactory.py +++ b/azurelinuxagent/distro/ubuntu/handlerFactory.py @@ -24,6 +24,6 @@ from azurelinuxagent.distro.default.handlerFactory import DefaultHandlerFactory class UbuntuHandlerFactory(DefaultHandlerFactory): def __init__(self): super(UbuntuHandlerFactory, self).__init__() - self.provision_handler = UbuntuProvisionHandler() - self.deprovision_handler = UbuntuDeprovisionHandler() + self.provision_handler = UbuntuProvisionHandler(self) + self.deprovision_handler = UbuntuDeprovisionHandler(self) diff --git a/azurelinuxagent/distro/ubuntu/provision.py b/azurelinuxagent/distro/ubuntu/provision.py index a68fe4d..655ad08 100644 --- a/azurelinuxagent/distro/ubuntu/provision.py +++ b/azurelinuxagent/distro/ubuntu/provision.py @@ -22,9 +22,10 @@ import time import azurelinuxagent.logger as logger from azurelinuxagent.future import text import azurelinuxagent.conf as conf -import azurelinuxagent.protocol as prot +from azurelinuxagent.protocol.factory import PROT_FACTORY +import azurelinuxagent.protocol.ovfenv as ovf from azurelinuxagent.event import add_event, WALAEventOperation -from azurelinuxagent.exception import * +from azurelinuxagent.exception import ProvisionError from azurelinuxagent.utils.osutil import OSUTIL import azurelinuxagent.utils.shellutil as shellutil import azurelinuxagent.utils.fileutil as fileutil @@ -34,6 +35,9 @@ from azurelinuxagent.distro.default.provision import ProvisionHandler On ubuntu image, provision could be disabled. """ class UbuntuProvisionHandler(ProvisionHandler): + def __init__(self, handlers): + self.handlers = handlers + def process(self): #If provision is enabled, run default provision handler if conf.get_switch("Provisioning.Enabled", False): @@ -45,37 +49,43 @@ class UbuntuProvisionHandler(ProvisionHandler): if os.path.isfile(provisioned): return - logger.info("Waiting cloud-init to finish provisioning.") - protocol = prot.FACTORY.get_default_protocol() + logger.info("Waiting cloud-init to copy ovf-env.xml.") + self.wait_for_ovfenv() + + protocol = PROT_FACTORY.detect_protocol_by_file() + self.report_not_ready(protocol, "Provisioning", "Starting") try: logger.info("Wait for ssh host key to be generated.") thumbprint = self.wait_for_ssh_host_key() fileutil.write_file(provisioned, "") - logger.info("Finished provisioning") - status = prot.ProvisionStatus(status="Ready") - status.properties.certificateThumbprint = thumbprint - try: - protocol.report_provision_status(status) - except prot.ProtocolError as pe: - add_event(name="WALA", is_success=False, message=text(pe), - op=WALAEventOperation.Provision) - + except ProvisionError as e: logger.error("Provision failed: {0}", e) - status = prot.ProvisionStatus(status="NotReady", - subStatus="ProvisioningFailed", - description= text(e)) - try: - protocol.report_provision_status(status) - except prot.ProtocolError as pe: - add_event(name="WALA", is_success=False, message=text(pe), - op=WALAEventOperation.Provision) + self.report_not_ready(protocol, "ProvisioningFailed", text(e)) + self.report_event(text(e)) + return + + self.report_ready(protocol, thumbprint) + self.report_event("Provision succeed", is_success=True) - add_event(name="WALA", is_success=False, message=text(e), - op=WALAEventOperation.Provision) + def wait_for_ovfenv(self, max_retry=60): + """ + Wait for cloud-init to copy ovf-env.xml file from provision ISO + """ + ovf_file_path = os.path.join(OSUTIL.get_lib_dir(), ovf.OVF_FILE_NAME) + for retry in range(0, max_retry): + if os.path.isfile(ovf_file_path): + return + if retry < max_retry - 1: + logger.info("Wait for cloud-init to copy ovf-env.xml") + time.sleep(5) + raise ProvisionError("ovf-env.xml is not copied") def wait_for_ssh_host_key(self, max_retry=60): + """ + Wait for cloud-init to generate ssh host key + """ kepair_type = conf.get("Provisioning.SshHostKeyPairType", "rsa") path = '/etc/ssh/ssh_host_{0}_key'.format(kepair_type) for retry in range(0, max_retry): diff --git a/azurelinuxagent/event.py b/azurelinuxagent/event.py index 02e8017..af783f6 100644 --- a/azurelinuxagent/event.py +++ b/azurelinuxagent/event.py @@ -27,6 +27,7 @@ import platform import azurelinuxagent.logger as logger from azurelinuxagent.future import text import azurelinuxagent.protocol as prot +from azurelinuxagent.protocol.factory import PROT_FACTORY from azurelinuxagent.metadata import DISTRO_NAME, DISTRO_VERSION, \ DISTRO_CODE_NAME, AGENT_VERSION from azurelinuxagent.utils.osutil import OSUTIL @@ -67,7 +68,7 @@ class EventMonitor(object): self.sysinfo.append(prot.TelemetryEventParam("Processors", OSUTIL.get_processor_cores())) try: - protocol = prot.FACTORY.get_default_protocol() + protocol = PROT_FACTORY.get_protocol() vminfo = protocol.get_vminfo() self.sysinfo.append(prot.TelemetryEventParam("VMName", vminfo.vmName)) @@ -109,7 +110,7 @@ class EventMonitor(object): data = json.loads(data_str) except ValueError as e: logger.verb(data_str) - logger.error("Failed to decode json event file: {0}", e) + logger.verb("Failed to decode json event file: {0}", e) continue event = prot.TelemetryEvent() @@ -120,7 +121,7 @@ class EventMonitor(object): return try: - protocol = prot.FACTORY.get_default_protocol() + protocol = PROT_FACTORY.get_protocol() protocol.report_event(event_list) except prot.ProtocolError as e: logger.error("{0}", e) diff --git a/azurelinuxagent/protocol/__init__.py b/azurelinuxagent/protocol/__init__.py index a4572e6..b605ff7 100644 --- a/azurelinuxagent/protocol/__init__.py +++ b/azurelinuxagent/protocol/__init__.py @@ -18,6 +18,3 @@ # from azurelinuxagent.protocol.common import * -from azurelinuxagent.protocol.protocolFactory import FACTORY, \ - detect_default_protocol - diff --git a/azurelinuxagent/distro/default/dhcp.py b/azurelinuxagent/protocol/dhcp.py similarity index 80% rename from azurelinuxagent/distro/default/dhcp.py rename to azurelinuxagent/protocol/dhcp.py index 4fd23ef..31a0544 100644 --- a/azurelinuxagent/distro/default/dhcp.py +++ b/azurelinuxagent/protocol/dhcp.py @@ -19,53 +19,27 @@ import os import socket import array import time +import threading import azurelinuxagent.logger as logger from azurelinuxagent.utils.osutil import OSUTIL -from azurelinuxagent.exception import AgentNetworkError import azurelinuxagent.utils.fileutil as fileutil import azurelinuxagent.utils.shellutil as shellutil from azurelinuxagent.utils.textutil import * +from azurelinuxagent.protocol.common import ProtocolError -WIRE_SERVER_ADDR_FILE_NAME="WireServer" - -class DhcpHandler(object): - def __init__(self): - self.endpoint = None - self.gateway = None - self.routes = None - - def wait_for_network(self): - ipv4 = OSUTIL.get_ip4_addr() - while ipv4 == '' or ipv4 == '0.0.0.0': - logger.info("Waiting for network.") - time.sleep(10) - OSUTIL.start_network() - ipv4 = OSUTIL.get_ip4_addr() - - def probe(self): - logger.info("Send dhcp request") - self.wait_for_network() - mac_addr = OSUTIL.get_mac_addr() - req = build_dhcp_request(mac_addr) - resp = send_dhcp_request(req) - if resp is None: - logger.warn("Failed to detect wire server.") - return +DHCP_FILE_NAME = "DHCP" + +class DhcpResponse(object): + def __init__(self, resp): endpoint, gateway, routes = parse_dhcp_resp(resp) self.endpoint = endpoint - logger.info("Wire server endpoint:{0}", endpoint) - logger.info("Gateway:{0}", gateway) - logger.info("Routes:{0}", routes) - if endpoint is not None: - path = os.path.join(OSUTIL.get_lib_dir(), WIRE_SERVER_ADDR_FILE_NAME) - fileutil.write_file(path, endpoint) + logger.verb("Wire server endpoint:{0}", endpoint) + logger.verb("Gateway:{0}", gateway) + logger.verb("Routes:{0}", routes) self.gateway = gateway self.routes = routes self.conf_routes() - def get_endpoint(self): - return self.endpoint - def conf_routes(self): logger.info("Configure routes") #Add default gateway @@ -75,6 +49,52 @@ class DhcpHandler(object): for route in self.routes: OSUTIL.route_add(route[0], route[1], route[2]) +def _load_dhcp_resp(): + dhcp_file_path = os.path.join(OSUTIL.get_lib_dir(), DHCP_FILE_NAME) + resp = fileutil.read_file(dhcp_file_path, asbin=True) + return DhcpResponse(resp) + +def _fetch_dhcp_resp(): + logger.info("Send dhcp request") + mac_addr = OSUTIL.get_mac_addr() + req = build_dhcp_request(mac_addr) + resp = send_dhcp_request(req) + if resp is None: + raise ProtocolError("Failed to receive dhcp response.") + dhcp_file_path = os.path.join(OSUTIL.get_lib_dir(), DHCP_FILE_NAME) + try: + fileutil.write_file(dhcp_file_path, resp, asbin=True) + except IOError as e: + logger.warn("Failed to save dhcp response: {0}", e) + return DhcpResponse(resp) + +class DhcpClient(object): + def __init__(self): + self._resp = None + self._lock = threading.Lock() + + def get_dhcp_resp(self): + self._lock.acquire() + try: + if self._resp is None: + try: + self._resp = _load_dhcp_resp() + except IOError: + self._resp = _fetch_dhcp_resp() + return self._resp + finally: + self._lock.release() + + def fetch_dhcp_resp(self): + self._lock.acquire() + try: + self._resp = _fetch_dhcp_resp() + return self._resp + finally: + self._lock.release() + +DHCPCLIENT = DhcpClient() + def validate_dhcp_resp(request, response): bytes_recv = len(response) if bytes_recv < 0xF6: @@ -92,28 +112,25 @@ def validate_dhcp_resp(request, response): logger.verb("Cookie not match:\nsend={0},\nreceive={1}", hex_dump3(request, 0xEC, 4), hex_dump3(response, 0xEC, 4)) - raise AgentNetworkError("Cookie in dhcp respones " - "doesn't match the request") + raise ProtocolError("Cookie in dhcp respones doesn't match the request") if not compare_bytes(request, response, 4, 4): logger.verb("TransactionID not match:\nsend={0},\nreceive={1}", hex_dump3(request, 4, 4), hex_dump3(response, 4, 4)) - raise AgentNetworkError("TransactionID in dhcp respones " - "doesn't match the request") + raise ProtocolError("TransactionID in dhcp respones " + "doesn't match the request") if not compare_bytes(request, response, 0x1C, 6): logger.verb("Mac Address not match:\nsend={0},\nreceive={1}", hex_dump3(request, 0x1C, 6), hex_dump3(response, 0x1C, 6)) - raise AgentNetworkError("Mac Addr in dhcp respones " - "doesn't match the request") + raise ProtocolError("Mac Addr in dhcp respones " + "doesn't match the request") def parse_route(response, option, i, length, bytes_recv): # http://msdn.microsoft.com/en-us/library/cc227282%28PROT.10%29.aspx - logger.verb("Routes at offset: {0} with length:{1}", - hex(i), - hex(length)) + logger.verb("Routes at offset: {0} with length:{1}", hex(i), hex(length)) routes = [] if length < 5: logger.error("Data too small for option:{0}", option) @@ -169,9 +186,7 @@ def parse_dhcp_resp(response): if (i + 1) < bytes_recv: length = str_to_ord(response[i + 1]) logger.verb("DHCP option {0} at offset:{1} with length:{2}", - hex(option), - hex(i), - hex(length)) + hex(option), hex(i), hex(length)) if option == 255: logger.verb("DHCP packet ended at offset:{0}", hex(i)) break @@ -179,19 +194,14 @@ def parse_dhcp_resp(response): routes = parse_route(response, option, i, length, bytes_recv) elif option == 3: gateway = parse_ip_addr(response, option, i, length, bytes_recv) - logger.verb("Default gateway:{0}, at {1}", - gateway, - hex(i)) + logger.verb("Default gateway:{0}, at {1}", gateway, hex(i)) elif option == 245: endpoint = parse_ip_addr(response, option, i, length, bytes_recv) - logger.verb("Azure wire protocol endpoint:{0}, at {1}", - gateway, - hex(i)) + logger.verb("Azure wire protocol endpoint:{0}, at {1}", gateway, + hex(i)) else: logger.verb("Skipping DHCP option:{0} at {1} with length {2}", - hex(option), - hex(i), - hex(length)) + hex(option), hex(i), hex(length)) i += length + 2 return endpoint, gateway, routes @@ -237,7 +247,7 @@ def send_dhcp_request(request): response = socket_send(request) validate_dhcp_resp(request, response) return response - except AgentNetworkError as e: + except ProtocolError as e: logger.warn("Failed to send DHCP request: {0}", e) time.sleep(duration) return None @@ -257,7 +267,7 @@ def socket_send(request): response = sock.recv(1024) return response except IOError as e: - raise AgentNetworkError("{0}".format(e)) + raise ProtocolError("{0}".format(e)) finally: if sock is not None: sock.close() diff --git a/azurelinuxagent/protocol/factory.py b/azurelinuxagent/protocol/factory.py new file mode 100644 index 0000000..6c98f2e --- /dev/null +++ b/azurelinuxagent/protocol/factory.py @@ -0,0 +1,124 @@ +# Microsoft Azure Linux Agent +# +# Copyright 2014 Microsoft Corporation +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# Requires Python 2.4+ and Openssl 1.0+ +# +import os +import traceback +import time +import threading +import azurelinuxagent.logger as logger +from azurelinuxagent.exception import * +from azurelinuxagent.future import text +import azurelinuxagent.utils.fileutil as fileutil +from azurelinuxagent.utils.osutil import OSUTIL +from azurelinuxagent.protocol.common import * +from azurelinuxagent.protocol.v1 import WireProtocol +from azurelinuxagent.protocol.v2 import MetadataProtocol +from azurelinuxagent.protocol.ovfenv import TAG_FILE_NAME + +PROTOCOL_FILE_NAME = "Protocol" +MAX_RETRY = 60 + + +def _detect_protocol(protocols=[WireProtocol, MetadataProtocol]): + protocol_file_path = os.path.join(OSUTIL.get_lib_dir(), PROTOCOL_FILE_NAME) + if os.path.isfile(protocol_file_path): + os.remove(protocol_file_path) + for retry in range(0, MAX_RETRY): + for protocol_cls in protocols: + try: + logger.info("Detecting protocol: {0}", protocol_cls.__name__) + protocol = protocol_cls() + protocol.initialize() + logger.info("Found protocol: {0}", protocol_cls.__name__) + fileutil.write_file(protocol_file_path, protocol_cls.__name__) + return protocol + except ProtocolError as e: + logger.info("Protocol endpoint not found: {0}, {1}", + protocol_cls.__name__, e) + if retry < MAX_RETRY -1: + logger.info("Retry detect protocols: retry={0}", retry) + time.sleep(10) + raise ProtocolNotFound("No protocol found.") + +def _get_protocol(): + protocol_file_path = os.path.join(OSUTIL.get_lib_dir(), + PROTOCOL_FILE_NAME) + if not os.path.isfile(protocol_file_path): + raise ProtocolError("No protocl found") + + protocol_name = fileutil.read_file(protocol_file_path) + if protocol_name == WireProtocol.__name__: + return WireProtocol() + else: + return MetadataProtocol() + +class ProtocolFactory(object): + def __init__(self): + self.protocol = None + self.lock = threading.Lock() + + def detect_protocol(self): + logger.info("Detect protocol endpoints") + self.lock.acquire() + try: + if self.protocol is None: + self.protocol = _detect_protocol() + return self.protocol + finally: + self.lock.release() + + def detect_protocol_by_file(self): + logger.info("Detect protocol by file") + self.lock.acquire() + try: + tag_file_path = os.path.join(OSUTIL.get_lib_dir(), TAG_FILE_NAME) + if self.protocol is None: + if os.path.isfile(tag_file_path): + protocol = _detect_protocol(protocols=[MetadataProtocol]) + else: + protocol = _detect_protocol(protocols=[WireProtocol]) + self.protocol = protocol + return self.protocol + finally: + self.lock.release() + + def get_protocol(self): + """ + Get protocol detected + """ + self.lock.acquire() + try: + if self.protocol is None: + self.protocol = _get_protocol() + return self.protocol + finally: + self.lock.release() + return self.protocol + + def wait_for_network(self): + """ + Wait for network stack to be initialized + """ + ipv4 = OSUTIL.get_ip4_addr() + while ipv4 == '' or ipv4 == '0.0.0.0': + logger.info("Waiting for network.") + time.sleep(10) + OSUTIL.start_network() + ipv4 = OSUTIL.get_ip4_addr() + +PROT_FACTORY = ProtocolFactory() diff --git a/azurelinuxagent/protocol/ovfenv.py b/azurelinuxagent/protocol/ovfenv.py index 906572d..a836aaa 100644 --- a/azurelinuxagent/protocol/ovfenv.py +++ b/azurelinuxagent/protocol/ovfenv.py @@ -17,19 +17,24 @@ # Requires Python 2.4+ and Openssl 1.0+ # """ -Copy and parse ovf-env.xml from provisiong ISO and local cache +Copy and parse ovf-env.xml from provisioning ISO and local cache """ import os import re +import shutil import xml.dom.minidom as minidom import azurelinuxagent.logger as logger from azurelinuxagent.future import text import azurelinuxagent.utils.fileutil as fileutil from azurelinuxagent.utils.textutil import parse_doc, findall, find, findtext from azurelinuxagent.utils.osutil import OSUTIL, OSUtilError -from azurelinuxagent.protocol import ProtocolError +from azurelinuxagent.protocol.common import ProtocolError OVF_FILE_NAME = "ovf-env.xml" + +#Tag file to indicate usage of metadata protocol +TAG_FILE_NAME = "useMetadataEndpoint.tag" + OVF_VERSION = "1.0" OVF_NAME_SPACE = "http://schemas.dmtf.org/ovf/environment/1" WA_NAME_SPACE = "http://schemas.microsoft.com/windowsazure" @@ -52,12 +57,21 @@ def copy_ovf_env(): """ try: OSUTIL.mount_dvd() + ovf_file_path_on_dvd = OSUTIL.get_ovf_env_file_path_on_dvd() ovfxml = fileutil.read_file(ovf_file_path_on_dvd, remove_bom=True) ovfenv = OvfEnv(ovfxml) ovfxml = re.sub(".*?<", "*<", ovfxml) ovf_file_path = os.path.join(OSUTIL.get_lib_dir(), OVF_FILE_NAME) fileutil.write_file(ovf_file_path, ovfxml) + + tag_file_path_on_dvd = os.path.join(OSUTIL.get_dvd_mount_point(), + TAG_FILE_NAME) + if os.path.isfile(tag_file_path_on_dvd): + logger.info("Found {0} in provisioning ISO", TAG_FILE_NAME) + tag_file_path = os.path.join(OSUTIL.get_lib_dir(), TAG_FILE_NAME) + shutil.copyfile(tag_file_path_on_dvd, tag_file_path) + OSUTIL.umount_dvd() OSUTIL.eject_dvd() except IOError as e: diff --git a/azurelinuxagent/protocol/protocolFactory.py b/azurelinuxagent/protocol/protocolFactory.py deleted file mode 100644 index 8652147..0000000 --- a/azurelinuxagent/protocol/protocolFactory.py +++ /dev/null @@ -1,113 +0,0 @@ -# Microsoft Azure Linux Agent -# -# Copyright 2014 Microsoft Corporation -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# -# Requires Python 2.4+ and Openssl 1.0+ -# -import os -import traceback -import threading -import azurelinuxagent.logger as logger -from azurelinuxagent.future import text -import azurelinuxagent.utils.fileutil as fileutil -from azurelinuxagent.utils.osutil import OSUTIL -from azurelinuxagent.protocol.common import * -from azurelinuxagent.protocol.v1 import WireProtocol -from azurelinuxagent.protocol.v2 import MetadataProtocol - -WIRE_SERVER_ADDR_FILE_NAME = "WireServer" - -def get_wire_protocol_endpoint(): - path = os.path.join(OSUTIL.get_lib_dir(), WIRE_SERVER_ADDR_FILE_NAME) - try: - endpoint = fileutil.read_file(path) - except IOError as e: - raise ProtocolNotFound("Wire server endpoint not found: {0}".format(e)) - - if endpoint is None: - raise ProtocolNotFound("Wire server endpoint is None") - - return endpoint - -def detect_wire_protocol(): - endpoint = get_wire_protocol_endpoint() - - protocol = WireProtocol(endpoint) - protocol.initialize() - logger.info("Protocol V1 found.") - return protocol - -def detect_metadata_protocol(): - protocol = MetadataProtocol() - protocol.initialize() - - logger.info("Protocol V2 found.") - return protocol - -def detect_available_protocols(prob_funcs=[detect_wire_protocol, - detect_metadata_protocol]): - available_protocols = [] - for probe_func in prob_funcs: - try: - protocol = probe_func() - available_protocols.append(protocol) - except ProtocolNotFound as e: - logger.info(text(e)) - return available_protocols - -def detect_default_protocol(): - logger.info("Detect default protocol.") - available_protocols = detect_available_protocols() - return choose_default_protocol(available_protocols) - -def choose_default_protocol(protocols): - if len(protocols) > 0: - return protocols[0] - else: - raise ProtocolNotFound("No available protocol detected.") - -def get_wire_protocol(): - endpoint = get_wire_protocol_endpoint() - return WireProtocol(endpoint) - -def get_metadata_protocol(): - return MetadataProtocol() - -def get_available_protocols(getters=[get_wire_protocol, get_metadata_protocol]): - available_protocols = [] - for getter in getters: - try: - protocol = getter() - available_protocols.append(protocol) - except ProtocolNotFound as e: - logger.info(text(e)) - return available_protocols - -class ProtocolFactory(object): - def __init__(self): - self._protocol = None - self._lock = threading.Lock() - - def get_default_protocol(self): - if self._protocol is None: - self._lock.acquire() - if self._protocol is None: - available_protocols = get_available_protocols() - self._protocol = choose_default_protocol(available_protocols) - self._lock.release() - - return self._protocol - -FACTORY = ProtocolFactory() diff --git a/azurelinuxagent/protocol/v1.py b/azurelinuxagent/protocol/v1.py index 2f959ba..772ee38 100644 --- a/azurelinuxagent/protocol/v1.py +++ b/azurelinuxagent/protocol/v1.py @@ -33,6 +33,7 @@ from azurelinuxagent.utils.osutil import OSUTIL import azurelinuxagent.utils.fileutil as fileutil import azurelinuxagent.utils.shellutil as shellutil from azurelinuxagent.protocol.common import * +import azurelinuxagent.protocol.dhcp as dhcp VERSION_INFO_URI = "http://{0}/?comp=versions" GOAL_STATE_URI = "http://{0}/machine/?comp=goalstate" @@ -54,19 +55,23 @@ TRANSPORT_CERT_FILE_NAME = "TransportCert.pem" TRANSPORT_PRV_FILE_NAME = "TransportPrivate.pem" PROTOCOL_VERSION = "2012-11-30" +ENDPOINT_FINE_NAME = "WireServer" class WireProtocolResourceGone(ProtocolError): pass class WireProtocol(Protocol): - def __init__(self, endpoint): - self.client = WireClient(endpoint) + def __init__(self): + dhcp_resp = dhcp.DHCPCLIENT.get_dhcp_resp() + self.client = WireClient(dhcp_resp.endpoint) def initialize(self): + dhcp_resp = dhcp.DHCPCLIENT.fetch_dhcp_resp() + self.client = WireClient(dhcp_resp.endpoint) + self.client.check_wire_protocol_version() OSUTIL.gen_transport_cert(TRANSPORT_PRV_FILE_NAME, TRANSPORT_CERT_FILE_NAME) - self.client.check_wire_protocol_version() self.client.update_goal_state(forced=True) def get_vminfo(self): @@ -462,6 +467,8 @@ def event_to_v1(event): class WireClient(object): def __init__(self, endpoint): + if endpoint is None: + raise ProtocolError("WireProtocl endpoint is None") self.endpoint = endpoint self.goal_state = None self.updated = None diff --git a/azurelinuxagent/protocol/v2.py b/azurelinuxagent/protocol/v2.py index d0c8d2f..f5adf86 100644 --- a/azurelinuxagent/protocol/v2.py +++ b/azurelinuxagent/protocol/v2.py @@ -108,6 +108,7 @@ class MetadataProtocol(Protocol): return textutil.get_bytes_from_pem(content) def initialize(self): + self.get_vminfo() trans_prv_file = os.path.join(OSUTIL.get_lib_dir(), TRANSPORT_PRV_FILE_NAME) trans_crt_file = os.path.join(OSUTIL.get_lib_dir(), @@ -122,18 +123,7 @@ class MetadataProtocol(Protocol): "{0}.crt".format(thumbprint)) shutil.copyfile(trans_prv_file, prv_file) shutil.copyfile(trans_crt_file, crt_file) - - #TODO remote workarround for azure stack test - for retry in range(0, MAX_PING): - try: - self.get_vminfo() - return - except ProtocolError as e: - logger.warn("Metadata server is not ready, retry = {0}", retry) - if retry < MAX_PING - 1: - time.sleep(RETRY_PING_INTERVAL) - raise ProtocolNotFound("Metadata server endpoint is not reachable") - + def get_vminfo(self): vminfo = VMInfo() data = self._get_data(self.identity_uri) diff --git a/tests/test_deprovision.py b/tests/test_deprovision.py index 8bad6b9..38da19e 100644 --- a/tests/test_deprovision.py +++ b/tests/test_deprovision.py @@ -36,7 +36,7 @@ def MockSetup(self, deluser): class TestDeprovisionHandler(unittest.TestCase): def test_setup(self): - handler = deprovision_handler.DeprovisionHandler() + handler = deprovision_handler.DeprovisionHandler(None) warnings, actions = handler.setup(False) self.assertNotEquals(None, warnings) self.assertNotEquals(0, len(warnings)) @@ -47,7 +47,7 @@ class TestDeprovisionHandler(unittest.TestCase): @mock(deprovision_handler.DeprovisionHandler, 'setup', MockSetup) def test_deprovision(self): - handler = deprovision_handler.DeprovisionHandler() + handler = deprovision_handler.DeprovisionHandler(None) handler.deprovision(force=True) if __name__ == '__main__': diff --git a/tests/test_dhcp.py b/tests/test_dhcp.py index 2206325..40ee54c 100644 --- a/tests/test_dhcp.py +++ b/tests/test_dhcp.py @@ -25,7 +25,7 @@ import unittest import os import json import azurelinuxagent.utils.fileutil as fileutil -import azurelinuxagent.distro.default.dhcp as dhcp_handler +import azurelinuxagent.protocol.dhcp as dhcp SampleDhcpResponse = None with open(os.path.join(env.test_root, "dhcp"), 'rb') as F: @@ -39,30 +39,28 @@ mock_send_dhcp_failed = MockFunc(retval=None) class TestdhcpHandler(unittest.TestCase): def test_build_dhcp_req(self): - req = dhcp_handler.build_dhcp_request(mock_get_mac_addr()) + req = dhcp.build_dhcp_request(mock_get_mac_addr()) self.assertNotEquals(None, req) - @mock(dhcp_handler, "gen_trans_id", mock_gen_trans_id) - @mock(dhcp_handler, "socket_send", mock_socket_send) + @mock(dhcp, "gen_trans_id", mock_gen_trans_id) + @mock(dhcp, "socket_send", mock_socket_send) def test_send_dhcp_req(self): - req = dhcp_handler.build_dhcp_request(mock_get_mac_addr()) - resp = dhcp_handler.send_dhcp_request(req) + req = dhcp.build_dhcp_request(mock_get_mac_addr()) + resp = dhcp.send_dhcp_request(req) self.assertNotEquals(None, resp) - @mock(dhcp_handler, "send_dhcp_request", mock_send_dhcp_failed) + @mock(dhcp, "send_dhcp_request", mock_send_dhcp_failed) def test_send_dhcp_failed(self): - dhcp = dhcp_handler.DhcpHandler() - dhcp.probe() + dhcp_resp = dhcp.DHCPCLIENT.get_dhcp_resp() - @mock(dhcp_handler, "socket_send", mock_socket_send) - @mock(dhcp_handler, "gen_trans_id", mock_gen_trans_id) - @mock(dhcp_handler.OSUTIL, "get_mac_addr", mock_get_mac_addr) - @mock(dhcp_handler.fileutil, "write_file", MockFunc()) + @mock(dhcp, "socket_send", mock_socket_send) + @mock(dhcp, "gen_trans_id", mock_gen_trans_id) + @mock(dhcp.OSUTIL, "get_mac_addr", mock_get_mac_addr) + @mock(dhcp.fileutil, "write_file", MockFunc()) def test_handle_dhcp(self): - dh = dhcp_handler.DhcpHandler() - dh.probe() - self.assertEquals("10.62.144.1", dh.gateway) - self.assertEquals("10.62.144.140", dh.endpoint) + dhcp_resp = dhcp.DHCPCLIENT.get_dhcp_resp() + self.assertEquals("10.62.144.1", dhcp_resp.gateway) + self.assertEquals("10.62.144.140", dhcp_resp.endpoint) if __name__ == '__main__': unittest.main() diff --git a/tests/test_distroLoader.py b/tests/test_distroLoader.py index 16987c5..8e522d3 100644 --- a/tests/test_distroLoader.py +++ b/tests/test_distroLoader.py @@ -31,7 +31,6 @@ class TestDistroLoader(unittest.TestCase): self.assertNotEquals(None, HANDLERS.init_handler) self.assertNotEquals(None, HANDLERS.main_handler) self.assertNotEquals(None, HANDLERS.scvmm_handler) - self.assertNotEquals(None, HANDLERS.dhcp_handler) self.assertNotEquals(None, HANDLERS.env_handler) self.assertNotEquals(None, HANDLERS.provision_handler) self.assertNotEquals(None, HANDLERS.resource_disk_handler) diff --git a/tests/test_envmon.py b/tests/test_envmon.py index 74b61ee..6941ec3 100644 --- a/tests/test_envmon.py +++ b/tests/test_envmon.py @@ -22,14 +22,18 @@ import tests.env from tests.tools import * import unittest import time +import azurelinuxagent.protocol.dhcp as dhcp from azurelinuxagent.future import text from azurelinuxagent.utils.osutil import OSUTIL from azurelinuxagent.distro.default.env import EnvMonitor -class MockDhcpHandler(object): +class MockDhcpResponse(object): def conf_routes(self): pass +def mock_get_dhcp_resp(self): + return MockDhcpResponse() + def mock_get_dhcp_pid(): return "1234" @@ -39,13 +43,15 @@ def mock_dhcp_pid_change(): class TestEnvMonitor(unittest.TestCase): @mock(OSUTIL, 'get_dhcp_pid', mock_get_dhcp_pid) + @mock(dhcp.DHCPCLIENT, 'get_dhcp_resp', mock_get_dhcp_resp) def test_dhcp_pid_not_change(self): - monitor = EnvMonitor(MockDhcpHandler()) + monitor = EnvMonitor() monitor.handle_dhclient_restart() @mock(OSUTIL, 'get_dhcp_pid', mock_dhcp_pid_change) + @mock(dhcp.DHCPCLIENT, 'get_dhcp_resp', mock_get_dhcp_resp) def test_dhcp_pid_change(self): - monitor = EnvMonitor(MockDhcpHandler()) + monitor = EnvMonitor() monitor.handle_dhclient_restart() if __name__ == '__main__': diff --git a/tests/test_event.py b/tests/test_event.py index fcf67c9..5b66b12 100644 --- a/tests/test_event.py +++ b/tests/test_event.py @@ -27,6 +27,7 @@ import shutil import azurelinuxagent.utils.fileutil as fileutil import azurelinuxagent.event as evt import azurelinuxagent.protocol as prot +from azurelinuxagent.protocol.factory import PROT_FACTORY class MockProtocol(object): def get_vminfo(self): @@ -42,8 +43,7 @@ class TestEvent(unittest.TestCase): self.assertNotEquals(0, len(eventsFile)) shutil.rmtree("/tmp/events") - @mock(evt.prot.FACTORY, 'get_default_protocol', - MockFunc(retval=MockProtocol())) + @mock(PROT_FACTORY, 'get_protocol', MockFunc(retval=MockProtocol())) def test_init_sys_info(self): monitor = evt.EventMonitor() monitor.init_sysinfo() diff --git a/tests/test_ext.py b/tests/test_ext.py index 355f099..7414d63 100644 --- a/tests/test_ext.py +++ b/tests/test_ext.py @@ -27,7 +27,7 @@ import json import azurelinuxagent.logger as logger from azurelinuxagent.utils.osutil import OSUTIL import azurelinuxagent.utils.fileutil as fileutil -import azurelinuxagent.protocol as prot +import azurelinuxagent.protocol.common as prot import azurelinuxagent.distro.default.extension as ext ext_sample_json = { diff --git a/tests/test_protocolFactory.py b/tests/test_protocolFactory.py deleted file mode 100644 index 9928b88..0000000 --- a/tests/test_protocolFactory.py +++ /dev/null @@ -1,37 +0,0 @@ -# Copyright 2014 Microsoft Corporation -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -# -# Requires Python 2.4+ and Openssl 1.0+ -# -# Implements parts of RFC 2131, 1541, 1497 and -# http://msdn.microsoft.com/en-us/library/cc227282%28PROT.10%29.aspx -# http://msdn.microsoft.com/en-us/library/cc227259%28PROT.13%29.aspx - -import tests.env -from tests.tools import * -import uuid -import unittest -import os -import azurelinuxagent.protocol as protocol -import azurelinuxagent.protocol.protocolFactory as protocolFactory - -class TestWireProtocolEndpoint(unittest.TestCase): - def test_get_available_protocols(self): - mockGetV1 = MockFunc(retval="Mock protocol") - protocols = protocolFactory.get_available_protocols([mockGetV1]) - self.assertNotEquals(None, protocols) - self.assertNotEquals(0, len(protocols)) - -if __name__ == '__main__': - unittest.main() diff --git a/tests/test_resourcedisk.py b/tests/test_resourcedisk.py index de54fd3..7c568b7 100644 --- a/tests/test_resourcedisk.py +++ b/tests/test_resourcedisk.py @@ -44,19 +44,19 @@ class TestResourceDisk(unittest.TestCase): @mock(rdh.shellutil, 'run_get_output', MockFunc(retval=(0, gpt_output_sample))) @mock(rdh.shellutil, 'run', MockFunc(retval=0)) def test_mountGPT(self): - handler = rdh.ResourceDiskHandler() + handler = rdh.ResourceDiskHandler(None) handler.mount_resource_disk('/tmp/foo', 'ext4') @mock(rdh.OSUTIL, 'device_for_ide_port', MockFunc(retval='foo')) @mock(rdh.shellutil, 'run_get_output', MockFunc(retval=(0, ""))) @mock(rdh.shellutil, 'run', MockFunc(retval=0)) def test_mountMBR(self): - handler = rdh.ResourceDiskHandler() + handler = rdh.ResourceDiskHandler(None) handler.mount_resource_disk('/tmp/foo', 'ext4') @mock(rdh.shellutil, 'run', MockFunc(retval=0)) def test_createSwapSpace(self): - handler = rdh.ResourceDiskHandler() + handler = rdh.ResourceDiskHandler(None) handler.create_swap_space('/tmp/foo', 512) if __name__ == '__main__':