Fix session title refinement after optimistic save
This commit is contained in:
@@ -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()
|
||||
|
||||
+60
-36
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user