Skip to content

Commit 2917288

Browse files
fix: recover durable gateway tasks by session
1 parent 359378b commit 2917288

4 files changed

Lines changed: 222 additions & 14 deletions

File tree

‎server.py‎

Lines changed: 57 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -171,11 +171,44 @@ def _execute_durable_task(task: dict[str, Any]) -> dict[str, Any]:
171171
}
172172

173173

174-
_task_worker = TaskWorker(
175-
TaskManager(_agent.memory_bank.store, workspace_id=_agent.memory_bank.workspace_id,
176-
user_id=_SERVER_ACTOR_ID),
177-
_execute_durable_task,
178-
)
174+
def _task_manager(session_id: str | None) -> TaskManager:
175+
"""Build a task manager bound to the authenticated deployment scope."""
176+
return TaskManager(
177+
_agent.memory_bank.store,
178+
workspace_id=_agent.memory_bank.workspace_id,
179+
session_id=session_id,
180+
user_id=_SERVER_ACTOR_ID,
181+
)
182+
183+
184+
def _task_worker(session_id: str | None) -> TaskWorker:
185+
"""Build a bounded worker that can only claim one exact session scope."""
186+
return TaskWorker(_task_manager(session_id), _execute_durable_task, max_tasks=1)
187+
188+
189+
def _task_scopes(statuses: set[str]) -> set[str | None]:
190+
rows = _agent.memory_bank.store.list_tasks(
191+
workspace_id=_agent.memory_bank.workspace_id,
192+
statuses=statuses,
193+
limit=10000,
194+
user_id=_SERVER_ACTOR_ID,
195+
)
196+
return {row.get("session_id") for row in rows}
197+
198+
199+
def _recover_task_scopes() -> list[dict[str, Any]]:
200+
"""Recover every unfinished task without collapsing session boundaries."""
201+
scopes = _task_scopes({"pending", "running", "failed", "blocked"})
202+
recovered: list[dict[str, Any]] = []
203+
for session_id in scopes:
204+
recovered.extend(_task_manager(session_id).recover())
205+
return recovered
206+
207+
208+
def _run_task_worker_cycle() -> None:
209+
"""Run at most one safe action per pending session in this scheduler tick."""
210+
for session_id in sorted(_task_scopes({"pending"}), key=lambda item: item or ""):
211+
_task_worker(session_id).run_once()
179212

180213

181214
def _run_scheduled_job(job: dict[str, Any]) -> None:
@@ -184,7 +217,7 @@ def _run_scheduled_job(job: dict[str, Any]) -> None:
184217
_agent.dispatch_learning_cycle(surface="web", trigger="scheduled", max_features=4)
185218
return
186219
if payload.get("type") == "task_worker":
187-
_task_worker.run_once()
220+
_run_task_worker_cycle()
188221
return
189222
if payload.get("type") == "chat":
190223
session_id = _normalise_session_id(payload.get("session_id"))
@@ -196,6 +229,7 @@ def _run_scheduled_job(job: dict[str, Any]) -> None:
196229

197230
_scheduler.register_callback("learning_cycle", _run_scheduled_job)
198231
_scheduler.register_callback("chat", _run_scheduled_job)
232+
_scheduler.register_callback("task_worker", _run_scheduled_job)
199233

200234

201235
def _normalise_session_id(raw_session_id: Any) -> str:
@@ -683,8 +717,7 @@ async def api_v2_memory_forget_claim(claim_id: str):
683717
@app.get("/api/v2/tasks", dependencies=[Depends(require_api_access)])
684718
async def api_v2_tasks(session_id: str | None = None):
685719
session = _normalise_session_id(session_id) if session_id else None
686-
manager = TaskManager(_agent.memory_bank.store, workspace_id=_agent.memory_bank.workspace_id,
687-
session_id=session, user_id=_SERVER_ACTOR_ID)
720+
manager = _task_manager(session)
688721
return {"tasks": manager.store.list_tasks(workspace_id=manager.workspace_id, session_id=session,
689722
user_id=_SERVER_ACTOR_ID)}
690723

@@ -695,8 +728,7 @@ async def api_v2_create_task(request: Request):
695728
if not isinstance(body, dict) or not str(body.get("description", "")).strip():
696729
raise HTTPException(400, "description is required")
697730
session = _normalise_session_id(body.get("session_id"))
698-
manager = TaskManager(_agent.memory_bank.store, workspace_id=_agent.memory_bank.workspace_id,
699-
session_id=session, user_id=_SERVER_ACTOR_ID)
731+
manager = _task_manager(session)
700732
checkpoint: dict[str, Any] = {}
701733
if body.get("action") is not None:
702734
action = _agent.TOOL_ALIASES.get(str(body.get("action")).strip(), str(body.get("action")).strip())
@@ -722,10 +754,23 @@ async def api_v2_create_task(request: Request):
722754

723755
@app.post("/api/v2/tasks/{task_id}/resume", dependencies=[Depends(require_api_access)])
724756
async def api_v2_resume_task(task_id: str):
725-
manager = _task_worker.manager
757+
candidates = _agent.memory_bank.store.list_tasks(
758+
workspace_id=_agent.memory_bank.workspace_id,
759+
statuses={"pending", "running", "failed", "blocked"},
760+
limit=10000,
761+
user_id=_SERVER_ACTOR_ID,
762+
)
763+
stored = next((item for item in candidates if item["id"] == str(task_id)), None)
764+
if stored is None:
765+
raise HTTPException(404, "Failed or blocked task not found in this scope")
766+
manager = _task_manager(stored.get("session_id"))
767+
manager.recover()
726768
task = manager.resume(task_id)
727769
if task is None:
728770
raise HTTPException(404, "Failed or blocked task not found in this scope")
771+
if not any(job.get("payload", {}).get("type") == "task_worker" for job in _scheduler.list_jobs()):
772+
_scheduler.schedule_every("durable-task-worker", 1.0, payload={"type": "task_worker"},
773+
job_id="job_durable_task_worker", delay_seconds=0)
729774
return {"task": task, "queued": True}
730775

731776

@@ -1431,7 +1476,7 @@ async def startup():
14311476
"on_startup", agent=_agent, surface="web", user_id=_SERVER_ACTOR_ID,
14321477
workspace_id=_agent.memory_bank.workspace_id,
14331478
)
1434-
recovered = _task_worker.recover()
1479+
recovered = _recover_task_scopes()
14351480
if recovered:
14361481
_audit("TASK_RECOVERY", f"recovered={len(recovered)} pending_or_resumable tasks")
14371482
if not any(job.get("payload", {}).get("type") == "task_worker" for job in _scheduler.list_jobs()):
Lines changed: 163 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,163 @@
1+
import json
2+
import os
3+
import socket
4+
import subprocess
5+
import sys
6+
import tempfile
7+
import time
8+
import unittest
9+
import urllib.error
10+
import urllib.request
11+
from pathlib import Path
12+
13+
from event_store import EventStore
14+
from task_engine import TaskManager
15+
16+
17+
class DurableTaskGatewayTests(unittest.TestCase):
18+
@staticmethod
19+
def _free_port() -> int:
20+
with socket.socket() as sock:
21+
sock.bind(("127.0.0.1", 0))
22+
return int(sock.getsockname()[1])
23+
24+
@staticmethod
25+
def _request(base_url: str, path: str, *, method: str = "GET", body: dict | None = None) -> dict:
26+
data = json.dumps(body).encode("utf-8") if body is not None else None
27+
request = urllib.request.Request(
28+
base_url + path,
29+
data=data,
30+
method=method,
31+
headers={"Content-Type": "application/json", "Authorization": "Bearer durable-test-token"},
32+
)
33+
with urllib.request.urlopen(request, timeout=5) as response:
34+
return json.loads(response.read().decode("utf-8"))
35+
36+
def _start_server(self, workspace: Path, db_path: Path, skills_path: Path) -> tuple[subprocess.Popen, str]:
37+
repository = Path(__file__).parents[1]
38+
port = self._free_port()
39+
env = os.environ.copy()
40+
env.update({
41+
"HOME": str(workspace / "home"),
42+
"KYROZEN_DB_PATH": str(db_path),
43+
"KYROZEN_DISABLE_VECTOR_INDEX": "1",
44+
"KYROZEN_PROVIDER": "ollama",
45+
"KYROZEN_BASE_URL": "http://127.0.0.1:11434/v1",
46+
"KYROZEN_SERVER_TOKEN": "durable-test-token",
47+
"KYROZEN_SKILLS_DIR": str(skills_path),
48+
})
49+
for variable in ("KYROZEN_API_KEY", "DEEPSEEK_API_KEY", "OPENAI_API_KEY", "ANTHROPIC_API_KEY"):
50+
env.pop(variable, None)
51+
process = subprocess.Popen(
52+
[sys.executable, str(repository / "server.py"), "--host", "127.0.0.1", "--port", str(port)],
53+
cwd=workspace,
54+
env=env,
55+
stdout=subprocess.PIPE,
56+
stderr=subprocess.STDOUT,
57+
text=True,
58+
)
59+
base_url = f"http://127.0.0.1:{port}"
60+
deadline = time.monotonic() + 20
61+
while time.monotonic() < deadline:
62+
if process.poll() is not None:
63+
output = process.stdout.read() if process.stdout else ""
64+
self.fail(f"Gateway exited before becoming ready: {output}")
65+
try:
66+
health = self._request(base_url, "/api/health")
67+
if health.get("status") in {"ok", "degraded"}:
68+
return process, base_url
69+
except (OSError, urllib.error.URLError):
70+
time.sleep(0.1)
71+
process.terminate()
72+
process.wait(timeout=5)
73+
self.fail("Gateway did not become ready")
74+
75+
@staticmethod
76+
def _stop_server(process: subprocess.Popen) -> None:
77+
if process.poll() is None:
78+
process.terminate()
79+
process.wait(timeout=5)
80+
if process.stdout:
81+
process.stdout.close()
82+
83+
def test_gateway_recovers_scoped_running_task_and_executes_api_task(self):
84+
with tempfile.TemporaryDirectory(prefix="openkyrozen-task-gateway-") as directory:
85+
root = Path(directory)
86+
workspace = root / "workspace"
87+
workspace.mkdir()
88+
(workspace / "home").mkdir()
89+
db_path = root / "state.sqlite3"
90+
skills_path = root / "skills"
91+
store = EventStore(db_path)
92+
recovery_manager = TaskManager(
93+
store, user_id="local", workspace_id="default", session_id="restart-session"
94+
)
95+
recovery_index = recovery_manager.add_task(
96+
"recover a running file task",
97+
checkpoint={"action": "write_file", "args": "recovered.txt|after restart"},
98+
)
99+
recovery_manager.set_status(recovery_index, "running")
100+
failed_manager = TaskManager(
101+
store, user_id="local", workspace_id="default", session_id="resume-session"
102+
)
103+
failed_index = failed_manager.add_task("resume only when requested")
104+
failed_manager.set_status(failed_index, "failed")
105+
failed_id = failed_manager.tasks[failed_index]["id"]
106+
107+
process, base_url = self._start_server(workspace, db_path, skills_path)
108+
try:
109+
deadline = time.monotonic() + 10
110+
recovered = self._request(base_url, "/api/v2/tasks?session_id=restart-session")
111+
while time.monotonic() < deadline:
112+
recovered = self._request(base_url, "/api/v2/tasks?session_id=restart-session")
113+
if recovered["tasks"][0]["status"] == "succeeded":
114+
break
115+
time.sleep(0.1)
116+
self.assertEqual(recovered["tasks"][0]["status"], "succeeded")
117+
self.assertEqual((workspace / "recovered.txt").read_text(encoding="utf-8"), "after restart")
118+
119+
created = self._request(
120+
base_url,
121+
"/api/v2/tasks",
122+
method="POST",
123+
body={
124+
"description": "write a live task marker",
125+
"session_id": "task-live",
126+
"action": "write_file",
127+
"args": "task-runtime.txt|created by durable worker",
128+
"acceptance": ["marker exists"],
129+
},
130+
)
131+
live_id = created["task"]["id"]
132+
deadline = time.monotonic() + 10
133+
live = created["task"]
134+
while time.monotonic() < deadline:
135+
live = self._request(base_url, "/api/v2/tasks?session_id=task-live")["tasks"][0]
136+
if live["status"] == "succeeded":
137+
break
138+
time.sleep(0.1)
139+
self.assertEqual(live["id"], live_id)
140+
self.assertEqual(live["status"], "succeeded")
141+
self.assertEqual(
142+
(workspace / "task-runtime.txt").read_text(encoding="utf-8"),
143+
"created by durable worker",
144+
)
145+
146+
before_resume = self._request(base_url, "/api/v2/tasks?session_id=resume-session")
147+
self.assertEqual(before_resume["tasks"][0]["status"], "failed")
148+
resumed = self._request(
149+
base_url, f"/api/v2/tasks/{failed_id}/resume", method="POST", body={}
150+
)
151+
self.assertEqual(resumed["task"]["id"], failed_id)
152+
self.assertEqual(resumed["task"]["status"], "pending")
153+
154+
events = self._request(base_url, "/api/v2/events?session_id=restart-session")["events"]
155+
event_types = {event["event_type"] for event in events}
156+
self.assertIn("task.recovered", event_types)
157+
self.assertIn("task.execution_completed", event_types)
158+
finally:
159+
self._stop_server(process)
160+
161+
162+
if __name__ == "__main__":
163+
unittest.main()

‎tests/test_learning_dispatcher.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -188,7 +188,7 @@ def test_web_startup_persists_a_learning_cycle_job(self):
188188
patch.object(server._agent, "_load_project_files_into_memory"), \
189189
patch.object(server, "_load_plugins"), \
190190
patch.object(server, "_trigger_hook"), \
191-
patch.object(server._task_worker, "recover", return_value=[]), \
191+
patch.object(server, "_recover_task_scopes", return_value=[]), \
192192
patch.object(server._scheduler, "list_jobs", return_value=[
193193
{"payload": {"type": "task_worker"}},
194194
]), \

‎tests/test_server.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -192,7 +192,7 @@ def test_server_startup_initializes_headlessly(self):
192192
with patch.object(server._agent, "_load_project_files_into_memory"):
193193
with patch.object(server, "_load_plugins"):
194194
with patch.object(server, "_trigger_hook"):
195-
with patch.object(server._task_worker, "recover", return_value=[]):
195+
with patch.object(server, "_recover_task_scopes", return_value=[]):
196196
with patch.object(server._scheduler, "list_jobs", return_value=[
197197
{"payload": {"type": "task_worker"}},
198198
]):

0 commit comments

Comments
 (0)