268 lines
9.2 KiB
Python
268 lines
9.2 KiB
Python
"""
|
|
PG agent_job 认领壳(与 Redis Stream 二选一真相时用)。
|
|
|
|
本文件职责:enqueue / claim / complete / requeue_expired;不含业务处理体。
|
|
禁止:一把大锁串行所有工单;禁止 Redis wecom:worker:singleton。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import threading
|
|
import uuid
|
|
from dataclasses import dataclass, field
|
|
from datetime import datetime, timezone
|
|
from typing import Any, Optional
|
|
|
|
from agent.runtime_db import assert_safe_database, connect_psycopg
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _utcnow() -> datetime:
|
|
return datetime.now(timezone.utc)
|
|
|
|
|
|
@dataclass
|
|
class JobRecord:
|
|
"""一条任务行。"""
|
|
|
|
id: str
|
|
job_type: str
|
|
payload: dict[str, Any] = field(default_factory=dict)
|
|
state: str = "pending"
|
|
idempotency_key: str = ""
|
|
attempts: int = 0
|
|
last_error: str = ""
|
|
|
|
|
|
class MemoryJobStore:
|
|
"""内存任务认领壳。"""
|
|
|
|
def __init__(self, *, claim_timeout_sec: float = 180.0) -> None:
|
|
self._rows: dict[str, JobRecord] = {}
|
|
self._claimed_at: dict[str, float] = {}
|
|
self._lock = threading.Lock()
|
|
self._claim_timeout = claim_timeout_sec
|
|
|
|
def enqueue(
|
|
self,
|
|
*,
|
|
job_type: str,
|
|
payload: dict[str, Any],
|
|
idempotency_key: str = "",
|
|
) -> tuple[bool, str]:
|
|
with self._lock:
|
|
if idempotency_key:
|
|
for j in self._rows.values():
|
|
if j.job_type == job_type and j.idempotency_key == idempotency_key:
|
|
return False, "duplicate_idempotency_key"
|
|
job_id = str(uuid.uuid4())
|
|
self._rows[job_id] = JobRecord(
|
|
id=job_id,
|
|
job_type=job_type,
|
|
payload=dict(payload),
|
|
idempotency_key=idempotency_key,
|
|
)
|
|
return True, job_id
|
|
|
|
def claim(self, *, worker_id: str, job_type: str = "") -> Optional[JobRecord]:
|
|
import time
|
|
|
|
now = time.time()
|
|
with self._lock:
|
|
for jid, job in self._rows.items():
|
|
if job_type and job.job_type != job_type:
|
|
continue
|
|
if job.state == "pending" or (
|
|
job.state == "claimed"
|
|
and now - self._claimed_at.get(jid, 0) > self._claim_timeout
|
|
):
|
|
job.state = "claimed"
|
|
job.attempts += 1
|
|
self._claimed_at[jid] = now
|
|
return JobRecord(
|
|
id=job.id,
|
|
job_type=job.job_type,
|
|
payload=dict(job.payload),
|
|
state=job.state,
|
|
idempotency_key=job.idempotency_key,
|
|
attempts=job.attempts,
|
|
last_error=job.last_error,
|
|
)
|
|
return None
|
|
|
|
def complete(self, job_id: str, *, error: str = "", result: Optional[dict] = None) -> None:
|
|
with self._lock:
|
|
job = self._rows.get(job_id)
|
|
if not job:
|
|
return
|
|
job.state = "failed" if error else "done"
|
|
job.last_error = error
|
|
logger.info("job.memory.complete id=%s state=%s", job_id, job.state)
|
|
|
|
|
|
class PostgresJobStore:
|
|
"""PG agent_job 认领壳。"""
|
|
|
|
def __init__(self, dsn: str, *, claim_timeout_sec: float = 180.0) -> None:
|
|
assert_safe_database(dsn)
|
|
self._dsn = dsn
|
|
self._claim_timeout = int(claim_timeout_sec)
|
|
|
|
def enqueue(
|
|
self,
|
|
*,
|
|
job_type: str,
|
|
payload: dict[str, Any],
|
|
idempotency_key: str = "",
|
|
) -> tuple[bool, str]:
|
|
job_id = str(uuid.uuid4())
|
|
try:
|
|
with connect_psycopg(self._dsn) as conn:
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO agent_job (id, job_type, idempotency_key, payload, state)
|
|
VALUES (%s, %s, %s, %s::jsonb, 'pending')
|
|
""",
|
|
(
|
|
job_id,
|
|
job_type,
|
|
idempotency_key,
|
|
json.dumps(payload, ensure_ascii=False),
|
|
),
|
|
)
|
|
conn.commit()
|
|
return True, job_id
|
|
except Exception as exc: # noqa: BLE001
|
|
if "unique" in str(exc).lower() or "duplicate" in str(exc).lower():
|
|
return False, "duplicate_idempotency_key"
|
|
raise
|
|
|
|
def claim(self, *, worker_id: str, job_type: str = "") -> Optional[JobRecord]:
|
|
type_filter = "AND job_type = %s" if job_type else ""
|
|
params: list[Any] = [self._claim_timeout, worker_id]
|
|
if job_type:
|
|
params.append(job_type)
|
|
sql = f"""
|
|
WITH cte AS (
|
|
SELECT id FROM agent_job
|
|
WHERE (
|
|
state = 'pending'
|
|
OR (state = 'claimed' AND claimed_at < NOW() - (%s || ' seconds')::interval)
|
|
)
|
|
{type_filter}
|
|
ORDER BY created_at
|
|
FOR UPDATE SKIP LOCKED
|
|
LIMIT 1
|
|
)
|
|
UPDATE agent_job j SET
|
|
state = 'claimed',
|
|
claimed_at = NOW(),
|
|
claimed_by = %s,
|
|
attempts = j.attempts + 1,
|
|
updated_at = NOW()
|
|
FROM cte WHERE j.id = cte.id
|
|
RETURNING j.id, j.job_type, j.payload, j.state, j.idempotency_key, j.attempts, j.last_error
|
|
"""
|
|
# 参数顺序:timeout, [job_type], worker_id — 上面拼法易错,改用固定两段
|
|
with connect_psycopg(self._dsn) as conn:
|
|
if job_type:
|
|
row = conn.execute(
|
|
"""
|
|
WITH cte AS (
|
|
SELECT id FROM agent_job
|
|
WHERE job_type = %s AND (
|
|
state = 'pending'
|
|
OR (state = 'claimed'
|
|
AND claimed_at < NOW() - (%s || ' seconds')::interval)
|
|
)
|
|
ORDER BY created_at
|
|
FOR UPDATE SKIP LOCKED
|
|
LIMIT 1
|
|
)
|
|
UPDATE agent_job j SET
|
|
state = 'claimed',
|
|
claimed_at = NOW(),
|
|
claimed_by = %s,
|
|
attempts = j.attempts + 1,
|
|
updated_at = NOW()
|
|
FROM cte WHERE j.id = cte.id
|
|
RETURNING j.id, j.job_type, j.payload, j.state,
|
|
j.idempotency_key, j.attempts, j.last_error
|
|
""",
|
|
(job_type, str(self._claim_timeout), worker_id),
|
|
).fetchone()
|
|
else:
|
|
row = conn.execute(
|
|
"""
|
|
WITH cte AS (
|
|
SELECT id FROM agent_job
|
|
WHERE state = 'pending'
|
|
OR (state = 'claimed'
|
|
AND claimed_at < NOW() - (%s || ' seconds')::interval)
|
|
ORDER BY created_at
|
|
FOR UPDATE SKIP LOCKED
|
|
LIMIT 1
|
|
)
|
|
UPDATE agent_job j SET
|
|
state = 'claimed',
|
|
claimed_at = NOW(),
|
|
claimed_by = %s,
|
|
attempts = j.attempts + 1,
|
|
updated_at = NOW()
|
|
FROM cte WHERE j.id = cte.id
|
|
RETURNING j.id, j.job_type, j.payload, j.state,
|
|
j.idempotency_key, j.attempts, j.last_error
|
|
""",
|
|
(str(self._claim_timeout), worker_id),
|
|
).fetchone()
|
|
conn.commit()
|
|
if not row:
|
|
return None
|
|
return JobRecord(
|
|
id=str(row[0]),
|
|
job_type=str(row[1]),
|
|
payload=dict(row[2] or {}),
|
|
state=str(row[3]),
|
|
idempotency_key=str(row[4] or ""),
|
|
attempts=int(row[5] or 0),
|
|
last_error=str(row[6] or ""),
|
|
)
|
|
|
|
def complete(
|
|
self,
|
|
job_id: str,
|
|
*,
|
|
error: str = "",
|
|
result: Optional[dict[str, Any]] = None,
|
|
) -> None:
|
|
state = "failed" if error else "done"
|
|
with connect_psycopg(self._dsn) as conn:
|
|
conn.execute(
|
|
"""
|
|
UPDATE agent_job SET
|
|
state = %s,
|
|
last_error = %s,
|
|
result = %s::jsonb,
|
|
updated_at = NOW()
|
|
WHERE id = %s
|
|
""",
|
|
(
|
|
state,
|
|
error,
|
|
json.dumps(result or {}, ensure_ascii=False),
|
|
job_id,
|
|
),
|
|
)
|
|
conn.commit()
|
|
logger.info("job.pg.complete id=%s state=%s", job_id, state)
|
|
|
|
|
|
def create_job_store(*, dsn: str = "", backend: str = "auto", claim_timeout_sec: float = 180.0):
|
|
"""工厂:Memory 或 PG。"""
|
|
if backend == "memory" or not dsn:
|
|
return MemoryJobStore(claim_timeout_sec=claim_timeout_sec)
|
|
return PostgresJobStore(dsn, claim_timeout_sec=claim_timeout_sec)
|