Files
creator/backend/pipeline.py
2026-07-02 03:05:57 +02:00

325 lines
12 KiB
Python

"""Pipeline building blocks: agent races (with grace), single-slot, schemas, prompts, guide status.
Holds the mutable pipeline state (generation semaphore, cancel set).
Access the cancel set ONLY through the functions here — copied references
in other modules would diverge on a re-assign.
"""
import asyncio
import logging
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable
from agents import run_agent, kill_process, cancel_scope, clear_scope
from config import MAX_CONCURRENT_GENERATIONS, TEMPLATES_DIR, TIMEOUTS
from database import update_guide
from jsonio import read_json_file as _json_file
from textkit import _STUFEN
log = logging.getLogger("creator.pipeline")
_semaphore = asyncio.Semaphore(MAX_CONCURRENT_GENERATIONS)
_cancelled: set[str] = set()
async def cancel_guide(guide_id: str) -> bool:
_cancelled.add(guide_id)
cancel_scope(f"{guide_id}-") # waiting agents bail before spawn
kill_process(guide_id) # kill running subprocesses
now = datetime.now(timezone.utc).isoformat()
await update_guide(guide_id, status="error", progress=None, error_msg="Cancelled — progress is preserved", updated_at=now)
return True
def is_guide_cancelled(guide_id: str) -> bool:
return guide_id in _cancelled
def clear_guide_cancelled(guide_id: str) -> None:
_cancelled.discard(guide_id)
clear_scope(f"{guide_id}-") # clear scope → restart not blocked
async def _set_progress(guide_id: str, progress: str) -> None:
now = datetime.now(timezone.utc).isoformat()
await update_guide(guide_id, progress=progress, updated_at=now)
async def _set_step(guide_id: str, step: int, progress: str) -> None:
now = datetime.now(timezone.utc).isoformat()
await update_guide(guide_id, step=step, progress=progress, updated_at=now)
async def _fail(guide_id: str, msg: str) -> None:
now = datetime.now(timezone.utc).isoformat()
await update_guide(guide_id, status="error", progress=None, error_msg=msg, updated_at=now)
def _prompt(name: str, **kwargs) -> str:
template = (TEMPLATES_DIR / "Prompt" / f"{name}.md").read_text(encoding="utf-8")
return template.format(**kwargs)
def _extra(instructions: str) -> str:
return f"\n\nADDITIONAL INSTRUCTIONS FROM THE USER:\n{instructions}\n" if instructions else ""
def _log(topic: str, msg: str) -> None:
log.info("[%s] %s", topic, msg)
def _claude_error(label: str, returncode: int, stdout: str, stderr: str) -> str:
stderr = (stderr or "").strip()
if stderr:
return f"{label}: {stderr[:1000]}"
tail = (stdout or "").strip()[-500:]
if tail:
return f"{label} (exit {returncode}, stderr empty): …{tail}"
return f"{label} (exit {returncode}, no output)"
def _gather_error(label: str, results: list) -> str:
for r in results:
if isinstance(r, BaseException):
return f"{label}: {type(r).__name__}: {r}"
returncode, stdout, stderr = r
if returncode != 0:
return _claude_error(label, returncode, stdout, stderr)
return f"{label}: no usable result"
def _timeout(step: str, n: int = 0) -> int:
base, per = TIMEOUTS[step]
return base + per * n
def _problems_schema(data):
"""{"ok": true} → [] · {"problems": [str]} → list · else None."""
if not isinstance(data, dict):
return None
if data.get("ok") is True:
return []
p = data.get("problems")
if not isinstance(p, list) or not p:
return None
out = [str(x).strip() for x in p if str(x).strip()]
return out or None
def _str_list(val) -> list[str] | None:
"""List of non-empty strings → stripped list (empty allowed) · else None."""
if not isinstance(val, list) or not all(isinstance(x, str) for x in val):
return None
out = [x.strip() for x in val]
return None if any(not x for x in out) else out
def _runde_schema(data, final: bool = False):
"""{"keep": [str], "rest": [str]} → (include, rest) · else None.
final=True: last clarification round — a non-empty rest is invalid.
"""
if not isinstance(data, dict):
return None
include = _str_list(data.get("keep"))
rest = _str_list(data.get("rest"))
if include is None or rest is None or (final and rest):
return None
return include, rest
_RELEVANCE = ("relevant", "peripheral")
_YESNO = ("ja", "nein")
def _enum_map_schema(key: str, allowed):
"""Factory for `{"<key>": {"1": value, …}}` → `{id: value}` parsers; value ∈ `allowed`
(casefolded). If `ids` are given, at least these must be covered (extras allowed). None
on any invalid id/value or wrong shape. The caller filters the result to `ids`."""
def parse(data, ids: set[int] | None = None):
if not isinstance(data, dict) or not isinstance(data.get(key), dict) or not data[key]:
return None
out: dict[int, str] = {}
for k, v in data[key].items():
try:
num = int(k)
except (ValueError, TypeError):
return None
value = str(v).strip().casefold()
if value not in allowed:
return None
out[num] = value
if ids is not None and not ids <= set(out):
return None
return out
return parse
_levels_schema = _enum_map_schema("levels", _STUFEN) # level ∈ beginner/advanced/expert
_relevance_schema = _enum_map_schema("relevance", _RELEVANCE) # relevance ∈ relevant/peripheral
_yesno_schema = _enum_map_schema("relevant", _YESNO) # triage gate ∈ ja/nein
_MAX_RESTARTS = 2
async def _race(topic: str, label: str, slots: list[dict], quorum: int, timeout: int, provider: str, on_update=None, cancelled=None, *, grace: int | None = None, min_runtime: int | None = None, max_runtime: int | None = None) -> list | None:
"""Starts all slots in parallel and collects `quorum` valid results.
Slot spec: {key, prompt, role, capabilities, payload}. `payload(result)`
checks validity and returns the slot result or None.
Error/timeout/invalid → slot restart (max. _MAX_RESTARTS). As soon as the
quorum stands, the remaining agents are killed. None = quorum missed.
`cancelled()` → True aborts (no restarts, returns None).
With `grace`, `quorum` becomes the minimum: the first valid result starts
a timer of `grace` seconds. After it expires, running agents are only
killed if the minimum stands — otherwise the race, including restarts,
keeps running until it stands. Returns: `quorum` to `len(slots)` results.
`min_runtime` (wall-clock from start): the race does not return before it
elapses while agents are still running — gives them time to search thoroughly.
`max_runtime` (wall-clock from start): hard cap — returns whatever is collected
(or None if nothing), killing the rest. Both default off; only Research sets them.
"""
attempts = {i: 0 for i in range(len(slots))}
tasks: dict[asyncio.Task, int] = {}
loop = asyncio.get_running_loop()
start = loop.time()
min_deadline = start + min_runtime if min_runtime else None
max_deadline = start + max_runtime if max_runtime else None
deadline: float | None = None
def spawn(i: int) -> None:
slot = slots[i]
task = asyncio.create_task(run_agent(
slot["key"], slot["prompt"], timeout,
provider=provider, role=slot["role"], capabilities=slot["capabilities"],
scope=topic, on_line=slot.get("on_line"),
))
tasks[task] = i
for i in range(len(slots)):
spawn(i)
results: list = []
try:
while tasks:
if cancelled and cancelled():
return None
# Hard wall-clock cap: return whatever we have (None if empty), kill the rest.
if max_deadline is not None and loop.time() >= max_deadline:
_log(topic, f"{label}: max runtime {max_runtime}s reached ({len(results)} valid)")
return results or None
min_ok = min_deadline is None or loop.time() >= min_deadline
if deadline is not None and len(results) >= quorum and loop.time() >= deadline and min_ok:
return results
# Wake up for the earliest relevant deadline (grace, min, or max).
waits = []
if deadline is not None and len(results) >= quorum:
waits.append(deadline - loop.time())
if min_deadline is not None:
waits.append(min_deadline - loop.time())
if max_deadline is not None:
waits.append(max_deadline - loop.time())
wait_timeout = max(0.0, min(waits)) if waits else None
done, _ = await asyncio.wait(tasks.keys(), return_when=asyncio.FIRST_COMPLETED, timeout=wait_timeout)
if not done:
continue
for task in done:
i = tasks.pop(task)
payload, err = None, None
try:
result = task.result()
if result[0] != 0:
err = _claude_error("Error", *result)
else:
payload = slots[i]["payload"](result)
if payload is None:
err = "result invalid/not parseable"
except asyncio.TimeoutError:
err = f"Timeout after {timeout}s"
except Exception as e:
err = f"{type(e).__name__}: {e}"
if payload is not None:
results.append(payload)
if grace is not None and deadline is None:
deadline = loop.time() + grace
_log(topic, f"{label}: first result — grace {grace}s running")
if on_update:
on_update(len(results))
if (len(results) >= quorum and (grace is None or loop.time() >= deadline)
and (min_deadline is None or loop.time() >= min_deadline)):
return results
continue
_log(topic, f"{label} {i + 1} (attempt {attempts[i] + 1}): {err}")
attempts[i] += 1
# If the minimum already stands, restarts are pointless — the restart
# would be killed at the grace end anyway.
enough = grace is not None and len(results) >= quorum
if attempts[i] <= _MAX_RESTARTS and not enough and not (cancelled and cancelled()):
spawn(i)
if len(results) >= quorum: # all slots done, minimum stands (only reachable with grace)
return results
_log(topic, f"{label}: quorum {quorum} not reached ({len(results)} valid)")
return None
finally:
for task, i in tasks.items():
kill_process(slots[i]["key"])
task.cancel()
if tasks:
await asyncio.gather(*tasks.keys(), return_exceptions=True)
@dataclass
class GenContext:
"""Pipeline parameters passed through — saves long argument signatures."""
topic: str
provider: str
is_cancelled: Callable[[], bool]
guide_id: str | None = None
# Result status of run_single_slot
OK, CANCELLED, FAILED = "ok", "cancelled", "failed"
async def run_single_slot(
ctx: GenContext, label: str, *,
key: str, prompt: str, role: str, capabilities: str, payload, timeout: int, on_line=None,
) -> tuple[str, object]:
"""One agent, one valid result (race with quorum 1).
→ (OK, value) | (CANCELLED, None) | (FAILED, None)
"""
slots = [{"key": key, "prompt": prompt, "role": role, "capabilities": capabilities, "payload": payload, "on_line": on_line}]
res = await _race(ctx.topic, label, slots, 1, timeout, ctx.provider, cancelled=ctx.is_cancelled)
if ctx.is_cancelled():
return CANCELLED, None
if res is None:
return FAILED, None
return OK, res[0]
async def _gather_progress(coros, total, report, start=0):
"""Runs `coros` concurrently and reports live progress: `await report(done, total)`
after each completion (and once initially). Results in order, return_exceptions=True."""
done = start
async def wrap(c):
nonlocal done
try:
return await c
finally:
done += 1
await report(done, total)
await report(done, total)
return await asyncio.gather(*[wrap(c) for c in coros], return_exceptions=True)