diff --git a/README.md b/README.md index d8b58ec..eea9566 100644 --- a/README.md +++ b/README.md @@ -73,4 +73,5 @@ This project is licensed under the MIT License. See the [LICENSE](LICENSE) file - Henry Birge-Lee - Grace Cimaszewski -- Dmitry Sharkov \ No newline at end of file +- Dmitry Sharkov +- Alan Hanafy \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 5322872..f9ba41d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -30,6 +30,7 @@ dependencies = [ "aiohttp==3.13.4", "cryptography==46.0.6", "uritools==6.0.1", + "opentelemetry-api==1.41.1", ] [project.optional-dependencies] diff --git a/src/open_mpic_core/__init__.py b/src/open_mpic_core/__init__.py index bcf256b..618fe24 100644 --- a/src/open_mpic_core/__init__.py +++ b/src/open_mpic_core/__init__.py @@ -40,6 +40,7 @@ from open_mpic_core.common_util.domain_encoder import DomainEncoder from open_mpic_core.common_util.trace_level_logger import get_logger from open_mpic_core.common_util.trace_level_logger import TRACE_LEVEL +from open_mpic_core.common_util.telemetry import get_meter, get_tracer from open_mpic_core.mpic_coordinator.domain.remote_perspective import RemotePerspective from open_mpic_core.mpic_coordinator.domain.mpic_orchestration_parameters import ( diff --git a/src/open_mpic_core/common_util/telemetry.py b/src/open_mpic_core/common_util/telemetry.py new file mode 100644 index 0000000..f6e75fc --- /dev/null +++ b/src/open_mpic_core/common_util/telemetry.py @@ -0,0 +1,33 @@ +"""OpenTelemetry accessor helpers for open-mpic-core. + +Depends only on ``opentelemetry-api``. When no SDK provider is registered the +returned Meter/Tracer objects are silent no-ops — the library is safe to use +without any OTEL configuration. + +Container services configure real providers at startup (before any requests are +processed), after which every ``get_meter`` / ``get_tracer`` call in the core +routes through those providers automatically. +""" + +from open_mpic_core.__about__ import __version__ +from opentelemetry import metrics, trace + +_INSTRUMENTATION_VERSION: str = __version__ + + +def get_meter(name: str) -> metrics.Meter: + """Return a :class:`~opentelemetry.metrics.Meter` scoped to *name*. + + Resolves through the global ``MeterProvider``; returns a no-op Meter when + no SDK provider has been registered. + """ + return metrics.get_meter(name, version=_INSTRUMENTATION_VERSION) + + +def get_tracer(name: str) -> trace.Tracer: + """Return a :class:`~opentelemetry.trace.Tracer` scoped to *name*. + + Resolves through the global ``TracerProvider``; returns a no-op Tracer when + no SDK provider has been registered. + """ + return trace.get_tracer(name, instrumenting_library_version=_INSTRUMENTATION_VERSION) diff --git a/src/open_mpic_core/mpic_caa_checker/mpic_caa_checker.py b/src/open_mpic_core/mpic_caa_checker/mpic_caa_checker.py index 58c2971..c86e7f9 100644 --- a/src/open_mpic_core/mpic_caa_checker/mpic_caa_checker.py +++ b/src/open_mpic_core/mpic_caa_checker/mpic_caa_checker.py @@ -6,10 +6,13 @@ from dns.name import Name from dns.rrset import RRset +from opentelemetry.trace import Status, StatusCode + from open_mpic_core import CaaCheckRequest, CaaCheckResponse, CaaCheckResponseDetails from open_mpic_core import MpicValidationError, ErrorMessages from open_mpic_core import DomainEncoder from open_mpic_core import get_logger +from open_mpic_core import get_meter, get_tracer from open_mpic_core import CertificateType ISSUE_TAG: Final[str] = "issue" @@ -47,18 +50,40 @@ def __init__( dns_resolution_lifetime if dns_resolution_lifetime is not None else self.resolver.lifetime ) + _meter = get_meter(__name__) + self._tracer = get_tracer(__name__) + self._request_counter = _meter.create_counter( + "mpic.caa.requests", + description="Total CAA check requests processed", + unit="1", + ) + self._duration_histogram = _meter.create_histogram( + "mpic.caa.duration", + description="CAA check request duration", + unit="ms", + ) + self._dns_duration_histogram = _meter.create_histogram( + "mpic.caa.dns_lookup.duration", + description="CAA DNS lookup duration", + unit="ms", + ) + async def find_caa_records_and_domain(self, caa_request) -> tuple[RRset, Name]: + _dns_start_ns = time.perf_counter_ns() rrset = None domain = dns.name.from_text(caa_request.domain_or_ip_target) - - while domain != dns.name.root: + with self._tracer.start_as_current_span("mpic.caa.dns_lookup"): try: - lookup = await self.resolver.resolve(domain, dns.rdatatype.CAA) - rrset = lookup.rrset - break - except (dns.resolver.NoAnswer, dns.resolver.NXDOMAIN): - domain = domain.parent() - # will raise other exceptions that we want to catch in the calling function + while domain != dns.name.root: + try: + lookup = await self.resolver.resolve(domain, dns.rdatatype.CAA) + rrset = lookup.rrset + break + except (dns.resolver.NoAnswer, dns.resolver.NXDOMAIN): + domain = domain.parent() + # will raise other exceptions that we want to catch in the calling function + finally: + self._dns_duration_histogram.record((time.perf_counter_ns() - _dns_start_ns) / 1_000_000) return rrset, domain @@ -94,46 +119,66 @@ async def check_caa(self, caa_request: CaaCheckRequest) -> CaaCheckResponse: timestamp_ns=None, ) - try: - # encode domain if needed - caa_request.domain_or_ip_target = DomainEncoder.prepare_target_for_lookup(caa_request.domain_or_ip_target) - - # noinspection PyUnresolvedReferences - async with self.logger.trace_timing(f"CAA lookup for target {caa_request.domain_or_ip_target}"): - rrset, domain = await self.find_caa_records_and_domain(caa_request) - caa_found = rrset is not None - except Exception as e: - error_encountered = True - caa_lookup_error = e - error_message = f"Error during CAA lookup for {caa_request.domain_or_ip_target}: {e}. Trace ID: {caa_request.trace_identifier}" - self.logger.error(error_message) - caa_check_response.errors = [MpicValidationError.create(ErrorMessages.CAA_LOOKUP_ERROR, error_message)] - caa_check_response.details.found_at = None - caa_check_response.details.records_seen = None - - if error_encountered: # if there was an error during lookup - # check if allow_lookup_failure is set to True, and allow issuance depending on error - if isinstance(caa_lookup_error, (dns.resolver.LifetimeTimeout, dns.resolver.NoNameservers)): - if caa_request.caa_check_parameters and caa_request.caa_check_parameters.allow_lookup_failure: - # if the error was from the lookup process itself (e.g. timeout), allow issuance - caa_check_response.check_completed = True - caa_check_response.check_passed = True - elif not caa_found: # if domain has no CAA records: valid for issuance - caa_check_response.check_completed = True - caa_check_response.check_passed = True - caa_check_response.details.caa_record_present = False - caa_check_response.details.found_at = None - caa_check_response.details.records_seen = None - else: - caa_check_response.check_completed = True - valid_for_issuance = MpicCaaChecker.is_valid_for_issuance( - caa_domains, certificate_type, is_wc_domain, rrset + _start_ns = time.perf_counter_ns() + with self._tracer.start_as_current_span("mpic.caa.check") as _span: + try: + # encode domain if needed + caa_request.domain_or_ip_target = DomainEncoder.prepare_target_for_lookup( + caa_request.domain_or_ip_target + ) + + # noinspection PyUnresolvedReferences + async with self.logger.trace_timing(f"CAA lookup for target {caa_request.domain_or_ip_target}"): + rrset, domain = await self.find_caa_records_and_domain(caa_request) + caa_found = rrset is not None + except Exception as e: + error_encountered = True + caa_lookup_error = e + error_message = f"Error during CAA lookup for {caa_request.domain_or_ip_target}: {e}. Trace ID: {caa_request.trace_identifier}" + self.logger.error(error_message) + caa_check_response.errors = [MpicValidationError.create(ErrorMessages.CAA_LOOKUP_ERROR, error_message)] + caa_check_response.details.found_at = None + caa_check_response.details.records_seen = None + + if error_encountered: # if there was an error during lookup + _span.record_exception(caa_lookup_error) + _span.set_status(Status(StatusCode.ERROR, description=type(caa_lookup_error).__name__)) + # check if allow_lookup_failure is set to True, and allow issuance depending on error + if isinstance(caa_lookup_error, (dns.resolver.LifetimeTimeout, dns.resolver.NoNameservers)): + if caa_request.caa_check_parameters and caa_request.caa_check_parameters.allow_lookup_failure: + # if the error was from the lookup process itself (e.g. timeout), allow issuance + caa_check_response.check_completed = True + caa_check_response.check_passed = True + elif not caa_found: # if domain has no CAA records: valid for issuance + caa_check_response.check_completed = True + caa_check_response.check_passed = True + caa_check_response.details.caa_record_present = False + caa_check_response.details.found_at = None + caa_check_response.details.records_seen = None + else: + caa_check_response.check_completed = True + valid_for_issuance = MpicCaaChecker.is_valid_for_issuance( + caa_domains, certificate_type, is_wc_domain, rrset + ) + caa_check_response.check_passed = valid_for_issuance + caa_check_response.details.caa_record_present = True + caa_check_response.details.found_at = domain.to_text(omit_final_dot=True) + caa_check_response.details.records_seen = [record_data.to_text() for record_data in rrset] + caa_check_response.timestamp_ns = time.time_ns() + + elapsed_ms = (time.perf_counter_ns() - _start_ns) / 1_000_000 + self._duration_histogram.record( + elapsed_ms, + {"check.passed": caa_check_response.check_passed}, + ) + self._request_counter.add( + 1, + { + "check.passed": caa_check_response.check_passed, + "check.completed": caa_check_response.check_completed, + "caa.lookup_error": error_encountered, + }, ) - caa_check_response.check_passed = valid_for_issuance - caa_check_response.details.caa_record_present = True - caa_check_response.details.found_at = domain.to_text(omit_final_dot=True) - caa_check_response.details.records_seen = [record_data.to_text() for record_data in rrset] - caa_check_response.timestamp_ns = time.time_ns() # noinspection PyUnresolvedReferences self.logger.trace(f"Completed CAA for {caa_request.domain_or_ip_target}") diff --git a/src/open_mpic_core/mpic_coordinator/mpic_coordinator.py b/src/open_mpic_core/mpic_coordinator/mpic_coordinator.py index 73116f1..c896e36 100644 --- a/src/open_mpic_core/mpic_coordinator/mpic_coordinator.py +++ b/src/open_mpic_core/mpic_coordinator/mpic_coordinator.py @@ -1,10 +1,11 @@ import asyncio import json +import time from itertools import cycle - from pprint import pformat -import time + import hashlib +from opentelemetry.trace import Status, StatusCode from open_mpic_core import CaaCheckResponse, DcvCheckResponse, CaaCheckResponseDetails from open_mpic_core import MpicRequest, MpicResponse, PerspectiveResponse @@ -22,6 +23,7 @@ from open_mpic_core import MpicRequestValidator from open_mpic_core import MpicResponseBuilder from open_mpic_core import get_logger +from open_mpic_core import get_meter, get_tracer logger = get_logger(__name__) @@ -59,104 +61,168 @@ def __init__( if log_level is not None: self.logger.setLevel(log_level) + _meter = get_meter(__name__) + self._tracer = get_tracer(__name__) + self._request_counter = _meter.create_counter( + "mpic.requests", + description="Total MPIC requests processed", + unit="1", + ) + self._coordinator_duration = _meter.create_histogram( + "mpic.coordinator.duration", + description="MPIC coordinator request processing duration", + unit="ms", + ) + self._perspective_response_counter = _meter.create_counter( + "mpic.perspective_responses", + description="Perspective check responses by outcome", + unit="1", + ) + self._remote_error_counter = _meter.create_counter( + "mpic.remote_check.errors", + description="Remote perspective check transport errors", + unit="1", + ) + # noinspection PyInconsistentReturns,PyTypeChecker async def coordinate_mpic(self, mpic_request: MpicRequest) -> MpicResponse: # noinspection PyUnresolvedReferences self.logger.trace(f"Coordinating MPIC request with trace ID {mpic_request.trace_identifier}") - self._raise_exception_on_invalid_request(mpic_request) - - orchestration_parameters = mpic_request.orchestration_parameters + check_type_value = mpic_request.check_type.value + _start_ns = time.perf_counter_ns() + _is_valid = False + _request_status = "error" - perspective_count = self.default_perspective_count - if orchestration_parameters is not None and orchestration_parameters.perspective_count is not None: - perspective_count = orchestration_parameters.perspective_count - - perspective_cohorts = self.shuffle_and_group_perspectives( - self.target_perspectives, perspective_count, mpic_request.domain_or_ip_target - ) + with self._tracer.start_as_current_span( + "mpic.coordinate", attributes={"check.type": check_type_value} + ) as _span: + try: + self._raise_exception_on_invalid_request(mpic_request) - if len(perspective_cohorts) == 0: - raise CohortCreationException(ErrorMessages.COHORT_CREATION_ERROR.message.format(perspective_count)) - - # check if a specific cohort is requested for single attempt - cohort_to_use = None - if orchestration_parameters is not None and orchestration_parameters.cohort_for_single_attempt is not None: - cohort_to_use = orchestration_parameters.cohort_for_single_attempt - if not MpicRequestValidator.is_requested_cohort_for_single_attempt_valid( - cohort_to_use, len(perspective_cohorts) - ): - raise CohortSelectionException(ErrorMessages.COHORT_SELECTION_ERROR.message.format(cohort_to_use)) - - quorum_count = self.determine_required_quorum_count(orchestration_parameters, perspective_count) - - if ( - orchestration_parameters is not None - and orchestration_parameters.max_attempts is not None - and orchestration_parameters.max_attempts > 0 - and orchestration_parameters.cohort_for_single_attempt is None - ): - max_attempts = orchestration_parameters.max_attempts - if self.global_max_attempts is not None and max_attempts > self.global_max_attempts: - max_attempts = self.global_max_attempts - else: - max_attempts = 1 - attempts = 1 - previous_attempt_results = None - cohort_cycle = cycle(perspective_cohorts) - - while attempts <= max_attempts: - if cohort_to_use is not None: - perspectives_to_use = perspective_cohorts[cohort_to_use - 1] # cohorts are 1-indexed for the user - else: - perspectives_to_use = next(cohort_cycle) + orchestration_parameters = mpic_request.orchestration_parameters - # Collect async calls to invoke for each perspective. - async_calls_to_issue = MpicCoordinator.collect_checker_calls_to_issue(mpic_request, perspectives_to_use) + perspective_count = self.default_perspective_count + if orchestration_parameters is not None and orchestration_parameters.perspective_count is not None: + perspective_count = orchestration_parameters.perspective_count - perspective_responses = await self.call_checkers_and_collect_responses( - mpic_request, perspectives_to_use, async_calls_to_issue - ) - - check_passed_per_perspective = { - response.perspective_code: response.check_response.check_passed for response in perspective_responses - } - - valid_perspective_count = sum(check_passed_per_perspective.values()) - is_valid_result = valid_perspective_count >= quorum_count - - # noinspection PyUnresolvedReferences - self.logger.trace(f"Perspectives used in attempt: \n%s", pformat(perspectives_to_use)) - # noinspection PyUnresolvedReferences - self.logger.trace(f"Check passed per perspective: \n%s", pformat(check_passed_per_perspective)) - - # if cohort size is larger than 2, then at least two RIRs must be represented in the SUCCESSFUL perspectives - if len(perspectives_to_use) > 2: - valid_perspectives = [ - perspective for perspective in perspectives_to_use if check_passed_per_perspective[perspective.code] - ] - rir_count = len(set(perspective.rir for perspective in valid_perspectives)) - is_valid_result = rir_count >= 2 and is_valid_result - - if is_valid_result or attempts == max_attempts: - response = MpicResponseBuilder.build_response( - mpic_request, - perspective_count, - quorum_count, - attempts, - perspective_responses, - is_valid_result, - previous_attempt_results, + perspective_cohorts = self.shuffle_and_group_perspectives( + self.target_perspectives, perspective_count, mpic_request.domain_or_ip_target ) - # noinspection PyUnresolvedReferences - self.logger.trace(f"Completed MPIC request with trace ID {mpic_request.trace_identifier}") - return response - else: - if previous_attempt_results is None: - previous_attempt_results = [] - previous_attempt_results.append(perspective_responses) - attempts += 1 + if len(perspective_cohorts) == 0: + raise CohortCreationException(ErrorMessages.COHORT_CREATION_ERROR.message.format(perspective_count)) + + # check if a specific cohort is requested for single attempt + cohort_to_use = None + if ( + orchestration_parameters is not None + and orchestration_parameters.cohort_for_single_attempt is not None + ): + cohort_to_use = orchestration_parameters.cohort_for_single_attempt + if not MpicRequestValidator.is_requested_cohort_for_single_attempt_valid( + cohort_to_use, len(perspective_cohorts) + ): + raise CohortSelectionException( + ErrorMessages.COHORT_SELECTION_ERROR.message.format(cohort_to_use) + ) + + quorum_count = self.determine_required_quorum_count(orchestration_parameters, perspective_count) + + if ( + orchestration_parameters is not None + and orchestration_parameters.max_attempts is not None + and orchestration_parameters.max_attempts > 0 + and orchestration_parameters.cohort_for_single_attempt is None + ): + max_attempts = orchestration_parameters.max_attempts + if self.global_max_attempts is not None and max_attempts > self.global_max_attempts: + max_attempts = self.global_max_attempts + else: + max_attempts = 1 + attempts = 1 + previous_attempt_results = None + cohort_cycle = cycle(perspective_cohorts) + + while attempts <= max_attempts: + if cohort_to_use is not None: + perspectives_to_use = perspective_cohorts[ + cohort_to_use - 1 + ] # cohorts are 1-indexed for the user + else: + perspectives_to_use = next(cohort_cycle) + + # Collect async calls to invoke for each perspective. + async_calls_to_issue = MpicCoordinator.collect_checker_calls_to_issue( + mpic_request, perspectives_to_use + ) + + perspective_responses = await self.call_checkers_and_collect_responses( + mpic_request, perspectives_to_use, async_calls_to_issue + ) + + check_passed_per_perspective = { + response.perspective_code: response.check_response.check_passed + for response in perspective_responses + } + + valid_perspective_count = sum(check_passed_per_perspective.values()) + is_valid_result = valid_perspective_count >= quorum_count + + # noinspection PyUnresolvedReferences + self.logger.trace(f"Perspectives used in attempt: \n%s", pformat(perspectives_to_use)) + # noinspection PyUnresolvedReferences + self.logger.trace(f"Check passed per perspective: \n%s", pformat(check_passed_per_perspective)) + + # if cohort size is larger than 2, then at least two RIRs must be represented in the SUCCESSFUL perspectives + if len(perspectives_to_use) > 2: + valid_perspectives = [ + perspective + for perspective in perspectives_to_use + if check_passed_per_perspective[perspective.code] + ] + rir_count = len(set(perspective.rir for perspective in valid_perspectives)) + is_valid_result = rir_count >= 2 and is_valid_result + + if is_valid_result or attempts == max_attempts: + response = MpicResponseBuilder.build_response( + mpic_request, + perspective_count, + quorum_count, + attempts, + perspective_responses, + is_valid_result, + previous_attempt_results, + ) + + _is_valid = is_valid_result + _request_status = "ok" + _span.set_attribute("mpic.is_valid", is_valid_result) + _span.set_attribute("mpic.attempts", attempts) + + # noinspection PyUnresolvedReferences + self.logger.trace(f"Completed MPIC request with trace ID {mpic_request.trace_identifier}") + return response + else: + if previous_attempt_results is None: + previous_attempt_results = [] + previous_attempt_results.append(perspective_responses) + attempts += 1 + except Exception as exc: + _span.record_exception(exc) + _span.set_status(Status(StatusCode.ERROR, description=type(exc).__name__)) + raise + finally: + elapsed_ms = (time.perf_counter_ns() - _start_ns) / 1_000_000 + self._coordinator_duration.record(elapsed_ms, {"check.type": check_type_value}) + self._request_counter.add( + 1, + { + "check.type": check_type_value, + "mpic.is_valid": _is_valid, + "request.status": _request_status, + }, + ) def _raise_exception_on_invalid_request(self, mpic_request): is_request_valid, validation_issues = MpicRequestValidator.is_request_valid( @@ -225,20 +291,29 @@ async def call_remote_perspective( This assumes the wrapper will provide an async version of call_remote_perspective_function, or that we'll wrap the sync function using asyncio.to_thread() if needed. """ - try: - # noinspection PyUnresolvedReferences - async with self.logger.trace_timing( - f"MPIC round-trip with perspective {call_config.perspective.code}; trace ID: {call_config.check_request.trace_identifier}" - ): - response = await call_remote_perspective_function( - call_config.perspective, call_config.check_type, call_config.check_request - ) - except Exception as exc: - error_message = str(exc) if str(exc) else exc.__class__.__name__ - raise RemoteCheckException( - f"Check failed for perspective {call_config.perspective.code}, target {call_config.check_request.domain_or_ip_target}: {error_message}; trace ID: {call_config.check_request.trace_identifier}", - call_config=call_config, - ) from exc + with self._tracer.start_as_current_span( + "mpic.call_remote_perspective", + attributes={ + "check.type": call_config.check_type.value, + "perspective.code": call_config.perspective.code, + }, + ) as span: + try: + # noinspection PyUnresolvedReferences + async with self.logger.trace_timing( + f"MPIC round-trip with perspective {call_config.perspective.code}; trace ID: {call_config.check_request.trace_identifier}" + ): + response = await call_remote_perspective_function( + call_config.perspective, call_config.check_type, call_config.check_request + ) + except Exception as exc: + span.record_exception(exc) + span.set_status(Status(StatusCode.ERROR, description=type(exc).__name__)) + error_message = str(exc) if str(exc) else exc.__class__.__name__ + raise RemoteCheckException( + f"Check failed for perspective {call_config.perspective.code}, target {call_config.check_request.domain_or_ip_target}: {error_message}; trace ID: {call_config.check_request.trace_identifier}", + call_config=call_config, + ) from exc return PerspectiveResponse(perspective_code=call_config.perspective.code, check_response=response) @staticmethod @@ -290,18 +365,50 @@ async def call_checkers_and_collect_responses( ): responses = await asyncio.gather(*tasks, return_exceptions=True) - for response in responses: + check_type_value = mpic_request.check_type.value + for call_config, response in zip(async_calls_to_issue, responses): # check for exception (return_exceptions=True above will return exceptions as responses) # every Exception should be rethrown as RemoteCheckException # (trying to handle other Exceptions should be unreachable code) if isinstance(response, RemoteCheckException): response_as_string = str(response) - log_msg = f"{response_as_string} - trace ID: {mpic_request.trace_identifier}" - logger.warning(log_msg) + logger.warning(response_as_string) error_response = MpicCoordinator.build_error_perspective_response_from_exception(response) perspective_responses.append(error_response) - else: - # Now we know it's a valid PerspectiveResponse + self._remote_error_counter.add( + 1, + { + "check.type": check_type_value, + "perspective.code": response.call_config.perspective.code, + }, + ) + elif isinstance(response, PerspectiveResponse): perspective_responses.append(response) + self._perspective_response_counter.add( + 1, + { + "check.type": check_type_value, + "perspective.code": response.perspective_code, + "check.passed": response.check_response.check_passed, + "check.completed": response.check_response.check_completed, + }, + ) + else: + # Defensive handling for unexpected exceptions returned by gather(..., return_exceptions=True). + response_error = response if isinstance(response, Exception) else Exception(str(response)) + wrapped_error = RemoteCheckException( + f"Unexpected error type {type(response_error).__name__} for perspective {call_config.perspective.code}, target {call_config.check_request.domain_or_ip_target}: {response_error}; trace ID: {mpic_request.trace_identifier}", + call_config=call_config, + ) + logger.warning(str(wrapped_error)) + error_response = MpicCoordinator.build_error_perspective_response_from_exception(wrapped_error) + perspective_responses.append(error_response) + self._remote_error_counter.add( + 1, + { + "check.type": check_type_value, + "perspective.code": call_config.perspective.code, + }, + ) return perspective_responses diff --git a/src/open_mpic_core/mpic_dcv_checker/mpic_dcv_checker.py b/src/open_mpic_core/mpic_dcv_checker/mpic_dcv_checker.py index db1949a..5055b0e 100644 --- a/src/open_mpic_core/mpic_dcv_checker/mpic_dcv_checker.py +++ b/src/open_mpic_core/mpic_dcv_checker/mpic_dcv_checker.py @@ -14,6 +14,8 @@ from aiohttp import ClientError from aiohttp.web import HTTPException +from opentelemetry.trace import Status, StatusCode + from open_mpic_core import DcvCheckRequest, DcvCheckResponse from open_mpic_core import RedirectResponse, DcvUtils from open_mpic_core import DcvValidationMethod, DnsRecordType @@ -21,6 +23,7 @@ from open_mpic_core import DomainEncoder from open_mpic_core import DcvTlsAlpnValidator from open_mpic_core import get_logger +from open_mpic_core import get_meter, get_tracer logger = get_logger(__name__) @@ -66,6 +69,29 @@ def __init__( self.acme_tls_alpn_validator = DcvTlsAlpnValidator(log_level=log_level) self._http_client_timeout = http_client_timeout + _meter = get_meter(__name__) + self._tracer = get_tracer(__name__) + self._request_counter = _meter.create_counter( + "mpic.dcv.requests", + description="Total DCV check requests processed", + unit="1", + ) + self._duration_histogram = _meter.create_histogram( + "mpic.dcv.duration", + description="DCV check request duration", + unit="ms", + ) + self._http_duration_histogram = _meter.create_histogram( + "mpic.dcv.http.duration", + description="DCV HTTP-based validation duration", + unit="ms", + ) + self._dns_duration_histogram = _meter.create_histogram( + "mpic.dcv.dns.duration", + description="DCV DNS-based validation duration", + unit="ms", + ) + @asynccontextmanager async def get_async_http_client(self): connector = aiohttp.TCPConnector(ssl=self.verify_ssl, limit=0, force_close=True) @@ -84,6 +110,7 @@ async def get_async_http_client(self): async def check_dcv(self, dcv_request: DcvCheckRequest) -> DcvCheckResponse: validation_method = dcv_request.dcv_check_parameters.validation_method + validation_method_value = validation_method.value # noinspection PyUnresolvedReferences self.logger.trace( "Checking DCV for %s with method %s. Trace ID: %s", @@ -95,14 +122,39 @@ async def check_dcv(self, dcv_request: DcvCheckRequest) -> DcvCheckResponse: # encode domain if needed dcv_request.domain_or_ip_target = DomainEncoder.prepare_target_for_lookup(dcv_request.domain_or_ip_target) - result = None - match validation_method: - case DcvValidationMethod.WEBSITE_CHANGE | DcvValidationMethod.ACME_HTTP_01: - result = await self.perform_http_based_validation(dcv_request) - case DcvValidationMethod.ACME_TLS_ALPN_01: - result = await self.acme_tls_alpn_validator.perform_tls_alpn_validation(dcv_request) - case _: # all DNS based methods - result = await self.perform_general_dns_validation(dcv_request) + _start_ns = time.perf_counter_ns() + with self._tracer.start_as_current_span( + "mpic.dcv.check", attributes={"validation.method": validation_method_value} + ) as _span: + try: + result = None + match validation_method: + case DcvValidationMethod.WEBSITE_CHANGE | DcvValidationMethod.ACME_HTTP_01: + result = await self.perform_http_based_validation(dcv_request) + case DcvValidationMethod.ACME_TLS_ALPN_01: + result = await self.acme_tls_alpn_validator.perform_tls_alpn_validation(dcv_request) + case _: # all DNS based methods + result = await self.perform_general_dns_validation(dcv_request) + except Exception as exc: + _span.record_exception(exc) + _span.set_status(Status(StatusCode.ERROR, description=type(exc).__name__)) + raise + finally: + elapsed_ms = (time.perf_counter_ns() - _start_ns) / 1_000_000 + check_passed = result.check_passed if result is not None else False + check_completed = result.check_completed if result is not None else False + self._duration_histogram.record( + elapsed_ms, + {"validation.method": validation_method_value, "check.passed": check_passed}, + ) + self._request_counter.add( + 1, + { + "validation.method": validation_method_value, + "check.passed": check_passed, + "check.completed": check_completed, + }, + ) # noinspection PyUnresolvedReferences self.logger.trace( @@ -116,6 +168,7 @@ async def check_dcv(self, dcv_request: DcvCheckRequest) -> DcvCheckResponse: async def perform_general_dns_validation(self, request: DcvCheckRequest) -> DcvCheckResponse: check_parameters = request.dcv_check_parameters validation_method = check_parameters.validation_method + validation_method_value = validation_method.value dns_name_prefix = check_parameters.dns_name_prefix dns_record_type = check_parameters.dns_record_type exact_match = True @@ -136,32 +189,43 @@ async def perform_general_dns_validation(self, request: DcvCheckRequest) -> DcvC dcv_check_response = DcvUtils.create_empty_check_response(validation_method) - try: - # noinspection PyUnresolvedReferences - async with self.logger.trace_timing( - f"DNS lookup for target {name_to_resolve}. Trace ID: {request.trace_identifier}" - ): - lookup = await self.perform_dns_resolution(name_to_resolve, validation_method, dns_record_type) - MpicDcvChecker.evaluate_dns_lookup_response( - dcv_check_response, - lookup, - validation_method, - dns_record_type, - expected_dns_record_content, - exact_match, - require_exact_case, - ) - except dns.exception.DNSException as e: - log_msg = f"DNS lookup error for {name_to_resolve}: {str(e)}. Trace ID: {request.trace_identifier}" - if isinstance(e, dns.resolver.NoAnswer) or isinstance(e, dns.resolver.NXDOMAIN): - dcv_check_response.check_completed = True # errors on the target domain, not the lookup + _dns_start_ns = time.perf_counter_ns() + with self._tracer.start_as_current_span( + "mpic.dcv.dns_validation", attributes={"validation.method": validation_method_value} + ) as _span: + try: # noinspection PyUnresolvedReferences - self.logger.trace(log_msg) - else: - self.logger.warning(log_msg) - dcv_check_response.errors = [ - MpicValidationError.create(ErrorMessages.DCV_LOOKUP_ERROR, e.__class__.__name__, e.msg) - ] + async with self.logger.trace_timing( + f"DNS lookup for target {name_to_resolve}. Trace ID: {request.trace_identifier}" + ): + lookup = await self.perform_dns_resolution(name_to_resolve, validation_method, dns_record_type) + MpicDcvChecker.evaluate_dns_lookup_response( + dcv_check_response, + lookup, + validation_method, + dns_record_type, + expected_dns_record_content, + exact_match, + require_exact_case, + ) + except dns.exception.DNSException as e: + log_msg = f"DNS lookup error for {name_to_resolve}: {str(e)}. Trace ID: {request.trace_identifier}" + if isinstance(e, dns.resolver.NoAnswer) or isinstance(e, dns.resolver.NXDOMAIN): + dcv_check_response.check_completed = True # errors on the target domain, not the lookup + # noinspection PyUnresolvedReferences + self.logger.trace(log_msg) + else: + self.logger.warning(log_msg) + _span.record_exception(e) + _span.set_status(Status(StatusCode.ERROR, description=type(e).__name__)) + dcv_check_response.errors = [ + MpicValidationError.create(ErrorMessages.DCV_LOOKUP_ERROR, e.__class__.__name__, e.msg) + ] + finally: + self._dns_duration_histogram.record( + (time.perf_counter_ns() - _dns_start_ns) / 1_000_000, + {"validation.method": validation_method_value}, + ) dcv_check_response.timestamp_ns = time.time_ns() return dcv_check_response @@ -202,6 +266,7 @@ def format_host_for_url(domain_or_ip_target: str) -> str: async def perform_http_based_validation(self, request: DcvCheckRequest) -> DcvCheckResponse: validation_method = request.dcv_check_parameters.validation_method + validation_method_value = validation_method.value domain_or_ip_target = request.domain_or_ip_target formatted_host = MpicDcvChecker.format_host_for_url(domain_or_ip_target) http_headers = request.dcv_check_parameters.http_headers @@ -218,31 +283,46 @@ async def perform_http_based_validation(self, request: DcvCheckRequest) -> DcvCh token = request.dcv_check_parameters.token token_url = f"http://{formatted_host}/{MpicDcvChecker.WELL_KNOWN_ACME_PATH}/{token}" # noqa E501 (http) dcv_check_response = DcvUtils.create_empty_check_response(DcvValidationMethod.ACME_HTTP_01) - try: - async with self.get_async_http_client() as async_http_client: - # noinspection PyUnresolvedReferences - async with self.logger.trace_timing( - f"HTTP lookup for target {token_url}, trace ID: {request.trace_identifier}" - ): - async with async_http_client.get(url=token_url, headers=http_headers, max_redirects=20) as response: - dcv_check_response = await MpicDcvChecker.evaluate_http_lookup_response( - request, dcv_check_response, response, token_url, expected_response_content - ) - except asyncio.TimeoutError as e: - dcv_check_response.timestamp_ns = time.time_ns() - log_message = f"Timeout connecting to {token_url}: {str(e)}. Trace ID: {request.trace_identifier}" - self.logger.warning(log_message) - message = f"Connection timed out while attempting to connect to {token_url}" - dcv_check_response.errors = [ - MpicValidationError.create(ErrorMessages.DCV_LOOKUP_ERROR, e.__class__.__name__, message) - ] - except (ClientError, HTTPException, OSError) as e: - log_message = f"Error connecting to {token_url}: {str(e)}. Trace ID: {request.trace_identifier}" - self.logger.error(log_message) - dcv_check_response.timestamp_ns = time.time_ns() - dcv_check_response.errors = [ - MpicValidationError.create(ErrorMessages.DCV_LOOKUP_ERROR, e.__class__.__name__, str(e)) - ] + _http_start_ns = time.perf_counter_ns() + with self._tracer.start_as_current_span( + "mpic.dcv.http_validation", attributes={"validation.method": validation_method_value} + ) as _span: + try: + async with self.get_async_http_client() as async_http_client: + # noinspection PyUnresolvedReferences + async with self.logger.trace_timing( + f"HTTP lookup for target {token_url}, trace ID: {request.trace_identifier}" + ): + async with async_http_client.get( + url=token_url, headers=http_headers, max_redirects=20 + ) as response: + dcv_check_response = await MpicDcvChecker.evaluate_http_lookup_response( + request, dcv_check_response, response, token_url, expected_response_content + ) + except asyncio.TimeoutError as e: + dcv_check_response.timestamp_ns = time.time_ns() + log_message = f"Timeout connecting to {token_url}: {str(e)}. Trace ID: {request.trace_identifier}" + self.logger.warning(log_message) + _span.record_exception(e) + _span.set_status(Status(StatusCode.ERROR, description=type(e).__name__)) + message = f"Connection timed out while attempting to connect to {token_url}" + dcv_check_response.errors = [ + MpicValidationError.create(ErrorMessages.DCV_LOOKUP_ERROR, e.__class__.__name__, message) + ] + except (ClientError, HTTPException, OSError) as e: + log_message = f"Error connecting to {token_url}: {str(e)}. Trace ID: {request.trace_identifier}" + self.logger.error(log_message) + _span.record_exception(e) + _span.set_status(Status(StatusCode.ERROR, description=type(e).__name__)) + dcv_check_response.timestamp_ns = time.time_ns() + dcv_check_response.errors = [ + MpicValidationError.create(ErrorMessages.DCV_LOOKUP_ERROR, e.__class__.__name__, str(e)) + ] + finally: + self._http_duration_histogram.record( + (time.perf_counter_ns() - _http_start_ns) / 1_000_000, + {"validation.method": validation_method_value}, + ) return dcv_check_response diff --git a/tests/unit/open_mpic_core/test_mpic_caa_checker.py b/tests/unit/open_mpic_core/test_mpic_caa_checker.py index de5cfea..9e58ca6 100644 --- a/tests/unit/open_mpic_core/test_mpic_caa_checker.py +++ b/tests/unit/open_mpic_core/test_mpic_caa_checker.py @@ -1,10 +1,11 @@ import logging from io import StringIO +from contextlib import nullcontext import dns import pytest -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, MagicMock from dns.asyncresolver import reset_default_resolver @@ -300,6 +301,26 @@ async def check_caa__should_return_failure_response_given_error_in_dns_lookup(se assert caa_response.check_completed is False assert caa_response.details == check_response_details + async def check_caa__should_record_exception_and_set_error_status_on_lookup_failure(self, mocker): + caa_checker = TestMpicCaaChecker.create_configured_caa_checker() + resolver = caa_checker.resolver + dns_lookup_error = dns.resolver.NoNameservers( + request=dns.message.make_query("example.com", "CAA", "IN"), + errors=[("192.0.2.1", True, 53, "SERVFAIL")], + ) + self.patch_resolver_with_answer_or_exception(mocker, resolver, dns_lookup_error) + + span_mock = MagicMock() + tracer_mock = MagicMock() + tracer_mock.start_as_current_span.return_value = nullcontext(span_mock) + caa_checker._tracer = tracer_mock + + caa_request = self.create_caa_check_request("example.com", ["ca111.com"]) + await caa_checker.check_caa(caa_request) + + span_mock.record_exception.assert_called_once_with(dns_lookup_error) + span_mock.set_status.assert_called_once() + @pytest.mark.parametrize("allow_lookup_failure", [True, False]) async def check_caa__should_return_errors_in_response_given_error_in_dns_lookup(self, allow_lookup_failure, mocker): caa_checker = TestMpicCaaChecker.create_configured_caa_checker() diff --git a/tests/unit/open_mpic_core/test_mpic_coordinator.py b/tests/unit/open_mpic_core/test_mpic_coordinator.py index aba2f36..c1ba05b 100644 --- a/tests/unit/open_mpic_core/test_mpic_coordinator.py +++ b/tests/unit/open_mpic_core/test_mpic_coordinator.py @@ -1,7 +1,8 @@ import logging from io import StringIO from itertools import cycle -from unittest.mock import AsyncMock +from contextlib import nullcontext +from unittest.mock import AsyncMock, MagicMock import pytest @@ -19,6 +20,7 @@ MpicResponse, MpicCoordinator, MpicCoordinatorConfiguration, + RemoteCheckException, TRACE_LEVEL, ) from open_mpic_core.common_domain.enum.regional_internet_registry import RegionalInternetRegistry @@ -581,6 +583,50 @@ async def coordinate_mpic__should_not_log_trace_timings_if_trace_level_logging_i log_contents = self.log_output.getvalue() assert "seconds" not in log_contents + async def coordinate_mpic__should_start_coordinate_span(self): + mpic_request = ValidMpicRequestCreator.create_valid_caa_mpic_request() + mpic_coordinator_config = self.create_mpic_coordinator_configuration() + mocked_call_remote_perspective_function = AsyncMock() + mocked_call_remote_perspective_function.side_effect = TestMpicCoordinator.SideEffectForMockedPayloads( + self.create_passing_caa_check_response + ) + mpic_coordinator = MpicCoordinator(mocked_call_remote_perspective_function, mpic_coordinator_config) + + span_mock = MagicMock() + tracer_mock = MagicMock() + tracer_mock.start_as_current_span.return_value = nullcontext(span_mock) + mpic_coordinator._tracer = tracer_mock + + await mpic_coordinator.coordinate_mpic(mpic_request) + + assert any( + call.args and call.args[0] == "mpic.coordinate" for call in tracer_mock.start_as_current_span.call_args_list + ) + + async def call_remote_perspective__should_record_exception_and_set_error_status_on_remote_failure(self): + mpic_request = ValidMpicRequestCreator.create_valid_caa_mpic_request() + mpic_coordinator_config = self.create_mpic_coordinator_configuration() + + failing_remote_call = AsyncMock(side_effect=Exception("test remote failure")) + mpic_coordinator = MpicCoordinator(failing_remote_call, mpic_coordinator_config) + + span_mock = MagicMock() + tracer_mock = MagicMock() + tracer_mock.start_as_current_span.return_value = nullcontext(span_mock) + mpic_coordinator._tracer = tracer_mock + + call_config = MpicCoordinator.collect_checker_calls_to_issue( + mpic_request, [mpic_coordinator_config.target_perspectives[0]] + )[0] + + with pytest.raises(RemoteCheckException): + await mpic_coordinator.call_remote_perspective( + mpic_coordinator.call_remote_perspective_function, call_config + ) + + span_mock.record_exception.assert_called_once() + span_mock.set_status.assert_called_once() + @pytest.mark.parametrize("should_complete_mpic", [True, False]) async def coordinate_mpic__should_set_mpic_completed_true_if_enough_perspectives_completed_check_otherwise_false( self, should_complete_mpic diff --git a/tests/unit/open_mpic_core/test_mpic_dcv_checker.py b/tests/unit/open_mpic_core/test_mpic_dcv_checker.py index e2dcf68..a05bdaf 100644 --- a/tests/unit/open_mpic_core/test_mpic_dcv_checker.py +++ b/tests/unit/open_mpic_core/test_mpic_dcv_checker.py @@ -2,6 +2,7 @@ import base64 import logging import time +from contextlib import nullcontext import dns import random @@ -305,6 +306,38 @@ async def check_dcv__should_be_able_to_trace_timing_of_http_and_dns_lookups(self log_contents = self.log_output.getvalue() assert all(text in log_contents for text in ["seconds", "TRACE", tracing_dcv_checker.logger.name]) + async def check_dcv__should_record_exception_and_set_error_status_on_dns_lookup_error(self, mocker): + dcv_request = ValidCheckCreator.create_valid_dcv_check_request(DcvValidationMethod.ACME_DNS_01) + dns_lookup_error = dns.exception.Timeout() + self._patch_resolver_with_answer_or_exception(mocker, dns_lookup_error) + + check_span = MagicMock() + dns_span = MagicMock() + tracer_mock = MagicMock() + tracer_mock.start_as_current_span.side_effect = [nullcontext(check_span), nullcontext(dns_span)] + self.dcv_checker._tracer = tracer_mock + + await self.dcv_checker.check_dcv(dcv_request) + + dns_span.record_exception.assert_called_once_with(dns_lookup_error) + dns_span.set_status.assert_called_once() + + async def check_dcv__should_record_exception_and_set_error_status_on_http_lookup_error(self, mocker): + dcv_request = ValidCheckCreator.create_valid_dcv_check_request(DcvValidationMethod.WEBSITE_CHANGE) + http_lookup_error = ClientConnectionError() + mocker.patch("aiohttp.ClientSession.get", side_effect=http_lookup_error) + + check_span = MagicMock() + http_span = MagicMock() + tracer_mock = MagicMock() + tracer_mock.start_as_current_span.side_effect = [nullcontext(check_span), nullcontext(http_span)] + self.dcv_checker._tracer = tracer_mock + + await self.dcv_checker.check_dcv(dcv_request) + + http_span.record_exception.assert_called_once_with(http_lookup_error) + http_span.set_status.assert_called_once() + async def check_dcv__should_include_trace_identifier_in_logs_if_included_in_request(self, mocker): dcv_request = ValidCheckCreator.create_valid_dcv_check_request(DcvValidationMethod.WEBSITE_CHANGE) dcv_request.trace_identifier = "test_trace_identifier"