Fix session title refinement after optimistic save

This commit is contained in:
wuyang6
2026-06-16 14:56:58 +08:00
parent 1d38753d08
commit c8b5b0659d
2 changed files with 62 additions and 36 deletions
+2
View File
@@ -3913,6 +3913,8 @@ def _read_session_title(directory: Path, session_id: str | None) -> str | None:
except FileNotFoundError: except FileNotFoundError:
return None return None
metadata = stored.session_metadata or {} metadata = stored.session_metadata or {}
if metadata.get('title_source') == 'first_message':
return None
title = metadata.get('title') title = metadata.get('title')
if isinstance(title, str) and title.strip(): if isinstance(title, str) and title.strip():
return title.strip() return title.strip()
+60 -36
View File
@@ -12,6 +12,7 @@ from __future__ import annotations
import tempfile import tempfile
import unittest import unittest
import json import json
from dataclasses import replace
from pathlib import Path from pathlib import Path
from unittest.mock import patch from unittest.mock import patch
@@ -24,6 +25,7 @@ from backend.api.server import (
_derive_initial_session_title, _derive_initial_session_title,
_ensure_session_title, _ensure_session_title,
_generate_session_title, _generate_session_title,
_read_session_title,
_runtime_event_stage, _runtime_event_stage,
_save_in_progress_session, _save_in_progress_session,
_sanitize_stored_session_for_resume, _sanitize_stored_session_for_resume,
@@ -795,26 +797,29 @@ class GuiServerTests(unittest.TestCase):
def test_update_session_title(self) -> None: def test_update_session_title(self) -> None:
with tempfile.TemporaryDirectory() as d: with tempfile.TemporaryDirectory() as d:
root = Path(d) root = Path(d)
client, _ = _build_client(root) client, state = _build_client(root)
session_root = root / 'accounts' / 'alice' / 'sessions' session_root = root / 'accounts' / 'alice' / 'sessions'
nested = session_root / 'same-id' session_id = 'same-id'
nested.mkdir(parents=True) agent = state.agent_for('alice', session_id)
session_file = nested / 'session.json' _save_in_progress_session(
session_file.write_text( directory=session_root,
json.dumps({'session_id': 'same-id', 'messages': []}), agent=agent,
encoding='utf-8', session_id=session_id,
prompt='帮我分析数据。',
) )
response = client.patch( response = client.patch(
'/api/sessions/same-id', f'/api/sessions/{session_id}',
params={'account_id': 'alice'}, params={'account_id': 'alice'},
json={'title': ' 手动标题 '}, json={'title': ' 手动标题 '},
) )
self.assertEqual(response.status_code, 200) self.assertEqual(response.status_code, 200)
self.assertEqual(response.json()['title'], '手动标题') self.assertEqual(response.json()['title'], '手动标题')
data = json.loads(session_file.read_text(encoding='utf-8')) metadata = load_agent_session(
self.assertEqual(data['title'], '手动标题') session_id, directory=session_root
self.assertEqual(data['title_source'], 'manual') ).session_metadata or {}
self.assertEqual(metadata['title'], '手动标题')
self.assertEqual(metadata['title_source'], 'manual')
def test_in_progress_session_writes_first_message_title(self) -> None: def test_in_progress_session_writes_first_message_title(self) -> None:
with tempfile.TemporaryDirectory() as d: with tempfile.TemporaryDirectory() as d:
@@ -834,8 +839,26 @@ class GuiServerTests(unittest.TestCase):
session_file = directory / session_id / 'session.json' session_file = directory / session_id / 'session.json'
data = json.loads(session_file.read_text(encoding='utf-8')) data = json.loads(session_file.read_text(encoding='utf-8'))
self.assertTrue(title.startswith('帮我分析一下线上 router')) self.assertTrue(title.startswith('帮我分析一下线上 router'))
self.assertEqual(data['title'], title) metadata = data['session_metadata']
self.assertEqual(data['title_source'], 'first_message') 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: def test_session_title_refines_first_message_title_after_run(self) -> None:
with tempfile.TemporaryDirectory() as d: with tempfile.TemporaryDirectory() as d:
@@ -850,12 +873,6 @@ class GuiServerTests(unittest.TestCase):
session_id=session_id, session_id=session_id,
prompt='帮我生成地图和餐饮服务的边界数据。', 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( with patch(
'backend.api.server._generate_session_title', 'backend.api.server._generate_session_title',
@@ -871,9 +888,11 @@ class GuiServerTests(unittest.TestCase):
) )
generate.assert_called_once() generate.assert_called_once()
updated = json.loads(session_file.read_text(encoding='utf-8')) metadata = load_agent_session(
self.assertEqual(updated['title'], '地图餐饮边界数据') session_id, directory=directory
self.assertEqual(updated['title_source'], 'llm') ).session_metadata or {}
self.assertEqual(metadata['title'], '地图餐饮边界数据')
self.assertEqual(metadata['title_source'], 'llm')
def test_session_title_preserves_existing_manual_title(self) -> None: def test_session_title_preserves_existing_manual_title(self) -> None:
with tempfile.TemporaryDirectory() as d: with tempfile.TemporaryDirectory() as d:
@@ -881,21 +900,24 @@ class GuiServerTests(unittest.TestCase):
_, state = _build_client(root) _, state = _build_client(root)
session_id = 'manual-preserved' session_id = 'manual-preserved'
directory = state.account_paths('alice')['sessions'] directory = state.account_paths('alice')['sessions']
session_dir = directory / session_id agent = state.agent_for('alice', session_id)
session_dir.mkdir(parents=True) _save_in_progress_session(
session_file = session_dir / 'session.json' directory=directory,
session_file.write_text( agent=agent,
json.dumps( session_id=session_id,
{ prompt='帮我分析数据。',
'session_id': session_id, )
stored = load_agent_session(session_id, directory=directory)
save_agent_session(
replace(
stored,
session_metadata={
**(stored.session_metadata or {}),
'title': '手动标题', 'title': '手动标题',
'title_source': 'manual', 'title_source': 'manual',
'messages': [ },
{'role': 'user', 'content': '帮我分析数据。'},
],
}
), ),
encoding='utf-8', directory=directory,
) )
with patch('backend.api.server._generate_session_title') as generate: with patch('backend.api.server._generate_session_title') as generate:
@@ -909,8 +931,10 @@ class GuiServerTests(unittest.TestCase):
) )
generate.assert_not_called() generate.assert_not_called()
updated = json.loads(session_file.read_text(encoding='utf-8')) metadata = load_agent_session(
self.assertEqual(updated['title'], '手动标题') session_id, directory=directory
).session_metadata or {}
self.assertEqual(metadata['title'], '手动标题')
def test_derive_initial_session_title_strips_session_context(self) -> None: def test_derive_initial_session_title_strips_session_context(self) -> None:
title = _derive_initial_session_title( title = _derive_initial_session_title(