103 lines
3.5 KiB
Python
103 lines
3.5 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
from dataclasses import dataclass
|
|
from datetime import UTC, datetime
|
|
|
|
from sqlalchemy.orm import Session
|
|
|
|
from govoplan_core.auth import ApiPrincipal
|
|
from govoplan_core.core.registry import PlatformRegistry
|
|
from govoplan_core.core.tasks import WorkItem, WorkItemPage, WorkItemQuery
|
|
from govoplan_tasks.backend.schemas import WorkProviderDiagnostic
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class WorkAggregation:
|
|
items: tuple[WorkItem, ...]
|
|
total: int
|
|
truncated: bool = False
|
|
diagnostics: tuple[WorkProviderDiagnostic, ...] = ()
|
|
|
|
|
|
def aggregate_work_items(
|
|
registry: PlatformRegistry,
|
|
session: Session,
|
|
principal: ApiPrincipal,
|
|
*,
|
|
query: WorkItemQuery,
|
|
) -> WorkAggregation:
|
|
"""Aggregate current work while isolating optional provider failures."""
|
|
|
|
items: list[WorkItem] = []
|
|
total = 0
|
|
truncated = False
|
|
diagnostics: list[WorkProviderDiagnostic] = []
|
|
for registered, provider in registry.work_item_providers():
|
|
provider_id = registered.registration.id
|
|
if query.provider_ids and provider_id not in query.provider_ids:
|
|
continue
|
|
if query.owner_modules and registered.module_id not in query.owner_modules:
|
|
continue
|
|
try:
|
|
page = provider.list_items(session, principal, query=query)
|
|
_validate_page(registered.module_id, provider_id, query, page)
|
|
except Exception: # provider isolation is part of the aggregation contract
|
|
logger.exception("Work-item provider failed provider=%s", provider_id)
|
|
diagnostics.append(
|
|
WorkProviderDiagnostic(
|
|
provider_id=provider_id,
|
|
owner_module=registered.module_id,
|
|
code="provider_unavailable",
|
|
message="This work source is temporarily unavailable.",
|
|
)
|
|
)
|
|
continue
|
|
items.extend(page.items)
|
|
total += page.total
|
|
truncated = truncated or page.truncated
|
|
items.sort(key=_sort_key)
|
|
if len(items) > query.limit:
|
|
truncated = True
|
|
items = items[: query.limit]
|
|
return WorkAggregation(
|
|
items=tuple(items),
|
|
total=total,
|
|
truncated=truncated or total > len(items),
|
|
diagnostics=tuple(diagnostics),
|
|
)
|
|
|
|
|
|
def _validate_page(
|
|
owner_module: str,
|
|
provider_id: str,
|
|
query: WorkItemQuery,
|
|
page: WorkItemPage,
|
|
) -> None:
|
|
if len(page.items) > query.limit:
|
|
raise ValueError("Work-item provider exceeded the requested limit.")
|
|
for item in page.items:
|
|
if item.provider_id != provider_id:
|
|
raise ValueError("Work-item provider returned another provider id.")
|
|
if item.owner_module != owner_module:
|
|
raise ValueError("Work-item provider returned another owner module.")
|
|
if item.tenant_id != query.tenant_id:
|
|
raise ValueError("Work-item provider returned another tenant.")
|
|
|
|
|
|
def _sort_key(item: WorkItem) -> tuple[object, ...]:
|
|
priorities = {"urgent": 0, "high": 1, "normal": 2, "low": 3}
|
|
due = _aware(item.due_at) if item.due_at else datetime.max.replace(tzinfo=UTC)
|
|
updated_rank = -_aware(item.updated_at).timestamp() if item.updated_at else 0.0
|
|
return (priorities[item.priority], due, updated_rank, item.provider_id, item.id)
|
|
|
|
|
|
def _aware(value: datetime) -> datetime:
|
|
return value.replace(tzinfo=UTC) if value.tzinfo is None else value.astimezone(UTC)
|
|
|
|
|
|
__all__ = ["WorkAggregation", "aggregate_work_items"]
|