This commit is contained in:
tao.chen
2026-07-30 20:52:46 +08:00
parent b6a029fe7e
commit 91767461d6
12 changed files with 380 additions and 236 deletions
+182 -175
View File
@@ -5,18 +5,17 @@ import hashlib
import json
import logging
import os
import socket
import traceback
from collections.abc import Awaitable, Callable
from contextlib import suppress
from datetime import timedelta
from datetime import UTC, datetime, timedelta
from pathlib import Path
from typing import Any
from zoneinfo import ZoneInfo
import boto3
import httpx
from redis.asyncio import Redis
from redis.exceptions import ResponseError
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from apscheduler.triggers.cron import CronTrigger
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
@@ -30,16 +29,7 @@ from common.db.models import (
Versions,
Workspaces,
)
from common.eventing import (
STREAM_BY_EVENT_TYPE,
add_outbox_event,
event_time,
utcnow,
)
from common.ids import new_ulid
from common.db.session import session_scope
from schedule.execution import ExecutionResult, execute_artifact
from schedule.storage_client import SchedulerStorageClient
from common.scheduler import build_sqlalchemy_jobstore
LOGGER = logging.getLogger(__name__)
@@ -58,203 +48,229 @@ TERMINAL_RUN_STATES = {
"timed_out",
}
_ACTIVE_SERVICE: "SchedulerService | None" = None
def _naive_utc(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
value = value.replace(tzinfo=UTC)
return value.astimezone(UTC).replace(tzinfo=None)
async def run_scheduled_job(schedule_id: str) -> None:
service = _ACTIVE_SERVICE
if service is None:
LOGGER.warning("scheduler job skipped because service is not ready")
return
await service.trigger_schedule(schedule_id)
class SchedulerService:
def __init__(
self,
*,
session_factory: async_sessionmaker[AsyncSession],
redis: Redis,
object_store: Any,
storage_client: SchedulerStorageClient,
backend_http_client: httpx.AsyncClient,
workspace_root: Path,
database_url: str,
) -> None:
self.session_factory = session_factory
self.redis = redis
self.object_store = object_store
self.storage_client = storage_client
self.backend_http_client = backend_http_client
self.workspace_root = workspace_root
self.consumer_name = (
os.getenv("SCHEDULER_CONSUMER_NAME")
or f"{socket.gethostname()}-{os.getpid()}"
)
self.tasks: list[asyncio.Task[Any]] = []
self.dispatch_lock = asyncio.Lock()
self.scheduler = AsyncIOScheduler(
jobstores={
"default": build_sqlalchemy_jobstore(database_url)
},
timezone=UTC,
)
async def start(self) -> None:
await self._ensure_group(
"stream:scheduler:commands",
"schedule-orchestrator",
)
await self._ensure_group("stream:jobs:execute", "job-workers")
await self._ensure_group("stream:jobs:results", "schedule-results")
global _ACTIVE_SERVICE
_ACTIVE_SERVICE = self
self.scheduler.start()
await self._sync_cron_jobs()
self.tasks = [
asyncio.create_task(
self._outbox_loop(),
name="scheduler-outbox-publisher",
self._database_event_loop(),
name="scheduler-database-events",
),
asyncio.create_task(
self._consumer_loop(
"stream:scheduler:commands",
"schedule-orchestrator",
self._handle_run_requested,
),
name="schedule-orchestrator",
),
asyncio.create_task(
self._consumer_loop(
"stream:jobs:execute",
"job-workers",
self._handle_node_execute,
),
name="job-worker",
),
asyncio.create_task(
self._consumer_loop(
"stream:jobs:results",
"schedule-results",
self._handle_node_finished,
),
name="schedule-results",
self._schedule_sync_loop(),
name="scheduler-cron-sync",
),
]
async def close(self) -> None:
global _ACTIVE_SERVICE
for task in self.tasks:
task.cancel()
for task in self.tasks:
with suppress(asyncio.CancelledError):
await task
self.tasks.clear()
if self.scheduler.running:
self.scheduler.shutdown(wait=False)
_ACTIVE_SERVICE = None
async def _ensure_group(self, stream: str, group: str) -> None:
try:
await self.redis.xgroup_create(
stream,
group,
id="0-0",
mkstream=True,
)
except ResponseError as exc:
if "BUSYGROUP" not in str(exc):
raise
async def _outbox_loop(self) -> None:
async def _database_event_loop(self) -> None:
while True:
try:
published = await self._publish_outbox_batch()
if not published:
await asyncio.sleep(0.35)
processed = await self.process_pending_events(limit=20)
if not processed:
await asyncio.sleep(0.25)
except asyncio.CancelledError:
raise
except Exception:
LOGGER.exception("outbox publisher iteration failed")
LOGGER.exception("database event loop failed")
await asyncio.sleep(1)
async def _publish_outbox_batch(self) -> int:
now = utcnow()
async def process_pending_events(
self,
*,
limit: int = 20,
aggregate_id: str | None = None,
) -> int:
async with self.dispatch_lock:
async with session_scope(self.session_factory) as session:
statement = (
select(OutboxEvents)
.where(
OutboxEvents.event_status == "pending",
OutboxEvents.available_at <= utcnow(),
)
.order_by(OutboxEvents.created_at)
.limit(limit)
)
if aggregate_id:
statement = statement.where(
OutboxEvents.aggregate_id == aggregate_id
)
events = list((await session.scalars(statement)).all())
for item in events:
try:
await self._process_outbox_event(item)
item.event_status = "published"
item.published_at = utcnow()
item.last_error = None
except Exception as exc:
item.retry_count += 1
item.last_error = str(exc)[:2000]
if item.retry_count >= 5:
item.event_status = "failed"
else:
item.available_at = utcnow() + timedelta(
seconds=min(30, 2 ** item.retry_count)
)
LOGGER.exception(
"failed to process database event %s",
item.event_id,
)
return len(events)
async def _process_outbox_event(self, item: OutboxEvents) -> None:
handlers = {
"schedule.run.requested": self._handle_run_requested,
"job.node.execute": self._handle_node_execute,
"job.node.finished": self._handle_node_finished,
}
handler = handlers.get(item.event_type)
if handler is None:
raise ValueError(f"unsupported event type: {item.event_type}")
await handler(item.payload_json, f"mysql:{item.event_id}")
async def dispatch_run(self, run_id: str) -> int:
return await self.process_pending_events(
limit=50,
aggregate_id=run_id,
)
async def _schedule_sync_loop(self) -> None:
while True:
try:
await self._sync_cron_jobs()
except asyncio.CancelledError:
raise
except Exception:
LOGGER.exception("cron job synchronization failed")
await asyncio.sleep(5)
async def _sync_cron_jobs(self) -> None:
async with session_scope(self.session_factory) as session:
events = list(
schedules = list(
(
await session.scalars(
select(OutboxEvents)
.where(
OutboxEvents.event_status == "pending",
OutboxEvents.available_at <= now,
select(Schedules).where(
Schedules.deleted_at.is_(None),
Schedules.enabled == 1,
Schedules.trigger_type == "cron",
Schedules.cron_expression.is_not(None),
)
.order_by(OutboxEvents.created_at)
.limit(20)
.with_for_update(skip_locked=True)
)
).all()
)
for item in events:
stream = STREAM_BY_EVENT_TYPE.get(item.event_type)
if stream is None:
item.event_status = "failed"
item.last_error = f"unsupported event type: {item.event_type}"
continue
try:
await self.redis.xadd(
stream,
{
"event": json.dumps(
item.payload_json,
ensure_ascii=False,
separators=(",", ":"),
)
},
)
item.event_status = "published"
item.published_at = utcnow()
item.last_error = None
except Exception as exc:
item.retry_count += 1
item.last_error = str(exc)[:2000]
raise
return len(events)
async def _consumer_loop(
self,
stream: str,
group: str,
handler: Callable[[dict[str, Any], str], Awaitable[None]],
) -> None:
while True:
try:
messages = await self.redis.xreadgroup(
group,
self.consumer_name,
{stream: ">"},
count=5,
block=1000,
active_job_ids: set[str] = set()
for item in schedules:
job_id = f"schedule:{item.schedule_id}"
active_job_ids.add(job_id)
expression = (item.cron_expression or "").strip()
trigger = CronTrigger.from_crontab(
expression,
timezone=ZoneInfo(item.timezone),
)
entries: list[tuple[str, dict[str, str]]] = []
for _, stream_messages in messages:
entries.extend(stream_messages)
if not entries:
claimed = await self.redis.xautoclaim(
stream,
group,
self.consumer_name,
min_idle_time=10_000,
start_id="0-0",
count=5,
)
if len(claimed) >= 2:
entries.extend(claimed[1])
for message_id, fields in entries:
try:
raw = fields.get("event")
if not raw:
raise ValueError("stream message has no event field")
event = json.loads(raw)
except asyncio.CancelledError:
raise
except (ValueError, TypeError, json.JSONDecodeError):
LOGGER.exception(
"discarding malformed message %s from %s",
message_id,
stream,
)
await self.redis.xack(stream, group, message_id)
continue
try:
await handler(event, message_id)
except asyncio.CancelledError:
raise
except Exception:
LOGGER.exception(
"consumer %s failed for message %s",
group,
message_id,
)
continue
await self.redis.xack(stream, group, message_id)
except asyncio.CancelledError:
raise
except Exception:
LOGGER.exception("consumer loop %s failed", group)
await asyncio.sleep(1)
job = self.scheduler.add_job(
run_scheduled_job,
trigger=trigger,
args=[item.schedule_id],
id=job_id,
replace_existing=True,
coalesce=True,
max_instances=max(1, item.max_concurrency),
misfire_grace_time=60,
)
item.next_run_at = _naive_utc(job.next_run_time)
for job in self.scheduler.get_jobs():
if job.id.startswith("schedule:") and job.id not in active_job_ids:
self.scheduler.remove_job(job.id)
async def trigger_schedule(self, schedule_id: str) -> None:
async with self.session_factory() as session:
item = await session.get(Schedules, schedule_id)
if (
item is None
or item.deleted_at is not None
or not item.enabled
or item.trigger_type != "cron"
):
return
user_id = item.created_by
workspace_id = item.workspace_id
now = datetime.now(UTC)
idempotency_key = (
f"cron:{schedule_id}:{now.strftime('%Y%m%d%H%M')}"
)
response = await self.backend_http_client.post(
f"/api/v1/schedules/{schedule_id}/run",
headers={
"X-User-ID": user_id,
"X-Workspace-ID": workspace_id,
"X-Request-ID": new_ulid(),
"Idempotency-Key": idempotency_key,
},
json={"reason": "cron"},
)
if response.is_error:
raise RuntimeError(
f"backend rejected cron run: {response.status_code} "
f"{response.text[:500]}"
)
async def _start_inbox(
self,
@@ -874,15 +890,6 @@ class SchedulerService:
self._finish_inbox(inbox)
def build_redis_client() -> Redis:
return Redis(
host=os.getenv("REDIS_HOST", "redis"),
port=int(os.getenv("REDIS_PORT", "6379")),
password=os.getenv("REDIS_PASSWORD") or None,
decode_responses=True,
)
def build_object_store() -> Any:
return boto3.client(
"s3",
@@ -898,6 +905,6 @@ def build_object_store() -> Any:
def build_storage_http_client() -> httpx.AsyncClient:
return httpx.AsyncClient(
base_url=os.getenv("STORAGE_API_URL", "http://storage_api:8000"),
base_url=os.getenv("BACKEND_API_URL", "http://backend:8000"),
timeout=httpx.Timeout(60.0),
)