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
+68 -31
View File
@@ -38,6 +38,7 @@ from src.agent_types import (
ModelConfig, ModelConfig,
) )
from src.bundled_skills import ALWAYS_ENABLED_HIDDEN_SKILL_NAMES, get_bundled_skills from src.bundled_skills import ALWAYS_ENABLED_HIDDEN_SKILL_NAMES, get_bundled_skills
from src.openai_compat import OpenAICompatClient, OpenAICompatError
from src.session_store import ( from src.session_store import (
DEFAULT_AGENT_SESSION_DIR, DEFAULT_AGENT_SESSION_DIR,
StoredAgentSession, StoredAgentSession,
@@ -1176,17 +1177,18 @@ def create_app(state: AgentState) -> FastAPI:
state.run_manager.record_event(run_record.run_id, event) state.run_manager.record_event(run_record.run_id, event)
_emit_runtime_event(event_sink, event) _emit_runtime_event(event_sink, event)
previous_title = _read_session_title(
session_directory,
requested_session_id,
)
fallback_title = None
if request.resume_session_id is None: if request.resume_session_id is None:
_save_in_progress_session( fallback_title = _save_in_progress_session(
directory=session_directory, directory=session_directory,
agent=agent, agent=agent,
session_id=requested_session_id, session_id=requested_session_id,
prompt=request.prompt.strip(), prompt=request.prompt.strip(),
) )
previous_title = _read_session_title(
session_directory,
requested_session_id,
)
if run_lock.locked(): if run_lock.locked():
queued_event = { queued_event = {
'type': 'run_queued', 'type': 'run_queued',
@@ -1297,6 +1299,7 @@ def create_app(state: AgentState) -> FastAPI:
base_url=config.base_url, base_url=config.base_url,
api_key=config.api_key, api_key=config.api_key,
previous_title=previous_title, previous_title=previous_title,
fallback_title=fallback_title,
) )
_annotate_transcript_elapsed(payload, elapsed_ms) _annotate_transcript_elapsed(payload, elapsed_ms)
state.run_manager.finish( state.run_manager.finish(
@@ -1728,14 +1731,15 @@ def _save_in_progress_session(
agent: LocalCodingAgent, agent: LocalCodingAgent,
session_id: str | None, session_id: str | None,
prompt: str, prompt: str,
) -> None: ) -> str | None:
safe_id = _safe_session_id(session_id) safe_id = _safe_session_id(session_id)
if safe_id is None or not prompt: if safe_id is None or not prompt:
return return None
if _session_json_path(directory, safe_id).exists(): if _session_json_path(directory, safe_id).exists():
return return None
# 新会话的第一轮执行可能很久。先写一个运行中占位,避免刷新页面后找不到会话。 # 新会话的第一轮执行可能很久。先写一个运行中占位,避免刷新页面后找不到会话。
initial_title = _derive_initial_session_title(prompt)
scratchpad_directory = ( scratchpad_directory = (
agent.runtime_config.scratchpad_root / safe_id / 'scratchpad' agent.runtime_config.scratchpad_root / safe_id / 'scratchpad'
).resolve() ).resolve()
@@ -1764,9 +1768,30 @@ def _save_in_progress_session(
plugin_state={}, plugin_state={},
scratchpad_directory=str(scratchpad_directory), scratchpad_directory=str(scratchpad_directory),
) )
save_agent_session(stored, directory=directory) saved_path = save_agent_session(stored, directory=directory)
if initial_title:
try:
data = json.loads(saved_path.read_text(encoding='utf-8'))
except (OSError, json.JSONDecodeError):
return initial_title
data['title'] = initial_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(saved_path, data)
return initial_title
except OSError: except OSError:
return return None
def _derive_initial_session_title(prompt: str) -> str | None:
# 第一条消息刚发出时先给侧边栏一个可读标题,后续再由模型摘要精修。
stripped = _strip_session_context(prompt)
stripped = re.sub(r'\s+', ' ', stripped).strip()
stripped = stripped.strip('"\'“”‘’`')
if not stripped:
return None
title = _clean_session_title(stripped)
return title or None
def _session_state_from_stored(stored: StoredAgentSession) -> AgentSessionState: def _session_state_from_stored(stored: StoredAgentSession) -> AgentSessionState:
@@ -1987,6 +2012,7 @@ def _ensure_session_title(
base_url: str, base_url: str,
api_key: str, api_key: str,
previous_title: str | None = None, previous_title: str | None = None,
fallback_title: str | None = None,
) -> None: ) -> None:
path = _session_json_path(directory, session_id) path = _session_json_path(directory, session_id)
try: try:
@@ -1998,13 +2024,28 @@ def _ensure_session_title(
_write_session_metadata(path, data) _write_session_metadata(path, data)
return return
title = data.get('title') title = data.get('title')
if isinstance(title, str) and title.strip(): title_source = data.get('title_source')
if (
isinstance(title, str)
and title.strip()
and title_source != 'first_message'
):
return return
messages = data.get('messages') messages = data.get('messages')
if not isinstance(messages, list): if not isinstance(messages, list):
if fallback_title and not (isinstance(title, str) and title.strip()):
data['title'] = fallback_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(path, data)
return return
user_messages = _session_user_messages(messages) user_messages = _session_user_messages(messages)
if not user_messages: if not user_messages:
if fallback_title and not (isinstance(title, str) and title.strip()):
data['title'] = fallback_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(path, data)
return return
generated = _generate_session_title( generated = _generate_session_title(
user_messages, user_messages,
@@ -2013,6 +2054,11 @@ def _ensure_session_title(
api_key=api_key, api_key=api_key,
) )
if not generated: if not generated:
if fallback_title and not (isinstance(title, str) and title.strip()):
data['title'] = fallback_title
data['title_source'] = 'first_message'
data['title_generated_at'] = int(time.time())
_write_session_metadata(path, data)
return return
data['title'] = generated data['title'] = generated
data['title_source'] = 'llm' data['title_source'] = 'llm'
@@ -2060,28 +2106,19 @@ def _generate_session_title(
{'role': 'user', 'content': conversation}, {'role': 'user', 'content': conversation},
], ],
} }
req = request.Request( client = OpenAICompatClient(
_join_url(base_url, '/chat/completions'), ModelConfig(
data=json.dumps(payload).encode('utf-8'), model=model,
headers={ base_url=base_url,
'Authorization': f'Bearer {api_key}', api_key=api_key,
'Content-Type': 'application/json', temperature=0.2,
}, )
method='POST',
) )
try: try:
with request.urlopen(req, timeout=15) as response: turn = client.complete(payload['messages'], tools=[])
response_payload = json.loads(response.read().decode('utf-8')) except OpenAICompatError:
except (error.HTTPError, error.URLError, OSError, json.JSONDecodeError):
return None return None
choices = response_payload.get('choices') return _clean_session_title(turn.content)
if not isinstance(choices, list) or not choices:
return None
message = choices[0].get('message') if isinstance(choices[0], dict) else None
content = message.get('content') if isinstance(message, dict) else None
if not isinstance(content, str):
return None
return _clean_session_title(content)
def _clean_session_title(value: str) -> str | None: def _clean_session_title(value: str) -> str | None:
@@ -2089,7 +2126,7 @@ def _clean_session_title(value: str) -> str | None:
title = title.rstrip('。.!?') title = title.rstrip('。.!?')
if not title: if not title:
return None return None
return title[:24] return title[:24].rstrip()
def _clean_manual_session_title(value: str) -> str | None: def _clean_manual_session_title(value: str) -> str | None:
@@ -129,6 +129,16 @@ type ContextBudget = {
token_counter_accurate?: boolean; token_counter_accurate?: boolean;
}; };
function scheduleSessionListRefresh() {
if (typeof window === "undefined") return;
window.dispatchEvent(new Event("claw-sessions-changed"));
for (const delay of [300, 1000, 2500]) {
window.setTimeout(() => {
window.dispatchEvent(new Event("claw-sessions-changed"));
}, delay);
}
}
export const Thread: FC = () => { export const Thread: FC = () => {
useRefreshReplayedRun(); useRefreshReplayedRun();
return ( return (
@@ -1145,7 +1155,9 @@ function SkillInsertDialog() {
setSkills(payload.skills); setSkills(payload.skills);
const changedCount = payload.changed_files?.length ?? 0; const changedCount = payload.changed_files?.length ?? 0;
if (payload.updated) { if (payload.updated) {
const shortHash = payload.after ? payload.after.slice(0, 7) : "最新提交"; const shortHash = payload.after
? payload.after.slice(0, 7)
: "最新提交";
setStatus(`Skill 已更新到 ${shortHash},变更 ${changedCount} 个文件。`); setStatus(`Skill 已更新到 ${shortHash},变更 ${changedCount} 个文件。`);
} else { } else {
setStatus("Skill 已是最新。"); setStatus("Skill 已是最新。");
@@ -1761,6 +1773,7 @@ const ImeComposerInput: FC<ComponentProps<typeof TextareaAutosize>> = ({
const clearAfterSubmit = () => { const clearAfterSubmit = () => {
justSubmittedRef.current = true; justSubmittedRef.current = true;
setLocalText(""); setLocalText("");
scheduleSessionListRefresh();
}; };
const unsubscribeComposer = aui.on("composer.send", clearAfterSubmit); const unsubscribeComposer = aui.on("composer.send", clearAfterSubmit);
const unsubscribeRun = aui.on("thread.runStart", clearAfterSubmit); const unsubscribeRun = aui.on("thread.runStart", clearAfterSubmit);
+126 -1
View File
@@ -13,16 +13,21 @@ import tempfile
import unittest import unittest
import json import json
from pathlib import Path from pathlib import Path
from unittest.mock import patch
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from backend.api.server import ( from backend.api.server import (
AgentState, AgentState,
create_app, create_app,
_derive_initial_session_title,
_ensure_session_title,
_generate_session_title,
_runtime_event_stage, _runtime_event_stage,
_save_in_progress_session,
_sanitize_stored_session_for_resume, _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 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'], '手动标题')
self.assertEqual(data['title_source'], 'manual') 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: def test_clear_runtime_state(self) -> None:
with tempfile.TemporaryDirectory() as d: with tempfile.TemporaryDirectory() as d:
client, _ = _build_client(Path(d)) client, _ = _build_client(Path(d))