agentdatabase 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.
@@ -0,0 +1,831 @@
1
+ """
2
+ Claude Agent SDK runtime adapter (middleware-style).
3
+
4
+ This adapter does not import claude_agent_sdk directly. Instead, it provides a
5
+ runtime contract wrapper around AgentDB so Claude SDK users can enforce:
6
+
7
+ 1) pre-response retrieval (always call AgentDB.retrieve)
8
+ 2) post-response persistence policy
9
+ 3) outcome recording for reinforcement
10
+
11
+ Use `run_turn(...)` for the simplest enforced flow, or call
12
+ `prepare_turn(...)` / `finalize_turn(...)` manually if your runtime has custom
13
+ callbacks.
14
+ """
15
+ from __future__ import annotations
16
+
17
+ from dataclasses import dataclass
18
+ from datetime import datetime, timezone
19
+ from typing import TYPE_CHECKING, Any, Awaitable, Callable, Optional
20
+ import json
21
+ import os
22
+
23
+ from ..core.directory_tracking import collect_directory_snapshots
24
+ from ..gateway.adapter import GatewayAdapter
25
+
26
+ if TYPE_CHECKING:
27
+ from ..interface.client import AgentDB, RetrievalResult
28
+
29
+
30
+ def _utcnow() -> datetime:
31
+ return datetime.now(timezone.utc)
32
+
33
+
34
+ def _extract_text_from_model_response(raw_response: Any) -> str:
35
+ if isinstance(raw_response, str):
36
+ return raw_response
37
+ if isinstance(raw_response, dict):
38
+ if "output_text" in raw_response:
39
+ return str(raw_response["output_text"])
40
+ content = raw_response.get("content")
41
+ if isinstance(content, str):
42
+ return content
43
+ if isinstance(content, list):
44
+ parts: list[str] = []
45
+ for part in content:
46
+ if isinstance(part, dict):
47
+ text = part.get("text")
48
+ if text is not None:
49
+ parts.append(str(text))
50
+ else:
51
+ text = getattr(part, "text", None)
52
+ if text is not None:
53
+ parts.append(str(text))
54
+ if parts:
55
+ return "\n".join(parts)
56
+ content_attr = getattr(raw_response, "content", None)
57
+ if isinstance(content_attr, list):
58
+ parts: list[str] = []
59
+ for part in content_attr:
60
+ if isinstance(part, dict):
61
+ text = part.get("text")
62
+ if text is not None:
63
+ parts.append(str(text))
64
+ else:
65
+ text = getattr(part, "text", None)
66
+ if text is not None:
67
+ parts.append(str(text))
68
+ if parts:
69
+ return "\n".join(parts)
70
+ return str(raw_response)
71
+
72
+
73
+ @dataclass(frozen=True)
74
+ class ClaudeRuntimeConfig:
75
+ model: str
76
+ system_prompt: Optional[str] = None
77
+ max_tokens: int = 1024
78
+
79
+
80
+ class AgentDBClaudeMiddleware(GatewayAdapter):
81
+ """
82
+ Middleware-style helper for Claude Agent SDK integrations.
83
+
84
+ Model calls are supplied as callbacks so this module stays runtime-agnostic.
85
+ Inherits GatewayAdapter for access to db.gateway (reconciliation, proposals).
86
+ """
87
+
88
+ def __init__(
89
+ self,
90
+ db: "AgentDB",
91
+ agent_id: str,
92
+ conversation_started_callback: Optional[Callable[..., Any]] = None,
93
+ ):
94
+ resolved = db.config_for(agent_id)
95
+ super().__init__(db, agent_id, profile=resolved["name"])
96
+ self._conversation_started_callback = conversation_started_callback
97
+ self._resolved = resolved
98
+ self._session_working_dirs: dict[str, list[str]] = {}
99
+ self.gateway.subscribe("conflict_detected", self._on_conflict)
100
+
101
+ def prepare_turn(
102
+ self,
103
+ *,
104
+ user_query: str,
105
+ session_id: str,
106
+ system_prompt: Optional[str] = None,
107
+ external_context: Optional[list[dict[str, Any]]] = None,
108
+ base_messages: Optional[list[dict[str, str]]] = None,
109
+ working_dirs: Optional[list[str]] = None,
110
+ ) -> tuple[list[dict[str, str]], "RetrievalResult"]:
111
+ """
112
+ Enforced pre-response hook: always retrieves governed memory first.
113
+ Returns model messages and raw retrieval result for optional inspection.
114
+ """
115
+ if working_dirs is not None:
116
+ self._session_working_dirs[session_id] = [
117
+ os.path.realpath(d) for d in working_dirs
118
+ ]
119
+ if self._conversation_started_callback is not None:
120
+ self._conversation_started_callback(
121
+ db=self._db,
122
+ agent_id=self._agent_id,
123
+ session_id=session_id,
124
+ user_query=user_query,
125
+ )
126
+ self.gateway.dispatch("conversation_started", session_id=session_id)
127
+
128
+ retrieval = self._db.retrieve(
129
+ query=user_query,
130
+ agent_id=self._agent_id,
131
+ limit=self._resolved["retrieval"]["retrieve_limit"],
132
+ external_context=external_context,
133
+ session_id=session_id,
134
+ budget_tokens=self._resolved["injection"]["context_token_limit"],
135
+ working_dirs=self._session_working_dirs.get(session_id),
136
+ )
137
+
138
+ messages: list[dict[str, str]] = []
139
+ if system_prompt:
140
+ messages.append({"role": "system", "content": system_prompt})
141
+
142
+ # Principle 9: conflict flags must be assembled BEFORE human-approved
143
+ # memories, not after — an agent must know a memory is contested
144
+ # before it receives the memory content itself.
145
+ override_prefix = self._resolved["injection"]["override_prefix"]
146
+ for instruction in retrieval.override_instructions:
147
+ messages.append({"role": "system", "content": f"{override_prefix} {instruction}"})
148
+
149
+ memory_context = self._build_memory_context(retrieval)
150
+ if memory_context:
151
+ messages.append(
152
+ {
153
+ "role": "system",
154
+ "content": f"{self._resolved['injection']['memory_header']}\n{memory_context}",
155
+ }
156
+ )
157
+
158
+ for msg in base_messages or []:
159
+ messages.append(msg)
160
+ messages.append({"role": "user", "content": user_query})
161
+ return messages, retrieval
162
+
163
+ def finalize_turn(
164
+ self,
165
+ *,
166
+ session_id: str,
167
+ model_response: str,
168
+ memory_writes: Optional[list[dict[str, Any]]] = None,
169
+ message_id: Optional[str] = None,
170
+ tracked_directories: Optional[list[str]] = None,
171
+ directory_snapshots: Optional[list[dict[str, Any]]] = None,
172
+ tracking_base_dir: Optional[str] = None,
173
+ max_files_per_directory: Optional[int] = None,
174
+ outcome_type: Optional[str] = None,
175
+ outcome_value: Optional[float] = None,
176
+ touched_entities: Optional[list[str]] = None,
177
+ user_query: Optional[str] = None,
178
+ ) -> None:
179
+ """
180
+ Post-response hook: persist memory updates and optional outcome signal.
181
+ """
182
+ session_dirs = self._session_working_dirs.get(session_id)
183
+
184
+ explicit_writes = list(memory_writes or [])
185
+ for item in explicit_writes:
186
+ self._db.write(
187
+ key=item["key"],
188
+ value=item["value"],
189
+ origin=item.get("origin", "agent_inferred"),
190
+ agent_id=self._agent_id,
191
+ session_id=session_id,
192
+ entities=touched_entities or [],
193
+ working_dirs=session_dirs,
194
+ )
195
+
196
+ if self._resolved["capture"]["auto_persist_response"] and model_response:
197
+ self._db.write(
198
+ key=self._response_key(session_id),
199
+ value={
200
+ "response": model_response,
201
+ "captured_at": _utcnow().isoformat(),
202
+ "profile": self._resolved["name"],
203
+ },
204
+ origin="agent_inferred",
205
+ agent_id=self._agent_id,
206
+ session_id=session_id,
207
+ entities=touched_entities or [],
208
+ working_dirs=session_dirs,
209
+ )
210
+
211
+ # Write conversation JSONL if enabled
212
+ if user_query is not None:
213
+ capture_cfg = self._resolved["capture"]
214
+ if capture_cfg.get("conversation_jsonl", False):
215
+ memory_dir = getattr(self._db, "memory_dir", None)
216
+ if memory_dir is not None:
217
+ import json as _json
218
+ from pathlib import Path as _Path
219
+ conv_dir = _Path(memory_dir) / "conversations"
220
+ conv_dir.mkdir(parents=True, exist_ok=True)
221
+ filename = f"{self._agent_id}--{session_id}.jsonl"
222
+ filepath = conv_dir / filename
223
+ entry = _json.dumps({
224
+ "session_id": session_id,
225
+ "message_id": message_id,
226
+ "user_query": user_query,
227
+ "model_response": model_response,
228
+ })
229
+ with filepath.open("a", encoding="utf-8") as f:
230
+ f.write(entry + "\n")
231
+
232
+ normalized_snapshots = list(directory_snapshots or [])
233
+ effective_tracked_directories = (
234
+ tracked_directories
235
+ if tracked_directories is not None
236
+ else getattr(self._db, "tracking_directories", [])
237
+ )
238
+ effective_tracking_base_dir = (
239
+ tracking_base_dir
240
+ if tracking_base_dir is not None
241
+ else str(getattr(self._db, "tracking_root", "."))
242
+ )
243
+ effective_max_files = (
244
+ max_files_per_directory
245
+ if max_files_per_directory is not None
246
+ else int(getattr(self._db, "tracking_max_files_per_directory", 200))
247
+ )
248
+ if effective_tracked_directories:
249
+ normalized_snapshots.extend(
250
+ collect_directory_snapshots(
251
+ effective_tracked_directories,
252
+ base_dir=effective_tracking_base_dir,
253
+ max_files_per_directory=effective_max_files,
254
+ )
255
+ )
256
+ if normalized_snapshots:
257
+ ts = _utcnow().strftime("%Y%m%dT%H%M%S%f")
258
+ for idx, snapshot in enumerate(normalized_snapshots):
259
+ self._db.write(
260
+ key=f"memory/episodic/claude/{self._agent_id}/{session_id}/directory_tracking/{ts}-{idx:04d}",
261
+ value={
262
+ "session_id": session_id,
263
+ "message_id": message_id,
264
+ "directory_snapshot": snapshot,
265
+ "captured_at": _utcnow().isoformat(),
266
+ },
267
+ origin="agent_inferred",
268
+ agent_id=self._agent_id,
269
+ session_id=session_id,
270
+ )
271
+
272
+ if outcome_type is not None and outcome_value is not None:
273
+ self._db.record_outcome(
274
+ session_id=session_id,
275
+ agent_id=self._agent_id,
276
+ outcome_type=outcome_type,
277
+ outcome_value=outcome_value,
278
+ )
279
+
280
+ def run_turn(
281
+ self,
282
+ *,
283
+ user_query: str,
284
+ session_id: str,
285
+ call_model: Callable[[list[dict[str, str]]], str],
286
+ system_prompt: Optional[str] = None,
287
+ external_context: Optional[list[dict[str, Any]]] = None,
288
+ base_messages: Optional[list[dict[str, str]]] = None,
289
+ memory_writes: Optional[list[dict[str, Any]]] = None,
290
+ message_id: Optional[str] = None,
291
+ tracked_directories: Optional[list[str]] = None,
292
+ directory_snapshots: Optional[list[dict[str, Any]]] = None,
293
+ tracking_base_dir: Optional[str] = None,
294
+ max_files_per_directory: Optional[int] = None,
295
+ outcome_type: Optional[str] = None,
296
+ outcome_value: Optional[float] = None,
297
+ ) -> str:
298
+ """
299
+ End-to-end enforced turn wrapper:
300
+ retrieve -> call_model -> persist -> record_outcome.
301
+ """
302
+ messages, _ = self.prepare_turn(
303
+ user_query=user_query,
304
+ session_id=session_id,
305
+ system_prompt=system_prompt,
306
+ external_context=external_context,
307
+ base_messages=base_messages,
308
+ )
309
+ response_text = call_model(messages)
310
+ self.finalize_turn(
311
+ session_id=session_id,
312
+ model_response=response_text,
313
+ memory_writes=memory_writes,
314
+ message_id=message_id,
315
+ tracked_directories=tracked_directories,
316
+ directory_snapshots=directory_snapshots,
317
+ tracking_base_dir=tracking_base_dir,
318
+ max_files_per_directory=max_files_per_directory,
319
+ outcome_type=outcome_type,
320
+ outcome_value=outcome_value,
321
+ )
322
+ return response_text
323
+
324
+ async def prepare_turn_async(
325
+ self,
326
+ *,
327
+ user_query: str,
328
+ session_id: str,
329
+ system_prompt: Optional[str] = None,
330
+ external_context: Optional[list[dict[str, Any]]] = None,
331
+ base_messages: Optional[list[dict[str, str]]] = None,
332
+ working_dirs: Optional[list[str]] = None,
333
+ ) -> tuple[list[dict[str, str]], "RetrievalResult"]:
334
+ return self.prepare_turn(
335
+ user_query=user_query,
336
+ session_id=session_id,
337
+ system_prompt=system_prompt,
338
+ external_context=external_context,
339
+ base_messages=base_messages,
340
+ working_dirs=working_dirs,
341
+ )
342
+
343
+ async def finalize_turn_async(
344
+ self,
345
+ *,
346
+ session_id: str,
347
+ model_response: str,
348
+ memory_writes: Optional[list[dict[str, Any]]] = None,
349
+ message_id: Optional[str] = None,
350
+ tracked_directories: Optional[list[str]] = None,
351
+ directory_snapshots: Optional[list[dict[str, Any]]] = None,
352
+ tracking_base_dir: Optional[str] = None,
353
+ max_files_per_directory: Optional[int] = None,
354
+ outcome_type: Optional[str] = None,
355
+ outcome_value: Optional[float] = None,
356
+ touched_entities: Optional[list[str]] = None,
357
+ user_query: Optional[str] = None,
358
+ ) -> None:
359
+ self.finalize_turn(
360
+ session_id=session_id,
361
+ model_response=model_response,
362
+ memory_writes=memory_writes,
363
+ message_id=message_id,
364
+ tracked_directories=tracked_directories,
365
+ directory_snapshots=directory_snapshots,
366
+ tracking_base_dir=tracking_base_dir,
367
+ max_files_per_directory=max_files_per_directory,
368
+ outcome_type=outcome_type,
369
+ outcome_value=outcome_value,
370
+ touched_entities=touched_entities,
371
+ user_query=user_query,
372
+ )
373
+
374
+ async def run_turn_async(
375
+ self,
376
+ *,
377
+ user_query: str,
378
+ session_id: str,
379
+ call_model_async: Callable[[list[dict[str, str]]], Awaitable[str]],
380
+ system_prompt: Optional[str] = None,
381
+ external_context: Optional[list[dict[str, Any]]] = None,
382
+ base_messages: Optional[list[dict[str, str]]] = None,
383
+ memory_writes: Optional[list[dict[str, Any]]] = None,
384
+ message_id: Optional[str] = None,
385
+ tracked_directories: Optional[list[str]] = None,
386
+ directory_snapshots: Optional[list[dict[str, Any]]] = None,
387
+ tracking_base_dir: Optional[str] = None,
388
+ max_files_per_directory: Optional[int] = None,
389
+ outcome_type: Optional[str] = None,
390
+ outcome_value: Optional[float] = None,
391
+ ) -> str:
392
+ messages, _ = await self.prepare_turn_async(
393
+ user_query=user_query,
394
+ session_id=session_id,
395
+ system_prompt=system_prompt,
396
+ external_context=external_context,
397
+ base_messages=base_messages,
398
+ )
399
+ response_text = await call_model_async(messages)
400
+ await self.finalize_turn_async(
401
+ session_id=session_id,
402
+ model_response=response_text,
403
+ memory_writes=memory_writes,
404
+ message_id=message_id,
405
+ tracked_directories=tracked_directories,
406
+ directory_snapshots=directory_snapshots,
407
+ tracking_base_dir=tracking_base_dir,
408
+ max_files_per_directory=max_files_per_directory,
409
+ outcome_type=outcome_type,
410
+ outcome_value=outcome_value,
411
+ )
412
+ return response_text
413
+
414
+ def _on_conflict(self, **payload) -> None:
415
+ """
416
+ Platform-agnostic conflict handler.
417
+ Concrete notification wiring is platform-specific (Slack, Telegram, etc.).
418
+ """
419
+ pass
420
+
421
+ def pending_conflicts(self, resolved: bool = False) -> list:
422
+ """Read-only view into db.gateway's conflict queue."""
423
+ return self.gateway.conflict_queue(resolved=resolved)
424
+
425
+ def propose_skill(
426
+ self,
427
+ name: str,
428
+ content: str,
429
+ scope: list[str],
430
+ entities: Optional[list[str]] = None,
431
+ session_id: Optional[str] = None,
432
+ ) -> str:
433
+ """Draft a skill proposal via db.gateway.propose(), scoped to this adapter's agent_id."""
434
+ return self.gateway.propose(
435
+ name=name,
436
+ content=content,
437
+ scope=scope,
438
+ agent_id=self._agent_id,
439
+ entities=entities,
440
+ session_id=session_id,
441
+ )
442
+
443
+ def approve_proposal(self, proposal_id: str, resolution: Optional[str] = None) -> None:
444
+ """Approve a pending proposal via db.gateway.approve()."""
445
+ self.gateway.approve(proposal_id, resolution=resolution)
446
+
447
+ def reject_proposal(self, proposal_id: str, reason: Optional[str] = None) -> None:
448
+ """Reject a pending proposal via db.gateway.reject()."""
449
+ self.gateway.reject(proposal_id, reason=reason)
450
+
451
+ def list_proposals(self, status: str = "pending") -> list[dict]:
452
+ """List proposals by status via db.gateway.list_proposals()."""
453
+ return self.gateway.list_proposals(status=status)
454
+
455
+ def _response_key(self, session_id: str) -> str:
456
+ ts = _utcnow().strftime("%Y%m%dT%H%M%S%f")
457
+ return f"memory/episodic/claude/{self._agent_id}/{session_id}/{ts}"
458
+
459
+ # Keys matching these substrings are stored for reconciliation but
460
+ # excluded from context injection — they're operational metadata,
461
+ # not useful memories for the model.
462
+ _INJECTION_EXCLUDED_KEY_PATTERNS = ("/directory_tracking/",)
463
+
464
+ # Episodic response records can be very large; cap them so one turn's
465
+ # response doesn't exhaust the whole memory budget.
466
+ _MAX_EPISODIC_VALUE_CHARS = 2000
467
+
468
+ def _build_memory_context(self, retrieval: "RetrievalResult") -> str:
469
+ if not retrieval.records:
470
+ return ""
471
+ remaining = self._resolved["injection"]["memory_budget_chars"]
472
+ line_format = self._resolved["injection"]["record_line_format"]
473
+ chunks: list[str] = []
474
+ for record in retrieval.records:
475
+ if any(pat in record.key for pat in self._INJECTION_EXCLUDED_KEY_PATTERNS):
476
+ continue
477
+ value = record.value
478
+ if isinstance(value, (dict, list)):
479
+ value_text = json.dumps(value, ensure_ascii=True)
480
+ else:
481
+ value_text = str(value)
482
+ if len(value_text) > self._MAX_EPISODIC_VALUE_CHARS:
483
+ value_text = value_text[: self._MAX_EPISODIC_VALUE_CHARS] + "…"
484
+ line = line_format.format(key=record.key, value=value_text)
485
+ if len(line) > remaining:
486
+ break
487
+ chunks.append(line)
488
+ remaining -= len(line)
489
+ return "\n".join(chunks)
490
+
491
+ def _extract_entities_from_tool(
492
+ self, tool_name: str, tool_input: dict[str, Any]
493
+ ) -> list[str]:
494
+ """
495
+ Extract entity strings from a tool call.
496
+
497
+ Edit -> file_path
498
+ Bash -> parse --select X flag or model names from command
499
+ """
500
+ entities: list[str] = []
501
+ if tool_name == "Edit":
502
+ file_path = tool_input.get("file_path")
503
+ if file_path:
504
+ entities.append(str(file_path))
505
+ elif tool_name == "Bash":
506
+ command = tool_input.get("command", "")
507
+ # Parse --select <value> (dbt-style)
508
+ import re
509
+ select_match = re.search(r"--select\s+(\S+)", command)
510
+ if select_match:
511
+ entities.append(select_match.group(1))
512
+ # Also look for model names in the command (words that look like identifiers)
513
+ # Extract words that are likely model/file references (not flags)
514
+ words = command.split()
515
+ for word in words:
516
+ if word.startswith("-"):
517
+ continue
518
+ # Skip common shell commands and dbt subcommands
519
+ if word in ("dbt", "build", "run", "test", "compile", "echo", "ls",
520
+ "cat", "grep", "python", "pip", "cd", "export"):
521
+ continue
522
+ # Accept words that look like identifiers (letters, digits, underscores, dots, slashes)
523
+ if re.match(r"^[a-zA-Z][a-zA-Z0-9_./]*$", word):
524
+ entities.append(word)
525
+ return entities
526
+
527
+ def build_action_context(
528
+ self, tool_name: str, tool_input: dict[str, Any], session_id: str
529
+ ) -> str:
530
+ """
531
+ Build governed action context for a tool call.
532
+
533
+ Returns a formatted context string if matching memories are found,
534
+ otherwise returns empty string.
535
+ """
536
+ entities = self._extract_entities_from_tool(tool_name, tool_input)
537
+ if not entities:
538
+ return ""
539
+
540
+ retrieval = self._db.retrieve_by_entities(
541
+ entities=entities,
542
+ agent_id=self._agent_id,
543
+ session_id=session_id,
544
+ )
545
+
546
+ if not retrieval.records:
547
+ return ""
548
+
549
+ lines = ["Governed action context:"]
550
+ for record in retrieval.records:
551
+ value = record.value
552
+ if isinstance(value, (dict, list)):
553
+ value_text = json.dumps(value, ensure_ascii=True)
554
+ else:
555
+ value_text = str(value)
556
+ lines.append(f"- {record.key}: {value_text}")
557
+
558
+ return "\n".join(lines)
559
+
560
+ def pre_tool_use(
561
+ self, tool_name: str, tool_input: dict[str, Any], session_id: str
562
+ ) -> dict[str, Any]:
563
+ """
564
+ Pre-tool-use hook: extract context and touched entities.
565
+
566
+ Returns:
567
+ {"additionalContext": str, "touchedEntities": list[str]}
568
+ or {} if no matching context.
569
+ """
570
+ entities = self._extract_entities_from_tool(tool_name, tool_input)
571
+ if not entities:
572
+ return {}
573
+
574
+ retrieval = self._db.retrieve_by_entities(
575
+ entities=entities,
576
+ agent_id=self._agent_id,
577
+ session_id=session_id,
578
+ )
579
+
580
+ if not retrieval.records:
581
+ return {}
582
+
583
+ lines = ["Governed action context:"]
584
+ touched: list[str] = []
585
+ for record in retrieval.records:
586
+ value = record.value
587
+ if isinstance(value, (dict, list)):
588
+ value_text = json.dumps(value, ensure_ascii=True)
589
+ else:
590
+ value_text = str(value)
591
+ lines.append(f"- {record.key}: {value_text}")
592
+ touched.extend(record.entities or [])
593
+
594
+ # Also include the original entities we searched for
595
+ all_touched = list(dict.fromkeys(entities + touched))
596
+
597
+ return {
598
+ "additionalContext": "\n".join(lines),
599
+ "touchedEntities": all_touched,
600
+ }
601
+
602
+ def build_hook_handlers(self) -> dict[str, Any]:
603
+ """
604
+ Build a dict of hook handlers for use with Claude SDK.
605
+
606
+ Returns:
607
+ {"PreToolUse": callable(tool_name, tool_input, session_id) -> dict}
608
+ """
609
+ return {
610
+ "PreToolUse": lambda tool_name, tool_input, session_id: self.pre_tool_use(
611
+ tool_name=tool_name,
612
+ tool_input=tool_input,
613
+ session_id=session_id,
614
+ )
615
+ }
616
+
617
+
618
+ class AgentDBClaudeRuntime:
619
+ """
620
+ Higher-level runtime wrapper for easy integration.
621
+
622
+ Developers can instantiate once and call `handle_turn(...)` from their
623
+ Slack/event loop without manually orchestrating middleware hooks.
624
+ """
625
+
626
+ def __init__(
627
+ self,
628
+ *,
629
+ middleware: AgentDBClaudeMiddleware,
630
+ invoke_model: Callable[[list[dict[str, str]]], Any],
631
+ default_system_prompt: Optional[str] = None,
632
+ response_to_text: Callable[[Any], str] = _extract_text_from_model_response,
633
+ ):
634
+ self._middleware = middleware
635
+ self._invoke_model = invoke_model
636
+ self._default_system_prompt = default_system_prompt
637
+ self._response_to_text = response_to_text
638
+ self._sdk_client_factory = None
639
+
640
+ @classmethod
641
+ def from_claude_client(
642
+ cls,
643
+ *,
644
+ db: "AgentDB",
645
+ agent_id: str,
646
+ client: Any,
647
+ model: str,
648
+ system_prompt: Optional[str] = None,
649
+ max_tokens: int = 1024,
650
+ ) -> "AgentDBClaudeRuntime":
651
+ middleware = AgentDBClaudeMiddleware(db=db, agent_id=agent_id)
652
+
653
+ def _invoke(messages: list[dict[str, str]]) -> Any:
654
+ system_blocks = [m["content"] for m in messages if m.get("role") == "system"]
655
+ chat_messages = [m for m in messages if m.get("role") != "system"]
656
+ if hasattr(client, "messages") and hasattr(client.messages, "create"):
657
+ payload: dict[str, Any] = {
658
+ "model": model,
659
+ "max_tokens": max_tokens,
660
+ "messages": chat_messages,
661
+ }
662
+ if system_blocks:
663
+ payload["system"] = "\n\n".join(system_blocks)
664
+ return client.messages.create(**payload)
665
+ if hasattr(client, "run"):
666
+ return client.run(model=model, messages=messages, max_tokens=max_tokens)
667
+ raise ValueError(
668
+ "Unsupported client: expected `messages.create(...)` or `run(...)` API."
669
+ )
670
+
671
+ return cls(
672
+ middleware=middleware,
673
+ invoke_model=_invoke,
674
+ default_system_prompt=system_prompt,
675
+ response_to_text=_extract_text_from_model_response,
676
+ )
677
+
678
+ @classmethod
679
+ def from_sdk_client(
680
+ cls,
681
+ *,
682
+ db: "AgentDB",
683
+ agent_id: str,
684
+ sdk_client_factory: Callable[[str], Any],
685
+ system_prompt: Optional[str] = None,
686
+ ) -> "AgentDBClaudeRuntime":
687
+ """
688
+ Build a runtime that wraps a ClaudeSDKClient (agentic subprocess).
689
+
690
+ sdk_client_factory: callable(thread_id: str) -> client
691
+ The client must support: await client.query(prompt) and
692
+ async-for client.receive_response() yielding messages with
693
+ .content (list of blocks with .text) or .is_error/.result.
694
+ """
695
+ middleware = AgentDBClaudeMiddleware(db=db, agent_id=agent_id)
696
+ instance = cls(
697
+ middleware=middleware,
698
+ invoke_model=lambda messages: "",
699
+ default_system_prompt=system_prompt,
700
+ )
701
+ instance._sdk_client_factory = sdk_client_factory
702
+ return instance
703
+
704
+ async def handle_turn_async(
705
+ self,
706
+ *,
707
+ user_query: str,
708
+ session_id: str,
709
+ message_id: Optional[str] = None,
710
+ system_prompt: Optional[str] = None,
711
+ external_context: Optional[list[dict[str, Any]]] = None,
712
+ base_messages: Optional[list[dict[str, str]]] = None,
713
+ memory_writes: Optional[list[dict[str, Any]]] = None,
714
+ tracked_directories: Optional[list[str]] = None,
715
+ directory_snapshots: Optional[list[dict[str, Any]]] = None,
716
+ tracking_base_dir: Optional[str] = None,
717
+ max_files_per_directory: Optional[int] = None,
718
+ outcome_type: Optional[str] = None,
719
+ outcome_value: Optional[float] = None,
720
+ ) -> dict[str, Any]:
721
+ sdk_factory = getattr(self, "_sdk_client_factory", None)
722
+ if sdk_factory is None:
723
+ raise RuntimeError(
724
+ "handle_turn_async requires a runtime built with from_sdk_client(). "
725
+ "Use handle_turn() for the messages.create() path."
726
+ )
727
+
728
+ messages, retrieval = await self._middleware.prepare_turn_async(
729
+ user_query=user_query,
730
+ session_id=session_id,
731
+ system_prompt=system_prompt if system_prompt is not None else self._default_system_prompt,
732
+ external_context=external_context,
733
+ base_messages=base_messages,
734
+ )
735
+
736
+ system_blocks = [m["content"] for m in messages if m.get("role") == "system"]
737
+ user_messages = [m for m in messages if m.get("role") == "user"]
738
+ prompt_parts = system_blocks + [m["content"] for m in user_messages]
739
+ prompt = "\n\n".join(prompt_parts)
740
+
741
+ client = sdk_factory(session_id)
742
+ await client.query(prompt)
743
+
744
+ text_parts: list[str] = []
745
+ async for message in client.receive_response():
746
+ is_error = getattr(message, "is_error", False)
747
+ if is_error:
748
+ error_msg = getattr(message, "result", "Claude SDK error")
749
+ raise RuntimeError(error_msg or "Claude SDK returned an error")
750
+ content = getattr(message, "content", None)
751
+ if content is not None:
752
+ for block in content:
753
+ text = getattr(block, "text", None)
754
+ if text is not None:
755
+ text_parts.append(text)
756
+
757
+ response_text = "\n".join(text_parts).strip()
758
+
759
+ await self._middleware.finalize_turn_async(
760
+ session_id=session_id,
761
+ model_response=response_text,
762
+ memory_writes=memory_writes,
763
+ message_id=message_id,
764
+ tracked_directories=tracked_directories,
765
+ directory_snapshots=directory_snapshots,
766
+ tracking_base_dir=tracking_base_dir,
767
+ max_files_per_directory=max_files_per_directory,
768
+ outcome_type=outcome_type,
769
+ outcome_value=outcome_value,
770
+ )
771
+
772
+ return {
773
+ "response_text": response_text,
774
+ "raw_response": None,
775
+ "retrieval": retrieval,
776
+ "messages": messages,
777
+ }
778
+
779
+ def handle_turn(
780
+ self,
781
+ *,
782
+ user_query: str,
783
+ session_id: str,
784
+ message_id: Optional[str] = None,
785
+ system_prompt: Optional[str] = None,
786
+ external_context: Optional[list[dict[str, Any]]] = None,
787
+ base_messages: Optional[list[dict[str, str]]] = None,
788
+ memory_writes: Optional[list[dict[str, Any]]] = None,
789
+ tracked_directories: Optional[list[str]] = None,
790
+ directory_snapshots: Optional[list[dict[str, Any]]] = None,
791
+ tracking_base_dir: Optional[str] = None,
792
+ max_files_per_directory: Optional[int] = None,
793
+ outcome_type: Optional[str] = None,
794
+ outcome_value: Optional[float] = None,
795
+ ) -> dict[str, Any]:
796
+ messages, retrieval = self._middleware.prepare_turn(
797
+ user_query=user_query,
798
+ session_id=session_id,
799
+ system_prompt=system_prompt if system_prompt is not None else self._default_system_prompt,
800
+ external_context=external_context,
801
+ base_messages=base_messages,
802
+ )
803
+ raw_response = self._invoke_model(messages)
804
+ response_text = self._response_to_text(raw_response)
805
+ self._middleware.finalize_turn(
806
+ session_id=session_id,
807
+ model_response=response_text,
808
+ memory_writes=memory_writes,
809
+ message_id=message_id,
810
+ tracked_directories=tracked_directories,
811
+ directory_snapshots=directory_snapshots,
812
+ tracking_base_dir=tracking_base_dir,
813
+ max_files_per_directory=max_files_per_directory,
814
+ outcome_type=outcome_type,
815
+ outcome_value=outcome_value,
816
+ )
817
+ return {
818
+ "response_text": response_text,
819
+ "raw_response": raw_response,
820
+ "retrieval": retrieval,
821
+ "messages": messages,
822
+ }
823
+
824
+ def get_sdk_hooks(self) -> dict[str, Any]:
825
+ """
826
+ Return hook handlers from the middleware for Claude SDK integration.
827
+
828
+ Returns:
829
+ {"PreToolUse": callable}
830
+ """
831
+ return self._middleware.build_hook_handlers()