Fix session title generation

This commit is contained in:
wuyang6
2026-05-11 16:22:19 +08:00
parent 82a444bfb2
commit 64ad7aa97a
3 changed files with 208 additions and 33 deletions
+126 -1
View File
@@ -13,16 +13,21 @@ import tempfile
import unittest
import json
from pathlib import Path
from unittest.mock import patch
from fastapi.testclient import TestClient
from backend.api.server import (
AgentState,
create_app,
_derive_initial_session_title,
_ensure_session_title,
_generate_session_title,
_runtime_event_stage,
_save_in_progress_session,
_sanitize_stored_session_for_resume,
)
from src.agent_types import AgentRunResult
from src.agent_types import AgentRunResult, AssistantTurn
from src.session_store import StoredAgentSession, load_agent_session, save_agent_session
@@ -724,6 +729,126 @@ class GuiServerTests(unittest.TestCase):
self.assertEqual(data['title'], '手动标题')
self.assertEqual(data['title_source'], 'manual')
def test_in_progress_session_writes_first_message_title(self) -> None:
with tempfile.TemporaryDirectory() as d:
root = Path(d)
_, state = _build_client(root)
session_id = 'first-title'
directory = state.account_paths('alice')['sessions']
agent = state.agent_for('alice', session_id)
title = _save_in_progress_session(
directory=directory,
agent=agent,
session_id=session_id,
prompt='帮我分析一下线上 router session 数据,然后统计 device 分布。',
)
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')
def test_session_title_refines_first_message_title_after_run(self) -> None:
with tempfile.TemporaryDirectory() as d:
root = Path(d)
_, state = _build_client(root)
session_id = 'refine-title'
directory = state.account_paths('alice')['sessions']
agent = state.agent_for('alice', session_id)
fallback_title = _save_in_progress_session(
directory=directory,
agent=agent,
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',
return_value='地图餐饮边界数据',
) as generate:
_ensure_session_title(
directory,
session_id,
model='test-model',
base_url='http://127.0.0.1:8000/v1',
api_key='local-token',
fallback_title=fallback_title,
)
generate.assert_called_once()
updated = json.loads(session_file.read_text(encoding='utf-8'))
self.assertEqual(updated['title'], '地图餐饮边界数据')
self.assertEqual(updated['title_source'], 'llm')
def test_session_title_preserves_existing_manual_title(self) -> None:
with tempfile.TemporaryDirectory() as d:
root = Path(d)
_, 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,
'title': '手动标题',
'title_source': 'manual',
'messages': [
{'role': 'user', 'content': '帮我分析数据。'},
],
}
),
encoding='utf-8',
)
with patch('backend.api.server._generate_session_title') as generate:
_ensure_session_title(
directory,
session_id,
model='test-model',
base_url='http://127.0.0.1:8000/v1',
api_key='local-token',
previous_title='手动标题',
)
generate.assert_not_called()
updated = json.loads(session_file.read_text(encoding='utf-8'))
self.assertEqual(updated['title'], '手动标题')
def test_derive_initial_session_title_strips_session_context(self) -> None:
title = _derive_initial_session_title(
'请帮我生成线上挖掘数据。\n\n[当前会话目录]\n/root/session'
)
self.assertEqual(title, '请帮我生成线上挖掘数据')
def test_generate_session_title_uses_model_compat_client(self) -> None:
with patch('backend.api.server.OpenAICompatClient') as client_cls:
client = client_cls.return_value
client.complete.return_value = AssistantTurn(content='“模型兼容标题”')
title = _generate_session_title(
['帮我分析一下线上 session 数据'],
model='ppio/pa/claude-opus-4-7',
base_url='http://model.mify.ai.srv/v1',
api_key='local-token',
)
self.assertEqual(title, '模型兼容标题')
config = client_cls.call_args.args[0]
self.assertEqual(config.model, 'ppio/pa/claude-opus-4-7')
self.assertEqual(config.base_url, 'http://model.mify.ai.srv/v1')
client.complete.assert_called_once()
def test_clear_runtime_state(self) -> None:
with tempfile.TemporaryDirectory() as d:
client, _ = _build_client(Path(d))