mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-11 14:38:38 +00:00
* fix(frontend): clarify reuse-thread scheduling behavior * fix(scheduler): enqueue overlapping scheduled runs * fix(scheduler): preserve queue lease fencing * fix(scheduler): close queue concurrency races * fix(scheduler): harden queue timeout bookkeeping * fix(scheduler): preserve manual failure schedule --------- Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
159 lines
5.5 KiB
Python
159 lines
5.5 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import os
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
|
|
|
|
import pytest
|
|
import pytest_asyncio
|
|
from sqlalchemy import text
|
|
|
|
from deerflow.config.database_config import DatabaseConfig
|
|
from deerflow.persistence.engine import close_engine, get_engine, get_session_factory, init_engine_from_config
|
|
from deerflow.persistence.run import RunRepository
|
|
from deerflow.persistence.scheduled_task_runs import ScheduledTaskRunRepository
|
|
from deerflow.persistence.scheduled_tasks import ScheduledTaskRepository
|
|
|
|
POSTGRES_URL = os.environ.get("TEST_POSTGRES_URI")
|
|
|
|
pytestmark = pytest.mark.skipif(
|
|
not POSTGRES_URL,
|
|
reason="requires TEST_POSTGRES_URI (real Postgres for scheduler concurrency)",
|
|
)
|
|
|
|
|
|
def _postgres_url(url: str) -> str:
|
|
parts = urlsplit(url)
|
|
query = urlencode([(key, value) for key, value in parse_qsl(parts.query, keep_blank_values=True) if key not in {"sslmode", "channel_binding"}])
|
|
return urlunsplit(parts._replace(query=query))
|
|
|
|
|
|
@pytest_asyncio.fixture()
|
|
async def postgres_repositories():
|
|
assert POSTGRES_URL is not None
|
|
schema = f"scheduler_{uuid.uuid4().hex}"
|
|
await init_engine_from_config(
|
|
DatabaseConfig(
|
|
backend="postgres",
|
|
postgres_url=_postgres_url(POSTGRES_URL),
|
|
postgres_schema=schema,
|
|
)
|
|
)
|
|
sf = get_session_factory()
|
|
assert sf is not None
|
|
run_repo = RunRepository(sf)
|
|
task_repo = ScheduledTaskRepository(sf, run_repository=run_repo)
|
|
task_run_repo = ScheduledTaskRunRepository(sf, run_repository=run_repo)
|
|
try:
|
|
yield task_repo, task_run_repo, run_repo
|
|
finally:
|
|
engine = get_engine()
|
|
async with engine.begin() as connection:
|
|
await connection.execute(text(f'DROP SCHEMA IF EXISTS "{schema}" CASCADE'))
|
|
await close_engine()
|
|
|
|
|
|
async def _create_cron_task(task_repo: ScheduledTaskRepository, task_id: str, *, next_run_at: datetime | None) -> None:
|
|
await task_repo.create(
|
|
task_id=task_id,
|
|
user_id="user-1",
|
|
thread_id=None,
|
|
context_mode="fresh_thread_per_run",
|
|
assistant_id="lead_agent",
|
|
title=task_id,
|
|
prompt="p",
|
|
schedule_type="cron",
|
|
schedule_spec={"cron": "* * * * *"},
|
|
timezone="UTC",
|
|
next_run_at=next_run_at,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_postgres_global_budget_serializes_cross_pod_claims(postgres_repositories):
|
|
task_repo, task_run_repo, _run_repo = postgres_repositories
|
|
now = datetime.now(UTC)
|
|
for suffix in ("a", "b"):
|
|
await _create_cron_task(task_repo, f"task-{suffix}", next_run_at=now)
|
|
await task_run_repo.create(
|
|
run_record_id=f"task-run-{suffix}",
|
|
task_id=f"task-{suffix}",
|
|
thread_id=f"thread-{suffix}",
|
|
scheduled_for=now,
|
|
trigger="scheduled",
|
|
status="queued",
|
|
)
|
|
|
|
claims = await asyncio.gather(
|
|
task_run_repo.claim_queued_run(
|
|
"task-run-a",
|
|
now=now,
|
|
lease_owner="pod-a",
|
|
lease_seconds=60,
|
|
global_max_concurrent_runs=1,
|
|
),
|
|
task_run_repo.claim_queued_run(
|
|
"task-run-b",
|
|
now=now,
|
|
lease_owner="pod-b",
|
|
lease_seconds=60,
|
|
global_max_concurrent_runs=1,
|
|
),
|
|
)
|
|
|
|
assert sum(claimed is not None for claimed in claims) == 1
|
|
statuses = {(await task_run_repo.list_by_task(task_id))[0]["status"] for task_id in ("task-a", "task-b")}
|
|
assert statuses == {"queued", "launching"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_postgres_reconciliation_uses_metadata_and_atomically_claims_expired_run(postgres_repositories):
|
|
task_repo, task_run_repo, run_repo = postgres_repositories
|
|
now = datetime.now(UTC)
|
|
for suffix in ("live", "expired"):
|
|
await _create_cron_task(task_repo, f"task-{suffix}", next_run_at=None)
|
|
await task_run_repo.create(
|
|
run_record_id=f"task-run-{suffix}",
|
|
task_id=f"task-{suffix}",
|
|
thread_id=f"thread-{suffix}",
|
|
scheduled_for=now,
|
|
trigger="scheduled",
|
|
status="running" if suffix == "expired" else "queued",
|
|
)
|
|
|
|
await task_run_repo.update_status("task-run-expired", status="running", run_id="run-expired")
|
|
await run_repo.put(
|
|
"run-live",
|
|
thread_id="thread-live",
|
|
user_id="user-1",
|
|
status="running",
|
|
metadata={
|
|
"scheduled_task_id": "task-live",
|
|
"scheduled_task_run_id": "task-run-live",
|
|
},
|
|
owner_worker_id="pod-a",
|
|
lease_expires_at=(now + timedelta(seconds=60)).isoformat(),
|
|
)
|
|
await run_repo.put(
|
|
"run-expired",
|
|
thread_id="thread-expired",
|
|
user_id="user-1",
|
|
status="running",
|
|
metadata={
|
|
"scheduled_task_id": "task-expired",
|
|
"scheduled_task_run_id": "task-run-expired",
|
|
},
|
|
owner_worker_id="pod-dead",
|
|
lease_expires_at=(now - timedelta(seconds=60)).isoformat(),
|
|
)
|
|
|
|
assert await task_run_repo.reconcile_active_runs(error="lease expired", now=now) == 1
|
|
assert (await task_run_repo.list_by_task("task-live"))[0]["status"] == "queued"
|
|
assert (await task_run_repo.list_by_task("task-expired"))[0]["status"] == "interrupted"
|
|
recovered = await run_repo.get("run-expired", user_id=None)
|
|
assert recovered is not None
|
|
assert recovered["status"] == "error"
|
|
assert recovered["stop_reason"] == "scheduled_task_orphan_recovered"
|