codebind 0.6.1__py3-none-any.whl → 0.6.3__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.
- codebind/conversation.py +166 -15
- codebind/execution.py +56 -11
- codebind/extension.py +102 -57
- codebind/images.py +127 -0
- codebind/instructions.md +9 -5
- codebind/jupyter.py +16 -45
- codebind/session.py +160 -96
- {codebind-0.6.1.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/package.json +2 -2
- codebind-0.6.3.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/590.dae4561c5427a855.js +1 -0
- codebind-0.6.3.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/remoteEntry.76c22319f766843f.js +1 -0
- {codebind-0.6.1.dist-info → codebind-0.6.3.dist-info}/METADATA +8 -4
- codebind-0.6.3.dist-info/RECORD +23 -0
- {codebind-0.6.1.dist-info → codebind-0.6.3.dist-info}/WHEEL +1 -1
- codebind-0.6.1.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/590.5b773af7a27693b3.js +0 -1
- codebind-0.6.1.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/remoteEntry.b171b12ff4a78939.js +0 -1
- codebind-0.6.1.dist-info/RECORD +0 -22
- {codebind-0.6.1.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/install.json +0 -0
- {codebind-0.6.1.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/static/style.js +0 -0
- {codebind-0.6.1.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/static/third-party-licenses.json +0 -0
- {codebind-0.6.1.dist-info → codebind-0.6.3.dist-info}/entry_points.txt +0 -0
- {codebind-0.6.1.dist-info → codebind-0.6.3.dist-info}/licenses/LICENSE +0 -0
codebind/conversation.py
CHANGED
|
@@ -9,9 +9,16 @@ from dataclasses import dataclass, field
|
|
|
9
9
|
from typing import Any, Protocol
|
|
10
10
|
from uuid import uuid4
|
|
11
11
|
|
|
12
|
-
from langchain_core.messages import
|
|
12
|
+
from langchain_core.messages import (
|
|
13
|
+
AIMessage,
|
|
14
|
+
BaseMessage,
|
|
15
|
+
HumanMessage,
|
|
16
|
+
ToolMessage,
|
|
17
|
+
message_to_dict,
|
|
18
|
+
messages_from_dict,
|
|
19
|
+
)
|
|
13
20
|
|
|
14
|
-
|
|
21
|
+
from .images import PreparedImage, prepare_mime_image
|
|
15
22
|
|
|
16
23
|
|
|
17
24
|
class ConversationStore(Protocol):
|
|
@@ -35,28 +42,39 @@ class Conversation:
|
|
|
35
42
|
|
|
36
43
|
@classmethod
|
|
37
44
|
def from_dict(cls, value: dict[str, Any]) -> Conversation:
|
|
38
|
-
if value
|
|
39
|
-
raise ValueError("
|
|
45
|
+
if not isinstance(value, dict):
|
|
46
|
+
raise ValueError("Invalid Codebind conversation")
|
|
40
47
|
identifier = value.get("id")
|
|
41
48
|
events = value.get("events")
|
|
42
49
|
revision = value.get("revision")
|
|
43
50
|
cwd = value.get("cwd")
|
|
44
|
-
if not isinstance(identifier, str) or not isinstance(events, list):
|
|
51
|
+
if not isinstance(identifier, str) or not identifier or not isinstance(events, list):
|
|
45
52
|
raise ValueError("Invalid Codebind conversation")
|
|
46
|
-
if
|
|
53
|
+
if type(revision) is not int or revision != len(events):
|
|
47
54
|
raise ValueError("Invalid Codebind conversation revision")
|
|
48
55
|
if any(
|
|
49
|
-
not isinstance(event, dict)
|
|
56
|
+
not isinstance(event, dict)
|
|
57
|
+
or type(event.get("sequence")) is not int
|
|
58
|
+
or event["sequence"] != sequence
|
|
59
|
+
or event.get("type") not in {"message", "notebook", "turn_started", "turn_finished"}
|
|
60
|
+
or not isinstance(event.get("turn_id"), str)
|
|
61
|
+
or (
|
|
62
|
+
event.get("type") == "turn_finished"
|
|
63
|
+
and event.get("status") not in {"completed", "cancelled", "failed", "interrupted"}
|
|
64
|
+
)
|
|
50
65
|
for sequence, event in enumerate(events, start=1)
|
|
51
66
|
):
|
|
52
|
-
raise ValueError("Invalid Codebind conversation
|
|
67
|
+
raise ValueError("Invalid Codebind conversation event")
|
|
53
68
|
if not isinstance(cwd, str):
|
|
54
|
-
|
|
55
|
-
|
|
69
|
+
raise ValueError("Invalid Codebind conversation directory")
|
|
70
|
+
conversation = cls(identifier, revision, deepcopy(events), cwd)
|
|
71
|
+
conversation._validate_tool_history()
|
|
72
|
+
conversation.notebook()
|
|
73
|
+
conversation.messages()
|
|
74
|
+
return conversation
|
|
56
75
|
|
|
57
76
|
def as_dict(self) -> dict[str, Any]:
|
|
58
77
|
return {
|
|
59
|
-
"version": _VERSION,
|
|
60
78
|
"id": self.id,
|
|
61
79
|
"revision": self.revision,
|
|
62
80
|
"cwd": self.cwd,
|
|
@@ -64,21 +82,83 @@ class Conversation:
|
|
|
64
82
|
}
|
|
65
83
|
|
|
66
84
|
def append(self, event_type: str, **data: Any) -> dict[str, Any]:
|
|
85
|
+
if event_type == "turn_finished" and self.pending_tool_calls({data.get("turn_id")}):
|
|
86
|
+
raise ValueError("Cannot finish a turn with unanswered tool calls")
|
|
67
87
|
event = {"sequence": self.revision + 1, "type": event_type, **data}
|
|
68
88
|
self.events.append(event)
|
|
69
89
|
self.revision += 1
|
|
70
90
|
return event
|
|
71
91
|
|
|
72
92
|
def append_message(self, message: BaseMessage, turn_id: str) -> dict[str, Any]:
|
|
93
|
+
if isinstance(message, ToolMessage):
|
|
94
|
+
if message.tool_call_id not in dict(self.pending_tool_calls({turn_id})):
|
|
95
|
+
raise ValueError("Tool result has no matching unanswered call")
|
|
96
|
+
elif isinstance(message, AIMessage):
|
|
97
|
+
identifiers = [call.get("id") for call in message.tool_calls]
|
|
98
|
+
if any(not isinstance(value, str) or not value for value in identifiers):
|
|
99
|
+
raise ValueError("Assistant tool calls require IDs")
|
|
100
|
+
if len(identifiers) != len(set(identifiers)):
|
|
101
|
+
raise ValueError("Assistant tool call IDs must be unique")
|
|
73
102
|
return self.append("message", turn_id=turn_id, message=message_to_dict(message))
|
|
74
103
|
|
|
104
|
+
def _validate_tool_history(self) -> None:
|
|
105
|
+
pending: dict[str, str] = {}
|
|
106
|
+
seen: set[str] = set()
|
|
107
|
+
for event in self.events:
|
|
108
|
+
turn_id = event.get("turn_id")
|
|
109
|
+
if event.get("type") == "turn_finished":
|
|
110
|
+
if turn_id in pending.values():
|
|
111
|
+
raise ValueError("Finished turn contains unanswered tool calls")
|
|
112
|
+
continue
|
|
113
|
+
if event.get("type") != "message":
|
|
114
|
+
continue
|
|
115
|
+
message = event.get("message")
|
|
116
|
+
if not isinstance(message, dict) or not isinstance(turn_id, str):
|
|
117
|
+
raise ValueError("Invalid Codebind message event")
|
|
118
|
+
data = message.get("data")
|
|
119
|
+
if not isinstance(data, dict):
|
|
120
|
+
raise ValueError("Invalid Codebind message data")
|
|
121
|
+
if message.get("type") == "ai":
|
|
122
|
+
calls = data.get("tool_calls") or []
|
|
123
|
+
if not isinstance(calls, list):
|
|
124
|
+
raise ValueError("Invalid assistant tool calls")
|
|
125
|
+
additional = data.get("additional_kwargs") or {}
|
|
126
|
+
if not isinstance(additional, dict):
|
|
127
|
+
raise ValueError("Invalid assistant response metadata")
|
|
128
|
+
response_items = additional.get("response_items")
|
|
129
|
+
if isinstance(response_items, list):
|
|
130
|
+
raw_calls = [
|
|
131
|
+
item.get("call_id")
|
|
132
|
+
for item in response_items
|
|
133
|
+
if isinstance(item, dict) and item.get("type") == "function_call"
|
|
134
|
+
]
|
|
135
|
+
parsed_calls = [
|
|
136
|
+
call.get("id") if isinstance(call, dict) else None for call in calls
|
|
137
|
+
]
|
|
138
|
+
if raw_calls != parsed_calls:
|
|
139
|
+
raise ValueError(
|
|
140
|
+
"Assistant tool calls disagree with provider response items"
|
|
141
|
+
)
|
|
142
|
+
for call in calls:
|
|
143
|
+
identifier = call.get("id") if isinstance(call, dict) else None
|
|
144
|
+
if not isinstance(identifier, str) or not identifier or identifier in seen:
|
|
145
|
+
raise ValueError("Invalid or repeated assistant tool call ID")
|
|
146
|
+
seen.add(identifier)
|
|
147
|
+
pending[identifier] = turn_id
|
|
148
|
+
elif message.get("type") == "tool":
|
|
149
|
+
identifier = data.get("tool_call_id")
|
|
150
|
+
if not isinstance(identifier, str) or pending.get(identifier) != turn_id:
|
|
151
|
+
raise ValueError("Tool result has no matching unanswered call")
|
|
152
|
+
del pending[identifier]
|
|
153
|
+
|
|
75
154
|
def append_notebook(
|
|
76
155
|
self,
|
|
77
156
|
cells: list[dict[str, Any]],
|
|
78
157
|
turn_id: str,
|
|
79
158
|
) -> dict[str, Any] | None:
|
|
80
159
|
current = self.notebook()
|
|
81
|
-
|
|
160
|
+
images: list[tuple[str, PreparedImage]] = []
|
|
161
|
+
incoming = _normalize_notebook(cells, images)
|
|
82
162
|
has_notebook = any(event.get("type") == "notebook" for event in self.events)
|
|
83
163
|
|
|
84
164
|
if not has_notebook:
|
|
@@ -109,11 +189,19 @@ class Conversation:
|
|
|
109
189
|
)
|
|
110
190
|
+ "</notebook-context>"
|
|
111
191
|
)
|
|
192
|
+
previous_images = _image_digests(current)
|
|
193
|
+
blocks: list[dict[str, str]] = [{"type": "text", "text": content}]
|
|
194
|
+
for label, image in images:
|
|
195
|
+
if previous_images.get(label) == image.source_sha256:
|
|
196
|
+
continue
|
|
197
|
+
blocks.append({"type": "text", "text": f"Displayed image from {label}:"})
|
|
198
|
+
blocks.append(image.content_block())
|
|
199
|
+
message = HumanMessage(content_blocks=blocks) if len(blocks) > 1 else HumanMessage(content)
|
|
112
200
|
return self.append(
|
|
113
201
|
"notebook",
|
|
114
202
|
turn_id=turn_id,
|
|
115
203
|
**payload,
|
|
116
|
-
message=message_to_dict(
|
|
204
|
+
message=message_to_dict(message),
|
|
117
205
|
)
|
|
118
206
|
|
|
119
207
|
def rollback(self, event: dict[str, Any]) -> None:
|
|
@@ -226,7 +314,43 @@ class NotebookConversationStore:
|
|
|
226
314
|
await self.bridge.save_conversation(value)
|
|
227
315
|
|
|
228
316
|
|
|
229
|
-
def
|
|
317
|
+
def _canonical_mime_bundle(
|
|
318
|
+
value: dict[str, Any],
|
|
319
|
+
label: str,
|
|
320
|
+
images: list[tuple[str, PreparedImage]] | None,
|
|
321
|
+
) -> dict[str, Any]:
|
|
322
|
+
data = {key: deepcopy(item) for key, item in value.items() if not key.startswith("image/")}
|
|
323
|
+
try:
|
|
324
|
+
image = prepare_mime_image(value)
|
|
325
|
+
except ValueError as error:
|
|
326
|
+
data["image_omitted"] = str(error)
|
|
327
|
+
else:
|
|
328
|
+
if image is not None:
|
|
329
|
+
data["image"] = image.descriptor()
|
|
330
|
+
if images is not None:
|
|
331
|
+
images.append((label, image))
|
|
332
|
+
return data
|
|
333
|
+
|
|
334
|
+
|
|
335
|
+
def _image_digests(cells: list[dict[str, Any]]) -> dict[str, str]:
|
|
336
|
+
digests: dict[str, str] = {}
|
|
337
|
+
for cell in cells:
|
|
338
|
+
identifier = cell["id"]
|
|
339
|
+
for index, output in enumerate(cell.get("outputs", [])):
|
|
340
|
+
descriptor = (output.get("data") or {}).get("image")
|
|
341
|
+
if isinstance(descriptor, dict) and isinstance(descriptor.get("source_sha256"), str):
|
|
342
|
+
digests[f"cell {identifier} output {index + 1}"] = descriptor["source_sha256"]
|
|
343
|
+
for name, bundle in cell.get("attachments", {}).items():
|
|
344
|
+
descriptor = bundle.get("image")
|
|
345
|
+
if isinstance(descriptor, dict) and isinstance(descriptor.get("source_sha256"), str):
|
|
346
|
+
digests[f"cell {identifier} attachment {name}"] = descriptor["source_sha256"]
|
|
347
|
+
return digests
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
def _normalize_notebook(
|
|
351
|
+
value: Any,
|
|
352
|
+
images: list[tuple[str, PreparedImage]] | None = None,
|
|
353
|
+
) -> list[dict[str, Any]]:
|
|
230
354
|
if not isinstance(value, list):
|
|
231
355
|
raise ValueError("Invalid Codebind notebook snapshot")
|
|
232
356
|
cells: list[dict[str, Any]] = []
|
|
@@ -255,7 +379,34 @@ def _normalize_notebook(value: Any) -> list[dict[str, Any]]:
|
|
|
255
379
|
):
|
|
256
380
|
raise ValueError("Invalid Codebind code cell")
|
|
257
381
|
cell["execution_count"] = execution_count
|
|
258
|
-
|
|
382
|
+
normalized_outputs: list[dict[str, Any]] = []
|
|
383
|
+
for index, output in enumerate(outputs):
|
|
384
|
+
if not isinstance(output, dict):
|
|
385
|
+
raise ValueError("Invalid Codebind cell output")
|
|
386
|
+
normalized = deepcopy(output)
|
|
387
|
+
data = output.get("data")
|
|
388
|
+
if isinstance(data, dict):
|
|
389
|
+
normalized["data"] = _canonical_mime_bundle(
|
|
390
|
+
data,
|
|
391
|
+
f"cell {identifier} output {index + 1}",
|
|
392
|
+
images,
|
|
393
|
+
)
|
|
394
|
+
normalized_outputs.append(normalized)
|
|
395
|
+
cell["outputs"] = normalized_outputs
|
|
396
|
+
attachments = raw_cell.get("attachments")
|
|
397
|
+
if attachments is not None:
|
|
398
|
+
if not isinstance(attachments, dict):
|
|
399
|
+
raise ValueError("Invalid Codebind cell attachments")
|
|
400
|
+
normalized_attachments: dict[str, Any] = {}
|
|
401
|
+
for name, bundle in attachments.items():
|
|
402
|
+
if not isinstance(name, str) or not isinstance(bundle, dict):
|
|
403
|
+
raise ValueError("Invalid Codebind cell attachment")
|
|
404
|
+
normalized_attachments[name] = _canonical_mime_bundle(
|
|
405
|
+
bundle,
|
|
406
|
+
f"cell {identifier} attachment {name}",
|
|
407
|
+
images,
|
|
408
|
+
)
|
|
409
|
+
cell["attachments"] = normalized_attachments
|
|
259
410
|
cells.append(cell)
|
|
260
411
|
return cells
|
|
261
412
|
|
codebind/execution.py
CHANGED
|
@@ -5,7 +5,7 @@ from __future__ import annotations
|
|
|
5
5
|
import sys
|
|
6
6
|
from collections.abc import Callable, Iterator
|
|
7
7
|
from contextlib import contextmanager
|
|
8
|
-
from dataclasses import
|
|
8
|
+
from dataclasses import dataclass
|
|
9
9
|
from typing import Any, cast
|
|
10
10
|
|
|
11
11
|
from IPython.core.displaypub import DisplayPublisher
|
|
@@ -13,9 +13,13 @@ from IPython.core.interactiveshell import InteractiveShell
|
|
|
13
13
|
from IPython.utils.capture import capture_output
|
|
14
14
|
|
|
15
15
|
from .display import display_cell
|
|
16
|
+
from .images import PreparedImage, prepare_mime_image
|
|
16
17
|
from .jupyter import JupyterLabBridge
|
|
17
18
|
|
|
18
19
|
|
|
20
|
+
_MAXIMUM_IMAGES_PER_REPORT = 32
|
|
21
|
+
|
|
22
|
+
|
|
19
23
|
@dataclass(frozen=True, slots=True)
|
|
20
24
|
class ExecutionReport:
|
|
21
25
|
"""The model-facing result of one IPython cell."""
|
|
@@ -26,10 +30,21 @@ class ExecutionReport:
|
|
|
26
30
|
result: str | None
|
|
27
31
|
displays: tuple[str, ...]
|
|
28
32
|
error: dict[str, str] | None
|
|
33
|
+
images: tuple[PreparedImage, ...] = ()
|
|
34
|
+
image_omissions: tuple[str, ...] = ()
|
|
29
35
|
|
|
30
36
|
def as_dict(self) -> dict[str, Any]:
|
|
31
|
-
"""Return
|
|
32
|
-
return
|
|
37
|
+
"""Return the text report and image descriptors without repeating image bytes."""
|
|
38
|
+
return {
|
|
39
|
+
"ok": self.ok,
|
|
40
|
+
"stdout": self.stdout,
|
|
41
|
+
"stderr": self.stderr,
|
|
42
|
+
"result": self.result,
|
|
43
|
+
"displays": self.displays,
|
|
44
|
+
"error": self.error,
|
|
45
|
+
"images": [image.descriptor() for image in self.images],
|
|
46
|
+
"image_omissions": self.image_omissions,
|
|
47
|
+
}
|
|
33
48
|
|
|
34
49
|
|
|
35
50
|
class _ExecutionDisplayHook:
|
|
@@ -156,6 +171,7 @@ def _captured_execution(
|
|
|
156
171
|
return
|
|
157
172
|
|
|
158
173
|
previous_displayhook = sys.displayhook
|
|
174
|
+
previous_trap_hook = shell.display_trap.hook
|
|
159
175
|
previous_showtraceback = shell.showtraceback
|
|
160
176
|
previous_showsyntaxerror = shell.showsyntaxerror
|
|
161
177
|
if on_output is not None:
|
|
@@ -169,12 +185,15 @@ def _captured_execution(
|
|
|
169
185
|
on_clear or (lambda wait: None),
|
|
170
186
|
),
|
|
171
187
|
)
|
|
172
|
-
|
|
188
|
+
execution_displayhook = _ExecutionDisplayHook(shell, expression_outputs, on_output)
|
|
189
|
+
shell.display_trap.hook = execution_displayhook
|
|
190
|
+
sys.displayhook = execution_displayhook
|
|
173
191
|
shell.showtraceback = lambda *args, **kwargs: None
|
|
174
192
|
shell.showsyntaxerror = lambda *args, **kwargs: None
|
|
175
193
|
try:
|
|
176
194
|
yield captured
|
|
177
195
|
finally:
|
|
196
|
+
shell.display_trap.hook = previous_trap_hook
|
|
178
197
|
sys.displayhook = previous_displayhook
|
|
179
198
|
shell.showtraceback = previous_showtraceback
|
|
180
199
|
shell.showsyntaxerror = previous_showsyntaxerror
|
|
@@ -221,7 +240,7 @@ class IPythonExecutor:
|
|
|
221
240
|
else:
|
|
222
241
|
captured.show()
|
|
223
242
|
|
|
224
|
-
return self._report(result, captured)
|
|
243
|
+
return self._report(result, captured, expression_outputs)
|
|
225
244
|
|
|
226
245
|
async def aexecute(self, cell: str) -> ExecutionReport:
|
|
227
246
|
"""Execute an async-capable cell and replay its native rich output."""
|
|
@@ -258,7 +277,7 @@ class IPythonExecutor:
|
|
|
258
277
|
else:
|
|
259
278
|
captured.show()
|
|
260
279
|
|
|
261
|
-
return self._report(result, captured)
|
|
280
|
+
return self._report(result, captured, expression_outputs)
|
|
262
281
|
|
|
263
282
|
def _stream_callbacks(
|
|
264
283
|
self,
|
|
@@ -279,15 +298,39 @@ class IPythonExecutor:
|
|
|
279
298
|
|
|
280
299
|
return on_output, on_clear
|
|
281
300
|
|
|
282
|
-
|
|
283
|
-
|
|
284
|
-
|
|
301
|
+
def _report(
|
|
302
|
+
self,
|
|
303
|
+
result: Any,
|
|
304
|
+
captured: Any,
|
|
305
|
+
expression_outputs: list[dict[str, Any]],
|
|
306
|
+
) -> ExecutionReport:
|
|
307
|
+
"""Build the model-facing text and image projection of an IPython execution."""
|
|
285
308
|
|
|
286
309
|
displays: list[str] = []
|
|
287
|
-
|
|
288
|
-
|
|
310
|
+
images: list[PreparedImage] = []
|
|
311
|
+
omissions: list[str] = []
|
|
312
|
+
bundles = [getattr(output, "data", None) for output in captured.outputs]
|
|
313
|
+
bundles.extend(output.get("data") for output in expression_outputs)
|
|
314
|
+
if not expression_outputs and result.result is not None:
|
|
315
|
+
data, _ = self.shell.display_formatter.format(result.result)
|
|
316
|
+
if any(key.startswith("image/") for key in data):
|
|
317
|
+
bundles.append(data)
|
|
318
|
+
for index, data in enumerate(bundles):
|
|
289
319
|
if isinstance(data, dict) and "text/plain" in data:
|
|
290
320
|
displays.append(str(data["text/plain"]))
|
|
321
|
+
if not isinstance(data, dict):
|
|
322
|
+
continue
|
|
323
|
+
if len(images) >= _MAXIMUM_IMAGES_PER_REPORT:
|
|
324
|
+
if any(key.startswith("image/") for key in data):
|
|
325
|
+
omissions.append(f"Display {index + 1}: image count limit reached")
|
|
326
|
+
continue
|
|
327
|
+
try:
|
|
328
|
+
image = prepare_mime_image(data)
|
|
329
|
+
except ValueError as error:
|
|
330
|
+
omissions.append(f"Display {index + 1}: {error}")
|
|
331
|
+
else:
|
|
332
|
+
if image is not None:
|
|
333
|
+
images.append(image)
|
|
291
334
|
|
|
292
335
|
exception = result.error_before_exec or result.error_in_exec
|
|
293
336
|
error = (
|
|
@@ -302,6 +345,8 @@ class IPythonExecutor:
|
|
|
302
345
|
result=repr(result.result) if result.result is not None else None,
|
|
303
346
|
displays=tuple(displays),
|
|
304
347
|
error=error,
|
|
348
|
+
images=tuple(images),
|
|
349
|
+
image_omissions=tuple(omissions),
|
|
305
350
|
)
|
|
306
351
|
|
|
307
352
|
def _notebook_outputs(
|
codebind/extension.py
CHANGED
|
@@ -10,23 +10,37 @@ from langchain_core.language_models import BaseChatModel
|
|
|
10
10
|
|
|
11
11
|
from .configuration import load_configuration, load_models
|
|
12
12
|
from .conversation import MemoryConversationStore, NotebookConversationStore
|
|
13
|
-
from .jupyter import JupyterLabBridge
|
|
13
|
+
from .jupyter import JupyterLabBridge, TARGET_NAME
|
|
14
14
|
from .session import Session
|
|
15
15
|
|
|
16
16
|
|
|
17
17
|
_STATE_ATTRIBUTE = "_codebind_extension_state"
|
|
18
18
|
|
|
19
19
|
|
|
20
|
+
def _close_state(state: dict[str, Any]) -> None:
|
|
21
|
+
manager = state.get("comm_manager")
|
|
22
|
+
target = state.get("comm_target")
|
|
23
|
+
if manager is not None and getattr(manager, "targets", {}).get(TARGET_NAME) is target:
|
|
24
|
+
manager.unregister_target(TARGET_NAME, target)
|
|
25
|
+
load_task = state.get("load_task")
|
|
26
|
+
if isinstance(load_task, asyncio.Task):
|
|
27
|
+
load_task.cancel()
|
|
28
|
+
session = state.get("session")
|
|
29
|
+
if isinstance(session, Session) and session.bridge is not None:
|
|
30
|
+
session.bridge.close()
|
|
31
|
+
discard_model = state.get("discard_model")
|
|
32
|
+
if callable(discard_model):
|
|
33
|
+
try:
|
|
34
|
+
asyncio.get_running_loop().create_task(discard_model())
|
|
35
|
+
except RuntimeError:
|
|
36
|
+
asyncio.run(discard_model())
|
|
37
|
+
|
|
38
|
+
|
|
20
39
|
def load_ipython_extension(ipython: InteractiveShell) -> None:
|
|
21
40
|
"""Load Codebind into the active IPython session."""
|
|
22
41
|
previous = getattr(ipython, _STATE_ATTRIBUTE, None)
|
|
23
42
|
if isinstance(previous, dict):
|
|
24
|
-
|
|
25
|
-
if isinstance(previous_load, asyncio.Task):
|
|
26
|
-
previous_load.cancel()
|
|
27
|
-
previous_session = previous.get("session")
|
|
28
|
-
if isinstance(previous_session, Session) and previous_session.bridge is not None:
|
|
29
|
-
previous_session.bridge.close()
|
|
43
|
+
_close_state(previous)
|
|
30
44
|
|
|
31
45
|
configuration = load_configuration()
|
|
32
46
|
models = load_models()
|
|
@@ -38,72 +52,103 @@ def load_ipython_extension(ipython: InteractiveShell) -> None:
|
|
|
38
52
|
model = models.chat(configuration.model, **configuration.parameters)
|
|
39
53
|
return model
|
|
40
54
|
|
|
41
|
-
def discard_model(
|
|
55
|
+
async def discard_model() -> None:
|
|
42
56
|
nonlocal model
|
|
43
|
-
|
|
44
|
-
|
|
45
|
-
|
|
46
|
-
|
|
47
|
-
|
|
48
|
-
|
|
49
|
-
|
|
50
|
-
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
|
|
56
|
-
|
|
57
|
+
selected = model
|
|
58
|
+
model = None
|
|
59
|
+
close = getattr(selected, "aclose", None)
|
|
60
|
+
if callable(close):
|
|
61
|
+
try:
|
|
62
|
+
await close()
|
|
63
|
+
except Exception:
|
|
64
|
+
pass
|
|
65
|
+
|
|
66
|
+
manager = getattr(getattr(ipython, "kernel", None), "comm_manager", None)
|
|
67
|
+
state: dict[str, Any] = {
|
|
68
|
+
"session": None,
|
|
69
|
+
"bridge": None,
|
|
70
|
+
"load_task": None,
|
|
71
|
+
"get_model": get_model,
|
|
72
|
+
"discard_model": discard_model,
|
|
73
|
+
"comm_manager": manager,
|
|
74
|
+
"comm_target": None,
|
|
75
|
+
}
|
|
76
|
+
|
|
77
|
+
if manager is None:
|
|
78
|
+
state["session"] = Session(shell=ipython, store=MemoryConversationStore())
|
|
79
|
+
else:
|
|
80
|
+
|
|
81
|
+
def accept_comm(comm: Any, message: dict[str, Any]) -> None:
|
|
82
|
+
data = message.get("content", {}).get("data", {})
|
|
83
|
+
if not isinstance(data, dict):
|
|
84
|
+
raise ValueError("Invalid Codebind notebook connection")
|
|
85
|
+
conversation = data.get("conversation")
|
|
86
|
+
if conversation is not None and not isinstance(conversation, dict):
|
|
87
|
+
raise ValueError("Invalid Codebind conversation")
|
|
88
|
+
|
|
89
|
+
previous_task = state.get("load_task")
|
|
90
|
+
if isinstance(previous_task, asyncio.Task):
|
|
91
|
+
previous_task.cancel()
|
|
92
|
+
previous_bridge = state.get("bridge")
|
|
93
|
+
if isinstance(previous_bridge, JupyterLabBridge):
|
|
94
|
+
previous_bridge.close()
|
|
95
|
+
|
|
96
|
+
bridge = JupyterLabBridge(comm, conversation)
|
|
97
|
+
session = Session(
|
|
98
|
+
shell=ipython,
|
|
99
|
+
bridge=bridge,
|
|
100
|
+
store=NotebookConversationStore(bridge),
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
async def answer_question(question: str, notebook: list[dict[str, Any]]) -> None:
|
|
104
|
+
selected = get_model()
|
|
105
|
+
try:
|
|
106
|
+
await session.asend(question, selected, notebook=notebook)
|
|
107
|
+
except BaseException:
|
|
108
|
+
await discard_model()
|
|
109
|
+
raise
|
|
110
|
+
|
|
111
|
+
async def prepare_session() -> None:
|
|
112
|
+
try:
|
|
113
|
+
await session.aload()
|
|
114
|
+
except asyncio.CancelledError:
|
|
115
|
+
raise
|
|
116
|
+
except Exception as error:
|
|
117
|
+
if bridge.ready:
|
|
118
|
+
bridge.report_session_ready(error)
|
|
119
|
+
else:
|
|
120
|
+
if bridge.ready:
|
|
121
|
+
bridge.report_session_ready()
|
|
122
|
+
|
|
123
|
+
bridge.handle_questions(answer_question)
|
|
124
|
+
state["bridge"] = bridge
|
|
125
|
+
state["session"] = session
|
|
126
|
+
state["load_task"] = asyncio.get_running_loop().create_task(prepare_session())
|
|
127
|
+
|
|
128
|
+
manager.register_target(TARGET_NAME, accept_comm)
|
|
129
|
+
state["comm_target"] = accept_comm
|
|
57
130
|
|
|
58
131
|
def question_magic(line: str, cell: str | None = None) -> None:
|
|
132
|
+
session = state.get("session")
|
|
133
|
+
if not isinstance(session, Session):
|
|
134
|
+
session = Session(shell=ipython, store=MemoryConversationStore())
|
|
135
|
+
state["session"] = session
|
|
59
136
|
selected = get_model()
|
|
60
137
|
try:
|
|
61
138
|
session.send(cell if cell is not None else line, selected)
|
|
62
139
|
except BaseException:
|
|
63
|
-
discard_model(
|
|
140
|
+
asyncio.run(discard_model())
|
|
64
141
|
raise
|
|
65
142
|
|
|
66
|
-
load_task: asyncio.Task[None] | None = None
|
|
67
|
-
if bridge is not None:
|
|
68
|
-
bridge.handle_questions(answer_question)
|
|
69
|
-
|
|
70
|
-
async def prepare_session() -> None:
|
|
71
|
-
try:
|
|
72
|
-
await session.aload()
|
|
73
|
-
except asyncio.CancelledError:
|
|
74
|
-
raise
|
|
75
|
-
except Exception as error:
|
|
76
|
-
bridge.report_session_ready(error)
|
|
77
|
-
else:
|
|
78
|
-
bridge.report_session_ready()
|
|
79
|
-
|
|
80
|
-
try:
|
|
81
|
-
load_task = asyncio.get_running_loop().create_task(prepare_session())
|
|
82
|
-
except RuntimeError:
|
|
83
|
-
pass
|
|
84
143
|
ipython.register_magic_function(question_magic, "line_cell", "question")
|
|
85
|
-
setattr(
|
|
86
|
-
ipython,
|
|
87
|
-
_STATE_ATTRIBUTE,
|
|
88
|
-
{
|
|
89
|
-
"session": session,
|
|
90
|
-
"bridge": bridge,
|
|
91
|
-
"load_task": load_task,
|
|
92
|
-
"get_model": get_model,
|
|
93
|
-
},
|
|
94
|
-
)
|
|
144
|
+
setattr(ipython, _STATE_ATTRIBUTE, state)
|
|
95
145
|
|
|
96
146
|
|
|
97
147
|
def unload_ipython_extension(ipython: InteractiveShell) -> None:
|
|
98
148
|
"""Unload Codebind without touching the user namespace."""
|
|
99
149
|
state: Any = getattr(ipython, _STATE_ATTRIBUTE, None)
|
|
100
150
|
if isinstance(state, dict):
|
|
101
|
-
|
|
102
|
-
if isinstance(load_task, asyncio.Task):
|
|
103
|
-
load_task.cancel()
|
|
104
|
-
session = state.get("session")
|
|
105
|
-
if isinstance(session, Session) and session.bridge is not None:
|
|
106
|
-
session.bridge.close()
|
|
151
|
+
_close_state(state)
|
|
107
152
|
delattr(ipython, _STATE_ATTRIBUTE)
|
|
108
153
|
ipython.magics_manager.magics["line"].pop("question", None)
|
|
109
154
|
ipython.magics_manager.magics["cell"].pop("question", None)
|