Fix session title generation
This commit is contained in:
+68
-31
@@ -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
@@ -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))
|
||||||
|
|||||||
Reference in New Issue
Block a user