codebind 0.6.2__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.
Files changed (21) hide show
  1. codebind/conversation.py +166 -15
  2. codebind/execution.py +56 -11
  3. codebind/extension.py +102 -57
  4. codebind/images.py +127 -0
  5. codebind/instructions.md +9 -5
  6. codebind/jupyter.py +16 -45
  7. codebind/session.py +142 -92
  8. {codebind-0.6.2.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/package.json +2 -2
  9. codebind-0.6.3.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/590.dae4561c5427a855.js +1 -0
  10. codebind-0.6.3.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/remoteEntry.76c22319f766843f.js +1 -0
  11. {codebind-0.6.2.dist-info → codebind-0.6.3.dist-info}/METADATA +7 -3
  12. codebind-0.6.3.dist-info/RECORD +23 -0
  13. {codebind-0.6.2.dist-info → codebind-0.6.3.dist-info}/WHEEL +1 -1
  14. codebind-0.6.2.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/590.5b773af7a27693b3.js +0 -1
  15. codebind-0.6.2.data/data/share/jupyter/labextensions/codebind-jupyterlab/static/remoteEntry.2cec6ffa067e0ab2.js +0 -1
  16. codebind-0.6.2.dist-info/RECORD +0 -22
  17. {codebind-0.6.2.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/install.json +0 -0
  18. {codebind-0.6.2.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/static/style.js +0 -0
  19. {codebind-0.6.2.data → codebind-0.6.3.data}/data/share/jupyter/labextensions/codebind-jupyterlab/static/third-party-licenses.json +0 -0
  20. {codebind-0.6.2.dist-info → codebind-0.6.3.dist-info}/entry_points.txt +0 -0
  21. {codebind-0.6.2.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 BaseMessage, HumanMessage, message_to_dict, messages_from_dict
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
- _VERSION = 2
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.get("version") != _VERSION:
39
- raise ValueError("Unsupported Codebind conversation version")
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 not isinstance(revision, int) or revision != len(events):
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) or event.get("sequence") != sequence
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 sequence")
67
+ raise ValueError("Invalid Codebind conversation event")
53
68
  if not isinstance(cwd, str):
54
- cwd = ""
55
- return cls(identifier, revision, deepcopy(events), cwd)
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
- incoming = _normalize_notebook(cells)
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(HumanMessage(content)),
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 _normalize_notebook(value: Any) -> list[dict[str, Any]]:
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
- cell["outputs"] = deepcopy(outputs)
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 asdict, dataclass
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 a JSON-serializable representation."""
32
- return asdict(self)
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
- sys.displayhook = _ExecutionDisplayHook(shell, expression_outputs, on_output)
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
- @staticmethod
283
- def _report(result: Any, captured: Any) -> ExecutionReport:
284
- """Build the model-facing text projection of an IPython execution."""
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
- for output in captured.outputs:
288
- data = getattr(output, "data", None)
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
- previous_load = previous.get("load_task")
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(selected: BaseChatModel) -> None:
55
+ async def discard_model() -> None:
42
56
  nonlocal model
43
- if model is selected:
44
- model = None
45
-
46
- bridge = JupyterLabBridge.connect(ipython)
47
- store = NotebookConversationStore(bridge) if bridge is not None else MemoryConversationStore()
48
- session = Session(shell=ipython, bridge=bridge, store=store)
49
-
50
- async def answer_question(question: str, notebook: list[dict[str, Any]]) -> None:
51
- selected = get_model()
52
- try:
53
- await session.asend(question, selected, notebook=notebook)
54
- except BaseException:
55
- discard_model(selected)
56
- raise
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(selected)
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
- load_task = state.get("load_task")
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)