diff --git a/ms_agent/session/session_log.py b/ms_agent/session/session_log.py index 72596d2a2..bfb199b91 100644 --- a/ms_agent/session/session_log.py +++ b/ms_agent/session/session_log.py @@ -38,7 +38,7 @@ import uuid from datetime import datetime, timezone from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any, Dict, Iterator, List, Optional from ms_agent.utils.atomic_file import atomic_write_json @@ -237,7 +237,7 @@ def get_all_messages(self) -> List[Dict[str, Any]]: if not self._path.exists(): self._messages = msgs return msgs - for line in self._path.read_text(encoding='utf-8').splitlines(): + for line in self._read_lines(): line = line.strip() if not line: continue @@ -304,7 +304,7 @@ def get_compaction_events(self) -> List[Dict[str, Any]]: events: List[Dict[str, Any]] = [] if not self._path.exists(): return events - for line in self._path.read_text(encoding='utf-8').splitlines(): + for line in self._read_lines(): line = line.strip() if not line: continue @@ -325,7 +325,7 @@ def get_errors(self) -> List[Dict[str, Any]]: errors: List[Dict[str, Any]] = [] if not self._path.exists(): return errors - for line in self._path.read_text(encoding='utf-8').splitlines(): + for line in self._read_lines(): line = line.strip() if not line: continue @@ -345,7 +345,7 @@ def get_permissions(self) -> List[Dict[str, Any]]: perms: List[Dict[str, Any]] = [] if not self._path.exists(): return perms - for line in self._path.read_text(encoding='utf-8').splitlines(): + for line in self._read_lines(): line = line.strip() if not line: continue @@ -366,7 +366,7 @@ def get_loop_ends(self) -> List[Dict[str, Any]]: out: List[Dict[str, Any]] = [] if not self._path.exists(): return out - for line in self._path.read_text(encoding='utf-8').splitlines(): + for line in self._read_lines(): line = line.strip() if not line: continue @@ -387,7 +387,7 @@ def get_skill_invocations(self) -> List[Dict[str, Any]]: out: List[Dict[str, Any]] = [] if not self._path.exists(): return out - for line in self._path.read_text(encoding='utf-8').splitlines(): + for line in self._read_lines(): line = line.strip() if not line: continue @@ -468,7 +468,7 @@ def _ensure_metadata(self) -> None: def _load_all_to_set_seq(self) -> None: """Scan the file to find the highest seq number.""" max_seq = -1 - for line in self._path.read_text(encoding='utf-8').splitlines(): + for line in self._read_lines(): line = line.strip() if not line: continue @@ -527,8 +527,11 @@ def _read_legacy_header(self) -> Optional[Dict[str, Any]]: """Return the first-line metadata record of the main log, if any.""" if not self._path.exists(): return None - with open(self._path, 'r', encoding='utf-8') as f: - first_line = f.readline().strip() + with open(self._path, 'rb') as f: + try: + first_line = f.readline().decode('utf-8').strip() + except UnicodeDecodeError: + return None if first_line: try: record = json.loads(first_line) @@ -538,10 +541,26 @@ def _read_legacy_header(self) -> Optional[Dict[str, Any]]: pass return None + def _read_lines(self) -> Iterator[str]: + """Read log lines, skipping records with invalid UTF-8.""" + with open(self._path, 'rb') as f: + for line in f: + try: + yield line.decode('utf-8') + except UnicodeDecodeError: + continue + def _append_line(self, record: Dict[str, Any]) -> None: """Append a single JSON line and flush.""" - with open(self._path, 'a', encoding='utf-8') as f: - f.write(json.dumps(record, ensure_ascii=False) + '\n') + data = (json.dumps(record, ensure_ascii=False) + '\n').encode('utf-8') + with open(self._path, 'ab+') as f: + f.seek(0, os.SEEK_END) + if f.tell(): + f.seek(-1, os.SEEK_END) + if f.read(1) != b'\n': + # Keep an interrupted record separate from the next write. + f.write(b'\n') + f.write(data) f.flush() os.fsync(f.fileno()) diff --git a/tests/session/test_session_log_recovery.py b/tests/session/test_session_log_recovery.py new file mode 100644 index 000000000..1ede1d7a1 --- /dev/null +++ b/tests/session/test_session_log_recovery.py @@ -0,0 +1,45 @@ +"""Resuming an interrupted JSONL write must preserve the next message.""" +import json + +import pytest + +from ms_agent.session.session_log import SessionLog + + +@pytest.mark.parametrize('tail', ['clean', 'partial_json', 'partial_utf8', 'complete']) +@pytest.mark.parametrize('keep_sidecar', [True, False]) +def test_resume_after_unterminated_record(tmp_path, tail, keep_sidecar): + log = SessionLog(tmp_path, 'recovery') + log.append({'role': 'user', 'content': '之前'}) + path = tmp_path / 'recovery.jsonl' + if tail == 'partial_json': + suffix = b'{"role":"assistant","content":"unfinished' + elif tail == 'partial_utf8': + suffix = b'{"role":"assistant","content":"\xe4\xb8' + elif tail == 'complete': + suffix = json.dumps({'role': 'assistant', 'content': 'Complete', 'seq': 1}).encode() + else: + suffix = b'' + with path.open('ab') as file: + file.write(suffix) + original = path.read_bytes() + if not keep_sidecar: + (tmp_path / 'recovery.meta.json').unlink() + + resumed = SessionLog(tmp_path, 'recovery') + expected_seq = 2 if tail == 'complete' else 1 + assert resumed.append({'role': 'user', 'content': 'After recovery'}) == expected_seq + resumed.record_compaction({'strategy': 'test'}) + + reopened = SessionLog(tmp_path, 'recovery') + messages = reopened.get_all_messages() + expected = ['之前', 'Complete', 'After recovery'] if tail == 'complete' else ['之前', 'After recovery'] + assert [message['content'] for message in messages] == expected + assert [message['seq'] for message in messages] == list(range(expected_seq + 1)) + assert reopened.get_compaction_events()[0]['strategy'] == 'test' + assert reopened.get_errors() == [] + assert reopened.get_permissions() == [] + assert reopened.get_loop_ends() == [] + assert reopened.get_skill_invocations() == [] + assert path.read_bytes().startswith(original) + assert reopened.append({'role': 'user', 'content': 'Next'}) == expected_seq + 2