Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions bases/renku_data_services/data_tasks/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,13 +33,32 @@ 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."""

db: DBConfig
solr: SolrClientConfig
posthog: PosthogConfig
metering: MeteringConfig
authz: AuthzConfig
keycloak: KeycloakConfig | None
k8s_config_root: str
Expand Down Expand Up @@ -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"))

Expand All @@ -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,
Expand Down
10 changes: 9 additions & 1 deletion bases/renku_data_services/data_tasks/dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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!")
Expand Down
2 changes: 1 addition & 1 deletion bases/renku_data_services/data_tasks/task_defs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
13 changes: 12 additions & 1 deletion components/renku_data_services/resource_usage/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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."""
Expand All @@ -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:
Expand Down
5 changes: 5 additions & 0 deletions components/renku_data_services/resource_usage/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
79 changes: 79 additions & 0 deletions components/renku_data_services/resource_usage/metering.py
Original file line number Diff line number Diff line change
@@ -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)
Loading