diff --git a/app/config.py b/app/config.py index f1c62c4..52a555c 100644 --- a/app/config.py +++ b/app/config.py @@ -46,7 +46,7 @@ class AdapterConfig(BaseModel): class ObservabilityConfig(BaseModel): service_name: str = "g2s-aggregator" - otlp_endpoint: str = "http://localhost:4317" + otlp_endpoint: str = "" log_level: str = "INFO" diff --git a/app/controllers/middleware.py b/app/controllers/middleware.py index baade96..c72b930 100644 --- a/app/controllers/middleware.py +++ b/app/controllers/middleware.py @@ -1,10 +1,55 @@ """HTTP middleware registration.""" -from fastapi import FastAPI +from uuid import uuid4 +from time import perf_counter + +from fastapi import FastAPI, Request +from starlette.middleware.base import BaseHTTPMiddleware +from starlette.responses import Response +from structlog.contextvars import bind_contextvars, clear_contextvars + +from app.observability.metrics import get_metrics + + +class RequestCorrelationMiddleware(BaseHTTPMiddleware): + """Assign and propagate request correlation identifiers.""" + + def __init__(self, app: FastAPI, *, request_id_header: str) -> None: + super().__init__(app) + self._request_id_header = request_id_header + + async def dispatch(self, request: Request, call_next) -> Response: + request_id = str(uuid4()) + request.state.request_id = request_id + bind_contextvars(request_id=request_id) + + started_at = perf_counter() + status_code = 500 + is_error = True + + try: + response = await call_next(request) + status_code = response.status_code + is_error = status_code >= 500 + response.headers[self._request_id_header] = request_id + return response + finally: + duration_ms = (perf_counter() - started_at) * 1000 + get_metrics().record_http_request( + method=request.method, + route=request.url.path, + status_code=status_code, + duration_ms=duration_ms, + is_error=is_error, + ) + clear_contextvars() def install_middleware(app: FastAPI) -> None: """Register middleware components for the API.""" - # Middleware stack is introduced in later tasks. - _ = app + settings = app.state.settings + app.add_middleware( + RequestCorrelationMiddleware, + request_id_header=settings.controller.request_id_header, + ) diff --git a/app/main.py b/app/main.py index f823766..75c39b0 100644 --- a/app/main.py +++ b/app/main.py @@ -5,6 +5,7 @@ from fastapi import FastAPI from app.config import Settings, get_settings from app.controllers.middleware import install_middleware from app.controllers.v1.delivery import router as delivery_router +from app.observability import setup_observability def create_app(settings: Settings | None = None) -> FastAPI: @@ -13,6 +14,7 @@ def create_app(settings: Settings | None = None) -> FastAPI: application = FastAPI(title="G2S Aggregator", version="0.1.0") application.state.settings = resolved_settings + setup_observability(application, resolved_settings.observability) install_middleware(application) application.include_router(delivery_router, prefix=resolved_settings.controller.api_prefix) return application diff --git a/app/observability/__init__.py b/app/observability/__init__.py new file mode 100644 index 0000000..0b572ee --- /dev/null +++ b/app/observability/__init__.py @@ -0,0 +1,5 @@ +"""Observability setup and helpers.""" + +from app.observability.setup import setup_observability + +__all__ = ["setup_observability"] diff --git a/app/observability/logging.py b/app/observability/logging.py new file mode 100644 index 0000000..31065be --- /dev/null +++ b/app/observability/logging.py @@ -0,0 +1,48 @@ +"""Structured logging setup with request and trace correlation.""" + +import logging +from typing import Any + +import structlog +from opentelemetry import trace +from opentelemetry.trace import SpanContext + + +def configure_structured_logging(log_level: str) -> None: + resolved_level = _resolve_log_level(log_level) + logging.basicConfig(level=resolved_level, format="%(message)s") + + structlog.configure( + processors=[ + structlog.contextvars.merge_contextvars, + structlog.stdlib.add_log_level, + add_trace_context, + structlog.processors.TimeStamper(fmt="iso", utc=True), + structlog.processors.JSONRenderer(), + ], + logger_factory=structlog.stdlib.LoggerFactory(), + wrapper_class=structlog.make_filtering_bound_logger(resolved_level), + cache_logger_on_first_use=True, + ) + + +def add_trace_context( + _: Any, + __: str, + event_dict: dict[str, Any], +) -> dict[str, Any]: + span = trace.get_current_span() + span_context = span.get_span_context() + if _is_valid_span_context(span_context): + event_dict["trace_id"] = format(span_context.trace_id, "032x") + event_dict["span_id"] = format(span_context.span_id, "016x") + return event_dict + + +def _resolve_log_level(log_level: str) -> int: + candidate = log_level.strip().upper() + return logging.getLevelNamesMapping().get(candidate, logging.INFO) + + +def _is_valid_span_context(span_context: SpanContext) -> bool: + return bool(span_context.is_valid and span_context.trace_id and span_context.span_id) diff --git a/app/observability/metrics.py b/app/observability/metrics.py new file mode 100644 index 0000000..8fb732c --- /dev/null +++ b/app/observability/metrics.py @@ -0,0 +1,129 @@ +"""OpenTelemetry metrics setup and recording helpers.""" + +from typing import Sequence + +from opentelemetry import metrics +from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import OTLPMetricExporter +from opentelemetry.sdk.metrics import MeterProvider +from opentelemetry.sdk.metrics.export import ( + MetricReader, + PeriodicExportingMetricReader, +) +from opentelemetry.sdk.resources import Resource + +_METER_PROVIDER: MeterProvider | None = None +_METRICS: "ObservabilityMetrics" | None = None +_METER_NAME = "g2s.observability" + + +class ObservabilityMetrics: + def __init__( + self, + meter_name: str, + *, + meter_provider: MeterProvider | None = None, + ) -> None: + if meter_provider is None: + meter = metrics.get_meter(meter_name) + else: + meter = meter_provider.get_meter(meter_name) + self._request_count = meter.create_counter( + name="http.server.request.count", + unit="1", + description="Total number of HTTP requests.", + ) + self._error_count = meter.create_counter( + name="http.server.error.count", + unit="1", + description="Total number of HTTP server errors.", + ) + self._request_latency = meter.create_histogram( + name="http.server.request.latency", + unit="ms", + description="HTTP request latency.", + ) + self._provider_availability = meter.create_histogram( + name="delivery.provider.availability", + unit="1", + description="Provider availability sample (1=available, 0=unavailable).", + ) + + def record_http_request( + self, + *, + method: str, + route: str, + status_code: int, + duration_ms: float, + is_error: bool, + ) -> None: + attributes = { + "http.method": method, + "http.route": route, + "http.status_code": status_code, + } + self._request_count.add(1, attributes) + self._request_latency.record(duration_ms, attributes) + if is_error: + self._error_count.add(1, attributes) + + def record_provider_availability( + self, + *, + provider: str, + is_available: bool, + ) -> None: + self._provider_availability.record( + 1.0 if is_available else 0.0, + {"provider": provider}, + ) + + +def setup_metrics( + *, + service_name: str, + otlp_endpoint: str, + meter_provider: MeterProvider | None = None, + metric_readers: Sequence[MetricReader] | None = None, +) -> MeterProvider: + global _METER_PROVIDER, _METRICS + + if meter_provider is None and _METER_PROVIDER is not None: + return _METER_PROVIDER + + if meter_provider is None: + readers = list(metric_readers or ()) + if otlp_endpoint: + metric_exporter = OTLPMetricExporter( + endpoint=otlp_endpoint, + insecure=otlp_endpoint.startswith("http://"), + ) + readers.append(PeriodicExportingMetricReader(metric_exporter)) + meter_provider = MeterProvider( + resource=Resource.create({"service.name": service_name}), + metric_readers=readers, + ) + + _METER_PROVIDER = meter_provider + _METRICS = ObservabilityMetrics( + _METER_NAME, + meter_provider=meter_provider, + ) + return meter_provider + + +def get_metrics() -> ObservabilityMetrics: + global _METRICS + + if _METRICS is None: + _METRICS = ObservabilityMetrics(_METER_NAME) + return _METRICS + + +def reset_metrics_for_tests() -> None: + global _METER_PROVIDER, _METRICS + + if _METER_PROVIDER is not None: + _METER_PROVIDER.shutdown() + _METER_PROVIDER = None + _METRICS = None diff --git a/app/observability/setup.py b/app/observability/setup.py new file mode 100644 index 0000000..99cff80 --- /dev/null +++ b/app/observability/setup.py @@ -0,0 +1,26 @@ +"""High-level observability setup wiring.""" + +from fastapi import FastAPI + +from app.config import ObservabilityConfig +from app.observability.logging import configure_structured_logging +from app.observability.metrics import setup_metrics +from app.observability.tracing import ( + instrument_fastapi_app, + instrument_httpx_client, + setup_tracing, +) + + +def setup_observability(app: FastAPI, config: ObservabilityConfig) -> None: + configure_structured_logging(config.log_level) + setup_tracing( + service_name=config.service_name, + otlp_endpoint=config.otlp_endpoint, + ) + setup_metrics( + service_name=config.service_name, + otlp_endpoint=config.otlp_endpoint, + ) + instrument_fastapi_app(app) + instrument_httpx_client() diff --git a/app/observability/tracing.py b/app/observability/tracing.py new file mode 100644 index 0000000..5586632 --- /dev/null +++ b/app/observability/tracing.py @@ -0,0 +1,80 @@ +"""OpenTelemetry tracing setup and instrumentation helpers.""" + +from typing import cast + +from fastapi import FastAPI +from opentelemetry import trace +from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import OTLPSpanExporter +from opentelemetry.instrumentation.fastapi import FastAPIInstrumentor +from opentelemetry.instrumentation.httpx import HTTPXClientInstrumentor +from opentelemetry.sdk.resources import Resource +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import BatchSpanProcessor +from opentelemetry.trace import Tracer, TracerProvider as TraceAPIProvider + +_TRACER_PROVIDER: TraceAPIProvider | None = None +_HTTPX_INSTRUMENTED = False + + +def setup_tracing( + *, + service_name: str, + otlp_endpoint: str, + tracer_provider: TraceAPIProvider | None = None, +) -> TraceAPIProvider: + global _TRACER_PROVIDER + + if tracer_provider is None and _TRACER_PROVIDER is not None: + return _TRACER_PROVIDER + + if tracer_provider is None: + tracer_provider = TracerProvider( + resource=Resource.create({"service.name": service_name}), + ) + if otlp_endpoint: + span_exporter = OTLPSpanExporter( + endpoint=otlp_endpoint, + insecure=otlp_endpoint.startswith("http://"), + ) + tracer_provider.add_span_processor(BatchSpanProcessor(span_exporter)) + + _TRACER_PROVIDER = tracer_provider + return tracer_provider + + +def instrument_fastapi_app(app: FastAPI) -> None: + tracer_provider = _TRACER_PROVIDER + if tracer_provider is None: + return + FastAPIInstrumentor.instrument_app(app, tracer_provider=tracer_provider) + + +def instrument_httpx_client() -> None: + global _HTTPX_INSTRUMENTED + + tracer_provider = _TRACER_PROVIDER + if tracer_provider is None: + return + + if _HTTPX_INSTRUMENTED: + return + HTTPXClientInstrumentor().instrument(tracer_provider=tracer_provider) + _HTTPX_INSTRUMENTED = True + + +def get_tracer(name: str) -> Tracer: + tracer_provider = _TRACER_PROVIDER + if tracer_provider is None: + return trace.get_tracer(name) + return cast(Tracer, tracer_provider.get_tracer(name)) + + +def reset_tracing_for_tests() -> None: + global _TRACER_PROVIDER, _HTTPX_INSTRUMENTED + + if _HTTPX_INSTRUMENTED: + HTTPXClientInstrumentor().uninstrument() + if isinstance(_TRACER_PROVIDER, TracerProvider): + _TRACER_PROVIDER.shutdown() + _TRACER_PROVIDER = None + _HTTPX_INSTRUMENTED = False diff --git a/app/services/aggregator.py b/app/services/aggregator.py index 90e836c..fcfa3e7 100644 --- a/app/services/aggregator.py +++ b/app/services/aggregator.py @@ -13,6 +13,8 @@ from app.domain.price import ( filter_and_sort_prices, normalize_delivery_request, ) +from app.observability.metrics import get_metrics +from app.observability.tracing import get_tracer from app.schemas.request import DeliveryRequest from app.schemas.response import DeliveryPrice @@ -44,33 +46,41 @@ class AggregatorService: self._filter_and_sort_prices = filter_and_sort_prices_fn async def get_all_prices(self, request: DeliveryRequest) -> list[DeliveryPrice]: - normalized_request = normalize_delivery_request( - request, weight_round_scale=self._weight_round_scale - ) - provider_request = self._to_provider_request(normalized_request) + tracer = get_tracer(__name__) + with tracer.start_as_current_span("AggregatorService.get_all_prices") as span: + span.set_attribute("from_city", request.from_city) + span.set_attribute("to_city", request.to_city) + span.set_attribute("weight_kg", float(request.weight_kg)) - provider_results = await asyncio.gather( - *( - self._get_provider_price( - provider=provider, - request=provider_request, - cache_key=self._build_cache_key( - provider_name=provider.name, - request=normalized_request, - ), - ) - for provider in self._providers - ), - return_exceptions=True, - ) + normalized_request = normalize_delivery_request( + request, weight_round_scale=self._weight_round_scale + ) + provider_request = self._to_provider_request(normalized_request) - successful_results = [ - price - for price in provider_results - if isinstance(price, DeliveryPrice) - ] - filtered_and_sorted = self._filter_and_sort_prices(successful_results) - return [self._coerce_delivery_price(price) for price in filtered_and_sorted] + provider_results = await asyncio.gather( + *( + self._get_provider_price( + provider=provider, + request=provider_request, + cache_key=self._build_cache_key( + provider_name=provider.name, + request=normalized_request, + ), + ) + for provider in self._providers + ), + return_exceptions=True, + ) + + successful_results = [ + price + for price in provider_results + if isinstance(price, DeliveryPrice) + ] + filtered_and_sorted = self._filter_and_sort_prices(successful_results) + result = [self._coerce_delivery_price(price) for price in filtered_and_sorted] + span.set_attribute("tariffs_found", len(result)) + return result async def _get_provider_price( self, @@ -79,31 +89,59 @@ class AggregatorService: request: DeliveryRequest, cache_key: str, ) -> DeliveryPrice: - cached_price = await self._get_cached_price(cache_key) - if cached_price is not None: - return cached_price + tracer = get_tracer(__name__) + with tracer.start_as_current_span("DeliveryProvider.get_price") as span: + span.set_attribute("provider", provider.name) + span.set_attribute("from_city", request.from_city) + span.set_attribute("to_city", request.to_city) + span.set_attribute("weight_kg", float(request.weight_kg)) - fresh_price = await provider.get_price(request) - await self._set_cached_price( - cache_key, - fresh_price, - ttl=getattr(provider, "cache_ttl_seconds", None), - ) - return fresh_price + cached_price = await self._get_cached_price(cache_key) + span.set_attribute("cache_hit", cached_price is not None) + if cached_price is not None: + return cached_price + + try: + fresh_price = await provider.get_price(request) + except Exception: + get_metrics().record_provider_availability( + provider=provider.name, + is_available=False, + ) + raise + + get_metrics().record_provider_availability( + provider=provider.name, + is_available=True, + ) + await self._set_cached_price( + cache_key, + fresh_price, + ttl=getattr(provider, "cache_ttl_seconds", None), + ) + return fresh_price async def _get_cached_price(self, cache_key: str) -> DeliveryPrice | None: - if self._cache is None: - return None - try: - payload = await self._cache.get(cache_key) - except Exception: - return None - if payload is None: - return None - try: - return self._coerce_delivery_price(payload) - except Exception: - return None + tracer = get_tracer(__name__) + with tracer.start_as_current_span("PriceCache.get") as span: + if self._cache is None: + span.set_attribute("cache_hit", False) + return None + try: + payload = await self._cache.get(cache_key) + except Exception: + span.set_attribute("cache_hit", False) + return None + if payload is None: + span.set_attribute("cache_hit", False) + return None + try: + price = self._coerce_delivery_price(payload) + except Exception: + span.set_attribute("cache_hit", False) + return None + span.set_attribute("cache_hit", True) + return price async def _set_cached_price( self, @@ -112,12 +150,14 @@ class AggregatorService: *, ttl: int | None, ) -> None: - if self._cache is None: - return - try: - await self._cache.set(cache_key, payload, ttl=ttl) - except Exception: - return + tracer = get_tracer(__name__) + with tracer.start_as_current_span("PriceCache.set"): + if self._cache is None: + return + try: + await self._cache.set(cache_key, payload, ttl=ttl) + except Exception: + return @staticmethod def _to_provider_request(request: NormalizedDeliveryRequest) -> DeliveryRequest: diff --git a/spec/index.md b/spec/index.md index 067fa97..996b2ba 100644 --- a/spec/index.md +++ b/spec/index.md @@ -1,7 +1,7 @@ # Spec Tasks Index > ⚠️ This file is generated. Do not edit manually. -> Generated at (UTC): `2026-03-07T20:54:20+00:00` +> Generated at (UTC): `2026-03-08T08:42:58+00:00` ## Tasks @@ -14,12 +14,12 @@ | 004 | DONE | 2026-03-07 | Add Redis price cache repository | `spec/tasks/004_add_redis_price_cache_repository.md` | | 005 | DONE | 2026-03-07 | Implement aggregator service workflow | `spec/tasks/005_implement_aggregator_service_workflow.md` | | 006 | DONE | 2026-03-07 | Add delivery price controller endpoint | `spec/tasks/006_add_delivery_price_controller.md` | -| 007 | TODO | 2026-03-07 | Add observability correlation and telemetry | `spec/tasks/007_add_observability_correlation.md` | +| 007 | DONE | 2026-03-07 | Add observability correlation and telemetry | `spec/tasks/007_add_observability_correlation.md` | | 008 | TODO | 2026-03-07 | Add SigNoz to Telegram alerting configuration | `spec/tasks/008_add_signoz_telegram_alerting.md` | | 009 | TODO | 2026-03-07 | Add local infrastructure stack | `spec/tasks/009_add_local_infra_stack.md` | ## Summary - Total: **10** -- TODO: **3** -- DONE: **7** +- TODO: **2** +- DONE: **8** diff --git a/spec/tasks/007_add_observability_correlation.md b/spec/tasks/007_add_observability_correlation.md index d37db0e..356df28 100644 --- a/spec/tasks/007_add_observability_correlation.md +++ b/spec/tasks/007_add_observability_correlation.md @@ -1,7 +1,7 @@ --- id: 007 title: Add observability correlation and telemetry -status: TODO +status: DONE created: 2026-03-07 --- diff --git a/tests/controllers/test_middleware_request_id.py b/tests/controllers/test_middleware_request_id.py new file mode 100644 index 0000000..c2624ee --- /dev/null +++ b/tests/controllers/test_middleware_request_id.py @@ -0,0 +1,101 @@ +import asyncio +from uuid import UUID + +import httpx + +from app.config import Settings +from app.controllers.v1.delivery import get_aggregator_service +from app.main import create_app +from app.schemas.request import DeliveryRequest + + +class StubAggregatorService: + def __init__(self) -> None: + self.calls: list[DeliveryRequest] = [] + + async def get_all_prices(self, request: DeliveryRequest) -> list[object]: + self.calls.append(request) + return [] + + +def _build_payload() -> dict[str, object]: + return { + "entity": "individual", + "from_city": "Moscow", + "to_city": "Kazan", + "weight_kg": 2.5, + "length_cm": 30.0, + "width_cm": 20.0, + "height_cm": 10.0, + } + + +def _build_app( + *, + request_id_header: str, + service: StubAggregatorService, +) -> tuple[object, StubAggregatorService]: + settings = Settings( + controller={"request_id_header": request_id_header}, + observability={ + "service_name": "g2s-tests", + "otlp_endpoint": "", + "log_level": "INFO", + }, + ) + app = create_app(settings=settings) + + async def override_service() -> StubAggregatorService: + return service + + app.dependency_overrides[get_aggregator_service] = override_service + return app, service + + +def test_request_id_is_generated_per_request_and_returned_in_header() -> None: + app, service = _build_app( + request_id_header="X-Request-ID", + service=StubAggregatorService(), + ) + + async def run_requests() -> tuple[httpx.Response, httpx.Response]: + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + ) as client: + first = await client.post("/api/v1/delivery/price", json=_build_payload()) + second = await client.post("/api/v1/delivery/price", json=_build_payload()) + return first, second + + first_response, second_response = asyncio.run(run_requests()) + + assert first_response.status_code == 200 + assert second_response.status_code == 200 + first_request_id = first_response.headers["X-Request-ID"] + second_request_id = second_response.headers["X-Request-ID"] + assert UUID(first_request_id).version == 4 + assert UUID(second_request_id).version == 4 + assert first_request_id != second_request_id + assert len(service.calls) == 2 + + +def test_request_id_uses_configured_header_name() -> None: + app, _ = _build_app( + request_id_header="X-Correlation-ID", + service=StubAggregatorService(), + ) + + async def run_request() -> httpx.Response: + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + ) as client: + return await client.post("/api/v1/delivery/price", json=_build_payload()) + + response = asyncio.run(run_request()) + + assert response.status_code == 200 + assert "X-Correlation-ID" in response.headers + assert "X-Request-ID" not in response.headers diff --git a/tests/observability/test_logging_context.py b/tests/observability/test_logging_context.py new file mode 100644 index 0000000..d38ba53 --- /dev/null +++ b/tests/observability/test_logging_context.py @@ -0,0 +1,50 @@ +import structlog +from opentelemetry.sdk.resources import Resource +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from structlog.testing import LogCapture +from structlog.contextvars import bind_contextvars, clear_contextvars + +from app.observability.logging import add_trace_context +from app.observability.tracing import get_tracer, reset_tracing_for_tests, setup_tracing + + +def test_logging_context_includes_request_id_and_trace_id() -> None: + reset_tracing_for_tests() + span_exporter = InMemorySpanExporter() + tracer_provider = TracerProvider(resource=Resource.create({"service.name": "tests"})) + tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) + setup_tracing( + service_name="tests", + otlp_endpoint="", + tracer_provider=tracer_provider, + ) + + log_capture = LogCapture() + structlog.configure( + processors=[ + structlog.contextvars.merge_contextvars, + add_trace_context, + log_capture, + ], + logger_factory=structlog.stdlib.LoggerFactory(), + wrapper_class=structlog.make_filtering_bound_logger(20), + cache_logger_on_first_use=False, + ) + logger = structlog.get_logger("test-logger") + + clear_contextvars() + bind_contextvars(request_id="req-test-id") + with get_tracer(__name__).start_as_current_span("test-span"): + logger.info("request_log") + + assert len(log_capture.entries) == 1 + entry = log_capture.entries[0] + assert entry["event"] == "request_log" + assert entry["request_id"] == "req-test-id" + assert len(entry["trace_id"]) == 32 + assert len(entry["span_id"]) == 16 + + clear_contextvars() + reset_tracing_for_tests() diff --git a/tests/observability/test_metrics.py b/tests/observability/test_metrics.py new file mode 100644 index 0000000..159a262 --- /dev/null +++ b/tests/observability/test_metrics.py @@ -0,0 +1,204 @@ +import asyncio +from decimal import Decimal +from typing import Any + +import httpx +from opentelemetry.sdk.metrics import MeterProvider +from opentelemetry.sdk.metrics.export import InMemoryMetricReader + +from app.config import Settings +from app.controllers.v1.delivery import get_aggregator_service +from app.observability.metrics import reset_metrics_for_tests, setup_metrics +from app.schemas.request import DeliveryEntity, DeliveryRequest +from app.schemas.response import DeliveryPrice +from app.services.aggregator import AggregatorService, AggregatorServiceError + + +class ToggleAggregatorService: + def __init__(self) -> None: + self.should_fail = False + + async def get_all_prices(self, request: DeliveryRequest) -> list[DeliveryPrice]: + _ = request + if self.should_fail: + raise AggregatorServiceError("failed") + return [] + + +class StubProvider: + def __init__(self, name: str, *, fail: bool) -> None: + self.name = name + self.fail = fail + self.cache_ttl_seconds = 120 + + async def get_price(self, request: DeliveryRequest) -> DeliveryPrice: + _ = request + if self.fail: + raise RuntimeError("provider unavailable") + return DeliveryPrice( + provider=self.name, + service_name="economy", + price=Decimal("99.90"), + currency="RUB", + delivery_days_min=2, + delivery_days_max=3, + ) + + +class StubCache: + async def get(self, key: str) -> object | None: + _ = key + return None + + async def set(self, key: str, value: object, ttl: int | None = None) -> None: + _ = key, value, ttl + + +class StubCacheHit: + def __init__(self, payload: object) -> None: + self.payload = payload + + async def get(self, key: str) -> object | None: + _ = key + return self.payload + + async def set(self, key: str, value: object, ttl: int | None = None) -> None: + _ = key, value, ttl + + +def _build_payload() -> dict[str, object]: + return { + "entity": "individual", + "from_city": "Moscow", + "to_city": "Kazan", + "weight_kg": 2.5, + "length_cm": 30.0, + "width_cm": 20.0, + "height_cm": 10.0, + } + + +def _build_request() -> DeliveryRequest: + return DeliveryRequest( + entity=DeliveryEntity.INDIVIDUAL, + from_city="Moscow", + to_city="Kazan", + weight_kg=2.5, + length_cm=30.0, + width_cm=20.0, + height_cm=10.0, + ) + + +def _iter_data_points(metrics_data: Any, metric_name: str) -> list[Any]: + points: list[Any] = [] + if metrics_data is None: + return points + for resource_metric in metrics_data.resource_metrics: + for scope_metric in resource_metric.scope_metrics: + for metric in scope_metric.metrics: + if metric.name == metric_name: + points.extend(metric.data.data_points) + return points + + +def test_metrics_export_requests_errors_latency_and_provider_availability() -> None: + reset_metrics_for_tests() + metric_reader = InMemoryMetricReader() + meter_provider = MeterProvider(metric_readers=[metric_reader]) + setup_metrics( + service_name="g2s-tests", + otlp_endpoint="", + meter_provider=meter_provider, + ) + + from app.main import create_app + + toggle_service = ToggleAggregatorService() + app = create_app( + settings=Settings( + observability={ + "service_name": "g2s-tests", + "otlp_endpoint": "", + "log_level": "INFO", + } + ) + ) + + async def override_service() -> ToggleAggregatorService: + return toggle_service + + app.dependency_overrides[get_aggregator_service] = override_service + + async def run_requests() -> tuple[httpx.Response, httpx.Response]: + transport = httpx.ASGITransport(app=app, raise_app_exceptions=False) + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + ) as client: + success = await client.post("/api/v1/delivery/price", json=_build_payload()) + toggle_service.should_fail = True + failure = await client.post("/api/v1/delivery/price", json=_build_payload()) + return success, failure + + success_response, failure_response = asyncio.run(run_requests()) + assert success_response.status_code == 200 + assert failure_response.status_code == 503 + + availability_service = AggregatorService( + providers=[ + StubProvider("available-provider", fail=False), + StubProvider("unavailable-provider", fail=True), + ], + cache=StubCache(), + ) + asyncio.run(availability_service.get_all_prices(_build_request())) + + metrics_data = metric_reader.get_metrics_data() + request_points = _iter_data_points(metrics_data, "http.server.request.count") + error_points = _iter_data_points(metrics_data, "http.server.error.count") + latency_points = _iter_data_points(metrics_data, "http.server.request.latency") + availability_points = _iter_data_points(metrics_data, "delivery.provider.availability") + + assert sum(point.value for point in request_points) == 2 + assert sum(point.value for point in error_points) == 1 + assert sum(point.count for point in latency_points) == 2 + assert sum(point.count for point in availability_points) == 2 + assert sorted(point.sum for point in availability_points) == [0.0, 1.0] + + reset_metrics_for_tests() + + +def test_provider_availability_not_recorded_on_cache_hit() -> None: + reset_metrics_for_tests() + metric_reader = InMemoryMetricReader() + meter_provider = MeterProvider(metric_readers=[metric_reader]) + setup_metrics( + service_name="g2s-tests", + otlp_endpoint="", + meter_provider=meter_provider, + ) + + cached_price = DeliveryPrice( + provider="cached-provider", + service_name="economy", + price=Decimal("95.50"), + currency="RUB", + delivery_days_min=2, + delivery_days_max=3, + ).model_dump(mode="json") + provider = StubProvider("cached-provider", fail=True) + service = AggregatorService( + providers=[provider], + cache=StubCacheHit(cached_price), + ) + + result = asyncio.run(service.get_all_prices(_build_request())) + assert len(result) == 1 + assert result[0].provider == "cached-provider" + + metrics_data = metric_reader.get_metrics_data() + availability_points = _iter_data_points(metrics_data, "delivery.provider.availability") + assert availability_points == [] + + reset_metrics_for_tests() diff --git a/tests/observability/test_tracing.py b/tests/observability/test_tracing.py new file mode 100644 index 0000000..c79cd3d --- /dev/null +++ b/tests/observability/test_tracing.py @@ -0,0 +1,160 @@ +import asyncio +from decimal import Decimal + +import httpx +from fastapi import FastAPI +from opentelemetry.instrumentation.httpx import HTTPXClientInstrumentor +from opentelemetry.sdk.resources import Resource +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.trace import SpanKind + +from app.observability.tracing import ( + instrument_fastapi_app, + instrument_httpx_client, + reset_tracing_for_tests, + setup_tracing, +) +from app.schemas.request import DeliveryEntity, DeliveryRequest +from app.schemas.response import DeliveryPrice +from app.services.aggregator import AggregatorService + + +class StubProvider: + name = "stub-provider" + cache_ttl_seconds = 120 + + async def get_price(self, request: DeliveryRequest) -> DeliveryPrice: + _ = request + return DeliveryPrice( + provider=self.name, + service_name="economy", + price=Decimal("100.10"), + currency="RUB", + delivery_days_min=2, + delivery_days_max=4, + ) + + +class StubCache: + async def get(self, key: str) -> object | None: + _ = key + return None + + async def set(self, key: str, value: object, ttl: int | None = None) -> None: + _ = key, value, ttl + + +def _build_request() -> DeliveryRequest: + return DeliveryRequest( + entity=DeliveryEntity.INDIVIDUAL, + from_city="Moscow", + to_city="Kazan", + weight_kg=1.5, + length_cm=10.0, + width_cm=20.0, + height_cm=30.0, + ) + + +def test_manual_spans_capture_required_attributes() -> None: + reset_tracing_for_tests() + span_exporter = InMemorySpanExporter() + tracer_provider = TracerProvider(resource=Resource.create({"service.name": "tests"})) + tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) + setup_tracing( + service_name="tests", + otlp_endpoint="", + tracer_provider=tracer_provider, + ) + + service = AggregatorService(providers=[StubProvider()], cache=StubCache()) + asyncio.run(service.get_all_prices(_build_request())) + + spans = span_exporter.get_finished_spans() + span_by_name = {span.name: span for span in spans} + + assert "AggregatorService.get_all_prices" in span_by_name + assert "DeliveryProvider.get_price" in span_by_name + assert "PriceCache.get" in span_by_name + assert "PriceCache.set" in span_by_name + + service_span = span_by_name["AggregatorService.get_all_prices"] + assert service_span.attributes["from_city"] == "Moscow" + assert service_span.attributes["to_city"] == "Kazan" + assert service_span.attributes["weight_kg"] == 1.5 + assert service_span.attributes["tariffs_found"] == 1 + + provider_span = span_by_name["DeliveryProvider.get_price"] + assert provider_span.attributes["provider"] == "stub-provider" + assert provider_span.attributes["cache_hit"] is False + assert provider_span.attributes["from_city"] == "Moscow" + assert provider_span.attributes["to_city"] == "Kazan" + assert provider_span.attributes["weight_kg"] == 1.5 + + cache_span = span_by_name["PriceCache.get"] + assert cache_span.attributes["cache_hit"] is False + + reset_tracing_for_tests() + + +def test_fastapi_and_httpx_instrumentation_emit_server_and_client_spans() -> None: + reset_tracing_for_tests() + span_exporter = InMemorySpanExporter() + tracer_provider = TracerProvider(resource=Resource.create({"service.name": "tests"})) + tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) + setup_tracing( + service_name="tests", + otlp_endpoint="", + tracer_provider=tracer_provider, + ) + + app = FastAPI() + + @app.get("/proxy") + async def proxy() -> dict[str, int]: + def handler(request: httpx.Request) -> httpx.Response: + _ = request + return httpx.Response(status_code=200, json={"ok": True}) + + transport = httpx.MockTransport(handler) + async with httpx.AsyncClient( + transport=transport, + base_url="https://provider.test", + ) as client: + HTTPXClientInstrumentor.instrument_client(client, tracer_provider=tracer_provider) + response = await client.get("/quote") + return {"status_code": response.status_code} + + instrument_fastapi_app(app) + instrument_httpx_client() + + async def run_request() -> httpx.Response: + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + ) as client: + return await client.get("/proxy") + + response = asyncio.run(run_request()) + assert response.status_code == 200 + + spans = span_exporter.get_finished_spans() + assert any(span.kind == SpanKind.SERVER for span in spans) + assert any(span.kind == SpanKind.CLIENT for span in spans) + assert any( + ( + "provider.test/quote" + in ( + span.attributes.get("http.url") + or span.attributes.get("url.full") + or "" + ) + ) + for span in spans + if span.kind == SpanKind.CLIENT + ) + + reset_tracing_for_tests()