Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions corecoder/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,7 @@ def load_session(session_id: str) -> tuple[list[dict], str] | None:
try:
data = json.loads(path.read_text(encoding="utf-8"))
return data["messages"], data["model"]
except (json.JSONDecodeError, KeyError, OSError):
except (json.JSONDecodeError, UnicodeDecodeError, KeyError, OSError):
# a corrupt or truncated session file shouldn't crash resume
return None

Expand All @@ -91,7 +91,7 @@ def list_sessions() -> list[dict]:
"saved_at": data.get("saved_at", "?"),
"preview": preview,
})
except (json.JSONDecodeError, KeyError):
except (json.JSONDecodeError, UnicodeDecodeError, KeyError):
continue

return sessions[:20] # cap at 20
31 changes: 30 additions & 1 deletion tests/test_session.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import pytest

from corecoder import session as session_module
from corecoder.session import load_session, save_session
from corecoder.session import list_sessions, load_session, save_session


def test_default_session_ids_do_not_collide(tmp_path, monkeypatch):
Expand Down Expand Up @@ -64,6 +66,33 @@ def test_corrupt_session_file_returns_none(tmp_path, monkeypatch):
assert load_session("broken") is None


@pytest.mark.parametrize("payload", [b"\xff", b'{"messages": [{"content": "\xe4\xbd'])
def test_invalid_utf8_session_returns_none(tmp_path, monkeypatch, payload):
monkeypatch.setattr(session_module, "SESSIONS_DIR", tmp_path)
path = tmp_path / "broken.json"
path.write_bytes(payload)

assert load_session("broken") is None
assert path.read_bytes() == payload


@pytest.mark.parametrize("payload", [b"\xff", b'{"messages": [{"content": "\xe4\xbd'])
def test_list_sessions_skips_invalid_utf8(tmp_path, monkeypatch, payload):
monkeypatch.setattr(session_module, "SESSIONS_DIR", tmp_path)
messages = [{"role": "user", "content": "你好"}]
sid = save_session(messages, "model-zh", "valid")
path = tmp_path / "z-broken.json"
path.write_bytes(payload)

sessions = list_sessions()

assert len(sessions) == 1
assert sessions[0]["id"] == sid
assert sessions[0]["preview"] == "你好"
assert load_session(sid) == (messages, "model-zh")
assert path.read_bytes() == payload


def test_session_roundtrips_unicode(tmp_path, monkeypatch):
monkeypatch.setattr(session_module, "SESSIONS_DIR", tmp_path)

Expand Down