deep-agent-cli 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (50) hide show
  1. agent/__init__.py +42 -0
  2. agent/attachments.py +303 -0
  3. agent/bootstrap.py +44 -0
  4. agent/cancel.py +107 -0
  5. agent/cli/__init__.py +5 -0
  6. agent/cli/app.py +1768 -0
  7. agent/cli/clipboard.py +224 -0
  8. agent/cli/commands.py +94 -0
  9. agent/cli/gitinfo.py +84 -0
  10. agent/cli/input.py +65 -0
  11. agent/cli/interactions.py +187 -0
  12. agent/cli/main.py +124 -0
  13. agent/cli/previews.py +710 -0
  14. agent/cli/rendering.py +770 -0
  15. agent/cli/session_controller.py +221 -0
  16. agent/cli/state.py +326 -0
  17. agent/config.example.yaml +76 -0
  18. agent/config.py +528 -0
  19. agent/control.py +171 -0
  20. agent/factory.py +232 -0
  21. agent/file_mutation.py +5 -0
  22. agent/llm.py +339 -0
  23. agent/middleware/__init__.py +9 -0
  24. agent/middleware/attachments.py +31 -0
  25. agent/middleware/cancel_tools.py +39 -0
  26. agent/middleware/pause.py +18 -0
  27. agent/middleware/recovery.py +65 -0
  28. agent/middleware/steering.py +35 -0
  29. agent/middleware/tool_arg_hints.py +128 -0
  30. agent/middleware/workspace_filesystem.py +38 -0
  31. agent/middleware/write_operation.py +60 -0
  32. agent/network.py +30 -0
  33. agent/permission.py +80 -0
  34. agent/runner.py +1393 -0
  35. agent/sandbox.py +699 -0
  36. agent/session.py +431 -0
  37. agent/session_lock.py +223 -0
  38. agent/session_runtime.py +209 -0
  39. agent/stream.py +168 -0
  40. agent/tools/__init__.py +9 -0
  41. agent/tools/examples.py +30 -0
  42. agent/tools/execute.py +73 -0
  43. agent/tools/human_input.py +170 -0
  44. agent/tools/human_interaction.py +101 -0
  45. agent/tools/web_search.py +131 -0
  46. deep_agent_cli-0.1.0.dist-info/METADATA +408 -0
  47. deep_agent_cli-0.1.0.dist-info/RECORD +50 -0
  48. deep_agent_cli-0.1.0.dist-info/WHEEL +4 -0
  49. deep_agent_cli-0.1.0.dist-info/entry_points.txt +2 -0
  50. deep_agent_cli-0.1.0.dist-info/licenses/LICENSE +21 -0
agent/session.py ADDED
@@ -0,0 +1,431 @@
1
+ """SQLite session catalog + LangGraph SqliteSaver for local persistence."""
2
+ from __future__ import annotations
3
+
4
+ from dataclasses import dataclass
5
+ from datetime import datetime, timezone
6
+ from enum import StrEnum
7
+ import hashlib
8
+ import os
9
+ from pathlib import Path
10
+ import sqlite3
11
+ from typing import Any, Callable, Literal
12
+ from uuid import uuid4
13
+
14
+ from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, ToolMessage
15
+ from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
16
+ from langgraph.checkpoint.sqlite import SqliteSaver
17
+
18
+ from agent.attachments import (
19
+ AttachmentCleanupResult,
20
+ AttachmentStore,
21
+ ImageAttachmentRef,
22
+ find_attachment_storage_keys,
23
+ refs_from_message,
24
+ )
25
+ from agent.session_lock import SessionLease, SessionLockManager
26
+ from agent.stream import message_text, reasoning_text, visible_text
27
+
28
+ SessionStatus = Literal[
29
+ "running", "completed", "waiting", "cancelled", "failed",
30
+ ]
31
+
32
+
33
+ class StopReason(StrEnum):
34
+ PENDING = "pending"
35
+ STOP = "stop"
36
+ ERROR = "error"
37
+ ABORTED = "aborted"
38
+ DEFERRED = "deferred"
39
+
40
+ _CATALOG_DDL = """
41
+ CREATE TABLE IF NOT EXISTS session_catalog (
42
+ id TEXT PRIMARY KEY,
43
+ title TEXT NOT NULL,
44
+ created_at TEXT NOT NULL,
45
+ updated_at TEXT NOT NULL,
46
+ model_id TEXT,
47
+ permission_mode TEXT,
48
+ last_run_status TEXT NOT NULL
49
+ );
50
+ CREATE INDEX IF NOT EXISTS session_catalog_updated_idx
51
+ ON session_catalog (updated_at DESC);
52
+ """
53
+
54
+
55
+ @dataclass(frozen=True)
56
+ class SessionInfo:
57
+ id: str
58
+ title: str
59
+ created_at: datetime
60
+ updated_at: datetime
61
+ status: SessionStatus
62
+ model_id: str | None = None
63
+ permission_mode: str | None = None
64
+ last_run_status: StopReason = StopReason.PENDING
65
+
66
+
67
+ @dataclass
68
+ class TranscriptBlock:
69
+ kind: str
70
+ content: str = ""
71
+ thinking: str = ""
72
+ tool_call_id: str = ""
73
+ name: str = ""
74
+ arguments: dict[str, Any] | None = None
75
+ is_error: bool = False
76
+ status: str = ""
77
+ attachments: tuple[ImageAttachmentRef, ...] = ()
78
+ exit_code: int | None = None
79
+ # Raw ToolMessage.artifact restored from the checkpoint; UI previews are
80
+ # re-derived from it instead of being persisted.
81
+ artifact: Any = None
82
+
83
+
84
+ def workspace_state_path(workspace: Path, *, override: str | Path | None = None) -> Path:
85
+ if override is not None:
86
+ return Path(override).expanduser().resolve()
87
+ env = os.environ.get("DEEP_AGENT_STATE_PATH")
88
+ if env:
89
+ return Path(env).expanduser().resolve()
90
+ digest = hashlib.sha256(str(workspace.expanduser().resolve()).encode("utf-8")).hexdigest()[:16]
91
+ return Path.home() / ".deep-agent" / "sessions" / f"{digest}.sqlite3"
92
+
93
+
94
+ def _utc_now() -> datetime:
95
+ return datetime.now(timezone.utc)
96
+
97
+
98
+ def _parse_dt(value: str) -> datetime:
99
+ return datetime.fromisoformat(value)
100
+
101
+
102
+ def _strict_serde() -> JsonPlusSerializer:
103
+ # Strict allowlist: only built-in safe msgpack types are reconstructed.
104
+ return JsonPlusSerializer(allowed_msgpack_modules=None, pickle_fallback=False)
105
+
106
+
107
+ class SessionStore:
108
+ """Presentation-neutral session API backed by one SQLite file per workspace."""
109
+
110
+ def __init__(self, path: Path) -> None:
111
+ self.path = path
112
+ self.path.parent.mkdir(parents=True, mode=0o700, exist_ok=True)
113
+ try:
114
+ os.chmod(self.path.parent, 0o700)
115
+ except OSError:
116
+ pass
117
+ self._conn = sqlite3.connect(str(self.path), check_same_thread=False)
118
+ self._conn.execute("PRAGMA journal_mode=WAL;")
119
+ self._conn.executescript(_CATALOG_DDL)
120
+ self._migrate_catalog()
121
+ if self.path.exists():
122
+ try:
123
+ os.chmod(self.path, 0o600)
124
+ except OSError:
125
+ pass
126
+ self.checkpointer = SqliteSaver(self._conn, serde=_strict_serde())
127
+ self.session_locks = SessionLockManager(self.path)
128
+ self.attachment_store = AttachmentStore.beside_database(self.path)
129
+ self.startup_cleanup: AttachmentCleanupResult | None = None
130
+ try:
131
+ self.startup_cleanup = self.cleanup_attachments()
132
+ except Exception: # noqa: BLE001 - persistence must remain usable if maintenance fails
133
+ self.startup_cleanup = None
134
+ self.attachment_store.acquire_runtime_lease()
135
+ self._closed = False
136
+
137
+ @classmethod
138
+ def for_workspace(
139
+ cls,
140
+ workspace: Path,
141
+ *,
142
+ override: str | Path | None = None,
143
+ ) -> "SessionStore":
144
+ return cls(workspace_state_path(workspace, override=override))
145
+
146
+ def try_acquire_session(self, thread_id: str) -> SessionLease | None:
147
+ """Take exclusive ownership of a persisted thread; None when busy."""
148
+ return self.session_locks.try_acquire(thread_id)
149
+
150
+ def wait_acquire_session(
151
+ self,
152
+ thread_id: str,
153
+ *,
154
+ cancelled: Callable[[], bool] | None = None,
155
+ ) -> SessionLease | None:
156
+ """Block until the thread is free, the wait is cancelled, or ownership lands."""
157
+ return self.session_locks.wait_acquire(thread_id, cancelled=cancelled)
158
+
159
+ def close(self) -> None:
160
+ if self._closed:
161
+ return
162
+ self._closed = True
163
+ self.session_locks.close()
164
+ self.attachment_store.close()
165
+ self._conn.close()
166
+
167
+ def _migrate_catalog(self) -> None:
168
+ # Hold a write reservation while checking the schema so concurrent
169
+ # processes cannot both decide to migrate the same workspace database.
170
+ self._conn.execute("BEGIN IMMEDIATE")
171
+ try:
172
+ columns = {row[1] for row in self._conn.execute("PRAGMA table_info(session_catalog)")}
173
+ if "status" not in columns:
174
+ self._conn.commit()
175
+ return
176
+ model = "model_id" if "model_id" in columns else "NULL"
177
+ permission = "permission_mode" if "permission_mode" in columns else "NULL"
178
+ last = "last_run_status" if "last_run_status" in columns else "NULL"
179
+ reason = (
180
+ f"CASE WHEN {last} IN ('pending', 'stop', 'error', 'aborted', 'deferred') THEN {last} "
181
+ "WHEN status = 'completed' OR status IN ('length', 'tool_use') THEN 'stop' "
182
+ "WHEN status = 'cancelled' THEN 'aborted' "
183
+ "WHEN status = 'failed' THEN 'error' "
184
+ "WHEN status IN ('waiting', 'interrupted') THEN 'deferred' "
185
+ "ELSE 'pending' END"
186
+ )
187
+ self._conn.execute(
188
+ "CREATE TABLE session_catalog_new ("
189
+ "id TEXT PRIMARY KEY, title TEXT NOT NULL, created_at TEXT NOT NULL, "
190
+ "updated_at TEXT NOT NULL, model_id TEXT, permission_mode TEXT, "
191
+ "last_run_status TEXT NOT NULL)"
192
+ )
193
+ self._conn.execute(
194
+ "INSERT INTO session_catalog_new "
195
+ "(id, title, created_at, updated_at, model_id, permission_mode, last_run_status) "
196
+ f"SELECT id, title, created_at, updated_at, {model}, {permission}, {reason} "
197
+ "FROM session_catalog"
198
+ )
199
+ self._conn.execute("DROP TABLE session_catalog")
200
+ self._conn.execute("ALTER TABLE session_catalog_new RENAME TO session_catalog")
201
+ self._conn.execute(
202
+ "CREATE INDEX session_catalog_updated_idx ON session_catalog (updated_at DESC)"
203
+ )
204
+ self._conn.commit()
205
+ except Exception:
206
+ self._conn.rollback()
207
+ raise
208
+
209
+ def attachment_storage_keys(self) -> set[str]:
210
+ """Return references from every retained checkpoint and pending write."""
211
+ live: set[str] = set()
212
+ for item in self.checkpointer.list(None):
213
+ live.update(find_attachment_storage_keys(item.checkpoint))
214
+ live.update(find_attachment_storage_keys(item.pending_writes))
215
+ return live
216
+
217
+ def cleanup_attachments(
218
+ self,
219
+ *,
220
+ protected: tuple[ImageAttachmentRef, ...] = (),
221
+ release_runtime_lease: bool = False,
222
+ ) -> AttachmentCleanupResult:
223
+ live = self.attachment_storage_keys()
224
+ return self.attachment_store.cleanup(
225
+ live,
226
+ protected_storage_keys=(ref.storage_key for ref in protected),
227
+ release_runtime_lease=release_runtime_lease,
228
+ )
229
+
230
+ def create_session(
231
+ self, *, title: str = "", session_id: str | None = None,
232
+ model_id: str | None = None, permission_mode: str = "ask",
233
+ ) -> SessionInfo:
234
+ now = _utc_now()
235
+ info = SessionInfo(
236
+ id=session_id or str(uuid4()),
237
+ title=title or "New session",
238
+ created_at=now,
239
+ updated_at=now,
240
+ status="running",
241
+ model_id=model_id,
242
+ permission_mode=permission_mode,
243
+ )
244
+ with self.checkpointer.lock:
245
+ self._conn.execute(
246
+ "INSERT INTO session_catalog (id, title, created_at, updated_at, model_id, permission_mode, last_run_status) VALUES (?, ?, ?, ?, ?, ?, ?)",
247
+ (info.id, info.title, info.created_at.isoformat(), info.updated_at.isoformat(),
248
+ info.model_id, info.permission_mode, info.last_run_status.value),
249
+ )
250
+ self._conn.commit()
251
+ return info
252
+
253
+ def touch(
254
+ self,
255
+ session_id: str,
256
+ *,
257
+ title: str | None = None,
258
+ model_id: str | None = None,
259
+ permission_mode: str | None = None,
260
+ last_run_status: StopReason | None = None,
261
+ ) -> None:
262
+ now = _utc_now().isoformat()
263
+ with self.checkpointer.lock:
264
+ row = self._conn.execute(
265
+ "SELECT title, model_id, permission_mode, last_run_status FROM session_catalog WHERE id = ?", (session_id,),
266
+ ).fetchone()
267
+ if row is None:
268
+ return
269
+ new_title = title if title is not None else row[0]
270
+ self._conn.execute(
271
+ "UPDATE session_catalog SET title = ?, updated_at = ?, model_id = ?, permission_mode = ?, last_run_status = ? WHERE id = ?",
272
+ (new_title, now, model_id if model_id is not None else row[1],
273
+ permission_mode if permission_mode is not None else row[2],
274
+ last_run_status.value if last_run_status is not None else row[3], session_id),
275
+ )
276
+ self._conn.commit()
277
+
278
+ def list_sessions(self, *, limit: int = 50) -> list[SessionInfo]:
279
+ with self.checkpointer.lock:
280
+ rows = self._conn.execute(
281
+ "SELECT id, title, created_at, updated_at, model_id, permission_mode, last_run_status "
282
+ "FROM session_catalog ORDER BY updated_at DESC LIMIT ?",
283
+ (limit,),
284
+ ).fetchall()
285
+ return [_session_info(row) for row in rows]
286
+
287
+ def get(self, session_id: str) -> SessionInfo | None:
288
+ with self.checkpointer.lock:
289
+ row = self._conn.execute(
290
+ "SELECT id, title, created_at, updated_at, model_id, permission_mode, last_run_status FROM session_catalog WHERE id = ?",
291
+ (session_id,),
292
+ ).fetchone()
293
+ if row is None:
294
+ return None
295
+ return _session_info(row)
296
+
297
+ def resolve_prefix(self, prefix: str) -> SessionInfo | None:
298
+ prefix = prefix.strip()
299
+ if not prefix:
300
+ return None
301
+ exact = self.get(prefix)
302
+ if exact is not None:
303
+ return exact
304
+ with self.checkpointer.lock:
305
+ rows = self._conn.execute(
306
+ "SELECT id, title, created_at, updated_at, model_id, permission_mode, last_run_status FROM session_catalog WHERE id LIKE ?",
307
+ (f"{prefix}%",),
308
+ ).fetchall()
309
+ if len(rows) != 1:
310
+ return None
311
+ return _session_info(rows[0])
312
+
313
+
314
+ def _session_info(row: Any) -> SessionInfo:
315
+ try:
316
+ reason = StopReason(row[6])
317
+ except ValueError:
318
+ reason = StopReason.PENDING
319
+ return SessionInfo(
320
+ id=row[0], title=row[1], created_at=_parse_dt(row[2]), updated_at=_parse_dt(row[3]),
321
+ status=_status_for_reason(reason), model_id=row[4], permission_mode=row[5], last_run_status=reason,
322
+ )
323
+
324
+
325
+ def _status_for_reason(reason: StopReason) -> SessionStatus:
326
+ return {
327
+ StopReason.PENDING: "running",
328
+ StopReason.STOP: "completed",
329
+ StopReason.ERROR: "failed",
330
+ StopReason.ABORTED: "cancelled",
331
+ StopReason.DEFERRED: "waiting",
332
+ }[reason]
333
+
334
+
335
+ def tool_message_is_error(message: ToolMessage) -> bool:
336
+ """Match live tool_completed error classification, including execute artifacts."""
337
+ status = str(getattr(message, "status", "") or "")
338
+ content = message_text(message)
339
+ name = str(getattr(message, "name", "") or "")
340
+ artifact = getattr(message, "artifact", None)
341
+ if name == "execute" and isinstance(artifact, dict):
342
+ exit_code = artifact.get("exit_code")
343
+ return (
344
+ status == "error"
345
+ or artifact.get("termination_reason") is not None
346
+ or (isinstance(exit_code, int) and exit_code != 0)
347
+ )
348
+ return (
349
+ status == "error"
350
+ or content.lower().startswith("error")
351
+ or "cancelled by user" in content.lower()
352
+ )
353
+
354
+
355
+ def messages_to_transcript(messages: list[BaseMessage]) -> list[TranscriptBlock]:
356
+ """Rebuild a presentation-neutral transcript from LangGraph checkpoint messages."""
357
+ blocks: list[TranscriptBlock] = []
358
+ todo_call_ids: set[str] = set()
359
+ for message in messages:
360
+ if isinstance(message, HumanMessage):
361
+ blocks.append(TranscriptBlock(
362
+ kind="user",
363
+ content=message_text(message),
364
+ attachments=refs_from_message(message),
365
+ ))
366
+ elif isinstance(message, AIMessage):
367
+ thinking = reasoning_text(message)
368
+ text = visible_text(message)
369
+ if thinking or text:
370
+ blocks.append(TranscriptBlock(kind="assistant", content=text, thinking=thinking))
371
+ for call in message.tool_calls or []:
372
+ if call.get("name") == "write_todos":
373
+ todo_call_ids.add(str(call.get("id") or ""))
374
+ continue
375
+ blocks.append(TranscriptBlock(
376
+ kind="tool",
377
+ tool_call_id=str(call.get("id") or ""),
378
+ name=str(call.get("name") or "tool"),
379
+ arguments=call.get("args") if isinstance(call.get("args"), dict) else {},
380
+ status="running",
381
+ ))
382
+ elif isinstance(message, ToolMessage):
383
+ content = message_text(message)
384
+ tool_call_id = str(getattr(message, "tool_call_id", "") or "")
385
+ if getattr(message, "name", None) == "write_todos" or tool_call_id in todo_call_ids:
386
+ continue
387
+ is_error = tool_message_is_error(message)
388
+ artifact = getattr(message, "artifact", None)
389
+ code = (
390
+ artifact.get("exit_code")
391
+ if isinstance(artifact, dict) and getattr(message, "name", None) == "execute"
392
+ else None
393
+ )
394
+ exit_code = code if isinstance(code, int) and not isinstance(code, bool) else None
395
+ updated = False
396
+ for block in reversed(blocks):
397
+ if block.kind == "tool" and block.tool_call_id == tool_call_id:
398
+ block.content = content
399
+ block.is_error = is_error
400
+ block.status = "error" if is_error else "completed"
401
+ block.exit_code = exit_code
402
+ block.artifact = artifact
403
+ updated = True
404
+ break
405
+ if not updated:
406
+ blocks.append(TranscriptBlock(
407
+ kind="tool",
408
+ tool_call_id=tool_call_id,
409
+ name=str(getattr(message, "name", "") or "tool"),
410
+ content=content,
411
+ is_error=is_error,
412
+ status="error" if is_error else "completed",
413
+ exit_code=exit_code,
414
+ artifact=artifact,
415
+ ))
416
+ return blocks
417
+
418
+
419
+ def settle_restored_tools(
420
+ blocks: list[TranscriptBlock], *, waiting_ids: set[str] | None = None,
421
+ waiting_human: bool = False,
422
+ ) -> None:
423
+ """Freeze tool calls without results when showing a saved checkpoint."""
424
+ pending = waiting_ids or set()
425
+ for block in blocks:
426
+ if block.kind != "tool" or block.status != "running":
427
+ continue
428
+ if block.tool_call_id in pending or (waiting_human and block.name == "request_human_input"):
429
+ block.status = "waiting"
430
+ else:
431
+ block.status = "interrupted"
agent/session_lock.py ADDED
@@ -0,0 +1,223 @@
1
+ """Exclusive thread ownership for one session database.
2
+
3
+ A lease is an advisory ``flock`` on ``<database>.locks/<sha256(thread_id)>.lock``
4
+ so two processes can never write the same persisted thread at once. Lock files
5
+ are created once and never deleted; releasing only drops the kernel hold. The
6
+ locks directory and files stay private to the database owner and every file
7
+ descriptor is close-on-exec so agent child processes cannot inherit a hold.
8
+
9
+ Runners without a persisted store share nothing across processes; for those,
10
+ ``InProcessSessionLockManager`` guards threads that flow through one
11
+ checkpointer object within this process.
12
+ """
13
+ from __future__ import annotations
14
+
15
+ import fcntl
16
+ import hashlib
17
+ import os
18
+ from pathlib import Path
19
+ import threading
20
+ import time
21
+ from typing import Any, Callable
22
+ from weakref import WeakKeyDictionary, ref
23
+
24
+ WAIT_POLL_SECONDS = 0.05
25
+
26
+
27
+ class SessionLockBusyError(RuntimeError):
28
+ """The thread is already owned by another runner or process."""
29
+
30
+
31
+ class SessionLease:
32
+ """Ownership of one persisted thread; ``release`` is idempotent."""
33
+
34
+ def __init__(
35
+ self,
36
+ thread_id: str,
37
+ handle: Any,
38
+ on_release: Callable[["SessionLease"], None],
39
+ ) -> None:
40
+ self._thread_id = thread_id
41
+ self._handle = handle
42
+ self._on_release = on_release
43
+ self._released = False
44
+
45
+ @property
46
+ def thread_id(self) -> str:
47
+ return self._thread_id
48
+
49
+ @property
50
+ def released(self) -> bool:
51
+ return self._released
52
+
53
+ def release(self) -> None:
54
+ if self._released:
55
+ return
56
+ self._released = True
57
+ try:
58
+ fcntl.flock(self._handle.fileno(), fcntl.LOCK_UN)
59
+ finally:
60
+ self._handle.close()
61
+ self._on_release(self)
62
+
63
+ def fileno(self) -> int:
64
+ return self._handle.fileno()
65
+
66
+
67
+ class SessionLockManager:
68
+ """flock-backed thread ownership beside one session database."""
69
+
70
+ def __init__(self, database: Path) -> None:
71
+ self._root = database.parent / f"{database.name}.locks"
72
+ self._root.mkdir(parents=True, mode=0o700, exist_ok=True)
73
+ try:
74
+ os.chmod(self._root, 0o700)
75
+ except OSError:
76
+ pass
77
+ self._mutex = threading.Lock()
78
+ self._leases: dict[str, "ref[SessionLease]"] = {}
79
+
80
+ def lock_path(self, thread_id: str) -> Path:
81
+ digest = hashlib.sha256(thread_id.encode("utf-8")).hexdigest()
82
+ return self._root / f"{digest}.lock"
83
+
84
+ def try_acquire(self, thread_id: str) -> SessionLease | None:
85
+ path = self.lock_path(thread_id)
86
+ fd = os.open(path, os.O_RDWR | os.O_CREAT | os.O_CLOEXEC, 0o600)
87
+ handle = os.fdopen(fd, "r+b")
88
+ try:
89
+ fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
90
+ except BlockingIOError:
91
+ handle.close()
92
+ return None
93
+ try:
94
+ os.chmod(path, 0o600)
95
+ except OSError:
96
+ pass
97
+ lease = SessionLease(thread_id, handle, self._forget)
98
+ with self._mutex:
99
+ self._leases[thread_id] = ref(lease)
100
+ return lease
101
+
102
+ def wait_acquire(
103
+ self,
104
+ thread_id: str,
105
+ *,
106
+ cancelled: Callable[[], bool] | None = None,
107
+ poll_seconds: float = WAIT_POLL_SECONDS,
108
+ ) -> SessionLease | None:
109
+ while True:
110
+ if cancelled is not None and cancelled():
111
+ return None
112
+ lease = self.try_acquire(thread_id)
113
+ if lease is not None:
114
+ # Cancellation arriving between the check and the acquisition
115
+ # must still win: hand the lock straight back.
116
+ if cancelled is not None and cancelled():
117
+ lease.release()
118
+ return None
119
+ return lease
120
+ time.sleep(poll_seconds)
121
+
122
+ def close(self) -> None:
123
+ """Release every lease this manager still tracks; idempotent."""
124
+ with self._mutex:
125
+ leases = [item() for item in self._leases.values()]
126
+ self._leases.clear()
127
+ for lease in leases:
128
+ if lease is not None:
129
+ lease.release()
130
+
131
+ def _forget(self, lease: SessionLease) -> None:
132
+ with self._mutex:
133
+ held = self._leases.get(lease.thread_id)
134
+ if held is not None and held() is lease:
135
+ del self._leases[lease.thread_id]
136
+
137
+
138
+ class InProcessSessionLease:
139
+ """Ownership of one in-memory thread; ``release`` is idempotent."""
140
+
141
+ def __init__(self, thread_id: str, owners: dict[str, "ref[InProcessSessionLease]"]) -> None:
142
+ self._thread_id = thread_id
143
+ self._owners = owners
144
+ self._released = False
145
+
146
+ @property
147
+ def thread_id(self) -> str:
148
+ return self._thread_id
149
+
150
+ @property
151
+ def released(self) -> bool:
152
+ return self._released
153
+
154
+ def release(self) -> None:
155
+ if self._released:
156
+ return
157
+ self._released = True
158
+ with _IN_PROCESS_GUARD:
159
+ held = self._owners.get(self._thread_id)
160
+ if held is not None and held() is self:
161
+ del self._owners[self._thread_id]
162
+
163
+
164
+ _IN_PROCESS_GUARD = threading.Lock()
165
+ _IN_PROCESS_OWNERS: "WeakKeyDictionary[Any, dict[str, ref[InProcessSessionLease]]]" = WeakKeyDictionary()
166
+
167
+
168
+ class InProcessSessionLockManager:
169
+ """Guards threads shared through one checkpointer object within this process."""
170
+
171
+ def __init__(self, checkpointer: Any) -> None:
172
+ self._checkpointer = checkpointer
173
+
174
+ def try_acquire(self, thread_id: str) -> InProcessSessionLease | None:
175
+ if self._checkpointer is None:
176
+ # Nothing is shared when the graph has no checkpointer at all.
177
+ return InProcessSessionLease(thread_id, {})
178
+ with _IN_PROCESS_GUARD:
179
+ owners = _IN_PROCESS_OWNERS.get(self._checkpointer)
180
+ if owners is None:
181
+ owners = {}
182
+ _IN_PROCESS_OWNERS[self._checkpointer] = owners
183
+ held = owners.get(thread_id)
184
+ if held is not None:
185
+ lease = held()
186
+ if lease is not None and not lease.released:
187
+ return None
188
+ lease = InProcessSessionLease(thread_id, owners)
189
+ owners[thread_id] = ref(lease)
190
+ return lease
191
+
192
+ def wait_acquire(
193
+ self,
194
+ thread_id: str,
195
+ *,
196
+ cancelled: Callable[[], bool] | None = None,
197
+ poll_seconds: float = WAIT_POLL_SECONDS,
198
+ ) -> InProcessSessionLease | None:
199
+ while True:
200
+ if cancelled is not None and cancelled():
201
+ return None
202
+ lease = self.try_acquire(thread_id)
203
+ if lease is not None:
204
+ # Cancellation arriving between the check and the acquisition
205
+ # must still win: hand the lock straight back.
206
+ if cancelled is not None and cancelled():
207
+ lease.release()
208
+ return None
209
+ return lease
210
+ time.sleep(poll_seconds)
211
+
212
+ def close(self) -> None:
213
+ if self._checkpointer is None:
214
+ return
215
+ with _IN_PROCESS_GUARD:
216
+ owners = _IN_PROCESS_OWNERS.get(self._checkpointer)
217
+ if owners is None:
218
+ return
219
+ leases = [item() for item in owners.values()]
220
+ owners.clear()
221
+ for lease in leases:
222
+ if lease is not None:
223
+ lease.release()