Skip to content
Open
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
4 changes: 3 additions & 1 deletion src/server/routers/v1/tasks.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,9 @@ async def list_tasks(
runtime: TaskRuntime = Depends(get_runtime),
logger: logging.Logger = Depends(get_logger),
) -> list[TaskInfo]:
tasks = runtime.list_tasks()
tasks = await asyncio.to_thread(
runtime.list_tasks, request.query_params.get("workflow_id")
)
allowed = await resolve_accessible_ids(
principal, ResourceKind.TASK, ResourceAction.READ, logger
)
Expand Down
3 changes: 2 additions & 1 deletion src/server/task/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -1355,11 +1355,12 @@ def describe_task(self, task_id: str) -> TaskInfo | None:
return None
return self._build_task_info_locked(task_id, record)

def list_tasks(self) -> list[TaskInfo]:
def list_tasks(self, workflow_id: str | None = None) -> list[TaskInfo]:
with self._lock:
return [
self._build_task_info_locked(task_id, record)
for task_id, record in self._tasks.items()
if workflow_id is None or record.workflow_id == workflow_id
]

# ------------------------------------------------------------------ #
Expand Down
99 changes: 99 additions & 0 deletions tests/server/task/test_runtime_list_tasks.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,99 @@
"""Tests for `TaskRuntime.list_tasks` workflow-scoped filtering."""

import asyncio
import logging
from typing import Any, cast

from server.task.runtime import TaskRuntime

from .merge_harness import WorkerRegistryStub, build_runtime

_PAYLOAD = """
apiVersion: flowmesh/v1
kind: Workflow
metadata:
name: list-tasks
spec:
graph:
nodes:
- name: a
spec:
taskType: echo
- name: b
spec:
taskType: echo
"""


class _BuildSpyRuntime(TaskRuntime):
"""Counts TaskInfo builds to prove non-matching tasks are never materialised."""

def __init__(self, *args: Any, **kwargs: Any) -> None:
super().__init__(*args, **kwargs)
self.build_count = 0

def _build_task_info_locked(self, task_id: str, record: Any) -> Any:
self.build_count += 1
return super()._build_task_info_locked(task_id, record)


def _spy_runtime() -> _BuildSpyRuntime:
async def _noop(*args: Any, **kwargs: Any) -> None:
return None

registry = type(
"Registry",
(),
{
"register_workflow_async": _noop,
"commit_transition": lambda *a, **k: None,
"save_task_states_async": _noop,
"save_workflow_sched_async": _noop,
},
)()
return _BuildSpyRuntime(
cast(Any, registry), cast(Any, WorkerRegistryStub()), logging.getLogger("t")
)


async def _register(runtime: TaskRuntime, payload: str) -> str:
workflow_id, _ = await runtime.register("owner", "org", payload, format="native")
return workflow_id


def test_list_tasks_filters_by_workflow_id() -> None:
runtime, _ = build_runtime()
wf_a = asyncio.run(_register(runtime, _PAYLOAD))
wf_b = asyncio.run(_register(runtime, _PAYLOAD))

tasks_a = runtime.list_tasks(workflow_id=wf_a)
assert {t.task_id for t in tasks_a} == set(
runtime._tasks[task_id].task_id
for task_id in runtime._tasks
if runtime._tasks[task_id].workflow_id == wf_a
)
assert all(t.workflow_id == wf_a for t in tasks_a)
assert len(tasks_a) == 2
assert wf_b not in {t.workflow_id for t in tasks_a}


def test_list_tasks_without_filter_returns_everything() -> None:
runtime, _ = build_runtime()
wf_a = asyncio.run(_register(runtime, _PAYLOAD))
wf_b = asyncio.run(_register(runtime, _PAYLOAD))

tasks = runtime.list_tasks()
assert {t.workflow_id for t in tasks} == {wf_a, wf_b}
assert len(tasks) == 4


def test_workflow_filter_does_not_build_non_matching_tasks() -> None:
runtime = _spy_runtime()
wf_a = asyncio.run(_register(runtime, _PAYLOAD))
asyncio.run(_register(runtime, _PAYLOAD)) # second workflow, 2 more tasks

runtime.build_count = 0
tasks = runtime.list_tasks(workflow_id=wf_a)

assert len(tasks) == 2
assert runtime.build_count == 2
166 changes: 166 additions & 0 deletions tests/server/test_tasks_router.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
"""Router tests for `GET /api/v1/tasks` workflow-scoped, non-blocking listing."""

import asyncio
import logging
import time
from collections.abc import Iterator
from typing import Any
from unittest.mock import MagicMock

import pytest
from fastapi import FastAPI
from httpx import ASGITransport, AsyncClient
from lumid_hooks import ResourceRef

from server.app_state import get_runtime
from server.auth.security import PrincipalContext, authenticate_connection
from server.hooks import PERMISSION_CHECKERS
from server.routers.v1 import tasks as tasks_router
from server.task.runtime import TaskRuntime
from tests.server.task.merge_harness import build_runtime

PREFIX = "/api/v1"

_PAYLOAD = """
apiVersion: flowmesh/v1
kind: Workflow
metadata:
name: list-tasks
spec:
graph:
nodes:
- name: a
spec:
taskType: echo
- name: b
spec:
taskType: echo
"""


class _ScopedChecker:
"""Scopes task access to a fixed set of task ids."""

name = "scoped"

def __init__(self, allowed: frozenset[str]) -> None:
self._allowed = allowed

async def require(
self,
principal: PrincipalContext,
resource: ResourceRef,
action: str,
logger: logging.Logger,
) -> None:
return None

async def accessible_ids(
self,
principal: PrincipalContext,
kind: str,
action: str,
logger: logging.Logger,
) -> frozenset[str] | None:
return self._allowed


@pytest.fixture
def no_checkers() -> Iterator[None]:
PERMISSION_CHECKERS.clear()
yield
PERMISSION_CHECKERS.clear()


def _client(runtime: TaskRuntime) -> AsyncClient:
app = FastAPI()
app.state.logger = logging.getLogger("test.tasks_router")
app.include_router(tasks_router.router, prefix=PREFIX)
app.dependency_overrides[get_runtime] = lambda: runtime
app.dependency_overrides[authenticate_connection] = lambda: MagicMock(
spec=PrincipalContext
)
return AsyncClient(transport=ASGITransport(app=app), base_url="http://t")


def _task_ids(runtime: TaskRuntime, workflow_id: str) -> list[str]:
return [
task_id
for task_id, record in runtime.tasks.items()
if record.workflow_id == workflow_id
]


async def _register(runtime: TaskRuntime, payload: str) -> str:
workflow_id, _ = await runtime.register("owner", "org", payload, format="native")
return workflow_id


@pytest.mark.anyio
async def test_list_tasks_filters_by_workflow_id_with_permissions(
no_checkers: None,
) -> None:
runtime, _ = build_runtime()
wf_a = await _register(runtime, _PAYLOAD)
wf_b = await _register(runtime, _PAYLOAD)
ids_a = _task_ids(runtime, wf_a)
ids_b = _task_ids(runtime, wf_b)

# Permission checker scopes access to only one task of workflow A.
PERMISSION_CHECKERS.append(_ScopedChecker(frozenset({ids_a[0]})))
try:
async with _client(runtime) as ac:
resp = await ac.get(f"{PREFIX}/tasks", params={"workflow_id": wf_a})
finally:
PERMISSION_CHECKERS.clear()

assert resp.status_code == 200
body = resp.json()
assert [t["task_id"] for t in body] == [ids_a[0]]
assert all(t["workflow_id"] == wf_a for t in body)
# No task from the other workflow leaks through.
assert all(t["task_id"] not in ids_b for t in body)


@pytest.mark.anyio
async def test_event_loop_stays_responsive_during_list(no_checkers: None) -> None:
runtime, _ = build_runtime()
wf_a = await _register(runtime, _PAYLOAD)
await _register(runtime, _PAYLOAD)

# Patch the store call to block, as a slow store would.
orig = runtime.list_tasks

def slow_list(*args: Any, **kwargs: Any) -> Any:
time.sleep(0.5)
return orig(*args, **kwargs)

runtime.list_tasks = slow_list # type: ignore[method-assign]

app = FastAPI()
app.state.logger = logging.getLogger("test.tasks_router")
app.include_router(tasks_router.router, prefix=PREFIX)
app.dependency_overrides[get_runtime] = lambda: runtime
app.dependency_overrides[authenticate_connection] = lambda: MagicMock(
spec=PrincipalContext
)

@app.get("/ping")
async def ping() -> dict[str, bool]:
return {"ok": True}

async with AsyncClient(transport=ASGITransport(app=app), base_url="http://t") as ac:
t0 = time.monotonic()
list_task = asyncio.create_task(
ac.get(f"{PREFIX}/tasks", params={"workflow_id": wf_a})
)
await asyncio.sleep(0.05) # let the list request start and block
ping_resp = await ac.get("/ping")
elapsed = time.monotonic() - t0
list_resp = await list_task

assert ping_resp.status_code == 200
# The trivial request completes while the slow list is still in flight, so the
# whole exchange is far shorter than the store call's 0.5s block.
assert elapsed < 0.3
assert list_resp.status_code == 200
Loading