hub-core/hub_core/routers/progress.py
tegwick 41a396e807
All checks were successful
CI Smoke / host-smoke (push) Successful in 0s
CI Smoke / container-smoke (push) Successful in 2s
Support legacy workstream_id body metering on progress POST
create_progress_router accepts optional body detection and metering
hooks so host apps can record usage when clients POST /progress/ with
the legacy workstream_id field.
2026-07-08 23:06:45 +02:00

173 lines
6.3 KiB
Python

import uuid
from collections.abc import Awaitable, Callable, Collection
from datetime import datetime
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from hub_core.events import ALERT_EVENT_TYPES, RISK_EVENT_TYPES
from hub_core.models.progress_event import ProgressEvent
from hub_core.schemas.progress_event import ProgressEventCreate, ProgressEventRead
from hub_core.utils.pagination import PageParams, apply_pagination
# List/create progress events. ``workplan_id`` is the preferred filter/body field;
# ``workstream_id`` remains a wire-compat alias until legacy-meter retires it.
MeterLegacyWorkstreamId = Callable[
[AsyncSession, Request, Response],
Awaitable[None],
]
ProgressBodyUsesLegacyWorkstreamId = Callable[[Any], bool]
def create_progress_router(
get_session: Callable[..., AsyncSession],
*,
progress_model: type[ProgressEvent] = ProgressEvent,
progress_create_schema: type[ProgressEventCreate] = ProgressEventCreate,
progress_read_schema: type[ProgressEventRead] = ProgressEventRead,
meter_legacy_workstream_id: MeterLegacyWorkstreamId | None = None,
progress_body_uses_legacy_workstream_id: ProgressBodyUsesLegacyWorkstreamId | None = None,
meter_legacy_workstream_id_body: MeterLegacyWorkstreamId | None = None,
) -> APIRouter:
router = APIRouter(prefix="/progress", tags=["progress"])
list_response_model = list[progress_read_schema]
async def _list_events(
session: AsyncSession,
*,
topic_id: uuid.UUID | None = None,
workstream_id: uuid.UUID | None = None,
workplan_id: uuid.UUID | None = None,
task_id: uuid.UUID | None = None,
decision_id: uuid.UUID | None = None,
event_type: str | None = None,
event_types: Collection[str] | None = None,
since: datetime | None = None,
limit: int = 100,
offset: int = 0,
) -> list[Any]:
q = select(progress_model)
scope_id = workplan_id if workplan_id is not None else workstream_id
for field, value in (
("topic_id", topic_id),
("task_id", task_id),
("decision_id", decision_id),
):
if value is not None:
column = getattr(progress_model, field, None)
if column is None:
raise HTTPException(
status_code=400,
detail=f"Progress events do not support filtering by {field}",
)
q = q.where(column == value)
if scope_id is not None:
column = getattr(progress_model, "workplan_id", None) or getattr(
progress_model, "workstream_id", None
)
if column is None:
raise HTTPException(
status_code=400,
detail="Progress events do not support filtering by workplan_id",
)
q = q.where(column == scope_id)
if event_type:
q = q.where(progress_model.event_type == event_type)
if event_types is not None:
q = q.where(progress_model.event_type.in_(sorted(event_types)))
if since:
q = q.where(progress_model.created_at >= since)
q = q.order_by(progress_model.created_at.desc())
q = apply_pagination(q, PageParams(limit=limit, offset=offset))
result = await session.execute(q)
return list(result.scalars().all())
@router.get("/", response_model=list_response_model)
async def list_progress(
request: Request,
response: Response,
topic_id: uuid.UUID | None = None,
workstream_id: uuid.UUID | None = None,
workplan_id: uuid.UUID | None = None,
task_id: uuid.UUID | None = None,
decision_id: uuid.UUID | None = None,
event_type: str | None = None,
since: datetime | None = None,
limit: int = Query(100, le=1000),
offset: int = Query(0, ge=0),
session: AsyncSession = Depends(get_session),
) -> list[Any]:
if (
meter_legacy_workstream_id is not None
and workstream_id is not None
and workplan_id is None
):
await meter_legacy_workstream_id(session, request, response)
return await _list_events(
session,
topic_id=topic_id,
workstream_id=workstream_id,
workplan_id=workplan_id,
task_id=task_id,
decision_id=decision_id,
event_type=event_type,
since=since,
limit=limit,
offset=offset,
)
@router.get("/risks", response_model=list_response_model)
async def get_risks(
since: datetime | None = None,
limit: int = Query(100, le=1000),
offset: int = Query(0, ge=0),
session: AsyncSession = Depends(get_session),
) -> list[Any]:
return await _list_events(
session,
event_types=RISK_EVENT_TYPES,
since=since,
limit=limit,
offset=offset,
)
@router.get("/alerts", response_model=list_response_model)
async def get_alerts(
since: datetime | None = None,
limit: int = Query(100, le=1000),
offset: int = Query(0, ge=0),
session: AsyncSession = Depends(get_session),
) -> list[Any]:
return await _list_events(
session,
event_types=ALERT_EVENT_TYPES,
since=since,
limit=limit,
offset=offset,
)
@router.post("/", response_model=progress_read_schema, status_code=status.HTTP_201_CREATED)
async def append_progress(
request: Request,
response: Response,
body: progress_create_schema,
session: AsyncSession = Depends(get_session),
) -> Any:
if (
meter_legacy_workstream_id_body is not None
and progress_body_uses_legacy_workstream_id is not None
and progress_body_uses_legacy_workstream_id(body)
):
await meter_legacy_workstream_id_body(session, request, response)
event = progress_model(**body.model_dump())
session.add(event)
await session.commit()
await session.refresh(event)
return event
return router