refactor
This commit is contained in:
+182
-175
@@ -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),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user