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:
|
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
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user