mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-15 00:19:14 +00:00
* feat(mcp): add durable task runtime foundation * fix(chart): sync embedded config version * fix(mcp): isolate task polls during shutdown * feat(mcp): track consecutive poll errors on mcp_tasks poll_attempt_count grows on every claim (successful polls included), so it cannot drive a failure backoff without misjudging normal long tasks. Add consecutive_poll_error_count: incremented when a claim is released after a poll error, reset to zero by any applied snapshot. The backoff/terminal policy that consumes it lands with the first concrete driver. * fix(mcp): harden durable task lifecycle * fix(mcp): preserve tracked task on dedup conflict --------- Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
214 lines
7.9 KiB
Python
214 lines
7.9 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
import socket
|
|
import uuid
|
|
from dataclasses import replace
|
|
from datetime import UTC, datetime, timedelta
|
|
|
|
from deerflow.mcp.tasks import McpTaskDriverRegistry, TaskReference, TaskSnapshot, TaskSubmitRequest
|
|
from deerflow.persistence.mcp_tasks import DuplicateMcpRemoteTaskError
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_MAX_POLL_ERROR_CHARS = 4000
|
|
|
|
|
|
class McpTaskService:
|
|
"""Persist and poll long-running MCP tasks outside the Agent loop."""
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
repository,
|
|
drivers: McpTaskDriverRegistry,
|
|
poll_interval_seconds: int,
|
|
lease_seconds: int,
|
|
max_concurrent_polls: int,
|
|
) -> None:
|
|
self._repository = repository
|
|
self._drivers = drivers
|
|
self._poll_interval_seconds = poll_interval_seconds
|
|
self._lease_seconds = lease_seconds
|
|
self._max_concurrent_polls = max_concurrent_polls
|
|
self._lease_owner = f"{socket.gethostname()}:{uuid.uuid4().hex}"
|
|
self._task: asyncio.Task[None] | None = None
|
|
self._stop = asyncio.Event()
|
|
|
|
@property
|
|
def drivers(self) -> McpTaskDriverRegistry:
|
|
return self._drivers
|
|
|
|
async def submit(
|
|
self,
|
|
*,
|
|
driver_name: str,
|
|
request: TaskSubmitRequest,
|
|
now: datetime | None = None,
|
|
) -> dict:
|
|
"""Submit through one driver and persist the remote handle before returning."""
|
|
driver = self._drivers.get(driver_name)
|
|
if driver is None:
|
|
raise LookupError(f"No MCP task driver registered as {driver_name!r}")
|
|
|
|
submitted_at = now or datetime.now(UTC)
|
|
local_task_id = request.local_task_id or f"mcp-task-{uuid.uuid4().hex}"
|
|
driver_request = replace(request, local_task_id=local_task_id)
|
|
submission = await driver.submit(driver_request)
|
|
snapshot = submission.snapshot
|
|
next_poll_at = self._next_poll_at(snapshot, now=submitted_at)
|
|
driver_data = {**request.driver_data, **submission.driver_data}
|
|
task_reference = TaskReference(
|
|
local_task_id=local_task_id,
|
|
user_id=request.user_id,
|
|
thread_id=request.thread_id,
|
|
server_name=request.server_name,
|
|
remote_task_id=submission.remote_task_id,
|
|
driver_data=driver_data,
|
|
)
|
|
try:
|
|
return await self._repository.create(
|
|
task_id=local_task_id,
|
|
user_id=request.user_id,
|
|
thread_id=request.thread_id,
|
|
run_id=request.run_id,
|
|
tool_call_id=request.tool_call_id,
|
|
server_name=request.server_name,
|
|
driver_name=driver_name,
|
|
remote_task_id=submission.remote_task_id,
|
|
task_name=request.task_name,
|
|
status=snapshot.status.value,
|
|
result=snapshot.result,
|
|
error=snapshot.error,
|
|
input_required=snapshot.input_required,
|
|
next_poll_at=next_poll_at,
|
|
driver_data=driver_data,
|
|
)
|
|
except DuplicateMcpRemoteTaskError:
|
|
# This handle already has a durable owner. Cancelling it as
|
|
# compensation would terminate the pre-existing tracked task.
|
|
raise
|
|
except Exception:
|
|
try:
|
|
await driver.cancel(task_reference)
|
|
except Exception: # noqa: BLE001 - preserve the original persistence failure
|
|
logger.exception(
|
|
"Failed to cancel untracked MCP task after persistence failure (task_id=%s, driver=%s, remote_task_id=%s)",
|
|
local_task_id,
|
|
driver_name,
|
|
submission.remote_task_id,
|
|
)
|
|
raise
|
|
|
|
async def run_once(self, *, now: datetime) -> None:
|
|
claimed = await self._repository.claim_due_tasks(
|
|
now=now,
|
|
lease_owner=self._lease_owner,
|
|
lease_seconds=self._lease_seconds,
|
|
limit=self._max_concurrent_polls,
|
|
)
|
|
if not claimed:
|
|
return
|
|
results = await asyncio.gather(
|
|
*(self._poll_one(task, now=now) for task in claimed),
|
|
return_exceptions=True,
|
|
)
|
|
for record, result in zip(claimed, results, strict=True):
|
|
if isinstance(result, BaseException):
|
|
logger.error(
|
|
"Unexpected MCP task poll failure (task_id=%s); the lease will expire for recovery",
|
|
record.get("id"),
|
|
exc_info=(type(result), result, result.__traceback__),
|
|
)
|
|
|
|
async def _poll_one(self, record: dict, *, now: datetime) -> None:
|
|
driver_name = str(record.get("driver_name") or "")
|
|
driver = self._drivers.get(driver_name)
|
|
if driver is None:
|
|
await self._release_after_error(
|
|
record,
|
|
now=now,
|
|
error=f"No MCP task driver registered as {driver_name!r}",
|
|
)
|
|
return
|
|
|
|
try:
|
|
snapshot = await driver.get_status(TaskReference.from_record(record))
|
|
except Exception as exc: # noqa: BLE001 - driver boundary; retry on the next poll
|
|
polled_at = datetime.now(UTC)
|
|
logger.warning(
|
|
"MCP task status poll failed (task_id=%s, driver=%s); retrying",
|
|
record.get("id"),
|
|
driver_name,
|
|
exc_info=True,
|
|
)
|
|
await self._release_after_error(record, now=polled_at, error=str(exc) or type(exc).__name__)
|
|
return
|
|
|
|
polled_at = datetime.now(UTC)
|
|
applied = await self._repository.apply_snapshot(
|
|
record["id"],
|
|
lease_owner=self._lease_owner,
|
|
status=snapshot.status.value,
|
|
result=snapshot.result,
|
|
error=snapshot.error,
|
|
input_required=snapshot.input_required,
|
|
next_poll_at=self._next_poll_at(snapshot, now=polled_at),
|
|
polled_at=polled_at,
|
|
)
|
|
if not applied:
|
|
logger.info(
|
|
"Discarded MCP task poll result after lease ownership changed or expired (task_id=%s)",
|
|
record.get("id"),
|
|
)
|
|
|
|
def _next_poll_at(self, snapshot: TaskSnapshot, *, now: datetime) -> datetime | None:
|
|
if not snapshot.is_pollable:
|
|
return None
|
|
interval = snapshot.poll_after_seconds or self._poll_interval_seconds
|
|
return now + timedelta(seconds=interval)
|
|
|
|
async def _release_after_error(self, record: dict, *, now: datetime, error: str) -> None:
|
|
await self._repository.release_claim(
|
|
record["id"],
|
|
lease_owner=self._lease_owner,
|
|
next_poll_at=now + timedelta(seconds=self._poll_interval_seconds),
|
|
error=error[:_MAX_POLL_ERROR_CHARS],
|
|
)
|
|
|
|
async def start(self) -> None:
|
|
if self._task is not None:
|
|
return
|
|
self._stop.clear()
|
|
self._task = asyncio.create_task(self._run_loop(), name="deerflow-mcp-task-poller")
|
|
|
|
async def stop(self) -> None:
|
|
task = self._task
|
|
if task is None:
|
|
return
|
|
self._stop.set()
|
|
task.cancel()
|
|
try:
|
|
await task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
finally:
|
|
self._task = None
|
|
|
|
async def _run_loop(self) -> None:
|
|
while not self._stop.is_set():
|
|
try:
|
|
# The first pass runs immediately. Expired leases therefore
|
|
# recover at startup without a separate destructive sweep.
|
|
await self.run_once(now=datetime.now(UTC))
|
|
except Exception:
|
|
logger.exception("MCP task poll failed; retrying next interval")
|
|
try:
|
|
await asyncio.wait_for(
|
|
self._stop.wait(),
|
|
timeout=self._poll_interval_seconds,
|
|
)
|
|
except TimeoutError:
|
|
continue
|