Skip to content
2 changes: 2 additions & 0 deletions sdk/src/flowmesh/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
APIGroupItem,
APIItem,
APIResult,
APIUsage,
BaseExecutorResult,
CostEstimates,
DataProfilingResult,
Expand Down Expand Up @@ -110,6 +111,7 @@
"APIGroupItem",
"APIItem",
"APIResult",
"APIUsage",
"ActiveWaitBreakdown",
"AgentBatchSummary",
"AgentItem",
Expand Down
2 changes: 2 additions & 0 deletions sdk/src/flowmesh/models/result/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@
AgentUsage,
APIGroupItem,
APIItem,
APIUsage,
CostEstimates,
DataRetrievalItem,
EchoItem,
Expand Down Expand Up @@ -95,6 +96,7 @@
"APIGroupItem",
"APIItem",
"APIResult",
"APIUsage",
"AgentBatchSummary",
"AgentItem",
"AgentMetadata",
Expand Down
3 changes: 2 additions & 1 deletion sdk/src/flowmesh/models/result/catalog.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
AgentUsage,
APIGroupItem,
APIItem,
APIUsage,
CostEstimates,
DataRetrievalItem,
EchoItem,
Expand Down Expand Up @@ -213,7 +214,7 @@ class APIResult(StrictExecutorResult):
truncated: bool = False
headers: dict[str, str] | None = None
response_json: Any = Field(default=None, alias="json")
usage: dict[str, Any] | None = None
usage: APIUsage | None = None
text: str | None = None
items: list[APIItem | APIGroupItem] = Field(default_factory=list)

Expand Down
11 changes: 11 additions & 0 deletions sdk/src/flowmesh/models/result/payloads.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,17 @@ class GenerationUsage(StrictModel):
latency_sec: float


class APIUsage(StrictModel):
prompt_tokens: int
completion_tokens: int
reasoning_tokens: int
calls: int
failures: int
retries: int
truncated_calls: int
wall_sec: float


class EmbeddingUsage(StrictModel):
prompt_tokens: int
total_tokens: int
Expand Down
2 changes: 2 additions & 0 deletions sdk/src/flowmesh/models/workflows.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from pydantic import BaseModel

from .common import TaskStatus, WorkflowStatus
from .result import APIUsage


class WorkflowSubmitTaskEntry(BaseModel):
Expand Down Expand Up @@ -46,3 +47,4 @@ class Workflow(BaseModel):
completed_tasks: list[str]
failed_tasks: list[str]
cancelled_tasks: list[str]
usage: APIUsage | None = None
4 changes: 4 additions & 0 deletions src/server/clients/redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,10 @@ def task_state_key(task_id: str) -> str:
return f"task:{task_id}:state"


def task_usage_key(task_id: str) -> str:
return f"task:{task_id}:usage"


def worker_key(worker_id: str) -> str:
return f"worker:{worker_id}"

Expand Down
1 change: 1 addition & 0 deletions src/server/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,7 @@
results_dir=RESULTS_DIR,
log_stream_ttl_sec=config.log_stream.ttl_sec,
server_base_url=config.identity.base_url,
workflow_registry=WORKFLOW_REGISTRY,
)

LOG_ARCHIVER = TaskLogArchiver(
Expand Down
72 changes: 72 additions & 0 deletions src/server/registries/workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,10 +13,13 @@
model_serializer,
)

from shared.schemas.result import APIUsage

from ..clients.redis import (
WORKFLOWS_SET_KEY,
RedisClient,
task_state_key,
task_usage_key,
workflow_cancelled_tasks_key,
workflow_dispatched_tasks_key,
workflow_failed_tasks_key,
Expand All @@ -28,6 +31,13 @@
from ..utils.time import now_iso


class UnknownUsage:
"""Sentinel for a model-calling task whose usage could not be mapped."""


UNKNOWN_USAGE = UnknownUsage()


class PersistedTask(BaseModel):
"""A durable per-task snapshot sufficient to rebuild scheduler state."""

Expand Down Expand Up @@ -103,6 +113,10 @@ class Workflow(BaseModel):
completed_tasks: list[str] = Field(description="Completed task identifiers.")
failed_tasks: list[str] = Field(description="Failed task identifiers.")
cancelled_tasks: list[str] = Field(description="Cancelled task identifiers.")
usage: APIUsage | None = Field(
default=None,
description="Token/call usage summed over the workflow's finished tasks.",
)


def _create_workflow_record(
Expand Down Expand Up @@ -175,6 +189,7 @@ def unregister_workflows(self, *workflow_ids: str) -> None:
pipe.delete(*(workflow_sched_key(wid) for wid in workflow_ids))
for task_id in task_ids:
pipe.delete(task_state_key(task_id))
pipe.delete(task_usage_key(task_id))
pipe.execute()

async def unregister_workflows_async(self, *workflow_ids: str) -> None:
Expand All @@ -189,6 +204,7 @@ async def unregister_workflows_async(self, *workflow_ids: str) -> None:
pipe.delete(*(workflow_sched_key(wid) for wid in workflow_ids))
for task_id in task_ids:
pipe.delete(task_state_key(task_id))
pipe.delete(task_usage_key(task_id))
await pipe.execute()

def get_workflow_ids(self) -> set[str]:
Expand Down Expand Up @@ -332,6 +348,62 @@ async def load_task_states_async(
PersistedTask.model_validate_json(blob) if blob else None for blob in blobs
]

def save_task_usage(
self, task_id: str, usage: APIUsage | None | UnknownUsage
) -> None:
"""Synchronous variant of ``save_task_usage_async`` for the event loop."""
key = task_usage_key(task_id)
if usage is None:
self._rds.sync.set_value(key, "null")
elif isinstance(usage, UnknownUsage):
self._rds.sync.set_value(key, "unknown")
else:
self._rds.sync.set_value(key, usage.model_dump_json())

async def save_task_usage_async(
self, task_id: str, usage: APIUsage | None | UnknownUsage
) -> None:
"""Persist one task's usage contribution.

``None`` marks a task type that makes no model calls (echo, lambda);
``UNKNOWN_USAGE`` marks a model-calling task whose usage could not be
mapped. Both are written so a completed task with no entry at all is
detectable as one whose result was never ingested.
"""
key = task_usage_key(task_id)
if usage is None:
await self._rds.asyncio.set_value(key, "null")
elif isinstance(usage, UnknownUsage):
await self._rds.asyncio.set_value(key, "unknown")
else:
await self._rds.asyncio.set_value(key, usage.model_dump_json())

async def load_task_usages_async(
self, *task_ids: str
) -> dict[str, APIUsage | None | UnknownUsage]:
"""Load persisted usage for the given tasks.

``None`` means the task type makes no model calls; ``UNKNOWN_USAGE``
means a model-calling task's usage could not be mapped; a task absent
from the result means its usage was never recorded.
"""
if not task_ids:
return {}
blobs = await self._rds.asyncio.mget(
[task_usage_key(task_id) for task_id in task_ids]
)
result: dict[str, APIUsage | None | UnknownUsage] = {}
for task_id, blob in zip(task_ids, blobs, strict=True):
if blob is None:
continue
if blob == "null":
result[task_id] = None
elif blob == "unknown":
result[task_id] = UNKNOWN_USAGE
else:
result[task_id] = APIUsage.model_validate_json(blob)
return result

async def save_workflow_sched_async(
self, workflow_id: str, in_epoch_order: bool, epoch_frontier: int
) -> None:
Expand Down
85 changes: 85 additions & 0 deletions src/server/routers/v1/results.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,16 @@

from shared.schemas.result import (
AnyExecutorResult,
APIResult,
APIUsage,
DataProfilingResult,
DataRetrievalResult,
EchoResult,
GenerationUsage,
InferenceResult,
PythonResult,
ResultEnvelope,
SSHResult,
read_result,
result_file_path,
write_result,
Expand All @@ -32,13 +41,15 @@
get_logger,
get_results_dir,
get_runtime,
get_workflow_registry,
)
from ...auth.security import (
PrincipalContext,
authenticate_connection,
require_permission,
)
from ...hooks import ResourceAction, ResourceKind
from ...registries.workflow import UNKNOWN_USAGE, UnknownUsage, WorkflowRegistry
from ...schemas.common import PathResponse
from ...services.monitoring import EventMonitor
from ...task.models import TERMINAL_TASK_STATUSES
Expand All @@ -52,6 +63,77 @@
router = APIRouter(prefix="/results", tags=["Results"])


# Result types that make no model calls and therefore never carry usage.
_NO_USAGE_RESULT_TYPES = (
EchoResult,
SSHResult,
DataProfilingResult,
DataRetrievalResult,
PythonResult,
)


def _task_usage_from_envelope(
envelope: ResultEnvelope,
) -> APIUsage | None | UnknownUsage:
"""Return a task's usage contribution from its result envelope.

API tasks carry an ``APIUsage``; vLLM inference tasks map their
``GenerationUsage`` token counts with reasoning 0, ``calls`` from
``num_requests``, and ``wall_sec`` from ``latency_sec``. Task types that
make no model calls (echo, ssh, data profiling, data retrieval, python)
contribute nothing (``None``). A model-calling task whose usage cannot be
mapped returns ``UNKNOWN_USAGE`` so the workflow fails closed.
"""
result = envelope.result
if isinstance(result, APIResult):
return result.usage if result.usage is not None else UNKNOWN_USAGE
if isinstance(result, InferenceResult):
if isinstance(result.usage, GenerationUsage):
return _inference_usage(result)
return UNKNOWN_USAGE
if isinstance(result, _NO_USAGE_RESULT_TYPES):
return None
# Any other result type either carries a usage field we do not map to
# APIUsage (embedding, agent, rag) or is an unknown model-calling type.
return UNKNOWN_USAGE


def _inference_usage(result: InferenceResult) -> APIUsage:
"""Map an inference result's usage, subtracting merged children's shares.

In a merged dispatch the parent's ``GenerationUsage`` is the whole batch
total, while each child carries its own share in ``result.children``. Each
child is also ingested separately, so the parent must record only its own
share (total minus the children's sum) or the workflow would count every
child's tokens and calls twice. ``wall_sec`` stays per task as reported.
"""
usage = result.usage
assert isinstance(usage, GenerationUsage)
prompt = usage.prompt_tokens
completion = usage.completion_tokens
calls = usage.num_requests
for child in result.children.values():
if not isinstance(child, InferenceResult):
continue
child_usage = child.usage
if not isinstance(child_usage, GenerationUsage):
continue
prompt -= child_usage.prompt_tokens
completion -= child_usage.completion_tokens
calls -= child_usage.num_requests
return APIUsage(
prompt_tokens=prompt,
completion_tokens=completion,
reasoning_tokens=0,
calls=calls,
failures=0,
retries=0,
truncated_calls=0,
wall_sec=usage.latency_sec,
)


def _resolve_artifact_path(filename: str) -> Path:
sanitized = Path(filename)
if (
Expand All @@ -77,6 +159,7 @@ async def ingest_result(
runtime: TaskRuntime = Depends(get_runtime),
event_monitor: EventMonitor = Depends(get_event_monitor),
results_dir: Path = Depends(get_results_dir),
registry: WorkflowRegistry = Depends(get_workflow_registry),
logger: logging.Logger = Depends(get_logger),
) -> PathResponse:
await require_permission(
Expand All @@ -97,6 +180,8 @@ async def ingest_result(
detail=f"Failed to store result: {exc}",
) from exc

await registry.save_task_usage_async(task_id, _task_usage_from_envelope(envelope))

expected_artifacts: list[str] = []
record = runtime.get_record(task_id)
if record:
Expand Down
Loading
Loading