diff --git a/backend/api/server.py b/backend/api/server.py index 7055000..779841d 100644 --- a/backend/api/server.py +++ b/backend/api/server.py @@ -38,6 +38,7 @@ from src.agent_types import ( ModelConfig, ) 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 ( DEFAULT_AGENT_SESSION_DIR, StoredAgentSession, @@ -1176,17 +1177,18 @@ def create_app(state: AgentState) -> FastAPI: state.run_manager.record_event(run_record.run_id, 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: - _save_in_progress_session( + fallback_title = _save_in_progress_session( directory=session_directory, agent=agent, session_id=requested_session_id, prompt=request.prompt.strip(), ) - previous_title = _read_session_title( - session_directory, - requested_session_id, - ) if run_lock.locked(): queued_event = { 'type': 'run_queued', @@ -1297,6 +1299,7 @@ def create_app(state: AgentState) -> FastAPI: base_url=config.base_url, api_key=config.api_key, previous_title=previous_title, + fallback_title=fallback_title, ) _annotate_transcript_elapsed(payload, elapsed_ms) state.run_manager.finish( @@ -1728,14 +1731,15 @@ def _save_in_progress_session( agent: LocalCodingAgent, session_id: str | None, prompt: str, -) -> None: +) -> str | None: safe_id = _safe_session_id(session_id) if safe_id is None or not prompt: - return + return None if _session_json_path(directory, safe_id).exists(): - return + return None # 新会话的第一轮执行可能很久。先写一个运行中占位,避免刷新页面后找不到会话。 + initial_title = _derive_initial_session_title(prompt) scratchpad_directory = ( agent.runtime_config.scratchpad_root / safe_id / 'scratchpad' ).resolve() @@ -1764,9 +1768,30 @@ def _save_in_progress_session( plugin_state={}, 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: - 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: @@ -1987,6 +2012,7 @@ def _ensure_session_title( base_url: str, api_key: str, previous_title: str | None = None, + fallback_title: str | None = None, ) -> None: path = _session_json_path(directory, session_id) try: @@ -1998,13 +2024,28 @@ def _ensure_session_title( _write_session_metadata(path, data) return 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 messages = data.get('messages') 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 user_messages = _session_user_messages(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 generated = _generate_session_title( user_messages, @@ -2013,6 +2054,11 @@ def _ensure_session_title( api_key=api_key, ) 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 data['title'] = generated data['title_source'] = 'llm' @@ -2060,28 +2106,19 @@ def _generate_session_title( {'role': 'user', 'content': conversation}, ], } - req = request.Request( - _join_url(base_url, '/chat/completions'), - data=json.dumps(payload).encode('utf-8'), - headers={ - 'Authorization': f'Bearer {api_key}', - 'Content-Type': 'application/json', - }, - method='POST', + client = OpenAICompatClient( + ModelConfig( + model=model, + base_url=base_url, + api_key=api_key, + temperature=0.2, + ) ) try: - with request.urlopen(req, timeout=15) as response: - response_payload = json.loads(response.read().decode('utf-8')) - except (error.HTTPError, error.URLError, OSError, json.JSONDecodeError): + turn = client.complete(payload['messages'], tools=[]) + except OpenAICompatError: return None - choices = response_payload.get('choices') - 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) + return _clean_session_title(turn.content) def _clean_session_title(value: str) -> str | None: @@ -2089,7 +2126,7 @@ def _clean_session_title(value: str) -> str | None: title = title.rstrip('。.!!??') if not title: return None - return title[:24] + return title[:24].rstrip() def _clean_manual_session_title(value: str) -> str | None: diff --git a/frontend/app/components/assistant-ui/thread.tsx b/frontend/app/components/assistant-ui/thread.tsx index bc871f5..91afd11 100644 --- a/frontend/app/components/assistant-ui/thread.tsx +++ b/frontend/app/components/assistant-ui/thread.tsx @@ -129,6 +129,16 @@ type ContextBudget = { 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 = () => { useRefreshReplayedRun(); return ( @@ -1145,7 +1155,9 @@ function SkillInsertDialog() { setSkills(payload.skills); const changedCount = payload.changed_files?.length ?? 0; 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} 个文件。`); } else { setStatus("Skill 已是最新。"); @@ -1761,6 +1773,7 @@ const ImeComposerInput: FC> = ({ const clearAfterSubmit = () => { justSubmittedRef.current = true; setLocalText(""); + scheduleSessionListRefresh(); }; const unsubscribeComposer = aui.on("composer.send", clearAfterSubmit); const unsubscribeRun = aui.on("thread.runStart", clearAfterSubmit); diff --git a/tests/test_gui_server.py b/tests/test_gui_server.py index 5146756..36131d1 100644 --- a/tests/test_gui_server.py +++ b/tests/test_gui_server.py @@ -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))