|
45 | 45 | CreateNotebookCommand, |
46 | 46 | ExecuteCellCommand, |
47 | 47 | ExecuteCellsCommand, |
| 48 | + ModelCommand, |
| 49 | + ModelCustomMessage, |
| 50 | + ModelUpdateMessage, |
48 | 51 | UpdateUIElementCommand, |
49 | 52 | ) |
50 | 53 | from marimo._session.state.session_view import ModelReplayState, SessionView |
@@ -351,6 +354,111 @@ def test_model_multiple_models(session_view: SessionView) -> None: |
351 | 354 | assert session_view.model_states[model_id2].state == {"key": "v2"} |
352 | 355 |
|
353 | 356 |
|
| 357 | +def _open_model( |
| 358 | + session_view: SessionView, |
| 359 | + model_id: WidgetModelId, |
| 360 | + state: dict[str, Any], |
| 361 | +) -> None: |
| 362 | + session_view.add_notification( |
| 363 | + ModelLifecycleNotification( |
| 364 | + model_id=model_id, |
| 365 | + message=ModelOpen(state=state, buffer_paths=[], buffers=[]), |
| 366 | + ) |
| 367 | + ) |
| 368 | + |
| 369 | + |
| 370 | +def test_model_command_merges_into_replay(session_view: SessionView) -> None: |
| 371 | + """A client's model write is recorded for reconnect replay.""" |
| 372 | + model_id = WidgetModelId("test_model") |
| 373 | + _open_model(session_view, model_id, {"count": 0, "label": "hi"}) |
| 374 | + |
| 375 | + session_view.add_control_request( |
| 376 | + ModelCommand( |
| 377 | + model_id=model_id, |
| 378 | + message=ModelUpdateMessage(state={"count": 5}, buffer_paths=[]), |
| 379 | + buffers=[], |
| 380 | + ) |
| 381 | + ) |
| 382 | + assert session_view.model_states[model_id].state == { |
| 383 | + "count": 5, |
| 384 | + "label": "hi", |
| 385 | + } |
| 386 | + |
| 387 | + |
| 388 | +def test_model_command_merges_buffers(session_view: SessionView) -> None: |
| 389 | + model_id = WidgetModelId("test_model") |
| 390 | + _open_model(session_view, model_id, {"img": None}) |
| 391 | + |
| 392 | + session_view.add_control_request( |
| 393 | + ModelCommand( |
| 394 | + model_id=model_id, |
| 395 | + message=ModelUpdateMessage( |
| 396 | + state={"img": None}, buffer_paths=[["img"]] |
| 397 | + ), |
| 398 | + buffers=[b"\x89PNG"], |
| 399 | + ) |
| 400 | + ) |
| 401 | + assert session_view.model_states[model_id].buffers == { |
| 402 | + ("img",): b"\x89PNG" |
| 403 | + } |
| 404 | + |
| 405 | + |
| 406 | +def test_model_command_strips_code_and_style( |
| 407 | + session_view: SessionView, |
| 408 | +) -> None: |
| 409 | + """Replayed state reaches future viewers, so a client must not |
| 410 | + be able to persist `_esm` or `_css` into it.""" |
| 411 | + model_id = WidgetModelId("test_model") |
| 412 | + _open_model(session_view, model_id, {"count": 0}) |
| 413 | + |
| 414 | + session_view.add_control_request( |
| 415 | + ModelCommand( |
| 416 | + model_id=model_id, |
| 417 | + message=ModelUpdateMessage( |
| 418 | + state={ |
| 419 | + "_esm": "alert('pwned')", |
| 420 | + "_css": "body { display: none }", |
| 421 | + "count": 2, |
| 422 | + }, |
| 423 | + buffer_paths=[], |
| 424 | + ), |
| 425 | + buffers=[], |
| 426 | + ) |
| 427 | + ) |
| 428 | + assert session_view.model_states[model_id].state == {"count": 2} |
| 429 | + |
| 430 | + |
| 431 | +def test_model_command_without_open_ignored( |
| 432 | + session_view: SessionView, |
| 433 | +) -> None: |
| 434 | + model_id = WidgetModelId("never_opened") |
| 435 | + session_view.add_control_request( |
| 436 | + ModelCommand( |
| 437 | + model_id=model_id, |
| 438 | + message=ModelUpdateMessage(state={"count": 1}, buffer_paths=[]), |
| 439 | + buffers=[], |
| 440 | + ) |
| 441 | + ) |
| 442 | + assert model_id not in session_view.model_states |
| 443 | + |
| 444 | + |
| 445 | +def test_model_command_custom_message_ignored( |
| 446 | + session_view: SessionView, |
| 447 | +) -> None: |
| 448 | + """Custom messages are ephemeral — they never mutate replay state.""" |
| 449 | + model_id = WidgetModelId("test_model") |
| 450 | + _open_model(session_view, model_id, {"count": 0}) |
| 451 | + |
| 452 | + session_view.add_control_request( |
| 453 | + ModelCommand( |
| 454 | + model_id=model_id, |
| 455 | + message=ModelCustomMessage(content={"foo": "bar"}), |
| 456 | + buffers=[], |
| 457 | + ) |
| 458 | + ) |
| 459 | + assert session_view.model_states[model_id].state == {"count": 0} |
| 460 | + |
| 461 | + |
354 | 462 | def test_get_model_notifications(session_view: SessionView) -> None: |
355 | 463 | # Empty initially |
356 | 464 | assert session_view.get_model_notifications() == [] |
|
0 commit comments