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