diff --git a/bases/renku_data_services/data_tasks/config.py b/bases/renku_data_services/data_tasks/config.py index 7d0934065..ea405bbea 100644 --- a/bases/renku_data_services/data_tasks/config.py +++ b/bases/renku_data_services/data_tasks/config.py @@ -33,6 +33,24 @@ def from_env( return cls(enabled, api_key, host, environment) +@dataclass +class MeteringConfig: + """Configuration for the Kong metering endpoint.""" + + enabled: bool + endpoint_url: str + token: str + + @classmethod + def from_env(cls) -> "MeteringConfig": + """Create metering config from environment variables.""" + return cls( + enabled=os.environ.get("METERING_ENABLED", "false").lower() == "true", + endpoint_url=os.environ.get("METERING_ENDPOINT_URL", ""), + token=os.environ.get("METERING_API_TOKEN", ""), + ) + + @dataclass class Config: """Configuration for data tasks.""" @@ -40,6 +58,7 @@ class Config: db: DBConfig solr: SolrClientConfig posthog: PosthogConfig + metering: MeteringConfig authz: AuthzConfig keycloak: KeycloakConfig | None k8s_config_root: str @@ -67,6 +86,7 @@ def from_env(cls) -> Config: main_tick = int(os.environ.get("MAIN_LOG_INTERVAL_SECONDS", "300")) solr_config = SolrClientConfig.from_env() posthog_config = PosthogConfig.from_env() + metering_config = MeteringConfig.from_env() tcp_host = os.environ.get("TCP_HOST", "127.0.0.1") tcp_port = int(os.environ.get("TCP_PORT", "8001")) @@ -89,6 +109,7 @@ def from_env(cls) -> Config: main_log_interval_seconds=main_tick, solr=solr_config, posthog=posthog_config, + metering=metering_config, authz=authz, keycloak=keycloak, k8s_config_root=k8s_config_root, diff --git a/bases/renku_data_services/data_tasks/dependencies.py b/bases/renku_data_services/data_tasks/dependencies.py index c7a9ac036..04ce957d1 100644 --- a/bases/renku_data_services/data_tasks/dependencies.py +++ b/bases/renku_data_services/data_tasks/dependencies.py @@ -25,6 +25,7 @@ ResourcesRequestRecorder, ResourceUsageService, ) +from renku_data_services.resource_usage.metering import MeteringClient from renku_data_services.resource_usage.db import ResourceRequestsRepo from renku_data_services.search.db import SearchUpdatesRepo from renku_data_services.session.db import SessionRepository @@ -123,8 +124,15 @@ def from_env(cls, cfg: Config | None = None) -> "DependencyManager": resource_requests_recorder: ResourcesRequestRecorder if cfg.enable_resource_request_tracking: + metering_client = ( + MeteringClient(endpoint_url=cfg.metering.endpoint_url, token=cfg.metering.token) + if cfg.metering.enabled + else None + ) resource_requests_recorder = DefaultResourcesRequestRecorder( - repo=resource_requests_repo, fetch=ResourceRequestsFetch(k8s_client) + repo=resource_requests_repo, + fetch=ResourceRequestsFetch(k8s_client), + metering=metering_client, ) else: logger.warning("Resource request tracking is disabled!") diff --git a/bases/renku_data_services/data_tasks/task_defs.py b/bases/renku_data_services/data_tasks/task_defs.py index 8f080e254..f1b7116dd 100644 --- a/bases/renku_data_services/data_tasks/task_defs.py +++ b/bases/renku_data_services/data_tasks/task_defs.py @@ -440,7 +440,7 @@ async def cleanup_orphaned_capacity_reservations(dm: DependencyManager) -> None: async def record_resource_requests(dm: DependencyManager) -> None: """Periodically record all resource requests.""" - interval_seconds = 600 + interval_seconds = 30 # increase before merging while True: await dm.resource_requests_recorder.record_resource_requests(timedelta(seconds=interval_seconds)) await asyncio.sleep(interval_seconds) diff --git a/components/renku_data_services/resource_usage/core.py b/components/renku_data_services/resource_usage/core.py index 91c106d39..c4dc31ec1 100644 --- a/components/renku_data_services/resource_usage/core.py +++ b/components/renku_data_services/resource_usage/core.py @@ -11,6 +11,7 @@ from renku_data_services.k8s.models import GVK, K8sObject, K8sObjectFilter, K8sObjectMeta from renku_data_services.resource_usage import apispec from renku_data_services.resource_usage.db import ResourceRequestsRepo +from renku_data_services.resource_usage.metering import MeteringClient, MetricCode from renku_data_services.resource_usage.model import ( Credit, ResourceClassCost, @@ -155,9 +156,15 @@ async def record_resource_requests(self, interval: timedelta) -> None: class DefaultResourcesRequestRecorder(ResourcesRequestRecorder): """Methods for recording resource requests.""" - def __init__(self, repo: ResourceRequestsRepo, fetch: ResourceRequestsFetchProto) -> None: + def __init__( + self, + repo: ResourceRequestsRepo, + fetch: ResourceRequestsFetchProto, + metering: MeteringClient | None = None, + ) -> None: self._repo = repo self._fetch = fetch + self._metering = metering async def record_resource_requests(self, interval: timedelta) -> None: """Fetches all resource requests in the given namespace and stores them.""" @@ -170,6 +177,10 @@ async def record_resource_requests(self, interval: timedelta) -> None: else: logger.info(f"Inserting {size} resource request records.") await self._repo.insert_many(result) + if self._metering is not None: + class_ids = {r.resource_class_id for r in result if r.resource_class_id is not None} + costs = await self._repo.get_costs_by_class_ids(class_ids) + await self._metering.emit(result, costs, MetricCode.session_resource_usage) class ResourceUsageService: diff --git a/components/renku_data_services/resource_usage/db.py b/components/renku_data_services/resource_usage/db.py index d402f2ac8..ac745e646 100644 --- a/components/renku_data_services/resource_usage/db.py +++ b/components/renku_data_services/resource_usage/db.py @@ -84,6 +84,11 @@ async def insert_many(self, reqs: Iterable[ResourcesRequest]) -> None: session.add_all(vals) await session.flush() + async def get_costs_by_class_ids(self, ids: set[int]) -> dict[int, Credit]: + """Return a mapping of resource_class_id to Credit for the given ids.""" + async with self.session_maker() as session: + return await self._get_all_costs(session, ids) + async def _get_all_costs(self, session: AsyncSession, ids: set[int]) -> dict[int, Credit]: stmt = sa.select(ResourceClassCostORM.id, ResourceClassCostORM.cost).where(ResourceClassCostORM.id.in_(ids)) rows = await session.execute(stmt) diff --git a/components/renku_data_services/resource_usage/metering.py b/components/renku_data_services/resource_usage/metering.py new file mode 100644 index 000000000..c0a7ae417 --- /dev/null +++ b/components/renku_data_services/resource_usage/metering.py @@ -0,0 +1,79 @@ +"""Meteroid event emission for session resource usage metering.""" + +from datetime import UTC +from enum import StrEnum + +import httpx + +from renku_data_services.app_config import logging +from renku_data_services.resource_usage.model import Credit, ResourcesRequest + +logger = logging.getLogger(__file__) + +class MetricCode(StrEnum): + session_resource_usage = "session_resource_usage" + + +def _to_meteroid_event(req: ResourcesRequest, costs: dict[int, Credit], metric_code: str) -> dict: + cost = costs.get(req.resource_class_id, Credit.zero()) # type: ignore[arg-type] + cu_cost = round(cost.value * (req.capture_interval.total_seconds() / 3600.0), 6) + + properties: dict[str, str] = { + "cu_cost": str(cu_cost), + "kind": req.kind, + "phase": req.phase, + "capture_interval_seconds": str(req.capture_interval.total_seconds()), + "resource_class_id": str(req.resource_class_id), + "user_id": str(req.user_id), + } + if req.resource_pool_id is not None: + properties["resource_pool_id"] = str(req.resource_pool_id) + if req.project_id is not None: + properties["project_id"] = str(req.project_id) + if req.launcher_id is not None: + properties["launcher_id"] = str(req.launcher_id) + if req.cluster_id is not None: + properties["cluster_id"] = str(req.cluster_id) + + return { + "event_id": f"{req.uid}/{req.capture_date.astimezone(UTC).isoformat()}", + "code": metric_code, + "customer_id": f"resource_pool_id-{req.resource_pool_id}", + "timestamp": req.capture_date.astimezone(UTC).isoformat(), + "properties": properties, + } + + +class MeteringClient: + """Emits session resource usage events to Meteroid.""" + + def __init__(self, endpoint_url: str, token: str) -> None: + self._endpoint_url = endpoint_url + self._headers = { + "Authorization": f"Bearer {token}", + "Content-Type": "application/json", + } + + async def emit(self, requests: list[ResourcesRequest], costs: dict[int, Credit], metric_code: MetricCode) -> None: + """POST all resource requests as a Meteroid ingest batch. Never raises.""" + events = [ + _to_meteroid_event(r, costs, metric_code) + for r in requests + if r.resource_class_id is not None and r.user_id is not None and r.resource_pool_id is not None + ] + if not events: + return + body = {"allow_partial_failures": True, "events": events} + try: + async with httpx.AsyncClient(timeout=30) as client: + resp = await client.post(self._endpoint_url, headers=self._headers, json=body) + if resp.status_code >= 300 or resp.status_code < 200: + logger.warning( + f"Metering endpoint returned unexpected status {resp.status_code}: {resp.text[:200]}" + ) + else: + logger.info(f"Emitted {len(events)} metering events, status={resp.status_code}: {body}") + except httpx.HTTPError as ex: + logger.warning(f"Failed to emit metering events: {ex}", exc_info=ex) + except Exception as ex: + logger.warning(f"Unexpected error emitting metering events: {ex}", exc_info=ex)