mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-11 14:38:38 +00:00
* fix(mcp): compensate cancelled task submissions * fix(mcp): shield submission compensation * docs(mcp): preserve notification lifecycle contract * fix(mcp): bound submission compensation wait --------- Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
1420 lines
45 KiB
Python
1420 lines
45 KiB
Python
import asyncio
|
|
import logging
|
|
from datetime import UTC, datetime, timedelta
|
|
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock
|
|
|
|
import pytest
|
|
|
|
import app.mcp_tasks.service as service_module
|
|
from app.mcp_tasks.errors import PermanentNotificationError
|
|
from app.mcp_tasks.service import McpTaskService
|
|
from deerflow.mcp.tasks import (
|
|
McpTaskDriverRegistry,
|
|
TaskSnapshot,
|
|
TaskStatus,
|
|
TaskSubmission,
|
|
TaskSubmitRequest,
|
|
)
|
|
from deerflow.mcp.tasks.ordinary import McpTaskProtocolError
|
|
from deerflow.persistence.mcp_tasks import DuplicateMcpRemoteTaskError
|
|
from deerflow.runtime.runs.manager import ConflictError
|
|
from deerflow.runtime.runs.schemas import RunStatus
|
|
|
|
|
|
class FakeRepository:
|
|
def __init__(self, rows=None):
|
|
self.rows = list(rows or [])
|
|
self.claimed = False
|
|
self.applied = []
|
|
self.released = []
|
|
self.created = []
|
|
|
|
async def create(self, **kwargs):
|
|
self.created.append(kwargs)
|
|
return {"id": kwargs["task_id"], **kwargs}
|
|
|
|
async def claim_due_tasks(self, **_kwargs):
|
|
if self.claimed:
|
|
return []
|
|
self.claimed = True
|
|
return [dict(row) for row in self.rows]
|
|
|
|
async def apply_snapshot(self, task_id, **kwargs):
|
|
self.applied.append((task_id, kwargs))
|
|
return True
|
|
|
|
async def release_claim(self, task_id, **kwargs):
|
|
self.released.append((task_id, kwargs))
|
|
return True
|
|
|
|
|
|
class FailingApplyRepository(FakeRepository):
|
|
async def apply_snapshot(self, task_id, **kwargs):
|
|
if task_id == "task-1":
|
|
raise RuntimeError("database unavailable")
|
|
return await super().apply_snapshot(task_id, **kwargs)
|
|
|
|
|
|
class FailingCreateRepository(FakeRepository):
|
|
async def create(self, **kwargs):
|
|
self.created.append(kwargs)
|
|
raise RuntimeError("database unavailable")
|
|
|
|
|
|
class BlockingCreateRepository(FakeRepository):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.create_started = asyncio.Event()
|
|
|
|
async def create(self, **kwargs):
|
|
self.created.append(kwargs)
|
|
self.create_started.set()
|
|
await asyncio.Event().wait()
|
|
|
|
|
|
class DuplicateCreateRepository(FakeRepository):
|
|
async def create(self, **kwargs):
|
|
self.created.append(kwargs)
|
|
raise DuplicateMcpRemoteTaskError("already tracked")
|
|
|
|
|
|
class FakeDriver:
|
|
def __init__(
|
|
self,
|
|
snapshots=None,
|
|
*,
|
|
submission=None,
|
|
error: Exception | None = None,
|
|
cancel_error: Exception | None = None,
|
|
):
|
|
self.snapshots = list(snapshots or [])
|
|
self.submission = submission
|
|
self.error = error
|
|
self.cancel_error = cancel_error
|
|
self.status_calls = []
|
|
self.submit_calls = []
|
|
self.cancel_calls = []
|
|
|
|
async def submit(self, request):
|
|
self.submit_calls.append(request)
|
|
if self.submission is None:
|
|
raise AssertionError(f"unexpected submit: {request}")
|
|
return self.submission
|
|
|
|
async def get_status(self, task):
|
|
self.status_calls.append(task)
|
|
if self.error is not None:
|
|
raise self.error
|
|
return self.snapshots.pop(0)
|
|
|
|
async def cancel(self, task):
|
|
self.cancel_calls.append(task)
|
|
if self.cancel_error is not None:
|
|
raise self.cancel_error
|
|
return TaskSnapshot(status=TaskStatus.CANCELLED)
|
|
|
|
|
|
class HangingDriver(FakeDriver):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.started = asyncio.Event()
|
|
self.cancelled = False
|
|
|
|
async def get_status(self, task):
|
|
self.status_calls.append(task)
|
|
self.started.set()
|
|
try:
|
|
await asyncio.Event().wait()
|
|
except asyncio.CancelledError:
|
|
self.cancelled = True
|
|
raise
|
|
|
|
|
|
class BlockingCancelDriver(FakeDriver):
|
|
def __init__(self, *, submission):
|
|
super().__init__(submission=submission)
|
|
self.cancel_started = asyncio.Event()
|
|
self.finish_cancel = asyncio.Event()
|
|
self.cancel_finished = asyncio.Event()
|
|
self.cancel_completed = False
|
|
self.cancel_interrupted = False
|
|
|
|
async def cancel(self, task):
|
|
self.cancel_calls.append(task)
|
|
self.cancel_started.set()
|
|
try:
|
|
await self.finish_cancel.wait()
|
|
except asyncio.CancelledError:
|
|
self.cancel_interrupted = True
|
|
self.cancel_finished.set()
|
|
raise
|
|
self.cancel_completed = True
|
|
self.cancel_finished.set()
|
|
return TaskSnapshot(status=TaskStatus.CANCELLED)
|
|
|
|
|
|
def _claimed_row(*, driver_name="fake"):
|
|
return {
|
|
"id": "task-1",
|
|
"user_id": "user-1",
|
|
"thread_id": "thread-1",
|
|
"run_id": "run-1",
|
|
"tool_call_id": "call-1",
|
|
"server_name": "reports",
|
|
"driver_name": driver_name,
|
|
"remote_task_id": "remote-1",
|
|
"task_name": "Generate report",
|
|
"status": "working",
|
|
"driver_data": {"status_tool": "status"},
|
|
"lease_owner": "ignored-by-service-fixture",
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_persists_remote_handle_before_returning():
|
|
now = datetime.now(UTC)
|
|
repo = FakeRepository()
|
|
driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED, poll_after_seconds=9),
|
|
driver_data={"status_tool": "status"},
|
|
)
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
request = TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={"topic": "MCP"},
|
|
driver_data={"submit_tool": "submit"},
|
|
)
|
|
|
|
created = await service.submit(driver_name="fake", request=request, now=now)
|
|
|
|
assert created["remote_task_id"] == "remote-1"
|
|
persisted = repo.created[0]
|
|
assert persisted["next_poll_at"] == now + timedelta(seconds=9)
|
|
assert persisted["driver_data"] == {"submit_tool": "submit", "status_tool": "status"}
|
|
assert driver.submit_calls[0].local_task_id == created["id"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_cancels_remote_task_when_persistence_fails():
|
|
repo = FailingCreateRepository()
|
|
driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
driver_data={"status_tool": "status", "cancel_tool": "cancel"},
|
|
)
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
request = TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={"topic": "MCP"},
|
|
driver_data={"submit_tool": "submit"},
|
|
local_task_id="task-1",
|
|
)
|
|
|
|
with pytest.raises(RuntimeError, match="database unavailable"):
|
|
await service.submit(driver_name="fake", request=request)
|
|
|
|
assert len(driver.cancel_calls) == 1
|
|
cancelled = driver.cancel_calls[0]
|
|
assert cancelled.local_task_id == "task-1"
|
|
assert cancelled.remote_task_id == "remote-1"
|
|
assert cancelled.driver_data == {
|
|
"submit_tool": "submit",
|
|
"status_tool": "status",
|
|
"cancel_tool": "cancel",
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_cancellation_during_persistence_cancels_remote_task():
|
|
repo = BlockingCreateRepository()
|
|
driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
driver_data={"cancel_tool": "cancel"},
|
|
)
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
request = TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={"topic": "MCP"},
|
|
local_task_id="task-1",
|
|
)
|
|
|
|
submit_task = asyncio.create_task(service.submit(driver_name="fake", request=request))
|
|
await repo.create_started.wait()
|
|
submit_task.cancel()
|
|
|
|
with pytest.raises(asyncio.CancelledError):
|
|
await submit_task
|
|
|
|
assert len(driver.cancel_calls) == 1
|
|
cancelled = driver.cancel_calls[0]
|
|
assert cancelled.local_task_id == "task-1"
|
|
assert cancelled.remote_task_id == "remote-1"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_repeated_cancellation_does_not_interrupt_compensation():
|
|
repo = BlockingCreateRepository()
|
|
driver = BlockingCancelDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
driver_data={"cancel_tool": "cancel"},
|
|
)
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
request = TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={"topic": "MCP"},
|
|
local_task_id="task-1",
|
|
)
|
|
|
|
submit_task = asyncio.create_task(service.submit(driver_name="fake", request=request))
|
|
await repo.create_started.wait()
|
|
submit_task.cancel()
|
|
await driver.cancel_started.wait()
|
|
|
|
submit_task.cancel()
|
|
driver.finish_cancel.set()
|
|
|
|
with pytest.raises(asyncio.CancelledError):
|
|
await submit_task
|
|
|
|
assert len(driver.cancel_calls) == 1
|
|
assert driver.cancel_completed
|
|
assert not driver.cancel_interrupted
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_stops_waiting_for_hung_compensation_without_cancelling_it(monkeypatch, caplog):
|
|
monkeypatch.setattr(service_module, "_UNTRACKED_TASK_COMPENSATION_WAIT_SECONDS", 0)
|
|
repo = BlockingCreateRepository()
|
|
driver = BlockingCancelDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
driver_data={"cancel_tool": "cancel"},
|
|
)
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
request = TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={"topic": "MCP"},
|
|
local_task_id="task-1",
|
|
)
|
|
|
|
submit_task = asyncio.create_task(service.submit(driver_name="fake", request=request))
|
|
await repo.create_started.wait()
|
|
submit_task.cancel()
|
|
|
|
with caplog.at_level(logging.WARNING), pytest.raises(asyncio.CancelledError):
|
|
await submit_task
|
|
|
|
assert "cancellation continues in the background" in caplog.text
|
|
await driver.cancel_started.wait()
|
|
assert not driver.cancel_interrupted
|
|
assert not driver.cancel_completed
|
|
|
|
driver.finish_cancel.set()
|
|
await driver.cancel_finished.wait()
|
|
|
|
assert len(driver.cancel_calls) == 1
|
|
assert driver.cancel_completed
|
|
assert not driver.cancel_interrupted
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_cancellation_preserves_cancelled_error_when_compensation_fails(caplog):
|
|
repo = BlockingCreateRepository()
|
|
driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
),
|
|
cancel_error=RuntimeError("cancel unavailable"),
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
request = TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={},
|
|
local_task_id="task-1",
|
|
)
|
|
|
|
submit_task = asyncio.create_task(service.submit(driver_name="fake", request=request))
|
|
await repo.create_started.wait()
|
|
submit_task.cancel()
|
|
|
|
with caplog.at_level(logging.ERROR), pytest.raises(asyncio.CancelledError):
|
|
await submit_task
|
|
|
|
assert len(driver.cancel_calls) == 1
|
|
assert "Failed to cancel untracked MCP task" in caplog.text
|
|
assert "cancel unavailable" in caplog.text
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_cancels_remote_task_when_its_id_exceeds_storage_limit():
|
|
repo = FakeRepository()
|
|
driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="r" * 256,
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
driver_data={"cancel_tool": "cancel"},
|
|
)
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
with pytest.raises(McpTaskProtocolError, match="remote_task_id.*255"):
|
|
await service.submit(
|
|
driver_name="fake",
|
|
request=TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={},
|
|
local_task_id="task-1",
|
|
),
|
|
)
|
|
|
|
assert repo.created == []
|
|
assert len(driver.cancel_calls) == 1
|
|
assert driver.cancel_calls[0].remote_task_id == "r" * 256
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_duplicate_remote_handle_is_rejected_without_cancelling_existing_task():
|
|
repo = DuplicateCreateRepository()
|
|
driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
driver_data={"cancel_tool": "cancel"},
|
|
)
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
with pytest.raises(DuplicateMcpRemoteTaskError, match="already tracked"):
|
|
await service.submit(
|
|
driver_name="fake",
|
|
request=TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-2",
|
|
run_id="run-2",
|
|
tool_call_id="call-2",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={},
|
|
),
|
|
)
|
|
|
|
assert driver.cancel_calls == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_cancel_task_persists_request_without_calling_remote():
|
|
record = {**_claimed_row(), "cancel_requested_at": datetime.now(UTC).isoformat()}
|
|
repo = SimpleNamespace(
|
|
request_cancel=AsyncMock(return_value=record),
|
|
claim_cancel_requests=AsyncMock(return_value=[{**record, "cancel_attempt_count": 1}]),
|
|
apply_cancel_snapshot=AsyncMock(return_value=True),
|
|
release_cancel_claim=AsyncMock(return_value=True),
|
|
get=AsyncMock(return_value={**record, "status": "cancelled"}),
|
|
)
|
|
driver = FakeDriver()
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
result = await service.cancel_task(
|
|
task_id="task-1",
|
|
thread_id="thread-1",
|
|
user_id="user-1",
|
|
)
|
|
|
|
assert result == record
|
|
assert driver.cancel_calls == []
|
|
repo.claim_cancel_requests.assert_not_awaited()
|
|
repo.apply_cancel_snapshot.assert_not_awaited()
|
|
repo.release_cancel_claim.assert_not_awaited()
|
|
repo.get.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_cancel_failure_schedules_retry_from_call_completion_time():
|
|
record = {**_claimed_row(), "cancel_attempt_count": 1}
|
|
repo = SimpleNamespace(
|
|
claim_cancel_requests=AsyncMock(return_value=[record]),
|
|
release_cancel_claim=AsyncMock(return_value=True),
|
|
claim_due_tasks=AsyncMock(return_value=[]),
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", FakeDriver(cancel_error=RuntimeError("cancel unavailable")))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
scan_started_at = datetime(2000, 1, 1, tzinfo=UTC)
|
|
|
|
await service.run_once(now=scan_started_at)
|
|
|
|
released = repo.release_cancel_claim.await_args.kwargs
|
|
retry_started_at = released["next_cancel_at"] - timedelta(seconds=5)
|
|
assert retry_started_at > scan_started_at
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_cancel_recovery_failures_are_isolated_and_later_phases_continue(caplog):
|
|
records = [
|
|
{**_claimed_row(), "id": "task-broken", "cancel_attempt_count": 1},
|
|
{**_claimed_row(), "id": "task-sibling", "remote_task_id": "remote-2", "cancel_attempt_count": 1},
|
|
]
|
|
|
|
async def release_cancel_claim(task_id, **_kwargs):
|
|
if task_id == "task-broken":
|
|
raise RuntimeError("cancel recovery store unavailable")
|
|
return True
|
|
|
|
repo = SimpleNamespace(
|
|
claim_cancel_requests=AsyncMock(return_value=records),
|
|
release_cancel_claim=AsyncMock(side_effect=release_cancel_claim),
|
|
claim_due_tasks=AsyncMock(return_value=[]),
|
|
claim_notification_work=AsyncMock(return_value=[]),
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", FakeDriver(cancel_error=RuntimeError("cancel unavailable")))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=AsyncMock(),
|
|
get_run=AsyncMock(),
|
|
)
|
|
|
|
with caplog.at_level(logging.ERROR):
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
assert repo.release_cancel_claim.await_count == 2
|
|
repo.claim_due_tasks.assert_awaited_once()
|
|
repo.claim_notification_work.assert_awaited_once()
|
|
assert "task-broken" in caplog.text
|
|
assert "cancel recovery store unavailable" in caplog.text
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_notification_delivery_waits_for_successful_agent_run():
|
|
repo = SimpleNamespace(
|
|
mark_notification_dispatched=AsyncMock(return_value=True),
|
|
finish_notification_run=AsyncMock(return_value=True),
|
|
release_notification_claim=AsyncMock(return_value=True),
|
|
defer_dispatched_notification=AsyncMock(return_value=True),
|
|
)
|
|
launch = AsyncMock(return_value={"run_id": "notify-run-1"})
|
|
get_run = AsyncMock(return_value=SimpleNamespace(status=RunStatus.running))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=launch,
|
|
get_run=get_run,
|
|
)
|
|
now = datetime.now(UTC)
|
|
claimed = {
|
|
**_claimed_row(),
|
|
"notification_status": "claimed",
|
|
"dispatch_version": 2,
|
|
"dispatch_attempt": 0,
|
|
"dispatch_event": {"status": "completed"},
|
|
}
|
|
|
|
await service._notify_one(claimed, now=now)
|
|
|
|
repo.mark_notification_dispatched.assert_awaited_once()
|
|
repo.finish_notification_run.assert_not_awaited()
|
|
|
|
get_run.return_value = SimpleNamespace(status=RunStatus.success)
|
|
await service._notify_one(
|
|
{
|
|
**claimed,
|
|
"notification_status": "dispatched",
|
|
"notification_run_id": "notify-run-1",
|
|
},
|
|
now=now,
|
|
)
|
|
repo.finish_notification_run.assert_awaited_once()
|
|
assert repo.finish_notification_run.await_args.kwargs["delivered"] is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_missing_dispatched_notification_run_retries_delivery():
|
|
repo = SimpleNamespace(
|
|
finish_notification_run=AsyncMock(return_value=True),
|
|
defer_dispatched_notification=AsyncMock(return_value=True),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=AsyncMock(),
|
|
get_run=AsyncMock(return_value=None),
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
await service._notify_one(
|
|
{
|
|
**_claimed_row(),
|
|
"notification_status": "dispatched",
|
|
"notification_run_id": "missing-run",
|
|
"dispatch_version": 2,
|
|
"notification_attempt_count": 2,
|
|
},
|
|
now=now,
|
|
)
|
|
|
|
repo.defer_dispatched_notification.assert_not_awaited()
|
|
repo.finish_notification_run.assert_awaited_once()
|
|
finished = repo.finish_notification_run.await_args.kwargs
|
|
assert finished["delivered"] is False
|
|
assert finished["next_notification_at"] == now + timedelta(seconds=20)
|
|
assert "missing-run" in finished["error"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_notification_failures_are_isolated_and_release_their_lease(caplog):
|
|
records = [
|
|
{
|
|
**_claimed_row(),
|
|
"id": "task-broken",
|
|
"notification_status": "dispatched",
|
|
"notification_run_id": "run-broken",
|
|
"dispatch_version": 2,
|
|
},
|
|
{
|
|
**_claimed_row(),
|
|
"id": "task-success",
|
|
"notification_status": "dispatched",
|
|
"notification_run_id": "run-success",
|
|
"dispatch_version": 3,
|
|
},
|
|
]
|
|
repo = SimpleNamespace(
|
|
claim_notification_work=AsyncMock(return_value=records),
|
|
finish_notification_run=AsyncMock(return_value=True),
|
|
defer_dispatched_notification=AsyncMock(return_value=True),
|
|
release_notification_lease=AsyncMock(return_value=True),
|
|
)
|
|
get_run = AsyncMock(
|
|
side_effect=[
|
|
RuntimeError("run store unavailable"),
|
|
SimpleNamespace(status=RunStatus.success),
|
|
]
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=AsyncMock(),
|
|
get_run=get_run,
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
with caplog.at_level(logging.ERROR):
|
|
await service._run_notifications(now=now)
|
|
|
|
repo.finish_notification_run.assert_awaited_once()
|
|
assert repo.finish_notification_run.await_args.args[0] == "task-success"
|
|
repo.release_notification_lease.assert_awaited_once()
|
|
released = repo.release_notification_lease.await_args
|
|
assert released.args[0] == "task-broken"
|
|
assert released.kwargs["next_notification_at"] == now + timedelta(seconds=5)
|
|
assert "run store unavailable" in released.kwargs["error"]
|
|
assert "task-broken" in caplog.text
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_notification_busy_thread_replaces_claim_with_latest_event():
|
|
repo = SimpleNamespace(
|
|
release_notification_claim=AsyncMock(return_value=True),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=AsyncMock(side_effect=ConflictError("thread busy")),
|
|
get_run=AsyncMock(return_value=SimpleNamespace(assistant_id="lead_agent")),
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
await service._notify_one(
|
|
{
|
|
**_claimed_row(),
|
|
"notification_status": "claimed",
|
|
"dispatch_version": 2,
|
|
"dispatch_attempt": 0,
|
|
"dispatch_event": {"status": "input_required"},
|
|
},
|
|
now=now,
|
|
)
|
|
|
|
released = repo.release_notification_claim.await_args.kwargs
|
|
assert released["replace_with_latest"] is True
|
|
assert released["next_notification_at"] == now + timedelta(seconds=5)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_notification_launch_failure_backs_off_and_replaces_with_latest_event():
|
|
repo = SimpleNamespace(
|
|
release_notification_claim=AsyncMock(return_value=True),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
max_poll_backoff_seconds=300,
|
|
launch_notification=AsyncMock(side_effect=RuntimeError("run store unavailable")),
|
|
get_run=AsyncMock(return_value=SimpleNamespace(assistant_id="lead_agent")),
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
await service._notify_one(
|
|
{
|
|
**_claimed_row(),
|
|
"notification_status": "claimed",
|
|
"dispatch_version": 2,
|
|
"dispatch_attempt": 0,
|
|
"notification_attempt_count": 3,
|
|
"dispatch_event": {"status": "input_required"},
|
|
},
|
|
now=now,
|
|
)
|
|
|
|
released = repo.release_notification_claim.await_args.kwargs
|
|
assert released["replace_with_latest"] is True
|
|
assert released["count_failure"] is True
|
|
assert released["next_notification_at"] == now + timedelta(seconds=40)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_permanently_rejected_notification_is_dead_lettered():
|
|
repo = SimpleNamespace(
|
|
dead_letter_notification=AsyncMock(return_value=True),
|
|
release_notification_claim=AsyncMock(return_value=True),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=AsyncMock(side_effect=PermanentNotificationError("Thread thread-1 not found")),
|
|
get_run=AsyncMock(return_value=SimpleNamespace(assistant_id="lead_agent")),
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
await service._notify_one(
|
|
{
|
|
**_claimed_row(),
|
|
"notification_status": "claimed",
|
|
"dispatch_version": 2,
|
|
"dispatch_attempt": 0,
|
|
"notification_attempt_count": 0,
|
|
"dispatch_event": {"status": "completed"},
|
|
},
|
|
now=now,
|
|
)
|
|
|
|
repo.dead_letter_notification.assert_awaited_once()
|
|
dead_lettered = repo.dead_letter_notification.await_args.kwargs
|
|
assert dead_lettered["dispatch_version"] == 2
|
|
assert "not found" in dead_lettered["error"]
|
|
assert dead_lettered["count_failure"] is True
|
|
repo.release_notification_claim.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_notification_retry_budget_dead_letters_before_creating_another_run():
|
|
repo = SimpleNamespace(
|
|
dead_letter_notification=AsyncMock(return_value=True),
|
|
)
|
|
launch_notification = AsyncMock()
|
|
get_run = AsyncMock()
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=launch_notification,
|
|
get_run=get_run,
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
await service._notify_one(
|
|
{
|
|
**_claimed_row(),
|
|
"notification_status": "retry",
|
|
"notification_error": "Agent run failed",
|
|
"dispatch_version": 2,
|
|
"dispatch_attempt": 5,
|
|
"notification_attempt_count": 5,
|
|
"dispatch_event": {"status": "completed"},
|
|
},
|
|
now=now,
|
|
)
|
|
|
|
launch_notification.assert_not_awaited()
|
|
get_run.assert_not_awaited()
|
|
dead_lettered = repo.dead_letter_notification.await_args.kwargs
|
|
assert dead_lettered["dispatch_version"] == 2
|
|
assert dead_lettered["count_failure"] is False
|
|
assert "5 failed attempts" in dead_lettered["error"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_dispatched_notification_retry_budget_dead_letters_before_hydrating_run():
|
|
repo = SimpleNamespace(
|
|
dead_letter_notification=AsyncMock(return_value=True),
|
|
)
|
|
get_run = AsyncMock()
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
launch_notification=AsyncMock(),
|
|
get_run=get_run,
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
await service._notify_one(
|
|
{
|
|
**_claimed_row(),
|
|
"notification_status": "dispatched",
|
|
"notification_run_id": "notify-run-1",
|
|
"notification_error": "run store unavailable",
|
|
"dispatch_version": 2,
|
|
"notification_attempt_count": 5,
|
|
},
|
|
now=now,
|
|
)
|
|
|
|
get_run.assert_not_awaited()
|
|
dead_lettered = repo.dead_letter_notification.await_args.kwargs
|
|
assert dead_lettered["dispatch_version"] == 2
|
|
assert dead_lettered["count_failure"] is False
|
|
assert "5 failed attempts" in dead_lettered["error"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_submit_preserves_persistence_error_when_compensation_cancel_fails(caplog):
|
|
repo = FailingCreateRepository()
|
|
driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.SUBMITTED),
|
|
),
|
|
cancel_error=RuntimeError("cancel unavailable"),
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
request = TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={"topic": "MCP"},
|
|
local_task_id="task-1",
|
|
)
|
|
|
|
with caplog.at_level(logging.ERROR), pytest.raises(RuntimeError, match="database unavailable"):
|
|
await service.submit(driver_name="fake", request=request)
|
|
|
|
assert "Failed to cancel untracked MCP task" in caplog.text
|
|
assert "cancel unavailable" in caplog.text
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_once_polls_without_an_llm_and_schedules_next_poll():
|
|
repo = FakeRepository([_claimed_row()])
|
|
driver = FakeDriver([TaskSnapshot(status=TaskStatus.WORKING, poll_after_seconds=12)])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
scan_started_at = datetime(2000, 1, 1, tzinfo=UTC)
|
|
|
|
await service.run_once(now=scan_started_at)
|
|
|
|
assert driver.status_calls[0].remote_task_id == "remote-1"
|
|
_, update = repo.applied[0]
|
|
assert update["status"] == "working"
|
|
assert update["next_poll_at"] == update["polled_at"] + timedelta(seconds=12)
|
|
assert update["polled_at"] > scan_started_at
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_once_caps_remote_poll_hint_to_one_day():
|
|
repo = FakeRepository([_claimed_row()])
|
|
driver = FakeDriver([TaskSnapshot(status=TaskStatus.WORKING, poll_after_seconds=1e20)])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
_, update = repo.applied[0]
|
|
assert update["next_poll_at"] == update["polled_at"] + timedelta(days=1)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_once_schedules_driver_error_retry_from_poll_completion_time():
|
|
repo = FakeRepository([_claimed_row()])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", FakeDriver(error=RuntimeError("network down")))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
scan_started_at = datetime(2000, 1, 1, tzinfo=UTC)
|
|
|
|
await service.run_once(now=scan_started_at)
|
|
|
|
_, released = repo.released[0]
|
|
retry_started_at = released["next_poll_at"] - timedelta(seconds=5)
|
|
assert retry_started_at > scan_started_at
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_once_stops_terminal_tasks_but_keeps_input_required_on_a_slow_poll():
|
|
rows = [_claimed_row(), {**_claimed_row(), "id": "task-2", "remote_task_id": "remote-2"}]
|
|
repo = FakeRepository(rows)
|
|
driver = FakeDriver(
|
|
[
|
|
TaskSnapshot(status=TaskStatus.COMPLETED, result={"report": "ready"}),
|
|
TaskSnapshot(status=TaskStatus.INPUT_REQUIRED, input_required={"prompt": "Approve?"}),
|
|
]
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
updates = {task_id: update for task_id, update in repo.applied}
|
|
assert updates["task-1"]["status"] == "completed"
|
|
assert updates["task-1"]["next_poll_at"] is None
|
|
assert updates["task-2"]["status"] == "input_required"
|
|
assert updates["task-2"]["input_required"] == {"prompt": "Approve?"}
|
|
assert updates["task-2"]["next_poll_at"] >= updates["task-2"]["polled_at"] + timedelta(seconds=60)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_once_uses_exponential_backoff_and_caps_transient_errors():
|
|
rows = [
|
|
{**_claimed_row(), "id": "task-1", "consecutive_poll_error_count": 0},
|
|
{**_claimed_row(), "id": "task-2", "consecutive_poll_error_count": 4},
|
|
]
|
|
repo = FakeRepository(rows)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", FakeDriver(error=RuntimeError("network down")))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
max_poll_backoff_seconds=30,
|
|
)
|
|
|
|
started_at = datetime.now(UTC)
|
|
await service.run_once(now=started_at)
|
|
finished_at = datetime.now(UTC)
|
|
|
|
released = {task_id: update for task_id, update in repo.released}
|
|
assert started_at + timedelta(seconds=5) <= released["task-1"]["next_poll_at"] <= finished_at + timedelta(seconds=5)
|
|
assert started_at + timedelta(seconds=30) <= released["task-2"]["next_poll_at"] <= finished_at + timedelta(seconds=30)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_protocol_error_terminalizes_instead_of_retrying_forever():
|
|
repo = FakeRepository([_claimed_row()])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", FakeDriver(error=McpTaskProtocolError("missing structuredContent")))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
assert repo.released == []
|
|
_, applied = repo.applied[0]
|
|
assert applied["status"] == "failed"
|
|
assert applied["error"] == "missing structuredContent"
|
|
assert applied["next_poll_at"] is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_protocol_error_message_is_bounded_before_terminal_persistence():
|
|
oversized_error = "e" * 5_000
|
|
repo = FakeRepository([_claimed_row()])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", FakeDriver(error=McpTaskProtocolError(oversized_error)))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
_, applied = repo.applied[0]
|
|
assert applied["status"] == "failed"
|
|
assert applied["error"] == oversized_error[:4_000]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_persisted_snapshot_errors_are_bounded_on_submit_and_poll():
|
|
oversized_error = "e" * 5_000
|
|
submit_repo = FakeRepository()
|
|
submit_driver = FakeDriver(
|
|
submission=TaskSubmission(
|
|
remote_task_id="remote-1",
|
|
snapshot=TaskSnapshot(status=TaskStatus.FAILED, error=oversized_error),
|
|
)
|
|
)
|
|
submit_registry = McpTaskDriverRegistry()
|
|
submit_registry.register("fake", submit_driver)
|
|
submit_service = McpTaskService(
|
|
repository=submit_repo,
|
|
drivers=submit_registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await submit_service.submit(
|
|
driver_name="fake",
|
|
request=TaskSubmitRequest(
|
|
user_id="user-1",
|
|
thread_id="thread-1",
|
|
run_id="run-1",
|
|
tool_call_id="call-1",
|
|
server_name="reports",
|
|
task_name="Generate report",
|
|
arguments={},
|
|
),
|
|
)
|
|
|
|
assert submit_repo.created[0]["error"] == oversized_error[:4_000]
|
|
|
|
poll_repo = FakeRepository([_claimed_row()])
|
|
poll_registry = McpTaskDriverRegistry()
|
|
poll_registry.register(
|
|
"fake",
|
|
FakeDriver([TaskSnapshot(status=TaskStatus.FAILED, error=oversized_error)]),
|
|
)
|
|
poll_service = McpTaskService(
|
|
repository=poll_repo,
|
|
drivers=poll_registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await poll_service.run_once(now=datetime.now(UTC))
|
|
|
|
_, applied = poll_repo.applied[0]
|
|
assert applied["error"] == oversized_error[:4_000]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_oversized_input_required_payload_terminalizes_without_persisting_it():
|
|
repo = FakeRepository([_claimed_row()])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register(
|
|
"fake",
|
|
FakeDriver(
|
|
[
|
|
TaskSnapshot(
|
|
status=TaskStatus.INPUT_REQUIRED,
|
|
input_required={"prompt": "x" * 65_536},
|
|
)
|
|
]
|
|
),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
assert repo.released == []
|
|
_, applied = repo.applied[0]
|
|
assert applied["status"] == "failed"
|
|
assert applied["input_required"] is None
|
|
assert "input_required payload exceeds the 65536-byte limit" in applied["error"]
|
|
assert applied["next_poll_at"] is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_oversized_result_stores_preview_without_invalid_truncated_json():
|
|
repo = FakeRepository([_claimed_row()])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register(
|
|
"fake",
|
|
FakeDriver(
|
|
[
|
|
TaskSnapshot(
|
|
status=TaskStatus.COMPLETED,
|
|
result={"report": "x" * 200},
|
|
result_artifact={"uri": "s3://reports/1.json", "mime_type": "application/json"},
|
|
)
|
|
]
|
|
),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
max_result_bytes=64,
|
|
result_preview_max_chars=24,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
_, applied = repo.applied[0]
|
|
assert applied["result"] is None
|
|
assert len(applied["result_preview"]) == 24
|
|
assert applied["result_truncated"] is True
|
|
assert applied["result_artifact"] == {
|
|
"uri": "s3://reports/1.json",
|
|
"mime_type": "application/json",
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_oversized_result_artifact_terminalizes_without_persisting_it():
|
|
repo = FakeRepository([_claimed_row()])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register(
|
|
"fake",
|
|
FakeDriver(
|
|
[
|
|
TaskSnapshot(
|
|
status=TaskStatus.COMPLETED,
|
|
result_artifact={
|
|
"uri": "https://example.test/" + "x" * 65_536,
|
|
"mime_type": "application/json",
|
|
},
|
|
)
|
|
]
|
|
),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
assert repo.released == []
|
|
_, applied = repo.applied[0]
|
|
assert applied["status"] == "failed"
|
|
assert applied["result_artifact"] is None
|
|
assert "result_artifact payload exceeds the 65536-byte limit" in applied["error"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_non_json_numeric_result_is_a_permanent_protocol_failure():
|
|
repo = FakeRepository([_claimed_row()])
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register(
|
|
"fake",
|
|
FakeDriver(
|
|
[
|
|
TaskSnapshot(
|
|
status=TaskStatus.COMPLETED,
|
|
result={"score": float("nan")},
|
|
)
|
|
]
|
|
),
|
|
)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
assert repo.released == []
|
|
_, applied = repo.applied[0]
|
|
assert applied["status"] == "failed"
|
|
assert "not valid JSON" in applied["error"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_once_releases_claim_when_driver_is_missing_or_fails():
|
|
rows = [_claimed_row(driver_name="missing"), {**_claimed_row(), "id": "task-2", "remote_task_id": "remote-2", "driver_name": "broken"}]
|
|
repo = FakeRepository(rows)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("broken", FakeDriver(error=RuntimeError("network down")))
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
now = datetime.now(UTC)
|
|
|
|
await service.run_once(now=now)
|
|
|
|
released = {task_id: update for task_id, update in repo.released}
|
|
assert "No MCP task driver registered" in released["task-1"]["error"]
|
|
assert released["task-2"]["error"] == "network down"
|
|
assert released["task-1"]["next_poll_at"] == now + timedelta(seconds=5)
|
|
assert released["task-2"]["next_poll_at"] > now + timedelta(seconds=5)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_once_isolates_unexpected_failure_to_its_claimed_task(caplog):
|
|
rows = [_claimed_row(), {**_claimed_row(), "id": "task-2", "remote_task_id": "remote-2"}]
|
|
repo = FailingApplyRepository(rows)
|
|
driver = FakeDriver(
|
|
[
|
|
TaskSnapshot(status=TaskStatus.COMPLETED, result={"report": "first"}),
|
|
TaskSnapshot(status=TaskStatus.COMPLETED, result={"report": "second"}),
|
|
]
|
|
)
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=5,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
with caplog.at_level(logging.ERROR):
|
|
await service.run_once(now=datetime.now(UTC))
|
|
|
|
assert [task_id for task_id, _update in repo.applied] == ["task-2"]
|
|
assert "task_id=task-1" in caplog.text
|
|
assert "database unavailable" in caplog.text
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_start_runs_recovery_poll_immediately_and_stop_is_clean():
|
|
repo = FakeRepository([])
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=McpTaskDriverRegistry(),
|
|
poll_interval_seconds=60,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.start()
|
|
for _ in range(20):
|
|
if repo.claimed:
|
|
break
|
|
await __import__("asyncio").sleep(0)
|
|
await service.stop()
|
|
|
|
assert repo.claimed is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_stop_cancels_a_hung_driver_poll():
|
|
repo = FakeRepository([_claimed_row()])
|
|
driver = HangingDriver()
|
|
registry = McpTaskDriverRegistry()
|
|
registry.register("fake", driver)
|
|
service = McpTaskService(
|
|
repository=repo,
|
|
drivers=registry,
|
|
poll_interval_seconds=60,
|
|
lease_seconds=120,
|
|
max_concurrent_polls=3,
|
|
)
|
|
|
|
await service.start()
|
|
await asyncio.wait_for(driver.started.wait(), timeout=1)
|
|
await asyncio.wait_for(service.stop(), timeout=1)
|
|
|
|
assert driver.cancelled is True
|