diff --git a/backend/api/server.py b/backend/api/server.py index 5eba318..6a712cc 100644 --- a/backend/api/server.py +++ b/backend/api/server.py @@ -3913,6 +3913,8 @@ def _read_session_title(directory: Path, session_id: str | None) -> str | None: except FileNotFoundError: return None metadata = stored.session_metadata or {} + if metadata.get('title_source') == 'first_message': + return None title = metadata.get('title') if isinstance(title, str) and title.strip(): return title.strip() diff --git a/tests/test_gui_server.py b/tests/test_gui_server.py index 91ede93..aadb983 100644 --- a/tests/test_gui_server.py +++ b/tests/test_gui_server.py @@ -12,6 +12,7 @@ from __future__ import annotations import tempfile import unittest import json +from dataclasses import replace from pathlib import Path from unittest.mock import patch @@ -24,6 +25,7 @@ from backend.api.server import ( _derive_initial_session_title, _ensure_session_title, _generate_session_title, + _read_session_title, _runtime_event_stage, _save_in_progress_session, _sanitize_stored_session_for_resume, @@ -795,26 +797,29 @@ class GuiServerTests(unittest.TestCase): def test_update_session_title(self) -> None: with tempfile.TemporaryDirectory() as d: root = Path(d) - client, _ = _build_client(root) + client, state = _build_client(root) session_root = root / 'accounts' / 'alice' / 'sessions' - nested = session_root / 'same-id' - nested.mkdir(parents=True) - session_file = nested / 'session.json' - session_file.write_text( - json.dumps({'session_id': 'same-id', 'messages': []}), - encoding='utf-8', + session_id = 'same-id' + agent = state.agent_for('alice', session_id) + _save_in_progress_session( + directory=session_root, + agent=agent, + session_id=session_id, + prompt='帮我分析数据。', ) response = client.patch( - '/api/sessions/same-id', + f'/api/sessions/{session_id}', params={'account_id': 'alice'}, json={'title': ' 手动标题 '}, ) self.assertEqual(response.status_code, 200) self.assertEqual(response.json()['title'], '手动标题') - data = json.loads(session_file.read_text(encoding='utf-8')) - self.assertEqual(data['title'], '手动标题') - self.assertEqual(data['title_source'], 'manual') + metadata = load_agent_session( + session_id, directory=session_root + ).session_metadata or {} + self.assertEqual(metadata['title'], '手动标题') + self.assertEqual(metadata['title_source'], 'manual') def test_in_progress_session_writes_first_message_title(self) -> None: with tempfile.TemporaryDirectory() as d: @@ -834,8 +839,26 @@ class GuiServerTests(unittest.TestCase): session_file = directory / session_id / 'session.json' data = json.loads(session_file.read_text(encoding='utf-8')) self.assertTrue(title.startswith('帮我分析一下线上 router')) - self.assertEqual(data['title'], title) - self.assertEqual(data['title_source'], 'first_message') + metadata = data['session_metadata'] + self.assertEqual(metadata['title'], title) + self.assertEqual(metadata['title_source'], 'first_message') + + def test_first_message_title_is_not_preserved_as_previous_title(self) -> None: + with tempfile.TemporaryDirectory() as d: + root = Path(d) + _, state = _build_client(root) + session_id = 'first-title-previous' + directory = state.account_paths('alice')['sessions'] + agent = state.agent_for('alice', session_id) + + _save_in_progress_session( + directory=directory, + agent=agent, + session_id=session_id, + prompt='帮我生成一批地图导航测试数据。', + ) + + self.assertIsNone(_read_session_title(directory, session_id)) def test_session_title_refines_first_message_title_after_run(self) -> None: with tempfile.TemporaryDirectory() as d: @@ -850,12 +873,6 @@ class GuiServerTests(unittest.TestCase): session_id=session_id, prompt='帮我生成地图和餐饮服务的边界数据。', ) - session_file = directory / session_id / 'session.json' - data = json.loads(session_file.read_text(encoding='utf-8')) - data['messages'] = [ - {'role': 'user', 'content': '帮我生成地图和餐饮服务的边界数据。'} - ] - session_file.write_text(json.dumps(data), encoding='utf-8') with patch( 'backend.api.server._generate_session_title', @@ -871,9 +888,11 @@ class GuiServerTests(unittest.TestCase): ) generate.assert_called_once() - updated = json.loads(session_file.read_text(encoding='utf-8')) - self.assertEqual(updated['title'], '地图餐饮边界数据') - self.assertEqual(updated['title_source'], 'llm') + metadata = load_agent_session( + session_id, directory=directory + ).session_metadata or {} + self.assertEqual(metadata['title'], '地图餐饮边界数据') + self.assertEqual(metadata['title_source'], 'llm') def test_session_title_preserves_existing_manual_title(self) -> None: with tempfile.TemporaryDirectory() as d: @@ -881,21 +900,24 @@ class GuiServerTests(unittest.TestCase): _, state = _build_client(root) session_id = 'manual-preserved' directory = state.account_paths('alice')['sessions'] - session_dir = directory / session_id - session_dir.mkdir(parents=True) - session_file = session_dir / 'session.json' - session_file.write_text( - json.dumps( - { - 'session_id': session_id, + agent = state.agent_for('alice', session_id) + _save_in_progress_session( + directory=directory, + agent=agent, + session_id=session_id, + prompt='帮我分析数据。', + ) + stored = load_agent_session(session_id, directory=directory) + save_agent_session( + replace( + stored, + session_metadata={ + **(stored.session_metadata or {}), 'title': '手动标题', 'title_source': 'manual', - 'messages': [ - {'role': 'user', 'content': '帮我分析数据。'}, - ], - } + }, ), - encoding='utf-8', + directory=directory, ) with patch('backend.api.server._generate_session_title') as generate: @@ -909,8 +931,10 @@ class GuiServerTests(unittest.TestCase): ) generate.assert_not_called() - updated = json.loads(session_file.read_text(encoding='utf-8')) - self.assertEqual(updated['title'], '手动标题') + metadata = load_agent_session( + session_id, directory=directory + ).session_metadata or {} + self.assertEqual(metadata['title'], '手动标题') def test_derive_initial_session_title_strips_session_context(self) -> None: title = _derive_initial_session_title(