refactor(code-quality): 清理 clippy 存量告警(unwrap/clone/redundant 等)

- 移除无用克隆与冗余引用,减少不必要内存分配
- 规范 unwrap/expect 使用,修复可提前失败路径
- 修复 anthropic provider llm_timeout_secs 死代码并补全超时日志
- cargo fmt 统一格式
This commit is contained in:
oudecheng 2026-08-16 23:22:22 +08:00
parent 3faefc74b9
commit 1019dbe8cc
71 changed files with 1338 additions and 1505 deletions

View File

@ -337,15 +337,14 @@ fn filter_images_by_age_and_count(
.count(); .count();
let content = if original_image_count > filtered_image_count { let content = if original_image_count > filtered_image_count {
let notice = if exceeds_age_limit { if exceeds_age_limit {
format!( format!(
"{} [图片已过期:超出 {} 条消息范围]", "{} [图片已过期:超出 {} 条消息范围]",
message.content, max_age_rounds message.content, max_age_rounds
) )
} else { } else {
format!("{} [图片已过期:超出最大图片数量限制]", message.content) format!("{} [图片已过期:超出最大图片数量限制]", message.content)
}; }
notice
} else { } else {
message.content.clone() message.content.clone()
}; };
@ -614,7 +613,7 @@ impl LoopDetector {
.count(); .count();
// Warn every warn_every times // Warn every warn_every times
if consecutive > 0 && consecutive % self.config.warn_every == 0 { if consecutive > 0 && consecutive.is_multiple_of(self.config.warn_every) {
LoopDetectionResult::Warning(format!( LoopDetectionResult::Warning(format!(
"注意: 工具 '{}' 已连续执行 {} 次,参数相同。如果任务没有进展,请尝试其他方法。", "注意: 工具 '{}' 已连续执行 {} 次,参数相同。如果任务没有进展,请尝试其他方法。",
last.name, consecutive last.name, consecutive
@ -1139,7 +1138,7 @@ impl AgentLoop {
// 避免每轮 serde_json::to_string 全量序列化工具定义。 // 避免每轮 serde_json::to_string 全量序列化工具定义。
let tools_tokens = tools let tools_tokens = tools
.as_ref() .as_ref()
.map(|t| estimate_tokens_from_serialized_json(t)) .map(estimate_tokens_from_serialized_json)
.unwrap_or_default(); .unwrap_or_default();
for iteration in 0..self.max_iterations { for iteration in 0..self.max_iterations {
@ -1513,71 +1512,67 @@ impl AgentLoop {
.and_then(|m| m.usage.as_ref()) .and_then(|m| m.usage.as_ref())
.map(|u| u.prompt_tokens); .map(|u| u.prompt_tokens);
if let Some(prompt_tokens) = last_prompt_tokens { if let Some(prompt_tokens) = last_prompt_tokens
if compressor.should_compress_by_usage(prompt_tokens) { && compressor.should_compress_by_usage(prompt_tokens)
// 阶段 1a工程化压缩截断非子代理 tool 结果,仅改内存) {
// 参数内聚到 ContextCompressorAgentLoop 不持有截断 token 数 // 阶段 1a工程化压缩截断非子代理 tool 结果,仅改内存)
compressor.truncate_tool_results(&mut messages); // 参数内聚到 ContextCompressorAgentLoop 不持有截断 token 数
engineering_compaction_applied = true; compressor.truncate_tool_results(&mut messages);
engineering_compaction_applied = true;
tracing::info!(
iteration,
prompt_tokens,
threshold = compressor.threshold(),
"Engineering compaction applied (tool results truncated)"
);
// 阶段 1b重新估算判断是否需要 LLM 压缩30% 阈值)
let estimated = crate::agent::context_compressor::estimate_tokens(&messages);
if estimated > compressor.llm_compaction_threshold() {
tracing::info!( tracing::info!(
iteration, iteration,
prompt_tokens, estimated_tokens = estimated,
threshold = compressor.threshold(), llm_threshold = compressor.llm_compaction_threshold(),
"Engineering compaction applied (tool results truncated)" "LLM compaction triggered (still above 30% after engineering compaction)"
); );
// LLM 压缩失败时降级为仅工程化压缩,不中断 agent loop
// 阶段 1b重新估算判断是否需要 LLM 压缩30% 阈值) match compressor
let estimated = .compress_two_segment_with_provider(&messages, self.provider.as_ref())
crate::agent::context_compressor::estimate_tokens(&messages); .await
if estimated > compressor.llm_compaction_threshold() { {
tracing::info!( Ok(compressed) => {
iteration, // sink 失败时记日志但不中断——内存已压缩DB 未更新
estimated_tokens = estimated, // 下次 process 从 DB 加载时会重新触发压缩
llm_threshold = compressor.llm_compaction_threshold(), if let Some(sink) = compaction_sink
"LLM compaction triggered (still above 30% after engineering compaction)" && let Err(e) = sink.compact(&compressed).await
); {
// LLM 压缩失败时降级为仅工程化压缩,不中断 agent loop tracing::error!(
match compressor
.compress_two_segment_with_provider(
&messages,
self.provider.as_ref(),
)
.await
{
Ok(compressed) => {
// sink 失败时记日志但不中断——内存已压缩DB 未更新
// 下次 process 从 DB 加载时会重新触发压缩
if let Some(sink) = compaction_sink {
if let Err(e) = sink.compact(&compressed).await {
tracing::error!(
error = %e,
iteration,
"CompactionSink compact failed; \
in-memory messages still replaced, DB will be re-compacted next round"
);
}
}
messages = compressed;
compaction_performed = true;
}
Err(e) => {
tracing::warn!(
error = %e, error = %e,
iteration, iteration,
"LLM compaction failed; \ "CompactionSink compact failed; \
falling back to engineering-only compaction (in-memory truncated messages retained)" in-memory messages still replaced, DB will be re-compacted next round"
); );
// 不设置 compaction_performedmessages 保持工程化压缩后的状态
} }
messages = compressed;
compaction_performed = true;
}
Err(e) => {
tracing::warn!(
error = %e,
iteration,
"LLM compaction failed; \
falling back to engineering-only compaction (in-memory truncated messages retained)"
);
// 不设置 compaction_performedmessages 保持工程化压缩后的状态
} }
} else {
tracing::info!(
iteration,
estimated_tokens = estimated,
llm_threshold = compressor.llm_compaction_threshold(),
"Engineering compaction sufficient (under 30%), skipping LLM compaction"
);
} }
} else {
tracing::info!(
iteration,
estimated_tokens = estimated,
llm_threshold = compressor.llm_compaction_threshold(),
"Engineering compaction sufficient (under 30%), skipping LLM compaction"
);
} }
} }
} }
@ -2319,14 +2314,14 @@ mod tests {
fn test_should_execute_in_parallel_single_tool() { fn test_should_execute_in_parallel_single_tool() {
// Would need a proper setup with AgentLoop to test fully // Would need a proper setup with AgentLoop to test fully
// For now, just verify the logic: single tool should return false // For now, just verify the logic: single tool should return false
let calls = vec![ToolCall { let calls = [ToolCall {
id: "1".to_string(), id: "1".to_string(),
name: "test".to_string(), name: "test".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
}]; }];
// If there's only 1 tool, should return false regardless // If there's only 1 tool, should return false regardless
assert_eq!(calls.len() <= 1, true); assert!(calls.len() <= 1);
} }
#[test] #[test]
@ -2619,9 +2614,15 @@ mod tests {
let filtered = filter_images_by_age_and_count(&messages, 10, 3); let filtered = filter_images_by_age_and_count(&messages, 10, 3);
// 检查结果 // 检查结果
assert!(filtered[19].media_refs.len() > 0, "最新消息应保留图片"); assert!(!filtered[19].media_refs.is_empty(), "最新消息应保留图片");
assert!(filtered[15].media_refs.len() > 0, "age=4 的消息应保留图片"); assert!(
assert!(filtered[10].media_refs.len() > 0, "age=9 的消息应保留图片"); !filtered[15].media_refs.is_empty(),
"age=4 的消息应保留图片"
);
assert!(
!filtered[10].media_refs.is_empty(),
"age=9 的消息应保留图片"
);
assert_eq!(filtered[5].media_refs.len(), 0, "age=14 的消息图片应被过滤"); assert_eq!(filtered[5].media_refs.len(), 0, "age=14 的消息图片应被过滤");
assert!(filtered[5].content.contains("超出 10 条消息范围")); assert!(filtered[5].content.contains("超出 10 条消息范围"));
assert_eq!(filtered[0].media_refs.len(), 0, "age=19 的消息图片应被过滤"); assert_eq!(filtered[0].media_refs.len(), 0, "age=19 的消息图片应被过滤");
@ -3117,7 +3118,7 @@ mod tests {
assert!( assert!(
messages messages
.iter() .iter()
.all(|m| m.tool_calls.as_ref().map_or(true, |c| c.is_empty())), .all(|m| m.tool_calls.as_ref().is_none_or(|c| c.is_empty())),
"no assistant should have tool_calls remaining" "no assistant should have tool_calls remaining"
); );
} }

View File

@ -54,7 +54,7 @@ fn is_assistant_with_tool_calls(msg: &ChatMessage) -> bool {
&& msg && msg
.tool_calls .tool_calls
.as_ref() .as_ref()
.map_or(false, |calls| !calls.is_empty()) .is_some_and(|calls| !calls.is_empty())
} }
/// Parse a flat message list into atomic units. Orphaned tool results /// Parse a flat message list into atomic units. Orphaned tool results
@ -713,10 +713,8 @@ OLDER SEGMENT (events from earlier in the session):
let middle_units = &compressible[preserve_count..split]; let middle_units = &compressible[preserve_count..split];
// Step 4: Build middle segment messages and transcript // Step 4: Build middle segment messages and transcript
let middle_messages: Vec<ChatMessage> = middle_units let middle_messages: Vec<ChatMessage> =
.iter() middle_units.iter().flat_map(unit_to_messages).collect();
.flat_map(unit_to_messages)
.collect();
let middle_transcript = Self::build_transcript(&middle_messages); let middle_transcript = Self::build_transcript(&middle_messages);
// Step 5: Summarize middle segment with LLM (heavy prompt) // Step 5: Summarize middle segment with LLM (heavy prompt)
@ -1116,8 +1114,8 @@ mod tests {
fn test_chinese_tokens_higher_than_english() { fn test_chinese_tokens_higher_than_english() {
// Use more characters to make the content difference significant // Use more characters to make the content difference significant
// compared to JSON overhead (50 tokens per message) // compared to JSON overhead (50 tokens per message)
let english = vec![ChatMessage::user(&"abcdefghij".repeat(20))]; // 200 English chars let english = vec![ChatMessage::user("abcdefghij".repeat(20))]; // 200 English chars
let chinese = vec![ChatMessage::user(&"这是一个测试消息字".repeat(20))]; // 200 CJK chars (10 chars * 20) let chinese = vec![ChatMessage::user("这是一个测试消息字".repeat(20))]; // 200 CJK chars (10 chars * 20)
let english_tokens = estimate_tokens(&english); let english_tokens = estimate_tokens(&english);
let chinese_tokens = estimate_tokens(&chinese); let chinese_tokens = estimate_tokens(&chinese);
@ -1153,7 +1151,7 @@ mod tests {
let compressor = ContextCompressor::new(20); let compressor = ContextCompressor::new(20);
// Need more content to trigger compression with new weighted calculation // Need more content to trigger compression with new weighted calculation
// 200 English chars / 4 = 50 tokens, plus overhead // 200 English chars / 4 = 50 tokens, plus overhead
let messages = vec![ChatMessage::user(&"x".repeat(400))]; let messages = vec![ChatMessage::user("x".repeat(400))];
assert!(compressor.should_compress(&messages)); assert!(compressor.should_compress(&messages));
} }
@ -1257,7 +1255,7 @@ mod tests {
#[test] #[test]
fn test_chunk_messages_for_summary_splits_oversized_message() { fn test_chunk_messages_for_summary_splits_oversized_message() {
let messages = vec![ChatMessage::user(&"x".repeat(25))]; let messages = vec![ChatMessage::user("x".repeat(25))];
let chunks = ContextCompressor::chunk_messages_for_summary(&messages, 10); let chunks = ContextCompressor::chunk_messages_for_summary(&messages, 10);

View File

@ -321,17 +321,17 @@ pub(crate) fn sanitize_incomplete_tool_call_sequences(messages: &mut Vec<ChatMes
for i in (0..messages.len()).rev() { for i in (0..messages.len()).rev() {
let msg = &messages[i]; let msg = &messages[i];
if msg.role == "tool" { if msg.role == "tool"
if let Some(ref tc_id) = msg.tool_call_id { && let Some(ref tc_id) = msg.tool_call_id
resolved_ids.insert(tc_id.clone()); {
} resolved_ids.insert(tc_id.clone());
} }
if msg.role == "assistant" if msg.role == "assistant"
&& msg && msg
.tool_calls .tool_calls
.as_ref() .as_ref()
.map_or(false, |calls| !calls.is_empty()) .is_some_and(|calls| !calls.is_empty())
{ {
let tool_calls = msg.tool_calls.as_ref().unwrap(); let tool_calls = msg.tool_calls.as_ref().unwrap();
let all_have_results = tool_calls.iter().all(|tc| resolved_ids.contains(&tc.id)); let all_have_results = tool_calls.iter().all(|tc| resolved_ids.contains(&tc.id));
@ -379,34 +379,31 @@ pub(crate) fn sanitize_incomplete_tool_call_sequences(messages: &mut Vec<ChatMes
// If we have pending tool_ids and encounter a non-tool message, // If we have pending tool_ids and encounter a non-tool message,
// the assistant's tool results were NOT immediately following. // the assistant's tool results were NOT immediately following.
if !pending_tool_ids.is_empty() && m.role != "tool" { if !pending_tool_ids.is_empty() && m.role != "tool" {
if let Some(idx) = pending_assistant_idx { if let Some(idx) = pending_assistant_idx
if !remove_indices.contains(&idx) { && !remove_indices.contains(&idx)
tracing::warn!( {
message_index = idx, tracing::warn!(
interrupted_by_index = i, message_index = idx,
interrupted_by_role = %m.role, interrupted_by_index = i,
pending_tool_call_count = pending_tool_ids.len(), interrupted_by_role = %m.role,
"Removing assistant with tool_calls — tool results \ pending_tool_call_count = pending_tool_ids.len(),
not immediately following (interrupted by non-tool message)" "Removing assistant with tool_calls — tool results \
); not immediately following (interrupted by non-tool message)"
// Remove this assistant's tool_call_ids from with_parent );
// so Phase 2 cleans up the now-orphaned tool messages // Remove this assistant's tool_call_ids from with_parent
if let Some(calls) = messages[idx].tool_calls.as_ref() { // so Phase 2 cleans up the now-orphaned tool messages
for tc in calls.iter() { if let Some(calls) = messages[idx].tool_calls.as_ref() {
with_parent.remove(&tc.id); for tc in calls.iter() {
} with_parent.remove(&tc.id);
} }
remove_indices.push(idx);
} }
remove_indices.push(idx);
} }
pending_tool_ids.clear(); pending_tool_ids.clear();
pending_assistant_idx = None; pending_assistant_idx = None;
} }
if m.role == "assistant" if m.role == "assistant" && m.tool_calls.as_ref().is_some_and(|calls| !calls.is_empty())
&& m.tool_calls
.as_ref()
.map_or(false, |calls| !calls.is_empty())
{ {
let already_marked = remove_indices.contains(&i); let already_marked = remove_indices.contains(&i);
if !already_marked { if !already_marked {
@ -419,32 +416,31 @@ pub(crate) fn sanitize_incomplete_tool_call_sequences(messages: &mut Vec<ChatMes
.collect(); .collect();
pending_assistant_idx = Some(i); pending_assistant_idx = Some(i);
} }
} else if m.role == "tool" { } else if m.role == "tool"
if let Some(ref tc_id) = m.tool_call_id { && let Some(ref tc_id) = m.tool_call_id
pending_tool_ids.remove(tc_id); {
if pending_tool_ids.is_empty() { pending_tool_ids.remove(tc_id);
pending_assistant_idx = None; if pending_tool_ids.is_empty() {
} pending_assistant_idx = None;
} }
} }
} }
// Handle trailing assistant with unresolved immediate tool results // Handle trailing assistant with unresolved immediate tool results
if !pending_tool_ids.is_empty() { if !pending_tool_ids.is_empty()
if let Some(idx) = pending_assistant_idx { && let Some(idx) = pending_assistant_idx
if !remove_indices.contains(&idx) { && !remove_indices.contains(&idx)
tracing::warn!( {
message_index = idx, tracing::warn!(
"Removing trailing assistant with incomplete immediate tool results" message_index = idx,
); "Removing trailing assistant with incomplete immediate tool results"
if let Some(calls) = messages[idx].tool_calls.as_ref() { );
for tc in calls.iter() { if let Some(calls) = messages[idx].tool_calls.as_ref() {
with_parent.remove(&tc.id); for tc in calls.iter() {
} with_parent.remove(&tc.id);
}
remove_indices.push(idx);
} }
} }
remove_indices.push(idx);
} }
} }
@ -939,7 +935,7 @@ fn format_tool_arguments_json(value: &serde_json::Value) -> String {
match value { match value {
serde_json::Value::Object(map) => { serde_json::Value::Object(map) => {
let mut entries: Vec<_> = map.iter().collect(); let mut entries: Vec<_> = map.iter().collect();
entries.sort_by(|(left, _), (right, _)| left.cmp(right)); entries.sort_by_key(|(left, _)| *left);
let body = entries let body = entries
.into_iter() .into_iter()
.map(|(key, value)| { .map(|(key, value)| {

View File

@ -234,10 +234,10 @@ impl FeishuChannel {
// 1. Check cache // 1. Check cache
{ {
let cached = self.tenant_token.read().await; let cached = self.tenant_token.read().await;
if let Some(ref token) = *cached { if let Some(ref token) = *cached
if Instant::now() < token.refresh_after { && Instant::now() < token.refresh_after
return Ok(token.value.clone()); {
} return Ok(token.value.clone());
} }
} }
@ -1076,10 +1076,10 @@ impl FeishuChannel {
.await?; .await?;
// Fetch and prepend quoted message content if this is a reply // Fetch and prepend quoted message content if this is a reply
if let Some(ref pid) = parent_id { if let Some(ref pid) = parent_id
if let Some(reply_ctx) = self.get_message_content(pid).await { && let Some(reply_ctx) = self.get_message_content(pid).await
content = format!("{}\n{}", reply_ctx, content); {
} content = format!("{}\n{}", reply_ctx, content);
} }
#[cfg(debug_assertions)] #[cfg(debug_assertions)]
@ -1532,15 +1532,15 @@ fn parse_post_content(content: &str) -> String {
// Fall back: try any dict child // Fall back: try any dict child
if let Some(root_obj) = root.as_object() { if let Some(root_obj) = root.as_object() {
for (_key, val) in root_obj { for (_key, val) in root_obj {
if let Some(obj) = val.as_object() { if let Some(obj) = val.as_object()
if obj.get("content").and_then(|c| c.as_array()).is_some() { && obj.get("content").and_then(|c| c.as_array()).is_some()
parse_block(val, &mut texts); {
let result = texts.join(""); parse_block(val, &mut texts);
if !result.trim().is_empty() { let result = texts.join("");
return result.trim().to_string(); if !result.trim().is_empty() {
} return result.trim().to_string();
texts.clear();
} }
texts.clear();
} }
} }
} }
@ -1565,21 +1565,20 @@ fn extract_interactive_content(content: &str) -> Result<(String, Option<MediaIte
} }
// Extract from card object // Extract from card object
if let Some(card) = parsed.get("card").and_then(|c| c.as_object()) { if let Some(card) = parsed.get("card").and_then(|c| c.as_object())
if let Some(elements) = card.get("elements").and_then(|e| e.as_array()) { && let Some(elements) = card.get("elements").and_then(|e| e.as_array())
for el in elements { {
extract_element_content(el, &mut texts); for el in elements {
} extract_element_content(el, &mut texts);
} }
} }
// Extract from header // Extract from header
if let Some(header) = parsed.get("header").and_then(|h| h.as_object()) { if let Some(header) = parsed.get("header").and_then(|h| h.as_object())
if let Some(title) = header.get("title").and_then(|t| t.as_object()) { && let Some(title) = header.get("title").and_then(|t| t.as_object())
if let Some(text) = title.get("content").and_then(|c| c.as_str()) { && let Some(text) = title.get("content").and_then(|c| c.as_str())
texts.push(format!("title: {}\n", text)); {
} texts.push(format!("title: {}\n", text));
}
} }
let result = texts.join("").trim().to_string(); let result = texts.join("").trim().to_string();
@ -1734,14 +1733,12 @@ fn collect_list_items(items: &[serde_json::Value], lines: &mut Vec<String>, dept
None None
} }
}) })
}) { }) && let Some(children) = children_arr
if let Some(children) = children_arr .as_object()
.as_object() .and_then(|o| o.get("children"))
.and_then(|o| o.get("children")) .and_then(|c| c.as_array())
.and_then(|c| c.as_array()) {
{ collect_list_items(children, lines, depth + 1);
collect_list_items(children, lines, depth + 1);
}
} }
} }
} }
@ -2269,138 +2266,6 @@ fn sanitize_download_file_name(file_name: &str) -> String {
.to_string() .to_string()
} }
#[cfg(test)]
mod tests {
use super::{
FeishuChannel, MsgFormat, extract_file_name_from_content_disposition,
infer_download_filename, parse_post_content, sanitize_download_file_name,
};
#[test]
fn markdown_post_uses_md_tag() {
let content = "**bold**\n1. item1\n2. item2\n[link](https://open.feishu.cn)";
let post = FeishuChannel::markdown_to_post(content);
let parsed: serde_json::Value = serde_json::from_str(&post).unwrap();
assert_eq!(parsed["zh_cn"]["content"][0][0]["tag"], "md");
assert_eq!(parsed["zh_cn"]["content"][0][0]["text"], content);
}
#[test]
fn multiline_markdown_is_not_misclassified_as_plain_post() {
let content = "intro\n1. item1\n2. item2";
assert_eq!(FeishuChannel::detect_msg_format(content), MsgFormat::Post);
}
#[test]
fn headings_still_use_interactive() {
let content = "intro\n## heading";
assert_eq!(
FeishuChannel::detect_msg_format(content),
MsgFormat::Interactive
);
}
#[test]
fn infer_download_filename_prefers_original_file_name() {
let content = serde_json::json!({
"file_key": "file_key_123",
"file_name": "demo-archive.zip"
});
let headers = reqwest::header::HeaderMap::new();
let filename =
infer_download_filename(&content, &headers, "om_123", "file_key_123", "file");
assert_eq!(filename, "om_123_demo-archive.zip");
}
#[test]
fn infer_download_filename_uses_content_disposition_when_message_lacks_name() {
let content = serde_json::json!({
"file_key": "file_key_123"
});
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::CONTENT_DISPOSITION,
reqwest::header::HeaderValue::from_static("attachment; filename=meeting-notes.zip"),
);
let filename =
infer_download_filename(&content, &headers, "om_123", "file_key_123", "file");
assert_eq!(filename, "om_123_meeting-notes.zip");
}
#[test]
fn infer_download_filename_falls_back_to_bin_without_name() {
let content = serde_json::json!({
"file_key": "file_key_123"
});
let headers = reqwest::header::HeaderMap::new();
let filename =
infer_download_filename(&content, &headers, "om_123", "file_key_123", "file");
assert_eq!(filename, "om_123_file_key.bin");
}
#[test]
fn sanitize_download_file_name_replaces_path_separators() {
let sanitized = sanitize_download_file_name("../../demo/archive.zip");
assert_eq!(sanitized, "_.._demo_archive.zip");
}
#[test]
fn extract_file_name_from_content_disposition_supports_filename_star() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::CONTENT_DISPOSITION,
reqwest::header::HeaderValue::from_static("attachment; filename*=UTF-8''archive.zip"),
);
let file_name = extract_file_name_from_content_disposition(&headers);
assert_eq!(file_name.as_deref(), Some("archive.zip"));
}
#[test]
fn parse_post_content_handles_code_block_with_content_array() {
// Test parsing code_block with content array (standard Feishu format)
let post_json = r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","language":"python","content":[{"tag":"text","text":"def hello():"},{"tag":"text","text":" print('world')"}]}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```python"));
assert!(result.contains("def hello():"));
assert!(result.contains("print('world')"));
}
#[test]
fn parse_post_content_handles_code_block_with_fallback_text() {
// Backwards compatibility: some formats might use text field directly
let post_json = r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","language":"rust","text":"fn main() {}"}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```rust"));
assert!(result.contains("fn main() {}"));
}
#[test]
fn parse_post_content_handles_code_block_without_language() {
// Test code_block without language field
let post_json = r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","content":[{"tag":"text","text":"plain text"}]}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```"));
assert!(result.contains("plain text"));
}
#[test]
fn parse_post_content_handles_empty_code_block() {
// Test code_block with empty content
let post_json =
r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","language":"go"}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```go"));
}
}
#[async_trait] #[async_trait]
impl Channel for FeishuChannel { impl Channel for FeishuChannel {
fn name(&self) -> &str { fn name(&self) -> &str {
@ -2502,7 +2367,7 @@ impl Channel for FeishuChannel {
let receive_id = if msg.chat_id.starts_with("oc_") { let receive_id = if msg.chat_id.starts_with("oc_") {
&msg.chat_id &msg.chat_id
} else { } else {
&msg.reply_to.as_ref().unwrap_or(&msg.chat_id) msg.reply_to.as_ref().unwrap_or(&msg.chat_id)
}; };
let receive_id_type = if msg.chat_id.starts_with("oc_") { let receive_id_type = if msg.chat_id.starts_with("oc_") {
"chat_id" "chat_id"
@ -2671,3 +2536,135 @@ impl Channel for FeishuChannel {
Ok(()) Ok(())
} }
} }
#[cfg(test)]
mod tests {
use super::{
FeishuChannel, MsgFormat, extract_file_name_from_content_disposition,
infer_download_filename, parse_post_content, sanitize_download_file_name,
};
#[test]
fn markdown_post_uses_md_tag() {
let content = "**bold**\n1. item1\n2. item2\n[link](https://open.feishu.cn)";
let post = FeishuChannel::markdown_to_post(content);
let parsed: serde_json::Value = serde_json::from_str(&post).unwrap();
assert_eq!(parsed["zh_cn"]["content"][0][0]["tag"], "md");
assert_eq!(parsed["zh_cn"]["content"][0][0]["text"], content);
}
#[test]
fn multiline_markdown_is_not_misclassified_as_plain_post() {
let content = "intro\n1. item1\n2. item2";
assert_eq!(FeishuChannel::detect_msg_format(content), MsgFormat::Post);
}
#[test]
fn headings_still_use_interactive() {
let content = "intro\n## heading";
assert_eq!(
FeishuChannel::detect_msg_format(content),
MsgFormat::Interactive
);
}
#[test]
fn infer_download_filename_prefers_original_file_name() {
let content = serde_json::json!({
"file_key": "file_key_123",
"file_name": "demo-archive.zip"
});
let headers = reqwest::header::HeaderMap::new();
let filename =
infer_download_filename(&content, &headers, "om_123", "file_key_123", "file");
assert_eq!(filename, "om_123_demo-archive.zip");
}
#[test]
fn infer_download_filename_uses_content_disposition_when_message_lacks_name() {
let content = serde_json::json!({
"file_key": "file_key_123"
});
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::CONTENT_DISPOSITION,
reqwest::header::HeaderValue::from_static("attachment; filename=meeting-notes.zip"),
);
let filename =
infer_download_filename(&content, &headers, "om_123", "file_key_123", "file");
assert_eq!(filename, "om_123_meeting-notes.zip");
}
#[test]
fn infer_download_filename_falls_back_to_bin_without_name() {
let content = serde_json::json!({
"file_key": "file_key_123"
});
let headers = reqwest::header::HeaderMap::new();
let filename =
infer_download_filename(&content, &headers, "om_123", "file_key_123", "file");
assert_eq!(filename, "om_123_file_key.bin");
}
#[test]
fn sanitize_download_file_name_replaces_path_separators() {
let sanitized = sanitize_download_file_name("../../demo/archive.zip");
assert_eq!(sanitized, "_.._demo_archive.zip");
}
#[test]
fn extract_file_name_from_content_disposition_supports_filename_star() {
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
reqwest::header::CONTENT_DISPOSITION,
reqwest::header::HeaderValue::from_static("attachment; filename*=UTF-8''archive.zip"),
);
let file_name = extract_file_name_from_content_disposition(&headers);
assert_eq!(file_name.as_deref(), Some("archive.zip"));
}
#[test]
fn parse_post_content_handles_code_block_with_content_array() {
// Test parsing code_block with content array (standard Feishu format)
let post_json = r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","language":"python","content":[{"tag":"text","text":"def hello():"},{"tag":"text","text":" print('world')"}]}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```python"));
assert!(result.contains("def hello():"));
assert!(result.contains("print('world')"));
}
#[test]
fn parse_post_content_handles_code_block_with_fallback_text() {
// Backwards compatibility: some formats might use text field directly
let post_json = r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","language":"rust","text":"fn main() {}"}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```rust"));
assert!(result.contains("fn main() {}"));
}
#[test]
fn parse_post_content_handles_code_block_without_language() {
// Test code_block without language field
let post_json = r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","content":[{"tag":"text","text":"plain text"}]}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```"));
assert!(result.contains("plain text"));
}
#[test]
fn parse_post_content_handles_empty_code_block() {
// Test code_block with empty content
let post_json =
r#"{"post":{"zh_cn":{"content":[[{"tag":"code_block","language":"go"}]]}}}"#;
let result = parse_post_content(post_json);
assert!(result.contains("```go"));
}
}

View File

@ -18,6 +18,12 @@ pub struct ChannelManager {
websocket_channel: Arc<CliChannel>, websocket_channel: Arc<CliChannel>,
} }
impl Default for ChannelManager {
fn default() -> Self {
Self::new()
}
}
impl ChannelManager { impl ChannelManager {
pub fn new() -> Self { pub fn new() -> Self {
let websocket_channel = Arc::new(CliChannel::new()); let websocket_channel = Arc::new(CliChannel::new());

View File

@ -69,9 +69,7 @@ impl WechatChannel {
let path = media.path.clone(); let path = media.path.clone();
let data = tokio::task::spawn_blocking(move || std::fs::read(&path)) let data = tokio::task::spawn_blocking(move || std::fs::read(&path))
.await .await
.map_err(|e| { .map_err(|e| ChannelError::SendError(format!("WeChat media read task failed: {}", e)))?
ChannelError::SendError(format!("WeChat media read task failed: {}", e))
})?
.map_err(|error| { .map_err(|error| {
ChannelError::SendError(format!( ChannelError::SendError(format!(
"WeChat media read failed for '{}': {}", "WeChat media read failed for '{}': {}",
@ -419,7 +417,9 @@ mod tests {
std::fs::rename(file.path(), &image_path).unwrap(); std::fs::rename(file.path(), &image_path).unwrap();
let media = MediaItem::new(image_path.to_string_lossy().to_string(), "image"); let media = MediaItem::new(image_path.to_string_lossy().to_string(), "image");
let content = WechatChannel::media_to_send_content(&media, None).await.unwrap(); let content = WechatChannel::media_to_send_content(&media, None)
.await
.unwrap();
assert!(matches!(content, SendContent::Image { .. })); assert!(matches!(content, SendContent::Image { .. }));
} }
@ -432,8 +432,9 @@ mod tests {
std::fs::rename(file.path(), &doc_path).unwrap(); std::fs::rename(file.path(), &doc_path).unwrap();
let media = MediaItem::new(doc_path.to_string_lossy().to_string(), "file"); let media = MediaItem::new(doc_path.to_string_lossy().to_string(), "file");
let content = let content = WechatChannel::media_to_send_content(&media, Some("note".to_string()))
WechatChannel::media_to_send_content(&media, Some("note".to_string())).await.unwrap(); .await
.unwrap();
match content { match content {
SendContent::File { SendContent::File {

View File

@ -209,11 +209,11 @@ impl InitWizard {
"2" => return self.modify_provider(existing).await, "2" => return self.modify_provider(existing).await,
"3" => { "3" => {
println!("Keeping existing providers."); println!("Keeping existing providers.");
return Ok(existing.providers.clone()); Ok(existing.providers.clone())
} }
"4" => { "4" => {
println!("Skipping provider configuration."); println!("Skipping provider configuration.");
return Ok(existing.providers.clone()); Ok(existing.providers.clone())
} }
_ => { _ => {
println!("Invalid option, adding new provider."); println!("Invalid option, adding new provider.");
@ -378,16 +378,16 @@ impl InitWizard {
match choice.as_str() { match choice.as_str() {
"1" => { "1" => {
println!("Keeping existing models."); println!("Keeping existing models.");
return Ok(existing.models.clone()); Ok(existing.models.clone())
} }
"2" => return self.add_model(existing).await, "2" => return self.add_model(existing).await,
"3" => { "3" => {
println!("Skipping model configuration."); println!("Skipping model configuration.");
return Ok(existing.models.clone()); Ok(existing.models.clone())
} }
_ => { _ => {
println!("Invalid option, keeping existing models."); println!("Invalid option, keeping existing models.");
return Ok(existing.models.clone()); Ok(existing.models.clone())
} }
} }
} else { } else {
@ -505,11 +505,11 @@ impl InitWizard {
"2" => return self.modify_agent(existing, providers, models).await, "2" => return self.modify_agent(existing, providers, models).await,
"3" => { "3" => {
println!("Keeping existing agents."); println!("Keeping existing agents.");
return Ok(existing.agents.clone()); Ok(existing.agents.clone())
} }
"4" => { "4" => {
println!("Skipping agent configuration."); println!("Skipping agent configuration.");
return Ok(existing.agents.clone()); Ok(existing.agents.clone())
} }
_ => { _ => {
println!("Invalid option, adding new agent."); println!("Invalid option, adding new agent.");

View File

@ -42,12 +42,11 @@ pub async fn run(gateway_url: &str) -> Result<(), Box<dyn std::error::Error>> {
let text = text.to_string(); let text = text.to_string();
if let Ok(outbound) = parse_message(&text) { if let Ok(outbound) = parse_message(&text) {
match outbound { match outbound {
WsOutbound::AssistantResponse { id, content, .. } => { WsOutbound::AssistantResponse { id, content, .. }
// Skip if already fully streamed via StreamDelta // Skip if already fully streamed via StreamDelta
if !streamed_message_ids.remove(&id) { if !streamed_message_ids.remove(&id) => {
input.write_response(&content).await?; input.write_response(&content).await?;
} }
}
WsOutbound::ToolCall { tool_name, arguments, .. } => { WsOutbound::ToolCall { tool_name, arguments, .. } => {
input.write_output(&format!("Tool call: {}\n{}\n", tool_name, format_json(&arguments))).await?; input.write_output(&format!("Tool call: {}\n{}\n", tool_name, format_json(&arguments))).await?;
} }
@ -235,12 +234,11 @@ pub async fn run(gateway_url: &str) -> Result<(), Box<dyn std::error::Error>> {
chat_id: current_session_id.clone(), chat_id: current_session_id.clone(),
sender_id: None, sender_id: None,
}; };
if let Ok(text) = serialize_inbound(&inbound) { if let Ok(text) = serialize_inbound(&inbound)
if sender.send(Message::Text(text.into())).await.is_err() { && sender.send(Message::Text(text.into())).await.is_err() {
tracing::error!("Failed to send message to gateway"); tracing::error!("Failed to send message to gateway");
break; break;
} }
}
} }
} }
} }

View File

@ -138,10 +138,10 @@ async fn handle_get_current_session(
.with_message(MessageKind::Notification, &message) .with_message(MessageKind::Notification, &message)
.with_metadata("topic_id", &topic.id) .with_metadata("topic_id", &topic.id)
.with_metadata("title", &topic.title) .with_metadata("title", &topic.title)
.with_metadata("message_count", &actual_message_count.to_string()) .with_metadata("message_count", actual_message_count.to_string())
.with_metadata("estimated_tokens", &total_tokens.to_string()) .with_metadata("estimated_tokens", total_tokens.to_string())
.with_metadata("system_prompt_tokens", &system_prompt_tokens.to_string()) .with_metadata("system_prompt_tokens", system_prompt_tokens.to_string())
.with_metadata("message_tokens", &message_tokens.to_string())) .with_metadata("message_tokens", message_tokens.to_string()))
} }
fn format_time_ago(timestamp_ms: i64) -> String { fn format_time_ago(timestamp_ms: i64) -> String {

View File

@ -57,5 +57,5 @@ async fn handle_list_channels(
Ok(CommandResponse::success(ctx.request_id) Ok(CommandResponse::success(ctx.request_id)
.with_message(MessageKind::Notification, &message) .with_message(MessageKind::Notification, &message)
.with_metadata("channels", &channels_json) .with_metadata("channels", &channels_json)
.with_metadata("count", &channels.len().to_string())) .with_metadata("count", channels.len().to_string()))
} }

View File

@ -85,10 +85,10 @@ async fn handle_list_sessions(
)); ));
// 显示描述(如果有) // 显示描述(如果有)
if let Some(ref desc) = topic.description { if let Some(ref desc) = topic.description
if !desc.is_empty() { && !desc.is_empty()
lines.push(format!(" {}", desc)); {
} lines.push(format!(" {}", desc));
} }
} }
@ -105,6 +105,6 @@ async fn handle_list_sessions(
Ok(CommandResponse::success(ctx.request_id) Ok(CommandResponse::success(ctx.request_id)
.with_message(MessageKind::Notification, &message) .with_message(MessageKind::Notification, &message)
.with_metadata("topics", &topics_json) .with_metadata("topics", &topics_json)
.with_metadata("count", &topics.len().to_string()) .with_metadata("count", topics.len().to_string())
.with_metadata("current_topic_id", current_topic_id)) .with_metadata("current_topic_id", current_topic_id))
} }

View File

@ -84,5 +84,5 @@ async fn handle_list_sessions_by_channel(
.with_message(MessageKind::Notification, &message) .with_message(MessageKind::Notification, &message)
.with_metadata("sessions", &sessions_json) .with_metadata("sessions", &sessions_json)
.with_metadata("channel_name", &channel_name) .with_metadata("channel_name", &channel_name)
.with_metadata("count", &summaries.len().to_string())) .with_metadata("count", summaries.len().to_string()))
} }

View File

@ -159,5 +159,5 @@ async fn handle_list_topics(
.with_message(MessageKind::Notification, &message) .with_message(MessageKind::Notification, &message)
.with_metadata("topics", &topics_json) .with_metadata("topics", &topics_json)
.with_metadata("session_id", &session_id) .with_metadata("session_id", &session_id)
.with_metadata("count", &summaries.len().to_string())) .with_metadata("count", summaries.len().to_string()))
} }

View File

@ -197,12 +197,12 @@ fn reconstruct_task_from_db(
/// New format: "Subagent [type]: description" /// New format: "Subagent [type]: description"
/// Legacy format: "Subagent: description" (defaults to "general") /// Legacy format: "Subagent: description" (defaults to "general")
fn parse_subagent_title(title: &str) -> (String, String) { fn parse_subagent_title(title: &str) -> (String, String) {
if let Some(rest) = title.strip_prefix("Subagent [") { if let Some(rest) = title.strip_prefix("Subagent [")
if let Some(bracket_pos) = rest.find("]: ") { && let Some(bracket_pos) = rest.find("]: ")
let agent_type = rest[..bracket_pos].to_string(); {
let desc = rest[bracket_pos + 3..].to_string(); let agent_type = rest[..bracket_pos].to_string();
return (agent_type, desc); let desc = rest[bracket_pos + 3..].to_string();
} return (agent_type, desc);
} }
let desc = title let desc = title
.strip_prefix("Subagent: ") .strip_prefix("Subagent: ")

View File

@ -60,5 +60,5 @@ async fn handle_load_topic(
.with_message(MessageKind::Notification, &topic.title) .with_message(MessageKind::Notification, &topic.title)
.with_metadata("topic_id", &topic.id) .with_metadata("topic_id", &topic.id)
.with_metadata("title", &topic.title) .with_metadata("title", &topic.title)
.with_metadata("message_count", &topic.message_count.to_string())) .with_metadata("message_count", topic.message_count.to_string()))
} }

View File

@ -95,7 +95,7 @@ async fn handle_rename_topic(
return Ok(CommandResponse::success(ctx.request_id) return Ok(CommandResponse::success(ctx.request_id)
.with_message( .with_message(
MessageKind::Notification, MessageKind::Notification,
&format!("✓ 话题标题未变化: {}", trimmed_title), format!("✓ 话题标题未变化: {}", trimmed_title),
) )
.with_metadata("topics", &topic_summaries_json) .with_metadata("topics", &topic_summaries_json)
.with_metadata("topic_id", &topic_id) .with_metadata("topic_id", &topic_id)

View File

@ -72,11 +72,12 @@ pub async fn save_session_to_file(
let output_path = resolve_filepath(filepath, &record); let output_path = resolve_filepath(filepath, &record);
// 创建父目录 // 创建父目录
if let Some(parent) = output_path.parent() { if let Some(parent) = output_path.parent()
if !parent.as_os_str().is_empty() && !parent.exists() { && !parent.as_os_str().is_empty()
std::fs::create_dir_all(parent) && !parent.exists()
.map_err(|e| format!("Failed to create directory: {}", e))?; {
} std::fs::create_dir_all(parent)
.map_err(|e| format!("Failed to create directory: {}", e))?;
} }
// 写入文件 // 写入文件
@ -192,7 +193,7 @@ async fn handle_save_session(
filepath, filepath,
include_all, include_all,
include_subagents, include_subagents,
&*handler.store, &handler.store,
Some(handler.task_repository.as_ref()), Some(handler.task_repository.as_ref()),
&*handler.system_prompt_provider, &*handler.system_prompt_provider,
) )
@ -213,16 +214,16 @@ async fn handle_save_session(
MessageKind::Notification, MessageKind::Notification,
// 路径中的反斜杠在 Markdown 渲染时会被当作转义符吃掉, // 路径中的反斜杠在 Markdown 渲染时会被当作转义符吃掉,
// 统一转换为正斜杠以保证显示完整(跨平台兼容) // 统一转换为正斜杠以保证显示完整(跨平台兼容)
&format!( format!(
"Session saved to: {}", "Session saved to: {}",
output_path.display().to_string().replace('\\', "/") output_path.display().to_string().replace('\\', "/")
), ),
) )
.with_metadata( .with_metadata(
"filepath", "filepath",
&output_path.display().to_string().replace('\\', "/"), output_path.display().to_string().replace('\\', "/"),
) )
.with_metadata("message_count", &message_count.to_string())) .with_metadata("message_count", message_count.to_string()))
} }
/// 子智能体任务数据 /// 子智能体任务数据
@ -391,21 +392,21 @@ pub fn generate_subagent_tasks_markdown(subagent_data: &[SubagentTaskData]) -> S
} }
// 工具调用 // 工具调用
if let Some(ref calls) = msg.tool_calls { if let Some(ref calls) = msg.tool_calls
if !calls.is_empty() { && !calls.is_empty()
output.push_str("**Tool Calls:**\n\n"); {
for call in calls { output.push_str("**Tool Calls:**\n\n");
output.push_str(&format!("- **{}** (`{}`)\n", call.name, call.id)); for call in calls {
output.push_str(" ```json\n"); output.push_str(&format!("- **{}** (`{}`)\n", call.name, call.id));
let args_json = serde_json::to_string_pretty(&call.arguments) output.push_str(" ```json\n");
.unwrap_or_else(|_| call.arguments.to_string()); let args_json = serde_json::to_string_pretty(&call.arguments)
for line in args_json.lines() { .unwrap_or_else(|_| call.arguments.to_string());
output.push_str(&format!(" {}\n", line)); for line in args_json.lines() {
} output.push_str(&format!(" {}\n", line));
output.push_str(" ```\n");
} }
output.push('\n'); output.push_str(" ```\n");
} }
output.push('\n');
} }
output.push_str("---\n\n"); output.push_str("---\n\n");
@ -560,21 +561,21 @@ pub fn generate_messages_markdown(messages: &[crate::bus::ChatMessage]) -> Strin
} }
// Tool calls // Tool calls
if let Some(ref calls) = msg.tool_calls { if let Some(ref calls) = msg.tool_calls
if !calls.is_empty() { && !calls.is_empty()
output.push_str("### Tool Calls\n\n"); {
for call in calls { output.push_str("### Tool Calls\n\n");
output.push_str(&format!("- **{}** (`{}`)\n", call.name, call.id)); for call in calls {
output.push_str(" ```json\n"); output.push_str(&format!("- **{}** (`{}`)\n", call.name, call.id));
let args_json = serde_json::to_string_pretty(&call.arguments) output.push_str(" ```json\n");
.unwrap_or_else(|_| call.arguments.to_string()); let args_json = serde_json::to_string_pretty(&call.arguments)
for line in args_json.lines() { .unwrap_or_else(|_| call.arguments.to_string());
output.push_str(&format!(" {}\n", line)); for line in args_json.lines() {
} output.push_str(&format!(" {}\n", line));
output.push_str(" ```\n");
} }
output.push('\n'); output.push_str(" ```\n");
} }
output.push('\n');
} }
// Media refs // Media refs
@ -621,16 +622,7 @@ pub fn resolve_filepath(filepath: Option<String>, record: &SessionRecord) -> Pat
// 生成安全标题(替换特殊字符) // 生成安全标题(替换特殊字符)
let safe_title = record let safe_title = record
.title .title
.replace(' ', "_") .replace([' ', '/', '\\', ':', '<', '>', '|', '?', '*', '"'], "_");
.replace('/', "_")
.replace('\\', "_")
.replace(':', "_")
.replace('<', "_")
.replace('>', "_")
.replace('|', "_")
.replace('?', "_")
.replace('*', "_")
.replace('"', "_");
// 使用标题或 session_id 作为文件名 // 使用标题或 session_id 作为文件名
let base_name = if safe_title.is_empty() { let base_name = if safe_title.is_empty() {
@ -716,7 +708,7 @@ impl InChatCommandHandler for SaveSessionInChatHandler {
filepath, filepath,
include_all, include_all,
include_subagents, include_subagents,
&*self.store, &self.store,
Some(self.task_repository.as_ref()), Some(self.task_repository.as_ref()),
&*self.system_prompt_provider, &*self.system_prompt_provider,
) )

View File

@ -54,11 +54,12 @@ pub async fn save_topic_to_file(
let output_path = resolve_topic_filepath(filepath, &topic); let output_path = resolve_topic_filepath(filepath, &topic);
// 创建父目录 // 创建父目录
if let Some(parent) = output_path.parent() { if let Some(parent) = output_path.parent()
if !parent.as_os_str().is_empty() && !parent.exists() { && !parent.as_os_str().is_empty()
std::fs::create_dir_all(parent) && !parent.exists()
.map_err(|e| format!("Failed to create directory: {}", e))?; {
} std::fs::create_dir_all(parent)
.map_err(|e| format!("Failed to create directory: {}", e))?;
} }
// 写入文件 // 写入文件
@ -138,16 +139,7 @@ fn resolve_topic_filepath(filepath: Option<String>, topic: &TopicRecord) -> Path
None => { None => {
let safe_title = topic let safe_title = topic
.title .title
.replace(' ', "_") .replace([' ', '/', '\\', ':', '<', '>', '|', '?', '*', '"'], "_");
.replace('/', "_")
.replace('\\', "_")
.replace(':', "_")
.replace('<', "_")
.replace('>', "_")
.replace('|', "_")
.replace('?', "_")
.replace('*', "_")
.replace('"', "_");
let base_name = if safe_title.is_empty() { let base_name = if safe_title.is_empty() {
format!("topic_{}", &topic.id[..8.min(topic.id.len())]) format!("topic_{}", &topic.id[..8.min(topic.id.len())])
@ -267,7 +259,7 @@ async fn handle_save_topic(
topic_id, topic_id,
filepath, filepath,
include_subagents, include_subagents,
&*handler.store, &handler.store,
Some(handler.task_repository.as_ref()), Some(handler.task_repository.as_ref()),
&*handler.system_prompt_provider, &*handler.system_prompt_provider,
&messages, &messages,
@ -282,14 +274,14 @@ async fn handle_save_topic(
MessageKind::Notification, MessageKind::Notification,
// 路径中的反斜杠在 Markdown 渲染时会被当作转义符吃掉, // 路径中的反斜杠在 Markdown 渲染时会被当作转义符吃掉,
// 统一转换为正斜杠以保证显示完整(跨平台兼容) // 统一转换为正斜杠以保证显示完整(跨平台兼容)
&format!( format!(
"Topic saved to: {}", "Topic saved to: {}",
output_path.display().to_string().replace('\\', "/") output_path.display().to_string().replace('\\', "/")
), ),
) )
.with_metadata( .with_metadata(
"filepath", "filepath",
&output_path.display().to_string().replace('\\', "/"), output_path.display().to_string().replace('\\', "/"),
) )
.with_metadata("message_count", &message_count.to_string())) .with_metadata("message_count", message_count.to_string()))
} }

View File

@ -94,13 +94,13 @@ async fn handle_create_session(
.ok_or_else(|| CommandError::new("NO_CHAT_ID", "No chat_id in context"))?; .ok_or_else(|| CommandError::new("NO_CHAT_ID", "No chat_id in context"))?;
// 如果有 SessionManager自动切换到新话题 // 如果有 SessionManager自动切换到新话题
if let Some(ref session_manager) = handler.session_manager { if let Some(ref session_manager) = handler.session_manager
if let Some(session) = session_manager.get(&ctx.channel_name).await { && let Some(session) = session_manager.get(&ctx.channel_name).await
let mut session_guard = session.lock().await; {
session_guard let mut session_guard = session.lock().await;
.switch_topic(chat_id, &topic.id) session_guard
.map_err(|e| CommandError::new("SWITCH_TOPIC_ERROR", e.to_string()))?; .switch_topic(chat_id, &topic.id)
} .map_err(|e| CommandError::new("SWITCH_TOPIC_ERROR", e.to_string()))?;
} }
// Query the full topic list so the frontend sidebar can update // Query the full topic list so the frontend sidebar can update
@ -119,7 +119,7 @@ async fn handle_create_session(
.with_metadata("topics", &topics_json) .with_metadata("topics", &topics_json)
.with_metadata("topic_id", &topic.id) .with_metadata("topic_id", &topic.id)
.with_metadata("session_id", &topic.session_id) .with_metadata("session_id", &topic.session_id)
.with_metadata("message_count", &topic.message_count.to_string())) .with_metadata("message_count", topic.message_count.to_string()))
} }
#[cfg(test)] #[cfg(test)]

View File

@ -108,10 +108,7 @@ impl CommandHandler for StopExecutionCommandHandler {
if cancelled || cancelled_subagents > 0 { if cancelled || cancelled_subagents > 0 {
let msg = if cancelled && cancelled_subagents > 0 { let msg = if cancelled && cancelled_subagents > 0 {
format!( format!("正在停止当前任务及 {} 个后台子代理...", cancelled_subagents)
"正在停止当前任务及 {} 个后台子代理...",
cancelled_subagents
)
} else if cancelled { } else if cancelled {
"正在停止当前任务...".to_string() "正在停止当前任务...".to_string()
} else { } else {

View File

@ -103,13 +103,13 @@ async fn handle_switch_topic(
})?; })?;
// 如果有 SessionManager实际切换话题历史 // 如果有 SessionManager实际切换话题历史
if let Some(ref session_manager) = handler.session_manager { if let Some(ref session_manager) = handler.session_manager
if let Some(session) = session_manager.get(&ctx.channel_name).await { && let Some(session) = session_manager.get(&ctx.channel_name).await
let mut session_guard = session.lock().await; {
session_guard let mut session_guard = session.lock().await;
.switch_topic(chat_id, &target_topic_id) session_guard
.map_err(|e| CommandError::new("SWITCH_TOPIC_ERROR", e.to_string()))?; .switch_topic(chat_id, &target_topic_id)
} .map_err(|e| CommandError::new("SWITCH_TOPIC_ERROR", e.to_string()))?;
} }
// 使用辅助方法获取消息数量 // 使用辅助方法获取消息数量
@ -127,5 +127,5 @@ async fn handle_switch_topic(
.with_message(MessageKind::Notification, &message) .with_message(MessageKind::Notification, &message)
.with_metadata("topic_id", &topic.id) .with_metadata("topic_id", &topic.id)
.with_metadata("title", &topic.title) .with_metadata("title", &topic.title)
.with_metadata("message_count", &msg_count.to_string())) .with_metadata("message_count", msg_count.to_string()))
} }

View File

@ -128,7 +128,7 @@ impl Default for CompactionConfig {
} }
/// 可观测性配置日志格式、metrics 开关等) /// 可观测性配置日志格式、metrics 开关等)
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize, Default)]
pub struct ObservabilityConfig { pub struct ObservabilityConfig {
/// 日志输出格式text默认或 json。 /// 日志输出格式text默认或 json。
/// json 格式便于接入 ELK/Loki 等日志聚合系统。 /// json 格式便于接入 ELK/Loki 等日志聚合系统。
@ -136,14 +136,6 @@ pub struct ObservabilityConfig {
pub log_format: LogFormat, pub log_format: LogFormat,
} }
impl Default for ObservabilityConfig {
fn default() -> Self {
Self {
log_format: LogFormat::default(),
}
}
}
/// 日志输出格式 /// 日志输出格式
#[derive(Debug, Clone, Deserialize, Serialize, Default, PartialEq, Eq)] #[derive(Debug, Clone, Deserialize, Serialize, Default, PartialEq, Eq)]
#[serde(rename_all = "lowercase")] #[serde(rename_all = "lowercase")]
@ -1984,7 +1976,7 @@ mod tests {
timezone: "Asia/Shanghai".to_string(), timezone: "Asia/Shanghai".to_string(),
}); });
assert_eq!(effective_jobs.len(), 3); // 2个内置 + 1个自定义 assert_eq!(effective_jobs.len(), 3); // 2个内置 + 1个自定义
// 第一个作业:内存维护(被覆盖为禁用) // 第一个作业:内存维护(被覆盖为禁用)
assert_eq!(effective_jobs[0].id, BUILTIN_MEMORY_MAINTENANCE_JOB_ID); assert_eq!(effective_jobs[0].id, BUILTIN_MEMORY_MAINTENANCE_JOB_ID);
assert!(!effective_jobs[0].enabled); assert!(!effective_jobs[0].enabled);
assert_eq!( assert_eq!(
@ -2305,25 +2297,33 @@ mod tests {
#[test] #[test]
fn test_scheduler_schedule_validation_rejects_invalid_values() { fn test_scheduler_schedule_validation_rejects_invalid_values() {
assert!(SchedulerSchedule::Delay { seconds: 0 } assert!(
.validate("delay.job") SchedulerSchedule::Delay { seconds: 0 }
.is_err()); .validate("delay.job")
assert!(SchedulerSchedule::Interval { .is_err()
seconds: 0, );
startup_delay_secs: 0, assert!(
} SchedulerSchedule::Interval {
.validate("interval.job") seconds: 0,
.is_err()); startup_delay_secs: 0,
assert!(SchedulerSchedule::At { }
timestamp: "bad timestamp".to_string(), .validate("interval.job")
} .is_err()
.validate("at.job") );
.is_err()); assert!(
assert!(SchedulerSchedule::Cron { SchedulerSchedule::At {
expression: "bad cron".to_string(), timestamp: "bad timestamp".to_string(),
} }
.validate("cron.job") .validate("at.job")
.is_err()); .is_err()
);
assert!(
SchedulerSchedule::Cron {
expression: "bad cron".to_string(),
}
.validate("cron.job")
.is_err()
);
} }
#[test] #[test]

View File

@ -63,13 +63,13 @@ impl CapabilityPolicy {
/// 校验指定子代理是否被允许。返回 Err 时附带拒绝原因。 /// 校验指定子代理是否被允许。返回 Err 时附带拒绝原因。
pub fn check_subagent_allowed(&self, name: &str) -> Result<(), String> { pub fn check_subagent_allowed(&self, name: &str) -> Result<(), String> {
if let Some(list) = &self.allowed_subagents { if let Some(list) = &self.allowed_subagents
if !list.iter().any(|s| s == name) { && !list.iter().any(|s| s == name)
return Err(format!( {
"subagent '{}' is not in the allowed_subagents whitelist", return Err(format!(
name "subagent '{}' is not in the allowed_subagents whitelist",
)); name
} ));
} }
if self.denied_subagents.iter().any(|s| s == name) { if self.denied_subagents.iter().any(|s| s == name) {
return Err(format!( return Err(format!(

View File

@ -1,12 +1,12 @@
use crate::config::ExpertsConfig; use crate::config::ExpertsConfig;
use crate::domain::CapabilityPolicy; use crate::domain::CapabilityPolicy;
use crate::platform::{atomic_rename, home_dir as platform_home_dir}; use crate::platform::{atomic_rename, home_dir as platform_home_dir};
use parking_lot::RwLock;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
use std::fs; use std::fs;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::Arc; use std::sync::Arc;
use parking_lot::RwLock;
#[cfg(test)] #[cfg(test)]
static EXPERT_TEST_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(()); static EXPERT_TEST_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
@ -294,18 +294,13 @@ impl ExpertRuntime {
/// Re-discover experts from the filesystem. /// Re-discover experts from the filesystem.
pub fn reload(&self) -> Result<ExpertCatalog, String> { pub fn reload(&self) -> Result<ExpertCatalog, String> {
let config = self let config = self.config.read().clone();
.config
.read()
.clone();
let catalog = ExpertCatalog::discover_with_state( let catalog = ExpertCatalog::discover_with_state(
&config, &config,
&self.cwd, &self.cwd,
Some(&load_expert_disable_state(&self.cwd)), Some(&load_expert_disable_state(&self.cwd)),
); );
let mut guard = self let mut guard = self.catalog.write();
.catalog
.write();
*guard = catalog.clone(); *guard = catalog.clone();
Ok(catalog) Ok(catalog)
} }
@ -323,18 +318,12 @@ impl ExpertRuntime {
/// List enabled experts (disabled ones are filtered out). /// List enabled experts (disabled ones are filtered out).
pub fn list_experts(&self) -> Vec<Expert> { pub fn list_experts(&self) -> Vec<Expert> {
self.catalog self.catalog.read().experts.clone()
.read()
.experts
.clone()
} }
/// List all discovered experts including disabled ones, with their disabled scopes. /// List all discovered experts including disabled ones, with their disabled scopes.
pub fn list_experts_with_status(&self) -> Vec<ExpertWithStatus> { pub fn list_experts_with_status(&self) -> Vec<ExpertWithStatus> {
let config = self let config = self.config.read().clone();
.config
.read()
.clone();
let catalog = ExpertCatalog::discover_without_state(&config, &self.cwd); let catalog = ExpertCatalog::discover_without_state(&config, &self.cwd);
let disable_state = load_expert_disable_state(&self.cwd); let disable_state = load_expert_disable_state(&self.cwd);
@ -361,10 +350,7 @@ impl ExpertRuntime {
} }
pub fn get_expert(&self, name: &str) -> Option<Expert> { pub fn get_expert(&self, name: &str) -> Option<Expert> {
self.catalog self.catalog.read().find_expert(name).cloned()
.read()
.find_expert(name)
.cloned()
} }
pub fn create_expert( pub fn create_expert(
@ -474,10 +460,7 @@ impl ExpertRuntime {
pub fn has_expert_definition(&self, name: &str) -> Result<bool, String> { pub fn has_expert_definition(&self, name: &str) -> Result<bool, String> {
validate_expert_name(name)?; validate_expert_name(name)?;
let config = self let config = self.config.read().clone();
.config
.read()
.clone();
let catalog = ExpertCatalog::discover_without_state(&config, &self.cwd); let catalog = ExpertCatalog::discover_without_state(&config, &self.cwd);
Ok(catalog.find_expert(name).is_some()) Ok(catalog.find_expert(name).is_some())
} }
@ -509,9 +492,7 @@ impl ExpertRuntime {
// update in-memory disable_state // update in-memory disable_state
{ {
let mut state = self let mut state = self.disable_state.write();
.disable_state
.write();
match scope { match scope {
ExpertScope::User => { ExpertScope::User => {
if enabled { if enabled {
@ -533,9 +514,7 @@ impl ExpertRuntime {
// refresh catalog so list_experts / get_expert reflect the change // refresh catalog so list_experts / get_expert reflect the change
let _ = self.reload()?; let _ = self.reload()?;
let state = self let state = self.disable_state.read();
.disable_state
.read();
let disabled_in_scopes = state.disabled_scopes_for(name); let disabled_in_scopes = state.disabled_scopes_for(name);
Ok(ExpertAvailabilityChange { Ok(ExpertAvailabilityChange {
@ -558,9 +537,7 @@ impl ExpertRuntime {
} }
{ {
let mut sessions = self let mut sessions = self.session_experts.write();
.session_experts
.write();
sessions.insert(session_id.to_string(), expert_name.to_string()); sessions.insert(session_id.to_string(), expert_name.to_string());
} }
persist_session_experts(&self.cwd, |state| { persist_session_experts(&self.cwd, |state| {
@ -573,9 +550,7 @@ impl ExpertRuntime {
/// Clear the selected expert for a session. /// Clear the selected expert for a session.
pub fn clear_expert(&self, session_id: &str) -> Result<(), String> { pub fn clear_expert(&self, session_id: &str) -> Result<(), String> {
{ {
let mut sessions = self let mut sessions = self.session_experts.write();
.session_experts
.write();
sessions.remove(session_id); sessions.remove(session_id);
} }
persist_session_experts(&self.cwd, |state| { persist_session_experts(&self.cwd, |state| {
@ -586,16 +561,12 @@ impl ExpertRuntime {
/// Returns the expert selected for a session, or None if none selected / disabled / not found. /// Returns the expert selected for a session, or None if none selected / disabled / not found.
pub fn selected_expert_for(&self, session_id: &str) -> Option<Expert> { pub fn selected_expert_for(&self, session_id: &str) -> Option<Expert> {
let name = { let name = {
let sessions = self let sessions = self.session_experts.read();
.session_experts
.read();
sessions.get(session_id).cloned() sessions.get(session_id).cloned()
}?; }?;
// Filter out disabled experts. // Filter out disabled experts.
let state = self let state = self.disable_state.read();
.disable_state
.read();
if state.is_disabled(&name) { if state.is_disabled(&name) {
return None; return None;
} }

View File

@ -3,7 +3,9 @@ use std::sync::Arc;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use crate::agent::context_compressor::ContextCompressor; use crate::agent::context_compressor::ContextCompressor;
use crate::agent::{AgentError, AgentLoop, AgentRuntimeConfig, CompositeSystemPromptProvider, SystemPromptProvider}; use crate::agent::{
AgentError, AgentLoop, AgentRuntimeConfig, CompositeSystemPromptProvider, SystemPromptProvider,
};
use crate::config::{CompactionConfig, LLMProviderConfig, ModelResolver}; use crate::config::{CompactionConfig, LLMProviderConfig, ModelResolver};
use crate::domain::CapabilityPolicy; use crate::domain::CapabilityPolicy;
use crate::experts::ExpertPromptProvider; use crate::experts::ExpertPromptProvider;
@ -14,10 +16,10 @@ use crate::gateway::tool_prompt_provider::ToolPromptProvider;
use crate::observability::Observer; use crate::observability::Observer;
use crate::skills::{SkillPromptProvider, SkillRuntime}; use crate::skills::{SkillPromptProvider, SkillRuntime};
use crate::storage::PromptInjectionRepository; use crate::storage::PromptInjectionRepository;
use crate::storage::persistent_session_id;
use crate::storage::SessionStore; use crate::storage::SessionStore;
use crate::tools::task::runtime::{SubagentPromptProvider, SubagentRuntime}; use crate::storage::persistent_session_id;
use crate::tools::task::SubagentResult; use crate::tools::task::SubagentResult;
use crate::tools::task::runtime::{SubagentPromptProvider, SubagentRuntime};
use crate::tools::{ToolContext, ToolRegistry, WaitCoordinator}; use crate::tools::{ToolContext, ToolRegistry, WaitCoordinator};
/// 构建与 Agent 实际使用的完全一致的组合系统提示词 Provider。 /// 构建与 Agent 实际使用的完全一致的组合系统提示词 Provider。
@ -133,7 +135,10 @@ impl AgentFactory {
/// 构造 ContextCompressor参数内聚到 ContextCompressorCompactionConfig 注入)。 /// 构造 ContextCompressor参数内聚到 ContextCompressorCompactionConfig 注入)。
/// AgentLoopin-loop 压缩)和 Sessionsync 兜底压缩)共用此方法, /// AgentLoopin-loop 压缩)和 Sessionsync 兜底压缩)共用此方法,
/// 确保两条压缩路径使用同一套用户配置的压缩参数。 /// 确保两条压缩路径使用同一套用户配置的压缩参数。
pub(crate) fn build_compressor(&self, runtime_config: &AgentRuntimeConfig) -> ContextCompressor { pub(crate) fn build_compressor(
&self,
runtime_config: &AgentRuntimeConfig,
) -> ContextCompressor {
ContextCompressor::with_compaction_config( ContextCompressor::with_compaction_config(
runtime_config.context_window_tokens, runtime_config.context_window_tokens,
runtime_config.context_summary_char_budget, runtime_config.context_summary_char_budget,
@ -201,21 +206,23 @@ impl AgentFactory {
// 物化:命中 session 级选择且话题无固化值时,将解析后的具体 // 物化:命中 session 级选择且话题无固化值时,将解析后的具体
// (provider, model) 写入 topics 行(持久化 + 内存缓存) // (provider, model) 写入 topics 行(持久化 + 内存缓存)
if !from_topic { if !from_topic && let Some(tid) = request.topic_id.as_deref() {
if let Some(tid) = request.topic_id.as_deref() { let provider = resolved.name.clone();
let provider = resolved.name.clone(); let model = resolved.model_id.clone();
let model = resolved.model_id.clone(); self.topic_model_selections.set(
self.topic_model_selections tid,
.set(tid, Some(provider.clone()), Some(model.clone())); Some(provider.clone()),
if let Err(err) = Some(model.clone()),
self.store.update_topic_model(tid, Some(&provider), Some(&model)) );
{ if let Err(err) =
tracing::warn!( self.store
error = %err, .update_topic_model(tid, Some(&provider), Some(&model))
topic_id = %tid, {
"AgentFactory: failed to materialize topic model selection" tracing::warn!(
); error = %err,
} topic_id = %tid,
"AgentFactory: failed to materialize topic model selection"
);
} }
} }
@ -289,7 +296,7 @@ impl AgentFactory {
// 供 wait_for_subagents 工具传递给 coordinator.wait() 的 select!。 // 供 wait_for_subagents 工具传递给 coordinator.wait() 的 select!。
// watch::Receiver::clone() 创建共享同一 sender 的新 receiver // watch::Receiver::clone() 创建共享同一 sender 的新 receiver
// 各 receiver 的 has_changed()/changed() 状态独立,互不影响。 // 各 receiver 的 has_changed()/changed() 状态独立,互不影响。
let cancel_rx_for_context = request.cancel_token.as_ref().map(|rx| rx.clone()); let cancel_rx_for_context = request.cancel_token.clone();
let runtime_config = AgentRuntimeConfig::from(effective_provider_config.clone()); let runtime_config = AgentRuntimeConfig::from(effective_provider_config.clone());
let compressor = Arc::new(self.build_compressor(&runtime_config)); let compressor = Arc::new(self.build_compressor(&runtime_config));

View File

@ -42,14 +42,14 @@ impl AgentPromptProvider {
/// 记录注入事件 /// 记录注入事件
fn record_injection(&self, context: &SystemPromptContext) { fn record_injection(&self, context: &SystemPromptContext) {
if let Some(session_id) = &context.session_id { if let Some(session_id) = &context.session_id
if let Err(e) = self.repository.mark_agent_prompt_reinjected(session_id) { && let Err(e) = self.repository.mark_agent_prompt_reinjected(session_id)
tracing::warn!( {
session_id = ?session_id, tracing::warn!(
error = %e, session_id = ?session_id,
"Failed to mark agent prompt reinjected; injection counter may be inaccurate" error = %e,
); "Failed to mark agent prompt reinjected; injection counter may be inaccurate"
} );
} }
} }
} }

View File

@ -6,11 +6,11 @@
//! - token 通过 `Authorization: Bearer <token>`HTTP或 `?token=<token>`WS传递。 //! - token 通过 `Authorization: Bearer <token>`HTTP或 `?token=<token>`WS传递。
//! - 校验使用常量时间比较,避免计时侧信道。 //! - 校验使用常量时间比较,避免计时侧信道。
use axum::Json;
use axum::extract::Request; use axum::extract::Request;
use axum::http::{HeaderMap, StatusCode}; use axum::http::{HeaderMap, StatusCode};
use axum::middleware::Next; use axum::middleware::Next;
use axum::response::{IntoResponse, Response}; use axum::response::{IntoResponse, Response};
use axum::Json;
use serde_json::json; use serde_json::json;
use subtle::ConstantTimeEq; use subtle::ConstantTimeEq;
@ -88,11 +88,7 @@ pub fn extract_bearer_token(headers: &HeaderMap) -> Option<&str> {
/// 仅在 `requires_auth` 为 true 时挂载。 /// 仅在 `requires_auth` 为 true 时挂载。
/// `/health`、`/ws`、静态资源放行;`/ws` 的 token 校验在 ws_handler 内完成。 /// `/health`、`/ws`、静态资源放行;`/ws` 的 token 校验在 ws_handler 内完成。
/// `/metrics` 包含运行时指标provider/model/耗时/token 用量),远程部署时需保护。 /// `/metrics` 包含运行时指标provider/model/耗时/token 用量),远程部署时需保护。
pub async fn require_bearer_auth( pub async fn require_bearer_auth(headers: HeaderMap, request: Request, next: Next) -> Response {
headers: HeaderMap,
request: Request,
next: Next,
) -> Response {
let path = request.uri().path(); let path = request.uri().path();
// /api/* 和 /metrics 需要认证;其余放行 // /api/* 和 /metrics 需要认证;其余放行

View File

@ -155,7 +155,7 @@ impl AgentExecutionService {
// 直接比较 current_topic(chat_id) 与 original_topic_id // 直接比较 current_topic(chat_id) 与 original_topic_id
// 这比"检查内存历史最新消息"更可靠,且天然处理"切走又切回"的 case // 这比"检查内存历史最新消息"更可靠,且天然处理"切走又切回"的 case
let is_current_turn = match request.original_topic_id.as_deref() { let is_current_turn = match request.original_topic_id.as_deref() {
Some(orig_tid) => session.current_topic(request.chat_id).as_deref() == Some(orig_tid), Some(orig_tid) => session.current_topic(request.chat_id) == Some(orig_tid),
None => true, // 无 topic 时总是视为当前回合 None => true, // 无 topic 时总是视为当前回合
}; };
@ -419,7 +419,11 @@ impl AgentExecutionService {
); );
let result = agent let result = agent
.process(history, Some(&system_prompt_context), Some(&compaction_sink)) .process(
history,
Some(&system_prompt_context),
Some(&compaction_sink),
)
.await?; .await?;
let mut metadata = HashMap::new(); let mut metadata = HashMap::new();
// 把用户消息的 UUID 回传给前端,前端用此更新本地消息 ID使 todo 点击跳转能匹配 // 把用户消息的 UUID 回传给前端,前端用此更新本地消息 ID使 todo 点击跳转能匹配
@ -605,7 +609,11 @@ impl AgentExecutionService {
); );
let result = agent let result = agent
.process(history, Some(&system_prompt_context), Some(&compaction_sink)) .process(
history,
Some(&system_prompt_context),
Some(&compaction_sink),
)
.await?; .await?;
let outbound_messages = self let outbound_messages = self

View File

@ -65,19 +65,19 @@ fn mask_config(config: &Config) -> Config {
} }
} }
for channel in masked.channels.values_mut() { for channel in masked.channels.values_mut() {
if let Some(feishu) = channel.as_feishu_mut() { if let Some(feishu) = channel.as_feishu_mut()
if !feishu.app_secret.is_empty() { && !feishu.app_secret.is_empty()
let visible: String = feishu.app_secret.chars().take(4).collect(); {
feishu.app_secret = format!("{}{}", visible, API_KEY_MASK); let visible: String = feishu.app_secret.chars().take(4).collect();
} feishu.app_secret = format!("{}{}", visible, API_KEY_MASK);
} }
} }
// 掩码网关认证 token避免通过 /api/config 泄露) // 掩码网关认证 token避免通过 /api/config 泄露)
if let Some(ref token) = masked.gateway.auth_token { if let Some(ref token) = masked.gateway.auth_token
if !token.is_empty() { && !token.is_empty()
let visible: String = token.chars().take(4).collect(); {
masked.gateway.auth_token = Some(format!("{}{}", visible, API_KEY_MASK)); let visible: String = token.chars().take(4).collect();
} masked.gateway.auth_token = Some(format!("{}{}", visible, API_KEY_MASK));
} }
masked masked
} }
@ -116,28 +116,26 @@ pub async fn save_config(
{ {
let cfg = state.config.read().await; let cfg = state.config.read().await;
for (name, provider) in new_config.providers.iter_mut() { for (name, provider) in new_config.providers.iter_mut() {
if is_masked_key(&provider.api_key) { if is_masked_key(&provider.api_key)
if let Some(original) = cfg.providers.get(name) { && let Some(original) = cfg.providers.get(name)
provider.api_key = original.api_key.clone(); {
} provider.api_key = original.api_key.clone();
} }
} }
for (name, channel) in new_config.channels.iter_mut() { for (name, channel) in new_config.channels.iter_mut() {
if let Some(feishu) = channel.as_feishu_mut() { if let Some(feishu) = channel.as_feishu_mut()
if is_masked_key(&feishu.app_secret) { && is_masked_key(&feishu.app_secret)
if let Some(original_channel) = cfg.channels.get(name) { && let Some(original_channel) = cfg.channels.get(name)
if let Some(original_feishu) = original_channel.as_feishu() { && let Some(original_feishu) = original_channel.as_feishu()
feishu.app_secret = original_feishu.app_secret.clone(); {
} feishu.app_secret = original_feishu.app_secret.clone();
}
}
} }
} }
// 保留原始 auth_token若提交的是掩码值 // 保留原始 auth_token若提交的是掩码值
if let Some(ref submitted) = new_config.gateway.auth_token { if let Some(ref submitted) = new_config.gateway.auth_token
if is_masked_key(submitted) { && is_masked_key(submitted)
new_config.gateway.auth_token = cfg.gateway.auth_token.clone(); {
} new_config.gateway.auth_token = cfg.gateway.auth_token.clone();
} }
} // read lock released here } // read lock released here
@ -243,9 +241,7 @@ pub async fn list_executions(State(state): State<Arc<GatewayState>>) -> Json<Exe
/// GET /metrics — Prometheus metrics 端点 /// GET /metrics — Prometheus metrics 端点
/// ///
/// 返回 Prometheus 格式的 metrics 文本。若 recorder 未安装则返回 503。 /// 返回 Prometheus 格式的 metrics 文本。若 recorder 未安装则返回 503。
pub async fn metrics_handler( pub async fn metrics_handler(State(state): State<Arc<GatewayState>>) -> (StatusCode, String) {
State(state): State<Arc<GatewayState>>,
) -> (StatusCode, String) {
match &state.prometheus_handle { match &state.prometheus_handle {
Some(handle) => (StatusCode::OK, handle.render()), Some(handle) => (StatusCode::OK, handle.render()),
None => ( None => (
@ -1160,27 +1156,27 @@ pub async fn session_select_model(
// 校验provider/model 名必须在 config 的 providers/models 表中存在 // 校验provider/model 名必须在 config 的 providers/models 表中存在
// (与 AgentFactory::create 中的解析失败行为对齐,提前反馈错误) // (与 AgentFactory::create 中的解析失败行为对齐,提前反馈错误)
let config = state.config.read().await; let config = state.config.read().await;
if let Some(name) = provider.as_ref() { if let Some(name) = provider.as_ref()
if !config.providers.contains_key(name) { && !config.providers.contains_key(name)
return ( {
StatusCode::BAD_REQUEST, return (
Json(SelectModelResponse { StatusCode::BAD_REQUEST,
success: false, Json(SelectModelResponse {
error: Some(format!("provider '{}' not found in config", name)), success: false,
}), error: Some(format!("provider '{}' not found in config", name)),
); }),
} );
} }
if let Some(name) = model.as_ref() { if let Some(name) = model.as_ref()
if !config.models.contains_key(name) { && !config.models.contains_key(name)
return ( {
StatusCode::BAD_REQUEST, return (
Json(SelectModelResponse { StatusCode::BAD_REQUEST,
success: false, Json(SelectModelResponse {
error: Some(format!("model '{}' not found in config", name)), success: false,
}), error: Some(format!("model '{}' not found in config", name)),
); }),
} );
} }
drop(config); drop(config);
@ -1274,27 +1270,27 @@ pub async fn topic_select_model(
// 校验provider/model 名必须在 config 的 providers/models 表中存在 // 校验provider/model 名必须在 config 的 providers/models 表中存在
let config = state.config.read().await; let config = state.config.read().await;
if let Some(name) = provider.as_ref() { if let Some(name) = provider.as_ref()
if !config.providers.contains_key(name) { && !config.providers.contains_key(name)
return ( {
StatusCode::BAD_REQUEST, return (
Json(SelectModelResponse { StatusCode::BAD_REQUEST,
success: false, Json(SelectModelResponse {
error: Some(format!("provider '{}' not found in config", name)), success: false,
}), error: Some(format!("provider '{}' not found in config", name)),
); }),
} );
} }
if let Some(name) = model.as_ref() { if let Some(name) = model.as_ref()
if !config.models.contains_key(name) { && !config.models.contains_key(name)
return ( {
StatusCode::BAD_REQUEST, return (
Json(SelectModelResponse { StatusCode::BAD_REQUEST,
success: false, Json(SelectModelResponse {
error: Some(format!("model '{}' not found in config", name)), success: false,
}), error: Some(format!("model '{}' not found in config", name)),
); }),
} );
} }
drop(config); drop(config);
@ -1320,7 +1316,9 @@ pub async fn topic_select_model(
if is_clear { if is_clear {
state.model_selections.set(&topic_session_id, None, None); state.model_selections.set(&topic_session_id, None, None);
} else { } else {
state.model_selections.set(&topic_session_id, provider, model); state
.model_selections
.set(&topic_session_id, provider, model);
} }
( (

View File

@ -707,13 +707,13 @@ pub(crate) fn validate_memory_maintenance_output(
} }
// 检查目标 namespace 是否与源一致 // 检查目标 namespace 是否与源一致
if let Some(src_ns) = source_namespaces.iter().next() { if let Some(src_ns) = source_namespaces.iter().next()
if *src_ns != merge.namespace { && *src_ns != merge.namespace
return Err(format!( {
"跨 namespace 合并被禁止: {} → {}", return Err(format!(
src_ns, merge.namespace "跨 namespace 合并被禁止: {} → {}",
)); src_ns, merge.namespace
} ));
} }
} }
@ -768,7 +768,7 @@ pub(crate) fn apply_memory_maintenance_output(
min_memories_to_keep, min_memories_to_keep,
max_merge_per_group, max_merge_per_group,
) )
.map_err(|e| AgentError::Other(e))?; .map_err(AgentError::Other)?;
let all_candidates = plan.candidates.clone(); let all_candidates = plan.candidates.clone();
@ -834,14 +834,14 @@ pub(crate) fn apply_memory_maintenance_output(
} }
for memory_id in &output.low_value_ids { for memory_id in &output.low_value_ids {
if let Some(candidate) = candidates_by_id.get(memory_id.as_str()) { if let Some(candidate) = candidates_by_id.get(memory_id.as_str())
if deleted_ids.insert(candidate.id.clone()) { && deleted_ids.insert(candidate.id.clone())
store {
.delete_memory("user", scope_key, &candidate.namespace, &candidate.key) store
.map_err(|err| { .delete_memory("user", scope_key, &candidate.namespace, &candidate.key)
AgentError::Other(format!("delete low value memory error: {}", err)) .map_err(|err| {
})?; AgentError::Other(format!("delete low value memory error: {}", err))
} })?;
} }
} }

View File

@ -110,27 +110,34 @@ impl GatewayState {
mcp_servers: config.mcp_servers.clone(), mcp_servers: config.mcp_servers.clone(),
}; };
let (session_manager, task_repository, mcp_manager, subagent_runtime, model_selections, topic_model_selections, subagent_executor) = let (
build_session_manager_with_sender( session_manager,
agent_prompt_reinject_every, task_repository,
show_tool_results, mcp_manager,
config.time.timezone.clone(), subagent_runtime,
provider_config, model_selections,
provider_configs, topic_model_selections,
skills.clone(), subagent_executor,
experts.clone(), ) = build_session_manager_with_sender(
Arc::new(BusSessionMessageSender::new(bus.clone())), agent_prompt_reinject_every,
std::collections::HashSet::new(), show_tool_results,
config.tools.task.clone(), config.time.timezone.clone(),
config.subagents.clone(), provider_config,
config.memory_maintenance.clone(), provider_configs,
session_ttl_hours, skills.clone(),
mcp_config, experts.clone(),
config.mcp_tool_timeout_secs, Arc::new(BusSessionMessageSender::new(bus.clone())),
Some(bus.clone()), std::collections::HashSet::new(),
Arc::new(crate::config::ModelResolver::from_config(&config)), config.tools.task.clone(),
config.compaction.clone(), config.subagents.clone(),
)?; config.memory_maintenance.clone(),
session_ttl_hours,
mcp_config,
config.mcp_tool_timeout_secs,
Some(bus.clone()),
Arc::new(crate::config::ModelResolver::from_config(&config)),
config.compaction.clone(),
)?;
// 诊断日志:记录新 GatewayState 的创建(用于排查重启后是否使用了新状态) // 诊断日志:记录新 GatewayState 的创建(用于排查重启后是否使用了新状态)
tracing::info!( tracing::info!(
@ -226,7 +233,11 @@ pub async fn run(
// 将 pending_subagents 表中所有 status='running' 的记录标记为 'interrupted'。 // 将 pending_subagents 表中所有 status='running' 的记录标记为 'interrupted'。
// 下次 wait_for_subagents 调用时,这些 task_id 不会出现在 pending 列表中, // 下次 wait_for_subagents 调用时,这些 task_id 不会出现在 pending 列表中,
// agent 可据此判断子代理未正常完成。 // agent 可据此判断子代理未正常完成。
match state.session_manager.store().mark_all_running_as_interrupted() { match state
.session_manager
.store()
.mark_all_running_as_interrupted()
{
Ok(0) => { Ok(0) => {
tracing::info!("Crash recovery: no interrupted subagents to recover"); tracing::info!("Crash recovery: no interrupted subagents to recover");
} }
@ -252,7 +263,7 @@ pub async fn run(
// Initialize and start channels // Initialize and start channels
state state
.channel_manager .channel_manager
.init(&*cfg, provider_config.clone()) .init(&cfg, provider_config.clone())
.await?; .await?;
drop(cfg); drop(cfg);
state.channel_manager.start_all().await?; state.channel_manager.start_all().await?;

View File

@ -1,5 +1,5 @@
use std::collections::HashMap;
use parking_lot::RwLock; use parking_lot::RwLock;
use std::collections::HashMap;
/// per-session 的用户模型覆盖选择存储。 /// per-session 的用户模型覆盖选择存储。
/// ///
@ -17,9 +17,7 @@ impl ModelSelectionStore {
/// 设置 session 的用户模型覆盖。provider 和 model 均为 None 时清除该 session 的选择。 /// 设置 session 的用户模型覆盖。provider 和 model 均为 None 时清除该 session 的选择。
pub fn set(&self, session_id: &str, provider: Option<String>, model: Option<String>) { pub fn set(&self, session_id: &str, provider: Option<String>, model: Option<String>) {
let mut selections = self let mut selections = self.selections.write();
.selections
.write();
if provider.is_none() && model.is_none() { if provider.is_none() && model.is_none() {
selections.remove(session_id); selections.remove(session_id);
} else { } else {
@ -29,10 +27,7 @@ impl ModelSelectionStore {
/// 读取 session 的用户模型覆盖。 /// 读取 session 的用户模型覆盖。
pub fn get(&self, session_id: &str) -> Option<(Option<String>, Option<String>)> { pub fn get(&self, session_id: &str) -> Option<(Option<String>, Option<String>)> {
self.selections self.selections.read().get(session_id).cloned()
.read()
.get(session_id)
.cloned()
} }
} }

View File

@ -77,37 +77,25 @@ impl OutboundDispatcher {
/// sender task 生命周期与 dispatcher 一致dispatcher `run()` 退出时 /// sender task 生命周期与 dispatcher 一致dispatcher `run()` 退出时
/// 通过 cancel token 终止所有 sender task。 /// 通过 cancel token 终止所有 sender task。
pub async fn register_channel(&self, name: &str, channel: Arc<dyn Channel + Send + Sync>) { pub async fn register_channel(&self, name: &str, channel: Arc<dyn Channel + Send + Sync>) {
let (high_tx, high_rx) = let (high_tx, high_rx) = mpsc::channel::<OutboundMessage>(HIGH_PRIORITY_QUEUE_CAPACITY);
mpsc::channel::<OutboundMessage>(HIGH_PRIORITY_QUEUE_CAPACITY); let (low_tx, low_rx) = mpsc::channel::<OutboundMessage>(LOW_PRIORITY_QUEUE_CAPACITY);
let (low_tx, low_rx) =
mpsc::channel::<OutboundMessage>(LOW_PRIORITY_QUEUE_CAPACITY);
let cancel = CancellationToken::new(); let cancel = CancellationToken::new();
let channel_name = name.to_string(); let channel_name = name.to_string();
let cancel_for_task = cancel.clone(); let cancel_for_task = cancel.clone();
tokio::spawn(async move { tokio::spawn(async move {
Self::run_sender_task( Self::run_sender_task(&channel_name, channel, high_rx, low_rx, cancel_for_task).await;
&channel_name,
channel,
high_rx,
low_rx,
cancel_for_task,
)
.await;
}); });
self.channels self.channels.write().await.insert(
.write() name.to_string(),
.await ChannelSink {
.insert( high_tx,
name.to_string(), low_tx,
ChannelSink { cancel,
high_tx, },
low_tx, );
cancel,
},
);
} }
/// sender task优先消费 high 队列(最终响应),再消费 low 队列(中间过程), /// sender task优先消费 high 队列(最终响应),再消费 low 队列(中间过程),
@ -166,11 +154,7 @@ impl OutboundDispatcher {
} }
/// 发送单条消息,处理重试结果日志。 /// 发送单条消息,处理重试结果日志。
async fn send_one( async fn send_one(channel: &dyn Channel, channel_name: &str, msg: OutboundMessage) {
channel: &dyn Channel,
channel_name: &str,
msg: OutboundMessage,
) {
let msg_chat_id = msg.chat_id.clone(); let msg_chat_id = msg.chat_id.clone();
let msg_trace_id = msg.trace_id.clone(); let msg_trace_id = msg.trace_id.clone();
match Self::send_with_retry(channel, msg).await { match Self::send_with_retry(channel, msg).await {
@ -419,7 +403,7 @@ mod tests {
return Err(ChannelError::ChannelFull); return Err(ChannelError::ChannelFull);
} }
if (count as u32) < self.fail_first_n { if count < self.fail_first_n {
return Err(ChannelError::SendError("simulated failure".to_string())); return Err(ChannelError::SendError("simulated failure".to_string()));
} }
@ -479,19 +463,42 @@ mod tests {
let error = make_error_message("c", "chat", "agent failed"); let error = make_error_message("c", "chat", "agent failed");
let tool_call = make_low_message("c", "chat", "calling tool"); let tool_call = make_low_message("c", "chat", "calling tool");
let tool_result = OutboundMessage::tool_result( let tool_result = OutboundMessage::tool_result(
"c", "chat", None, "id", "tool", "result", None, "c",
"chat",
None,
"id",
"tool",
"result",
None,
std::collections::HashMap::new(), std::collections::HashMap::new(),
); );
let exec_done = OutboundMessage::execution_completed( let exec_done = OutboundMessage::execution_completed(
"c", "chat", None, "c",
"chat",
None,
std::collections::HashMap::new(), std::collections::HashMap::new(),
); );
assert!(is_high_priority(&assistant), "AssistantResponse should be high priority"); assert!(
assert!(is_high_priority(&error), "ErrorNotification should be high priority"); is_high_priority(&assistant),
assert!(!is_high_priority(&tool_call), "ToolCall should be low priority"); "AssistantResponse should be high priority"
assert!(!is_high_priority(&tool_result), "ToolResult should be low priority"); );
assert!(!is_high_priority(&exec_done), "ExecutionCompleted should be low priority"); assert!(
is_high_priority(&error),
"ErrorNotification should be high priority"
);
assert!(
!is_high_priority(&tool_call),
"ToolCall should be low priority"
);
assert!(
!is_high_priority(&tool_result),
"ToolResult should be low priority"
);
assert!(
!is_high_priority(&exec_done),
"ExecutionCompleted should be low priority"
);
} }
#[tokio::test] #[tokio::test]
@ -514,8 +521,12 @@ mod tests {
}); });
// 先发一条 slow500ms 延迟),紧接着发一条 fast // 先发一条 slow500ms 延迟),紧接着发一条 fast
bus.publish_outbound(make_message("slow", "chat-1", "slow-msg")).await.unwrap(); bus.publish_outbound(make_message("slow", "chat-1", "slow-msg"))
bus.publish_outbound(make_message("fast", "chat-2", "fast-msg")).await.unwrap(); .await
.unwrap();
bus.publish_outbound(make_message("fast", "chat-2", "fast-msg"))
.await
.unwrap();
// 等待 fast 消息被投递(远早于 slow 完成) // 等待 fast 消息被投递(远早于 slow 完成)
tokio::time::timeout(Duration::from_millis(200), async { tokio::time::timeout(Duration::from_millis(200), async {
@ -524,7 +535,9 @@ mod tests {
} }
}) })
.await .await
.expect("fast channel should receive message within 200ms, but was blocked by slow channel"); .expect(
"fast channel should receive message within 200ms, but was blocked by slow channel",
);
// 等待 slow 消息完成 // 等待 slow 消息完成
tokio::time::timeout(Duration::from_secs(2), async { tokio::time::timeout(Duration::from_secs(2), async {
@ -568,8 +581,12 @@ mod tests {
}); });
// 先发 flaky会重试 3 秒),紧接着发 stable // 先发 flaky会重试 3 秒),紧接着发 stable
bus.publish_outbound(make_message("flaky", "chat-1", "flaky-msg")).await.unwrap(); bus.publish_outbound(make_message("flaky", "chat-1", "flaky-msg"))
bus.publish_outbound(make_message("stable", "chat-2", "stable-msg")).await.unwrap(); .await
.unwrap();
bus.publish_outbound(make_message("stable", "chat-2", "stable-msg"))
.await
.unwrap();
// stable 应在 200ms 内收到,远早于 flaky 的 3 秒重试完成 // stable 应在 200ms 内收到,远早于 flaky 的 3 秒重试完成
tokio::time::timeout(Duration::from_millis(200), async { tokio::time::timeout(Duration::from_millis(200), async {
@ -820,7 +837,10 @@ mod tests {
.await .await
.expect("high priority should succeed within extended retry budget"); .expect("high priority should succeed within extended retry budget");
assert!(result.is_ok(), "high priority should succeed after 4 attempts"); assert!(
result.is_ok(),
"high priority should succeed after 4 attempts"
);
assert_eq!( assert_eq!(
call_count.load(Ordering::SeqCst), call_count.load(Ordering::SeqCst),
4, 4,

View File

@ -1,7 +1,7 @@
use std::collections::HashSet;
use std::sync::Arc;
use futures_util::FutureExt; use futures_util::FutureExt;
use parking_lot::Mutex; use parking_lot::Mutex;
use std::collections::HashSet;
use std::sync::Arc;
use tokio::sync::Semaphore; use tokio::sync::Semaphore;
@ -28,8 +28,8 @@ use crate::providers::{ProviderRuntimeConfig, create_provider};
use crate::storage::persistent_session_id; use crate::storage::persistent_session_id;
use crate::topic_description::generate_topic_description; use crate::topic_description::generate_topic_description;
use super::session::{BusToolCallEmitter, SessionManager};
use super::message_prepare::enrich_user_content_with_media_refs; use super::message_prepare::enrich_user_content_with_media_refs;
use super::session::{BusToolCallEmitter, SessionManager};
#[derive(Clone)] #[derive(Clone)]
pub struct InboundProcessor { pub struct InboundProcessor {
@ -180,39 +180,37 @@ impl InboundProcessor {
let chat_id_for_span = inbound.chat_id.clone(); let chat_id_for_span = inbound.chat_id.clone();
let session_id_for_span = let session_id_for_span =
crate::storage::persistent_session_id(&inbound.channel, &inbound.chat_id); crate::storage::persistent_session_id(&inbound.channel, &inbound.chat_id);
tokio::spawn( tokio::spawn(crate::observability::tracing_ctx::traced(
crate::observability::tracing_ctx::traced( &trace_id,
&trace_id, &chat_id_for_span,
&chat_id_for_span, &session_id_for_span,
&session_id_for_span, async move {
async move { let _permit = permit; // 持有 permit 直到任务完成
let _permit = permit; // 持有 permit 直到任务完成 // catch_unwind 将 panic 归一化为错误:否则工具/历史清理中的
// catch_unwind 将 panic 归一化为错误:否则工具/历史清理中的 // panic 只会终止任务并打 panic hook 日志,跳过错误日志与指标,
// panic 只会终止任务并打 panic hook 日志,跳过错误日志与指标, // 用户消息被静默吞掉。参考 channels/wechat.rs 的同类用法。
// 用户消息被静默吞掉。参考 channels/wechat.rs 的同类用法。 let result = std::panic::AssertUnwindSafe(processor.process_one(inbound))
let result = std::panic::AssertUnwindSafe(processor.process_one(inbound)) .catch_unwind()
.catch_unwind() .await;
.await; match result {
match result { Ok(Ok(())) => {}
Ok(Ok(())) => {} Ok(Err(e)) => {
Ok(Err(e)) => { tracing::error!(
tracing::error!( error = %crate::utils::format_error_chain(&e),
error = %crate::utils::format_error_chain(&e), "Message processing failed"
"Message processing failed" );
); crate::observability::metrics::record_message_processing_error();
crate::observability::metrics::record_message_processing_error();
}
Err(payload) => {
tracing::error!(
error = %crate::utils::panic_payload_message(&payload),
"Message processing panicked"
);
crate::observability::metrics::record_message_processing_error();
}
} }
}, Err(payload) => {
), tracing::error!(
); error = %crate::utils::panic_payload_message(&payload),
"Message processing panicked"
);
crate::observability::metrics::record_message_processing_error();
}
}
},
));
} }
} }
@ -280,8 +278,8 @@ impl InboundProcessor {
} }
} }
} }
} else if let Some(error) = response.error { } else if let Some(error) = response.error
if let Err(e) = self && let Err(e) = self
.bus .bus
.publish_outbound( .publish_outbound(
OutboundMessage::assistant( OutboundMessage::assistant(
@ -295,14 +293,13 @@ impl InboundProcessor {
.with_trace_id(&inbound.trace_id), .with_trace_id(&inbound.trace_id),
) )
.await .await
{ {
match e { match e {
crate::bus::BusError::Dropped => { crate::bus::BusError::Dropped => {
tracing::warn!(error = %e, "Outbound dropped (bus full)"); tracing::warn!(error = %e, "Outbound dropped (bus full)");
} }
crate::bus::BusError::Closed => { crate::bus::BusError::Closed => {
tracing::error!(error = %e, "Failed to publish error response"); tracing::error!(error = %e, "Failed to publish error response");
}
} }
} }
} }
@ -326,73 +323,65 @@ impl InboundProcessor {
// //
// 安全性is_waiting 在持锁状态下检查wait_coordinator 清除 is_waiting 需先重获取锁, // 安全性is_waiting 在持锁状态下检查wait_coordinator 清除 is_waiting 需先重获取锁,
// 两者互斥,无 TOCTOU。 // 两者互斥,无 TOCTOU。
if let Some(ref topic_id) = current_topic { if let Some(ref topic_id) = current_topic
if let Some(session) = self.session_manager.get(&inbound.channel).await { && let Some(session) = self.session_manager.get(&inbound.channel).await
let lock_key = topic_id.clone(); {
let lock_key = topic_id.clone();
// 获取 serial_lock Arc短暂持有 session 锁) // 获取 serial_lock Arc短暂持有 session 锁)
let serial_lock = { let serial_lock = {
let mut g = session.lock().await;
g.ensure_sub_done_channel(&lock_key);
g.topic_serial_lock(&lock_key)
};
// 阻塞获取 serial_lock
// - agent 正常运行:阻塞至其完成(天然串行化)
// - agent 在 wait 中wait 已释放锁,可立即获取
let _inject_guard = serial_lock.clone().lock_owned().await;
// 检查 is_waiting持锁状态下安全
let is_waiting = {
let g = session.lock().await;
g.is_waiting(&lock_key)
};
if is_waiting {
// Agent 正在 wait_for_subagents 中等待 → 注入用户消息 + 唤醒
tracing::info!(
topic_id = %lock_key,
"Topic is in waiting state, injecting user message and waking up agent"
);
let wakeup = {
let mut g = session.lock().await; let mut g = session.lock().await;
g.ensure_sub_done_channel(&lock_key); // 确保 session 和 chat 已加载
g.topic_serial_lock(&lock_key) g.ensure_persistent_session(&inbound.chat_id)?;
g.ensure_chat_loaded(&inbound.chat_id, Some(&lock_key))?;
// 构造用户消息(与 prepare_and_execute_message 一致的处理流程)
let media_refs: Vec<String> =
inbound.media.iter().map(|m| m.path.clone()).collect();
let enriched_content =
enrich_user_content_with_media_refs(&inbound.content, &media_refs)?;
let user_message = g.create_user_message(&enriched_content, media_refs);
g.append_persisted_message(&inbound.chat_id, Some(&lock_key), user_message)?;
// 获取 wakeup 信号
g.wait_wakeup(&lock_key)
}; };
// 阻塞获取 serial_lock // 唤醒等待中的 agentwait_coordinator 的 select! 会捕获此通知)
// - agent 正常运行:阻塞至其完成(天然串行化) wakeup.notify_one();
// - agent 在 wait 中wait 已释放锁,可立即获取
let _inject_guard = serial_lock.clone().lock_owned().await;
// 检查 is_waiting持锁状态下安全 // _inject_guard 在此处 drop → 释放 serial_lock
let is_waiting = { // wait_coordinator 重获取锁后继续处理history 已包含新用户消息)
let g = session.lock().await; //
g.is_waiting(&lock_key) // 跳过 handle_message / cancel 注册 / execution_completed
}; // 因为等待中的 agent 会处理这条消息。
return Ok(());
if is_waiting {
// Agent 正在 wait_for_subagents 中等待 → 注入用户消息 + 唤醒
tracing::info!(
topic_id = %lock_key,
"Topic is in waiting state, injecting user message and waking up agent"
);
let wakeup = {
let mut g = session.lock().await;
// 确保 session 和 chat 已加载
g.ensure_persistent_session(&inbound.chat_id)?;
g.ensure_chat_loaded(&inbound.chat_id, Some(&lock_key))?;
// 构造用户消息(与 prepare_and_execute_message 一致的处理流程)
let media_refs: Vec<String> = inbound
.media
.iter()
.map(|m| m.path.clone())
.collect();
let enriched_content =
enrich_user_content_with_media_refs(&inbound.content, &media_refs)?;
let user_message =
g.create_user_message(&enriched_content, media_refs);
g.append_persisted_message(
&inbound.chat_id,
Some(&lock_key),
user_message,
)?;
// 获取 wakeup 信号
g.wait_wakeup(&lock_key)
};
// 唤醒等待中的 agentwait_coordinator 的 select! 会捕获此通知)
wakeup.notify_one();
// _inject_guard 在此处 drop → 释放 serial_lock
// wait_coordinator 重获取锁后继续处理history 已包含新用户消息)
//
// 跳过 handle_message / cancel 注册 / execution_completed
// 因为等待中的 agent 会处理这条消息。
return Ok(());
}
// is_waiting=false_inject_guard drop 释放锁,走正常 handle_message 路径
} }
// is_waiting=false_inject_guard drop 释放锁,走正常 handle_message 路径
} }
let live_emitter = Arc::new(PersistingEmittedMessageHandler::new( let live_emitter = Arc::new(PersistingEmittedMessageHandler::new(
@ -461,81 +450,74 @@ impl InboundProcessor {
// 异步生成 topic 描述(仅当描述为空且没有正在进行的生成任务时触发) // 异步生成 topic 描述(仅当描述为空且没有正在进行的生成任务时触发)
if let Some(ref topic_id) = current_topic { if let Some(ref topic_id) = current_topic {
let store = self.session_manager.store(); let store = self.session_manager.store();
if let Ok(Some(topic)) = store.get_topic(topic_id) { if let Ok(Some(topic)) = store.get_topic(topic_id)
if topic.description.is_none() && (topic.description.is_none()
|| topic || topic
.description .description
.as_ref() .as_ref()
.map(|d| d.is_empty()) .map(|d| d.is_empty())
.unwrap_or(true) .unwrap_or(true))
{ {
// 检查并设置"生成中"守卫,防止竞态条件导致重复生成 // 检查并设置"生成中"守卫,防止竞态条件导致重复生成
let should_generate = { let should_generate = {
let mut in_flight = let mut in_flight = self.description_generation_in_flight.lock();
self.description_generation_in_flight.lock(); if in_flight.contains(topic_id) {
if in_flight.contains(topic_id) { false
false } else {
} else { in_flight.insert(topic_id.clone());
in_flight.insert(topic_id.clone()); true
true }
} };
};
if should_generate { if should_generate {
let provider_config = self.provider_config.clone(); let provider_config = self.provider_config.clone();
let topic_id_clone = topic_id.clone(); let topic_id_clone = topic_id.clone();
let store_clone = store.clone(); let store_clone = store.clone();
let in_flight = self.description_generation_in_flight.clone(); let in_flight = self.description_generation_in_flight.clone();
tokio::spawn(async move { tokio::spawn(async move {
// 从 DB 查询该 topic 的第一条用户消息作为描述生成的依据 // 从 DB 查询该 topic 的第一条用户消息作为描述生成的依据
let first_user_message = store_clone let first_user_message = store_clone
.load_messages_for_topic_full(&topic_id_clone, None) .load_messages_for_topic_full(&topic_id_clone, None)
.ok() .ok()
.and_then(|msgs| { .and_then(|msgs| msgs.into_iter().find(|m| m.role == "user"))
msgs.into_iter().find(|m| m.role == "user") .map(|m| m.content);
})
.map(|m| m.content);
let message_content = match first_user_message { let message_content = match first_user_message {
Some(content) => content, Some(content) => content,
None => { None => {
tracing::warn!(topic_id = %topic_id_clone, "No user message found for topic, skipping description generation"); tracing::warn!(topic_id = %topic_id_clone, "No user message found for topic, skipping description generation");
in_flight.lock().remove(&topic_id_clone); in_flight.lock().remove(&topic_id_clone);
return; return;
}
};
let runtime_config: ProviderRuntimeConfig = provider_config.into();
if let Ok(provider) = create_provider(runtime_config) {
match generate_topic_description(
provider.as_ref(),
&message_content,
)
.await
{
Ok(description) => {
if let Err(e) = store_clone.update_topic_description(
&topic_id_clone,
&description,
) {
tracing::error!(error = %e, topic_id = %topic_id_clone, "Failed to update topic description");
} else {
tracing::info!(topic_id = %topic_id_clone, description = %description, "Topic description generated");
}
} }
}; Err(e) => {
tracing::error!(error = %e, topic_id = %topic_id_clone, "Failed to generate topic description");
let runtime_config: ProviderRuntimeConfig =
provider_config.into();
if let Ok(provider) = create_provider(runtime_config) {
match generate_topic_description(
provider.as_ref(),
&message_content,
)
.await
{
Ok(description) => {
if let Err(e) = store_clone
.update_topic_description(
&topic_id_clone,
&description,
)
{
tracing::error!(error = %e, topic_id = %topic_id_clone, "Failed to update topic description");
} else {
tracing::info!(topic_id = %topic_id_clone, description = %description, "Topic description generated");
}
}
Err(e) => {
tracing::error!(error = %e, topic_id = %topic_id_clone, "Failed to generate topic description");
}
} }
} }
// 无论成功失败,释放生成守卫 }
in_flight.lock().remove(&topic_id_clone); // 无论成功失败,释放生成守卫
}); in_flight.lock().remove(&topic_id_clone);
} });
} }
} }
} }

View File

@ -1,4 +1,6 @@
use crate::agent::{AgentError, AgentLoop, AgentRuntimeConfig, ContextCompressor, EmittedMessageHandler}; use crate::agent::{
AgentError, AgentLoop, AgentRuntimeConfig, ContextCompressor, EmittedMessageHandler,
};
#[cfg(test)] #[cfg(test)]
use crate::bus::SYSTEM_CONTEXT_SCHEDULED_PROMPT; use crate::bus::SYSTEM_CONTEXT_SCHEDULED_PROMPT;
use crate::bus::{ChatMessage, MessageBus, OutboundMessage}; use crate::bus::{ChatMessage, MessageBus, OutboundMessage};
@ -12,10 +14,10 @@ use crate::storage::{
SkillEventRepository, SkillEventRepository,
}; };
use crate::tools::ToolRegistry; use crate::tools::ToolRegistry;
use crate::tools::WaitCoordinator;
use crate::tools::task::SubagentResult;
use crate::tools::task::repository::TaskRepository; use crate::tools::task::repository::TaskRepository;
use crate::tools::task::runtime::SubagentRuntime; use crate::tools::task::runtime::SubagentRuntime;
use crate::tools::task::SubagentResult;
use crate::tools::WaitCoordinator;
use async_trait::async_trait; use async_trait::async_trait;
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
@ -557,10 +559,10 @@ impl Session {
} }
// 更新 topic 的最后活跃时间 // 更新 topic 的最后活跃时间
if let Some(ref topic_id) = topic_id { if let Some(ref topic_id) = topic_id
if let Err(e) = self.store.touch_topic(topic_id) { && let Err(e) = self.store.touch_topic(topic_id)
tracing::warn!(error = %e, topic_id = %topic_id, "Failed to touch topic"); {
} tracing::warn!(error = %e, topic_id = %topic_id, "Failed to touch topic");
} }
Ok(()) Ok(())
@ -1270,6 +1272,28 @@ impl SessionManager {
} }
} }
#[async_trait]
impl crate::scheduler::MaintenanceExecutor for SessionManager {
async fn cleanup_expired_sessions(&self) -> usize {
self.cleanup_expired_sessions().await
}
async fn run_memory_maintenance_for_all_scopes(
&self,
) -> anyhow::Result<Vec<crate::scheduler::MaintenanceRunSummary>> {
match self.run_memory_maintenance_for_all_scopes().await {
Ok(Some(result)) => Ok(vec![crate::scheduler::MaintenanceRunSummary {
scope_key: result.scope_key,
merges: result.output.merges.len(),
conflicts: result.output.conflicts.len(),
low_value: result.output.low_value_ids.len(),
}]),
Ok(None) => Ok(vec![]),
Err(error) => Err(anyhow::anyhow!(error.to_string())),
}
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@ -3033,25 +3057,3 @@ mod tests {
assert!(contents.contains(&"习惯先问方案再要代码".to_string())); assert!(contents.contains(&"习惯先问方案再要代码".to_string()));
} }
} }
#[async_trait]
impl crate::scheduler::MaintenanceExecutor for SessionManager {
async fn cleanup_expired_sessions(&self) -> usize {
self.cleanup_expired_sessions().await
}
async fn run_memory_maintenance_for_all_scopes(
&self,
) -> anyhow::Result<Vec<crate::scheduler::MaintenanceRunSummary>> {
match self.run_memory_maintenance_for_all_scopes().await {
Ok(Some(result)) => Ok(vec![crate::scheduler::MaintenanceRunSummary {
scope_key: result.scope_key,
merges: result.output.merges.len(),
conflicts: result.output.conflicts.len(),
low_value: result.output.low_value_ids.len(),
}]),
Ok(None) => Ok(vec![]),
Err(error) => Err(anyhow::anyhow!(error.to_string())),
}
}
}

View File

@ -76,11 +76,7 @@ impl SessionHistory {
} }
// 收集当前活跃 topic 集合 // 收集当前活跃 topic 集合
let active: HashSet<&str> = self let active: HashSet<&str> = self.chat_topic_ids.values().map(|s| s.as_str()).collect();
.chat_topic_ids
.values()
.map(|s| s.as_str())
.collect();
// 找一个非活跃 topic 驱逐 // 找一个非活跃 topic 驱逐
let to_evict = self.topic_histories.keys().find(|tid| { let to_evict = self.topic_histories.keys().find(|tid| {
@ -104,10 +100,10 @@ impl SessionHistory {
// 检查是否有活跃 agent 任务serial lock 被持有) // 检查是否有活跃 agent 任务serial lock 被持有)
// try_lock 成功 = 锁空闲 = 无活跃任务 = 可驱逐 // try_lock 成功 = 锁空闲 = 无活跃任务 = 可驱逐
// try_lock 失败 = 锁被持有 = 有活跃任务 = 不驱逐 // try_lock 失败 = 锁被持有 = 有活跃任务 = 不驱逐
if let Some(lock) = self.topic_serial_locks.get(*tid) { if let Some(lock) = self.topic_serial_locks.get(*tid)
if lock.try_lock().is_err() { && lock.try_lock().is_err()
return false; {
} return false;
} }
true true
}); });
@ -176,7 +172,10 @@ impl SessionHistory {
/// 获取该 topic 的 sub_done 队列 sender用于后台子代理发送结果 /// 获取该 topic 的 sub_done 队列 sender用于后台子代理发送结果
/// 调用前应已通过 `ensure_sub_done_channel` 创建队列。 /// 调用前应已通过 `ensure_sub_done_channel` 创建队列。
pub(crate) fn sub_done_sender(&mut self, topic_id: &str) -> Option<mpsc::Sender<SubagentResult>> { pub(crate) fn sub_done_sender(
&mut self,
topic_id: &str,
) -> Option<mpsc::Sender<SubagentResult>> {
self.ensure_sub_done_channel(topic_id); self.ensure_sub_done_channel(topic_id);
self.sub_done_senders.get(topic_id).cloned() self.sub_done_senders.get(topic_id).cloned()
} }
@ -334,14 +333,14 @@ impl SessionHistory {
chat_id: &str, chat_id: &str,
topic_id: Option<&str>, topic_id: Option<&str>,
) -> Result<(), AgentError> { ) -> Result<(), AgentError> {
if let Some(tid) = topic_id { if let Some(tid) = topic_id
if let Some(history) = self.topic_histories.get_mut(tid) { && let Some(history) = self.topic_histories.get_mut(tid)
#[cfg(debug_assertions)] {
let len = history.len(); #[cfg(debug_assertions)]
history.clear(); let len = history.len();
#[cfg(debug_assertions)] history.clear();
tracing::debug!(topic_id = %tid, previous_len = len, "Topic history cleared"); #[cfg(debug_assertions)]
} tracing::debug!(topic_id = %tid, previous_len = len, "Topic history cleared");
} }
self.conversations self.conversations

View File

@ -149,7 +149,7 @@ mod tests {
text: Some("hello".to_string()), text: Some("hello".to_string()),
// 使用临时目录确保跨平台兼容 // 使用临时目录确保跨平台兼容
attachments: vec![MediaItem::new( attachments: vec![MediaItem::new(
&std::env::temp_dir().join("demo.png").display().to_string(), std::env::temp_dir().join("demo.png").display().to_string(),
"image", "image",
)], )],
}, },

View File

@ -34,14 +34,14 @@ pub async fn static_handler(uri: Uri) -> Response<Body> {
None => { None => {
// 对于 SPA 应用,如果请求的是页面路由(不是静态资源),返回 index.html // 对于 SPA 应用,如果请求的是页面路由(不是静态资源),返回 index.html
// 静态资源通常包含 . (如 .js, .css, .png) // 静态资源通常包含 . (如 .js, .css, .png)
if !path.contains('.') { if !path.contains('.')
if let Some(index) = StaticAssets::get("index.html") { && let Some(index) = StaticAssets::get("index.html")
return Response::builder() {
.status(StatusCode::OK) return Response::builder()
.header(header::CONTENT_TYPE, "text/html") .status(StatusCode::OK)
.body(Body::from(index.data.into_owned())) .header(header::CONTENT_TYPE, "text/html")
.unwrap(); .body(Body::from(index.data.into_owned()))
} .unwrap();
} }
Response::builder() Response::builder()

View File

@ -11,6 +11,12 @@ use crate::agent::{SystemPrompt, SystemPromptContext, SystemPromptProvider};
/// - 两者独立演化:新增工具只需在此处加常量,不碰代理身份配置 /// - 两者独立演化:新增工具只需在此处加常量,不碰代理身份配置
pub struct ToolPromptProvider; pub struct ToolPromptProvider;
impl Default for ToolPromptProvider {
fn default() -> Self {
Self::new()
}
}
impl ToolPromptProvider { impl ToolPromptProvider {
pub fn new() -> Self { pub fn new() -> Self {
Self Self

View File

@ -111,17 +111,17 @@ impl ToolRegistryFactory {
if self.is_enabled("memory_manage") { if self.is_enabled("memory_manage") {
registry.register(MemoryManageTool::new(self.memories.clone())); registry.register(MemoryManageTool::new(self.memories.clone()));
} }
if self.is_enabled("todo_write") { if self.is_enabled("todo_write")
if let Some(ref state) = self.todo_state { && let Some(ref state) = self.todo_state
registry.register(TodoWriteTool::new( {
state.clone(), registry.register(TodoWriteTool::new(
self.todo_repository.clone(), state.clone(),
)); self.todo_repository.clone(),
registry.register(TodoReadTool::new( ));
state.clone(), registry.register(TodoReadTool::new(
self.todo_repository.clone(), state.clone(),
)); self.todo_repository.clone(),
} ));
} }
if self.is_enabled("session_send") { if self.is_enabled("session_send") {
registry.register(SessionSendTool::new(self.session_message_sender.clone())); registry.register(SessionSendTool::new(self.session_message_sender.clone()));
@ -157,15 +157,16 @@ impl ToolRegistryFactory {
} }
// 注册 Task 工具(如果启用且有 subagent_runtime // 注册 Task 工具(如果启用且有 subagent_runtime
if self.is_enabled("task") && self.task_config.enabled { if self.is_enabled("task")
if let Some(runtime) = &self.subagent_runtime { && self.task_config.enabled
registry.register(TaskTool::new(runtime.clone(), None)); && let Some(runtime) = &self.subagent_runtime
// 注册 wait_for_subagents 工具(仅主 agent用于等待异步子代理完成 {
// 默认超时从配置读取LLM 可通过 timeout_secs 参数覆盖 registry.register(TaskTool::new(runtime.clone(), None));
registry.register(WaitForSubagentsTool::new( // 注册 wait_for_subagents 工具(仅主 agent用于等待异步子代理完成
self.task_config.wait_default_timeout_secs, // 默认超时从配置读取LLM 可通过 timeout_secs 参数覆盖
)); registry.register(WaitForSubagentsTool::new(
} self.task_config.wait_default_timeout_secs,
));
} }
registry registry
@ -230,17 +231,17 @@ impl ToolRegistryFactory {
} }
// Todo 追踪工具 // Todo 追踪工具
if self.is_enabled("todo_write") { if self.is_enabled("todo_write")
if let Some(ref state) = self.todo_state { && let Some(ref state) = self.todo_state
registry.register(TodoWriteTool::new( {
state.clone(), registry.register(TodoWriteTool::new(
self.todo_repository.clone(), state.clone(),
)); self.todo_repository.clone(),
registry.register(TodoReadTool::new( ));
state.clone(), registry.register(TodoReadTool::new(
self.todo_repository.clone(), state.clone(),
)); self.todo_repository.clone(),
} ));
} }
// 注册 MCP 工具(如果提供) // 注册 MCP 工具(如果提供)

View File

@ -97,11 +97,7 @@ impl WaitCoordinator for SessionWaitCoordinator {
results results
} }
async fn wait( async fn wait(&self, timeout: Duration, cancel_rx: Option<watch::Receiver<()>>) -> WaitEvent {
&self,
timeout: Duration,
cancel_rx: Option<watch::Receiver<()>>,
) -> WaitEvent {
// 1. 设置 waiting=true // 1. 设置 waiting=true
{ {
let mut session = self.session.lock().await; let mut session = self.session.lock().await;

View File

@ -142,13 +142,13 @@ pub async fn ws_handler(
auth_cfg: Option<axum::Extension<crate::gateway::auth::AuthConfig>>, auth_cfg: Option<axum::Extension<crate::gateway::auth::AuthConfig>>,
) -> Response { ) -> Response {
// 若启用了认证auth_cfg 存在且 token 已配置),校验 query param 中的 token // 若启用了认证auth_cfg 存在且 token 已配置),校验 query param 中的 token
if let Some(axum::Extension(cfg)) = auth_cfg { if let Some(axum::Extension(cfg)) = auth_cfg
if let Some(ref expected) = cfg.token { && let Some(ref expected) = cfg.token
let provided = query.token.as_deref(); {
if !crate::gateway::auth::token_matches(provided, &Some(expected.clone())) { let provided = query.token.as_deref();
tracing::warn!("WebSocket connection rejected: missing or invalid token"); if !crate::gateway::auth::token_matches(provided, &Some(expected.clone())) {
return (StatusCode::UNAUTHORIZED, "missing or invalid token").into_response(); tracing::warn!("WebSocket connection rejected: missing or invalid token");
} return (StatusCode::UNAUTHORIZED, "missing or invalid token").into_response();
} }
} }
@ -653,66 +653,61 @@ async fn handle_inbound(
} }
// 处理定时任务列表 // 处理定时任务列表
if let Some(jobs_json) = response.metadata.get("scheduler_jobs") { if let Some(jobs_json) = response.metadata.get("scheduler_jobs")
if let Ok(jobs) = && let Ok(jobs) =
serde_json::from_str::<Vec<crate::protocol::SchedulerJobSummary>>(jobs_json) serde_json::from_str::<Vec<crate::protocol::SchedulerJobSummary>>(jobs_json)
{ {
let _ = sender.send(WsOutbound::SchedulerJobList { jobs }).await; let _ = sender.send(WsOutbound::SchedulerJobList { jobs }).await;
}
} }
// 处理技能列表 // 处理技能列表
if let Some(skills_json) = response.metadata.get("skills") { if let Some(skills_json) = response.metadata.get("skills")
if let Ok(skills) = && let Ok(skills) =
serde_json::from_str::<Vec<crate::protocol::SkillSummary>>(skills_json) serde_json::from_str::<Vec<crate::protocol::SkillSummary>>(skills_json)
{ {
let _ = sender.send(WsOutbound::SkillList { skills }).await; let _ = sender.send(WsOutbound::SkillList { skills }).await;
}
} }
// 处理 Todo 列表 // 处理 Todo 列表
if let Some(todos_json) = response.metadata.get("todos") { if let Some(todos_json) = response.metadata.get("todos")
if let Ok(todos) = && let Ok(todos) =
serde_json::from_str::<Vec<crate::protocol::TodoItemSummary>>(todos_json) serde_json::from_str::<Vec<crate::protocol::TodoItemSummary>>(todos_json)
{ {
let scope_key = response let scope_key = response
.metadata .metadata
.get("todos_scope_key") .get("todos_scope_key")
.cloned() .cloned()
.unwrap_or_default(); .unwrap_or_default();
tracing::debug!(todo_count = todos.len(), %scope_key, "list_todos command response"); tracing::debug!(todo_count = todos.len(), %scope_key, "list_todos command response");
let _ = sender.send(WsOutbound::TodoList { todos, scope_key }).await; let _ = sender.send(WsOutbound::TodoList { todos, scope_key }).await;
}
} }
// 处理记忆列表 // 处理记忆列表
if let Some(memories_json) = response.metadata.get("memories") { if let Some(memories_json) = response.metadata.get("memories")
if let Ok(memories) = && let Ok(memories) =
serde_json::from_str::<Vec<crate::protocol::MemorySummary>>(memories_json) serde_json::from_str::<Vec<crate::protocol::MemorySummary>>(memories_json)
{ {
let _ = sender.send(WsOutbound::MemoryList { memories }).await; let _ = sender.send(WsOutbound::MemoryList { memories }).await;
}
} }
// 记忆 CRUD 后自动刷新列表 // 记忆 CRUD 后自动刷新列表
if response.metadata.get("memory_updated").map(|v| v.as_str()) == Some("true") { if response.metadata.get("memory_updated").map(|v| v.as_str()) == Some("true")
if let Ok(records) = && let Ok(records) =
store.list_memories_for_scope("user", crate::storage::GLOBAL_SCOPE_KEY) store.list_memories_for_scope("user", crate::storage::GLOBAL_SCOPE_KEY)
{ {
let memories: Vec<crate::protocol::MemorySummary> = records let memories: Vec<crate::protocol::MemorySummary> = records
.into_iter() .into_iter()
.filter(|m| m.namespace != "_meta") .filter(|m| m.namespace != "_meta")
.map(|m| crate::protocol::MemorySummary { .map(|m| crate::protocol::MemorySummary {
id: m.id, id: m.id,
namespace: m.namespace, namespace: m.namespace,
memory_key: m.memory_key, memory_key: m.memory_key,
content: m.content, content: m.content,
created_at: m.created_at, created_at: m.created_at,
updated_at: m.updated_at, updated_at: m.updated_at,
}) })
.collect(); .collect();
let _ = sender.send(WsOutbound::MemoryList { memories }).await; let _ = sender.send(WsOutbound::MemoryList { memories }).await;
}
} }
// 处理加载聊天消息请求 // 处理加载聊天消息请求
@ -738,31 +733,29 @@ async fn handle_inbound(
} }
} }
if current_topic_id.is_none() { if current_topic_id.is_none()
if let Some(topics_json) = response.metadata.get("topics") { && let Some(topics_json) = response.metadata.get("topics")
match serde_json::from_str::<Vec<crate::protocol::TopicSummary>>( {
topics_json, match serde_json::from_str::<Vec<crate::protocol::TopicSummary>>(topics_json) {
) { Ok(topics) => {
Ok(topics) => { if let Some(first_topic) = topics.first() {
if let Some(first_topic) = topics.first() { let topic_id = first_topic.topic_id.clone();
let topic_id = first_topic.topic_id.clone(); *current_topic_id = Some(topic_id.clone());
*current_topic_id = Some(topic_id.clone()); if let Err(e) = send_topic_history(
if let Err(e) = send_topic_history( &store,
&store, current_session_id,
current_session_id, &topic_id,
&topic_id, sender,
sender, &state.task_repository,
&state.task_repository, )
) .await
.await {
{ tracing::warn!(error = %e, topic_id = %topic_id, "Failed to send initial topic history");
tracing::warn!(error = %e, topic_id = %topic_id, "Failed to send initial topic history");
}
} }
} }
Err(e) => { }
tracing::warn!(error = %e, "Failed to parse topics metadata for initial history"); Err(e) => {
} tracing::warn!(error = %e, "Failed to parse topics metadata for initial history");
} }
} }
} }
@ -820,10 +813,10 @@ async fn send_topic_history(
let mut tool_call_ids_with_results: std::collections::HashSet<String> = let mut tool_call_ids_with_results: std::collections::HashSet<String> =
std::collections::HashSet::new(); std::collections::HashSet::new();
for msg in &messages { for msg in &messages {
if msg.role == "tool" { if msg.role == "tool"
if let Some(ref tcid) = msg.tool_call_id { && let Some(ref tcid) = msg.tool_call_id
tool_call_ids_with_results.insert(tcid.clone()); {
} tool_call_ids_with_results.insert(tcid.clone());
} }
} }
@ -894,7 +887,8 @@ fn reconcile_running_in_messages(
topic_id: &str, topic_id: &str,
) { ) {
let has_running_placeholder = messages.iter().any(|m| { let has_running_placeholder = messages.iter().any(|m| {
m.role == "tool" && crate::gateway::session::extract_task_id_from_content(&m.content).is_some() m.role == "tool"
&& crate::gateway::session::extract_task_id_from_content(&m.content).is_some()
}); });
if !has_running_placeholder { if !has_running_placeholder {
return; // 无需查询 DB return; // 无需查询 DB
@ -916,13 +910,15 @@ fn reconcile_running_in_messages(
if msg.role != "tool" { if msg.role != "tool" {
continue; continue;
} }
let Some((task_id, is_json)) = crate::gateway::session::extract_task_id_from_content(&msg.content) let Some((task_id, is_json)) =
crate::gateway::session::extract_task_id_from_content(&msg.content)
else { else {
continue; continue;
}; };
match status_map.get(task_id.as_str()) { match status_map.get(task_id.as_str()) {
Some(&status) if status != "running" => { Some(&status) if status != "running" => {
msg.content = crate::gateway::session::format_reconciled_content(&task_id, status, is_json); msg.content =
crate::gateway::session::format_reconciled_content(&task_id, status, is_json);
} }
_ => {} // 不存在(已清理)或仍在运行:保留原占位 _ => {} // 不存在(已清理)或仍在运行:保留原占位
} }
@ -945,10 +941,10 @@ async fn send_task_messages(
let mut tool_call_ids_with_results: std::collections::HashSet<String> = let mut tool_call_ids_with_results: std::collections::HashSet<String> =
std::collections::HashSet::new(); std::collections::HashSet::new();
for msg in &messages { for msg in &messages {
if msg.role == "tool" { if msg.role == "tool"
if let Some(ref tcid) = msg.tool_call_id { && let Some(ref tcid) = msg.tool_call_id
tool_call_ids_with_results.insert(tcid.clone()); {
} tool_call_ids_with_results.insert(tcid.clone());
} }
} }
@ -1041,10 +1037,10 @@ fn set_subagent_task_id(outbound: &mut WsOutbound, task_id: &str) {
fn extract_parent_task_id(task: &crate::tools::task::types::TaskSession) -> Option<String> { fn extract_parent_task_id(task: &crate::tools::task::types::TaskSession) -> Option<String> {
let parent = &task.parent_session_id; let parent = &task.parent_session_id;
// 仅当父会话是子智能体会话时才提取(格式: "sub:...:task:{uuid}" // 仅当父会话是子智能体会话时才提取(格式: "sub:...:task:{uuid}"
if parent.starts_with("sub:") { if parent.starts_with("sub:")
if let Some(pos) = parent.find(":task:") { && let Some(pos) = parent.find(":task:")
return Some(parent[pos + 1..].to_string()); // "task:{uuid}" {
} return Some(parent[pos + 1..].to_string()); // "task:{uuid}"
} }
None None
} }

View File

@ -3,7 +3,7 @@ use chrono_tz::Tz;
use std::path::PathBuf; use std::path::PathBuf;
use tracing_appender::rolling::{RollingFileAppender, Rotation}; use tracing_appender::rolling::{RollingFileAppender, Rotation};
use tracing_subscriber::{ use tracing_subscriber::{
fmt, fmt::time::FormatTime, layer::SubscriberExt, util::SubscriberInitExt, EnvFilter, Layer, EnvFilter, Layer, fmt, fmt::time::FormatTime, layer::SubscriberExt, util::SubscriberInitExt,
}; };
use crate::config::LogFormat; use crate::config::LogFormat;
@ -61,14 +61,14 @@ pub fn init_logging(timezone: Tz, log_format: LogFormat) {
let log_dir = get_default_log_dir(); let log_dir = get_default_log_dir();
// Create log directory if it doesn't exist // Create log directory if it doesn't exist
if !log_dir.exists() { if !log_dir.exists()
if let Err(e) = std::fs::create_dir_all(&log_dir) { && let Err(e) = std::fs::create_dir_all(&log_dir)
eprintln!( {
"Warning: Failed to create log directory {}: {}", eprintln!(
log_dir.display(), "Warning: Failed to create log directory {}: {}",
e log_dir.display(),
); e
} );
} }
// Create file appender with daily rotation // Create file appender with daily rotation

View File

@ -6,9 +6,9 @@
//! - Connects to MCP servers asynchronously //! - Connects to MCP servers asynchronously
//! - Dynamically registers MCP tools via the Tool trait adapter //! - Dynamically registers MCP tools via the Tool trait adapter
use parking_lot::Mutex;
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use parking_lot::Mutex;
use tokio::sync::RwLock; use tokio::sync::RwLock;
use http::{HeaderName, HeaderValue}; use http::{HeaderName, HeaderValue};

View File

@ -102,7 +102,7 @@ impl McpServerConfig {
command, command,
args: self.args.clone().unwrap_or_default(), args: self.args.clone().unwrap_or_default(),
env: self.env.clone().unwrap_or_default(), env: self.env.clone().unwrap_or_default(),
cwd: self.cwd.as_ref().map(|s| std::path::PathBuf::from(s)), cwd: self.cwd.as_ref().map(std::path::PathBuf::from),
}) })
} }
"http" | "streamableHttp" => { "http" | "streamableHttp" => {

View File

@ -111,25 +111,22 @@ impl PicoBotTool for McpToolWrapper {
.call_tool(&self.server_key, &self.tool_name, args); .call_tool(&self.server_key, &self.tool_name, args);
let result = if self.timeout_secs > 0 { let result = if self.timeout_secs > 0 {
tokio::time::timeout( tokio::time::timeout(std::time::Duration::from_secs(self.timeout_secs), call)
std::time::Duration::from_secs(self.timeout_secs), .await
call, .map_err(|_| {
) tracing::warn!(
.await server_key = %self.server_key,
.map_err(|_| { tool = %self.tool_name,
tracing::warn!( timeout_secs = self.timeout_secs,
server_key = %self.server_key, "MCP tool call timed out"
tool = %self.tool_name, );
timeout_secs = self.timeout_secs, anyhow::anyhow!(
"MCP tool call timed out" "MCP tool '{}' on server '{}' timed out after {}s",
); self.tool_name,
anyhow::anyhow!( self.server_key,
"MCP tool '{}' on server '{}' timed out after {}s", self.timeout_secs
self.tool_name, )
self.server_key, })??
self.timeout_secs
)
})??
} else { } else {
call.await? call.await?
}; };
@ -183,12 +180,8 @@ pub async fn register_mcp_tools(
let all_tools = manager.all_tools().await; let all_tools = manager.all_tools().await;
for (server_key, tool_info) in all_tools { for (server_key, tool_info) in all_tools {
let wrapper = McpToolWrapper::new( let wrapper =
manager.clone(), McpToolWrapper::new(manager.clone(), server_key.clone(), tool_info, timeout_secs);
server_key.clone(),
tool_info,
timeout_secs,
);
tracing::info!( tracing::info!(
name = %wrapper.name(), name = %wrapper.name(),

View File

@ -40,7 +40,8 @@ pub const MESSAGE_PROCESSING_ERRORS: &str = "picobot_message_processing_errors_t
/// 幂等:首次调用安装 recorder 并缓存 handle后续调用含热重启返回缓存的 handle。 /// 幂等:首次调用安装 recorder 并缓存 handle后续调用含热重启返回缓存的 handle。
/// 这避免了热重启后 `install_recorder()` 因 recorder 已安装而失败、导致 `/metrics` 返回 503 的问题。 /// 这避免了热重启后 `install_recorder()` 因 recorder 已安装而失败、导致 `/metrics` 返回 503 的问题。
/// 返回 None 表示安装失败非致命metrics 静默降级)。 /// 返回 None 表示安装失败非致命metrics 静默降级)。
static PROMETHEUS_HANDLE: std::sync::OnceLock<Option<PrometheusHandle>> = std::sync::OnceLock::new(); static PROMETHEUS_HANDLE: std::sync::OnceLock<Option<PrometheusHandle>> =
std::sync::OnceLock::new();
pub fn init_recorder() -> Option<PrometheusHandle> { pub fn init_recorder() -> Option<PrometheusHandle> {
PROMETHEUS_HANDLE PROMETHEUS_HANDLE

View File

@ -342,7 +342,7 @@ pub fn home_dir() -> Option<PathBuf> {
// Windows: support USERPROFILE // Windows: support USERPROFILE
env::var_os("USERPROFILE").map(PathBuf::from) env::var_os("USERPROFILE").map(PathBuf::from)
}) })
.or_else(|| dirs::home_dir()) .or_else(dirs::home_dir)
} }
/// 返回 PicoBot 主目录,若无法确定则回退到当前目录 `"."`。 /// 返回 PicoBot 主目录,若无法确定则回退到当前目录 `"."`。

View File

@ -125,7 +125,6 @@ pub struct AnthropicProvider {
api_key: String, api_key: String,
base_url: String, base_url: String,
extra_headers: HashMap<String, String>, extra_headers: HashMap<String, String>,
#[cfg_attr(not(debug_assertions), allow(dead_code))]
llm_timeout_secs: u64, llm_timeout_secs: u64,
model_id: String, model_id: String,
temperature: Option<f32>, temperature: Option<f32>,
@ -316,15 +315,15 @@ impl LLMProvider for AnthropicProvider {
req_builder = req_builder.header(key.as_str(), value.as_str()); req_builder = req_builder.header(key.as_str(), value.as_str());
} }
let resp = req_builder.json(&body).send().await.map_err(|e| { let resp = req_builder.json(&body).send().await.inspect_err(|e| {
tracing::error!( tracing::error!(
provider = %self.name, provider = %self.name,
model = %self.model_id, model = %self.model_id,
url = %url, url = %url,
error = %format_error_chain(&e), timeout_secs = self.llm_timeout_secs,
error = %format_error_chain(e),
"Anthropic: HTTP request failed" "Anthropic: HTTP request failed"
); );
e
})?; })?;
let status = resp.status(); let status = resp.status();
let text = resp.text().await?; let text = resp.text().await?;
@ -635,7 +634,7 @@ mod tests {
#[test] #[test]
fn test_format_error_chain_single() { fn test_format_error_chain_single() {
let err = std::io::Error::new(std::io::ErrorKind::Other, "single error"); let err = std::io::Error::other("single error");
let chain = format_error_chain(&err); let chain = format_error_chain(&err);
assert_eq!(chain, "single error"); assert_eq!(chain, "single error");
} }
@ -649,7 +648,7 @@ mod tests {
#[test] #[test]
fn test_format_error_chain_nested() { fn test_format_error_chain_nested() {
let inner = std::io::Error::new(std::io::ErrorKind::Other, "root cause"); let inner = std::io::Error::other("root cause");
let outer = OuterError::Wrapped(inner); let outer = OuterError::Wrapped(inner);
let chain = format_error_chain(&outer); let chain = format_error_chain(&outer);
assert!(chain.contains("outer wrapper")); assert!(chain.contains("outer wrapper"));

View File

@ -63,22 +63,19 @@ impl StreamingAccumulator {
name: Option<&str>, name: Option<&str>,
arguments: Option<&str>, arguments: Option<&str>,
) { ) {
let entry = self let entry = self.tool_calls.entry(index).or_default();
.tool_calls
.entry(index)
.or_insert_with(StreamingToolCall::default);
// 只在 id 非空时才更新,防止流式响应中后续 chunk 的空 id 覆盖之前的值 // 只在 id 非空时才更新,防止流式响应中后续 chunk 的空 id 覆盖之前的值
if let Some(id) = id { if let Some(id) = id
if !id.is_empty() { && !id.is_empty()
entry.id = id.to_string(); {
} entry.id = id.to_string();
} }
// 只在 name 非空时才更新,防止流式响应中后续 chunk 的 None 覆盖之前的值 // 只在 name 非空时才更新,防止流式响应中后续 chunk 的 None 覆盖之前的值
if let Some(name) = name { if let Some(name) = name
if !name.is_empty() { && !name.is_empty()
entry.name = name.to_string(); {
} entry.name = name.to_string();
} }
if let Some(args) = arguments { if let Some(args) = arguments {
entry.arguments.push_str(args); entry.arguments.push_str(args);
@ -107,8 +104,8 @@ impl StreamingAccumulator {
.into_iter() .into_iter()
.filter(|(_, call)| !call.id.is_empty() && !call.name.is_empty()) .filter(|(_, call)| !call.id.is_empty() && !call.name.is_empty())
.map(|(_, call)| { .map(|(_, call)| {
let arguments = serde_json::from_str(&call.arguments) let arguments =
.unwrap_or_else(|_| serde_json::Value::Null); serde_json::from_str(&call.arguments).unwrap_or(serde_json::Value::Null);
ToolCall { ToolCall {
id: call.id, id: call.id,
name: call.name, name: call.name,
@ -218,14 +215,12 @@ fn convert_content_blocks(
} }
// 如果只有一个文本块且没有通知,返回字符串形式 // 如果只有一个文本块且没有通知,返回字符串形式
if converted_blocks.len() == 1 { if converted_blocks.len() == 1
if let Some(block) = converted_blocks.first() { && let Some(block) = converted_blocks.first()
if block.get("type").and_then(|t| t.as_str()) == Some("text") { && block.get("type").and_then(|t| t.as_str()) == Some("text")
if let Some(text) = block.get("text").and_then(|t| t.as_str()) { && let Some(text) = block.get("text").and_then(|t| t.as_str())
return Value::String(text.to_string()); {
} return Value::String(text.to_string());
}
}
} }
return Value::Array(converted_blocks); return Value::Array(converted_blocks);
@ -233,10 +228,10 @@ fn convert_content_blocks(
} }
// 原有逻辑 - 模型支持图片,正常转换 // 原有逻辑 - 模型支持图片,正常转换
if blocks.len() == 1 { if blocks.len() == 1
if let ContentBlock::Text { text } = &blocks[0] { && let ContentBlock::Text { text } = &blocks[0]
return Value::String(text.clone()); {
} return Value::String(text.clone());
} }
Value::Array( Value::Array(
blocks blocks
@ -481,14 +476,12 @@ impl OpenAIProvider {
} }
// 提取流式末帧的 usagestream_options.include_usage=true 时返回) // 提取流式末帧的 usagestream_options.include_usage=true 时返回)
if let Some(usage_val) = json.get("usage") { if let Some(usage_val) = json.get("usage")
if !usage_val.is_null() { && !usage_val.is_null()
if let Ok(u) = && let Ok(u) =
serde_json::from_value::<OpenAIUsage>(usage_val.clone()) serde_json::from_value::<OpenAIUsage>(usage_val.clone())
{ {
accumulator.set_usage(u); accumulator.set_usage(u);
}
}
} }
// 提取 choices // 提取 choices
@ -605,13 +598,11 @@ impl OpenAIProvider {
} }
// 提取流式末帧的 usage与主循环一致 // 提取流式末帧的 usage与主循环一致
if let Some(usage_val) = json.get("usage") { if let Some(usage_val) = json.get("usage")
if !usage_val.is_null() { && !usage_val.is_null()
if let Ok(u) = serde_json::from_value::<OpenAIUsage>(usage_val.clone()) && let Ok(u) = serde_json::from_value::<OpenAIUsage>(usage_val.clone())
{ {
accumulator.set_usage(u); accumulator.set_usage(u);
}
}
} }
if let Some(choices) = json.get("choices").and_then(|c| c.as_array()) { if let Some(choices) = json.get("choices").and_then(|c| c.as_array()) {
@ -684,53 +675,54 @@ impl OpenAIProvider {
// 回退:当流式解析未获取到任何内容且无 tool call 时, // 回退:当流式解析未获取到任何内容且无 tool call 时,
// 服务器可能返回的是非 SSE 格式的纯 JSON尝试直接反序列化整个响应体 // 服务器可能返回的是非 SSE 格式的纯 JSON尝试直接反序列化整个响应体
if response.content.is_empty() && response.tool_calls.is_empty() { if response.content.is_empty()
if let Ok(openai_resp) = serde_json::from_str::<OpenAIResponse>(&raw_body) { && response.tool_calls.is_empty()
let fallback_content = openai_resp && let Ok(openai_resp) = serde_json::from_str::<OpenAIResponse>(&raw_body)
{
let fallback_content = openai_resp
.choices
.first()
.and_then(|c| c.message.content.as_deref())
.unwrap_or("")
.to_string();
if !fallback_content.is_empty() {
tracing::debug!(
model = %self.model_id,
"Streaming accumulator empty, falling back to non-SSE JSON parsing"
);
response.content = fallback_content;
response.reasoning_content = openai_resp
.choices .choices
.first() .first()
.and_then(|c| c.message.content.as_deref()) .and_then(|c| c.message.reasoning_content.clone());
.unwrap_or("") response.tool_calls = openai_resp
.to_string(); .choices
if !fallback_content.is_empty() { .first()
tracing::debug!( .map(|c| {
model = %self.model_id, c.message
"Streaming accumulator empty, falling back to non-SSE JSON parsing" .tool_calls
); .iter()
response.content = fallback_content; .map(|tc| ToolCall {
response.reasoning_content = openai_resp id: tc.id.clone(),
.choices name: tc.function.name.clone(),
.first() arguments: match &tc.function.arguments {
.and_then(|c| c.message.reasoning_content.clone()); OAIFunctionArguments::Json(args) => args.clone(),
response.tool_calls = openai_resp OAIFunctionArguments::String(args) => {
.choices serde_json::from_str(args)
.first() .unwrap_or(serde_json::Value::Null)
.map(|c| { }
c.message },
.tool_calls })
.iter() .collect()
.map(|tc| ToolCall { })
id: tc.id.clone(), .unwrap_or_default();
name: tc.function.name.clone(), // 回退场景下也从非流式响应提取 usage
arguments: match &tc.function.arguments { response.usage = Usage {
OAIFunctionArguments::Json(args) => args.clone(), prompt_tokens: openai_resp.usage.prompt_tokens,
OAIFunctionArguments::String(args) => { completion_tokens: openai_resp.usage.completion_tokens,
serde_json::from_str(args) total_tokens: openai_resp.usage.total_tokens,
.unwrap_or(serde_json::Value::Null) cached_tokens: openai_resp.usage.cached_tokens(),
} };
},
})
.collect()
})
.unwrap_or_default();
// 回退场景下也从非流式响应提取 usage
response.usage = Usage {
prompt_tokens: openai_resp.usage.prompt_tokens,
completion_tokens: openai_resp.usage.completion_tokens,
total_tokens: openai_resp.usage.total_tokens,
cached_tokens: openai_resp.usage.cached_tokens(),
};
}
} }
} }
@ -761,26 +753,25 @@ impl OpenAIProvider {
std::collections::HashSet::new(); std::collections::HashSet::new();
for (i, m) in request.messages.iter().enumerate().rev() { for (i, m) in request.messages.iter().enumerate().rev() {
if m.role == "tool" { if m.role == "tool"
if let Some(ref tc_id) = m.tool_call_id { && let Some(ref tc_id) = m.tool_call_id
resolved_tool_ids.insert(tc_id.as_str()); {
} resolved_tool_ids.insert(tc_id.as_str());
} }
if m.role == "assistant" { if m.role == "assistant"
if let Some(ref calls) = m.tool_calls { && let Some(ref calls) = m.tool_calls
if !calls.is_empty() { && !calls.is_empty()
let all_resolved = calls {
.iter() let all_resolved = calls
.all(|tc| resolved_tool_ids.contains(tc.id.as_str())); .iter()
if all_resolved { .all(|tc| resolved_tool_ids.contains(tc.id.as_str()));
for tc in calls { if all_resolved {
with_parent.insert(tc.id.as_str()); for tc in calls {
} with_parent.insert(tc.id.as_str());
} else {
skip_assistant_indices.insert(i);
}
} }
} else {
skip_assistant_indices.insert(i);
} }
} }
} }
@ -827,36 +818,37 @@ impl OpenAIProvider {
} }
if m.role == "assistant" { if m.role == "assistant" {
if let Some(ref calls) = m.tool_calls { if let Some(ref calls) = m.tool_calls
if !calls.is_empty() && !skip_assistant_indices.contains(&i) { && !calls.is_empty()
pending_tool_ids = calls.iter().map(|tc| tc.id.as_str()).collect(); && !skip_assistant_indices.contains(&i)
pending_assistant_idx = Some(i); {
} pending_tool_ids = calls.iter().map(|tc| tc.id.as_str()).collect();
pending_assistant_idx = Some(i);
} }
} else if m.role == "tool" { } else if m.role == "tool"
if let Some(ref tc_id) = m.tool_call_id { && let Some(ref tc_id) = m.tool_call_id
pending_tool_ids.remove(tc_id.as_str()); {
if pending_tool_ids.is_empty() { pending_tool_ids.remove(tc_id.as_str());
pending_assistant_idx = None; if pending_tool_ids.is_empty() {
} pending_assistant_idx = None;
} }
} }
} }
// Handle trailing assistant with unresolved immediate tool results // Handle trailing assistant with unresolved immediate tool results
if !pending_tool_ids.is_empty() { if !pending_tool_ids.is_empty()
if let Some(idx) = pending_assistant_idx { && let Some(idx) = pending_assistant_idx
skip_assistant_indices.insert(idx); {
tracing::warn!( skip_assistant_indices.insert(idx);
message_index = idx, tracing::warn!(
pending_tool_call_count = pending_tool_ids.len(), message_index = idx,
"build_request_body: trailing assistant tool_calls without \ pending_tool_call_count = pending_tool_ids.len(),
"build_request_body: trailing assistant tool_calls without \
immediately following tool results stripping tool_calls" immediately following tool results stripping tool_calls"
); );
if let Some(calls) = &request.messages[idx].tool_calls { if let Some(calls) = &request.messages[idx].tool_calls {
for tc in calls.iter() { for tc in calls.iter() {
with_parent.remove(tc.id.as_str()); with_parent.remove(tc.id.as_str());
}
} }
} }
} }
@ -956,11 +948,10 @@ impl OpenAIProvider {
"content": convert_content_blocks(supports_images, &self.name, &self.model_id, &m.content, i) "content": convert_content_blocks(supports_images, &self.name, &self.model_id, &m.content, i)
}); });
if m.role == "assistant" { if m.role == "assistant"
if let Some(reasoning_content) = &m.reasoning_content { && let Some(reasoning_content) = &m.reasoning_content {
message["reasoning_content"] = Value::String(reasoning_content.clone()); message["reasoning_content"] = Value::String(reasoning_content.clone());
} }
}
Some(message) Some(message)
} }
@ -1150,15 +1141,14 @@ impl LLMProvider for OpenAIProvider {
for (i, msg) in msgs.iter().enumerate() { for (i, msg) in msgs.iter().enumerate() {
if let Some(content) = msg.get("content").and_then(|c| c.as_array()) { if let Some(content) = msg.get("content").and_then(|c| c.as_array()) {
for (j, item) in content.iter().enumerate() { for (j, item) in content.iter().enumerate() {
if item.get("type").and_then(|t| t.as_str()) == Some("image_url") { if item.get("type").and_then(|t| t.as_str()) == Some("image_url")
if let Some(url_str) = item && let Some(url_str) = item
.get("image_url") .get("image_url")
.and_then(|u| u.get("url")) .and_then(|u| u.get("url"))
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
{ {
let prefix: String = url_str.chars().take(20).collect(); let prefix: String = url_str.chars().take(20).collect();
tracing::debug!(msg_idx = i, item_idx = j, image_prefix = %prefix, image_url_len = %url_str.len(), "Image in LLM request (first 20 bytes shown)"); tracing::debug!(msg_idx = i, item_idx = j, image_prefix = %prefix, image_url_len = %url_str.len(), "Image in LLM request (first 20 bytes shown)");
}
} }
} }
} }

View File

@ -734,13 +734,13 @@ impl RuntimeJob {
return Ok(()); return Ok(());
} }
if let Some(max_runs) = self.max_runs { if let Some(max_runs) = self.max_runs
if self.run_count >= max_runs { && self.run_count >= max_runs
self.state = SchedulerJobState::Completed; {
self.next_fire_at = None; self.state = SchedulerJobState::Completed;
self.completed_at = Some(now.timestamp_millis()); self.next_fire_at = None;
return Ok(()); self.completed_at = Some(now.timestamp_millis());
} return Ok(());
} }
let reference_ms = self.next_fire_at.or(self.last_fired_at); let reference_ms = self.next_fire_at.or(self.last_fired_at);

View File

@ -1,13 +1,13 @@
use crate::platform::{ use crate::platform::{
atomic_rename, home_dir as platform_home_dir, path_to_uri, xml_escape as platform_xml_escape, atomic_rename, home_dir as platform_home_dir, path_to_uri, xml_escape as platform_xml_escape,
}; };
use parking_lot::RwLock;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use serde_json::json; use serde_json::json;
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
use std::fs; use std::fs;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::Arc; use std::sync::Arc;
use parking_lot::RwLock;
#[cfg(test)] #[cfg(test)]
static SKILL_TEST_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(()); static SKILL_TEST_ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
@ -143,9 +143,7 @@ impl SkillRuntime {
} }
pub fn is_empty(&self) -> bool { pub fn is_empty(&self) -> bool {
self.catalog self.catalog.read().is_empty()
.read()
.is_empty()
} }
pub fn len(&self) -> usize { pub fn len(&self) -> usize {
@ -153,9 +151,7 @@ impl SkillRuntime {
} }
pub fn system_index_prompt(&self) -> Option<String> { pub fn system_index_prompt(&self) -> Option<String> {
self.catalog self.catalog.read().system_index_prompt()
.read()
.system_index_prompt()
} }
/// 按白/黑名单过滤后的技能索引。供专家/子代理按 `CapabilityPolicy` 过滤技能可见性。 /// 按白/黑名单过滤后的技能索引。供专家/子代理按 `CapabilityPolicy` 过滤技能可见性。
@ -170,34 +166,23 @@ impl SkillRuntime {
} }
pub fn discovery_event_payload(&self) -> serde_json::Value { pub fn discovery_event_payload(&self) -> serde_json::Value {
self.catalog self.catalog.read().discovery_event_payload()
.read()
.discovery_event_payload()
} }
pub fn offered_event_payload(&self) -> serde_json::Value { pub fn offered_event_payload(&self) -> serde_json::Value {
self.catalog self.catalog.read().offered_event_payload()
.read()
.offered_event_payload()
} }
pub fn activation_payload(&self, name: &str) -> Result<String, String> { pub fn activation_payload(&self, name: &str) -> Result<String, String> {
self.catalog self.catalog.read().activation_payload(name)
.read()
.activation_payload(name)
} }
pub fn activation_event_payload(&self, name: &str) -> Result<serde_json::Value, String> { pub fn activation_event_payload(&self, name: &str) -> Result<serde_json::Value, String> {
self.catalog self.catalog.read().activation_event_payload(name)
.read()
.activation_event_payload(name)
} }
pub fn list_skills(&self) -> Vec<Skill> { pub fn list_skills(&self) -> Vec<Skill> {
self.catalog self.catalog.read().skills.clone()
.read()
.skills
.clone()
} }
/// List all discovered skills including disabled ones, with their disabled scopes. /// List all discovered skills including disabled ones, with their disabled scopes.
@ -226,10 +211,7 @@ impl SkillRuntime {
} }
pub fn get_skill(&self, name: &str) -> Option<Skill> { pub fn get_skill(&self, name: &str) -> Option<Skill> {
self.catalog self.catalog.read().find_skill(name).cloned()
.read()
.find_skill(name)
.cloned()
} }
pub fn create_skill( pub fn create_skill(
@ -450,7 +432,7 @@ impl SkillCatalog {
// Load from least specific to most specific so later sources win on conflicts. // Load from least specific to most specific so later sources win on conflicts.
for source in source_order(&config.sources) { for source in source_order(&config.sources) {
sources_seen += 1; sources_seen += 1;
let root = source_root(&source, &cwd); let root = source_root(&source, cwd);
let Some(root) = root else { continue }; let Some(root) = root else { continue };
for skill in load_skills_from_root(&root, source.clone()) { for skill in load_skills_from_root(&root, source.clone()) {
@ -519,7 +501,7 @@ impl SkillCatalog {
.filter(|s| { .filter(|s| {
allowed_set allowed_set
.as_ref() .as_ref()
.map_or(true, |set| set.contains(s.name.as_str())) .is_none_or(|set| set.contains(s.name.as_str()))
}) })
.collect(); .collect();

View File

@ -299,7 +299,7 @@ pub(super) fn ensure_todos_schema(conn: &Connection) -> Result<(), StorageError>
} }
// Column migration: add created_by_message_id if it doesn't exist // Column migration: add created_by_message_id if it doesn't exist
let has_column = has_column(&conn, "todos", "created_by_message_id")?; let has_column = has_column(conn, "todos", "created_by_message_id")?;
if !has_column { if !has_column {
tracing::info!("Adding created_by_message_id column to todos table"); tracing::info!("Adding created_by_message_id column to todos table");
conn.execute( conn.execute(

View File

@ -574,9 +574,7 @@ impl SessionStore {
let mut stmt = conn.prepare( let mut stmt = conn.prepare(
"SELECT id, provider, model FROM topics WHERE provider IS NOT NULL OR model IS NOT NULL", "SELECT id, provider, model FROM topics WHERE provider IS NOT NULL OR model IS NOT NULL",
)?; )?;
let rows = stmt.query_map([], |row| { let rows = stmt.query_map([], |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)))?;
Ok((row.get(0)?, row.get(1)?, row.get(2)?))
})?;
let mut result = Vec::new(); let mut result = Vec::new();
for row in rows { for row in rows {
result.push(row?); result.push(row?);
@ -1027,7 +1025,7 @@ impl SessionStore {
new_messages.iter().partition(|m| { new_messages.iter().partition(|m| {
m.system_context m.system_context
.as_deref() .as_deref()
.map_or(false, |sc| sc.starts_with("history_compaction")) .is_some_and(|sc| sc.starts_with("history_compaction"))
}); });
// 先删除该 topic 下已有的旧压缩摘要system_context LIKE 'history_compaction%')。 // 先删除该 topic 下已有的旧压缩摘要system_context LIKE 'history_compaction%')。

View File

@ -186,7 +186,9 @@ pub struct MemoryUpsert {
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
#[derive(Default)]
pub enum SchedulerJobState { pub enum SchedulerJobState {
#[default]
Scheduled, Scheduled,
Running, Running,
Paused, Paused,
@ -241,12 +243,6 @@ impl SchedulerJobStatus {
} }
} }
impl Default for SchedulerJobState {
fn default() -> Self {
Self::Scheduled
}
}
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SchedulerJobRecord { pub struct SchedulerJobRecord {
pub id: String, pub id: String,

View File

@ -92,7 +92,7 @@ impl ShellKind {
let info = self.to_info(); let info = self.to_info();
info.args info.args
.iter() .iter()
.map(|s| *s) .copied()
.chain(std::iter::once(command)) .chain(std::iter::once(command))
.collect() .collect()
} }
@ -385,7 +385,7 @@ impl Tool for BashTool {
let cwd = self let cwd = self
.working_dir .working_dir
.as_ref() .as_ref()
.map(|d| Path::new(d)) .map(Path::new)
.unwrap_or_else(|| Path::new(".")); .unwrap_or_else(|| Path::new("."));
match self match self
@ -629,7 +629,7 @@ fn format_command_output(stdout: &str, stderr: &str, exit_code: Option<i32>) ->
if !stderr.trim().is_empty() { if !stderr.trim().is_empty() {
if !output.is_empty() { if !output.is_empty() {
output.push_str("\n"); output.push('\n');
} }
output.push_str("STDERR:\n"); output.push_str("STDERR:\n");
output.push_str(stderr); output.push_str(stderr);

View File

@ -432,7 +432,9 @@ fn calc_evaluate(args: &serde_json::Value) -> Result<String, String> {
// 表达式可产生非有限结果(如 "1/0" → inf、"0/0" → NaN // 表达式可产生非有限结果(如 "1/0" → inf、"0/0" → NaN
// 与 extract_values/extract_f64 的边界策略保持一致:拒绝输出。 // 与 extract_values/extract_f64 的边界策略保持一致:拒绝输出。
if !n.is_finite() { if !n.is_finite() {
return Err(format!("Expression result is not a finite number: {expression}")); return Err(format!(
"Expression result is not a finite number: {expression}"
));
} }
Ok(format_num(n)) Ok(format_num(n))
}) })
@ -873,10 +875,7 @@ mod tests {
.await .await
.unwrap(); .unwrap();
assert!(ok.success); assert!(ok.success);
assert_eq!( assert_eq!(ok.output, "295232799039604140847618609643520000000");
ok.output,
"295232799039604140847618609643520000000"
);
} }
#[tokio::test] #[tokio::test]

View File

@ -140,16 +140,15 @@ impl Tool for FileWriteTool {
}; };
// Create parent directories if needed // Create parent directories if needed
if let Some(parent) = resolved.parent() { if let Some(parent) = resolved.parent()
if !parent.exists() { && !parent.exists()
if let Err(e) = std::fs::create_dir_all(parent) { && let Err(e) = std::fs::create_dir_all(parent)
return Ok(ToolResult { {
success: false, return Ok(ToolResult {
output: String::new(), success: false,
error: Some(format!("Failed to create parent directory: {}", e)), output: String::new(),
}); error: Some(format!("Failed to create parent directory: {}", e)),
} });
}
} }
match std::fs::write(&resolved, content) { match std::fs::write(&resolved, content) {

View File

@ -76,12 +76,11 @@ impl HttpRequestTool {
if let Some(obj) = headers.as_object() { if let Some(obj) = headers.as_object() {
for (key, value) in obj { for (key, value) in obj {
if let Some(str_val) = value.as_str() { if let Some(str_val) = value.as_str()
if let Ok(name) = reqwest::header::HeaderName::from_bytes(key.as_bytes()) { && let Ok(name) = reqwest::header::HeaderName::from_bytes(key.as_bytes())
if let Ok(val) = reqwest::header::HeaderValue::from_str(str_val) { && let Ok(val) = reqwest::header::HeaderValue::from_str(str_val)
header_map.insert(name, val); {
} header_map.insert(name, val);
}
} }
} }
} }

View File

@ -55,11 +55,8 @@ pub fn extract_string(args: &serde_json::Value, key: &str) -> Option<String> {
args.get(key).and_then(|v| { args.get(key).and_then(|v| {
if let Some(s) = v.as_str() { if let Some(s) = v.as_str() {
Some(s.to_string()) Some(s.to_string())
} else if let Some(n) = v.as_number() {
// Handle case where LLM sends a number but we need a string
Some(n.to_string())
} else { } else {
None v.as_number().map(|n| n.to_string())
} }
}) })
} }

View File

@ -1,6 +1,6 @@
use parking_lot::RwLock;
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use parking_lot::RwLock;
use crate::domain::tools::{Tool, ToolFunction}; use crate::domain::tools::{Tool, ToolFunction};
@ -24,20 +24,13 @@ impl ToolRegistry {
} }
pub fn get(&self, name: &str) -> Option<Arc<dyn ToolTrait>> { pub fn get(&self, name: &str) -> Option<Arc<dyn ToolTrait>> {
self.tools self.tools.read().get(name).cloned()
.read()
.get(name)
.cloned()
} }
/// Get all registered tools. /// Get all registered tools.
/// Used for concurrent tool execution when we need to look up tools by name. /// Used for concurrent tool execution when we need to look up tools by name.
pub fn get_all(&self) -> Vec<Arc<dyn ToolTrait>> { pub fn get_all(&self) -> Vec<Arc<dyn ToolTrait>> {
self.tools self.tools.read().values().cloned().collect()
.read()
.values()
.cloned()
.collect()
} }
pub fn get_definitions(&self) -> Vec<Tool> { pub fn get_definitions(&self) -> Vec<Tool> {
@ -56,18 +49,11 @@ impl ToolRegistry {
} }
pub fn has_tools(&self) -> bool { pub fn has_tools(&self) -> bool {
!self !self.tools.read().is_empty()
.tools
.read()
.is_empty()
} }
pub fn tool_names(&self) -> Vec<String> { pub fn tool_names(&self) -> Vec<String> {
self.tools self.tools.read().keys().cloned().collect()
.read()
.keys()
.cloned()
.collect()
} }
/// 创建一个排除指定工具的新 registry 副本 /// 创建一个排除指定工具的新 registry 副本
@ -80,9 +66,7 @@ impl ToolRegistry {
.map(|(k, v)| (k.clone(), v.clone())) .map(|(k, v)| (k.clone(), v.clone()))
.collect(); .collect();
let new_registry = ToolRegistry::new(); let new_registry = ToolRegistry::new();
*new_registry *new_registry.tools.write() = filtered;
.tools
.write() = filtered;
new_registry new_registry
} }
@ -97,9 +81,7 @@ impl ToolRegistry {
.map(|(k, v)| (k.clone(), v.clone())) .map(|(k, v)| (k.clone(), v.clone()))
.collect(); .collect();
let new_registry = ToolRegistry::new(); let new_registry = ToolRegistry::new();
*new_registry *new_registry.tools.write() = filtered;
.tools
.write() = filtered;
new_registry new_registry
} }
} }

View File

@ -338,30 +338,28 @@ fn enrich_target_from_context(
_ => return target, _ => return target,
}; };
if !has_non_empty_string(&object, "channel") { if !has_non_empty_string(&object, "channel")
if let Some(channel_name) = context && let Some(channel_name) = context
.channel_name .channel_name
.as_ref() .as_ref()
.filter(|value| !value.trim().is_empty()) .filter(|value| !value.trim().is_empty())
{ {
object.insert( object.insert(
"channel".to_string(), "channel".to_string(),
serde_json::Value::String(channel_name.clone()), serde_json::Value::String(channel_name.clone()),
); );
}
} }
if !has_non_empty_string(&object, "chat_id") { if !has_non_empty_string(&object, "chat_id")
if let Some(chat_id) = context && let Some(chat_id) = context
.chat_id .chat_id
.as_ref() .as_ref()
.filter(|value| !value.trim().is_empty()) .filter(|value| !value.trim().is_empty())
{ {
object.insert( object.insert(
"chat_id".to_string(), "chat_id".to_string(),
serde_json::Value::String(chat_id.clone()), serde_json::Value::String(chat_id.clone()),
); );
}
} }
serde_json::Value::Object(object) serde_json::Value::Object(object)

View File

@ -114,10 +114,11 @@ impl SchemaCleanr {
anyhow::bail!("Schema missing required 'type' field"); anyhow::bail!("Schema missing required 'type' field");
} }
if let Some(Value::String(t)) = obj.get("type") { if let Some(Value::String(t)) = obj.get("type")
if t == "object" && !obj.contains_key("properties") { && t == "object"
tracing::warn!("Object schema without 'properties' field may cause issues"); && !obj.contains_key("properties")
} {
tracing::warn!("Object schema without 'properties' field may cause issues");
} }
Ok(()) Ok(())
@ -173,10 +174,10 @@ impl SchemaCleanr {
} }
// Handle anyOf/oneOf simplification // Handle anyOf/oneOf simplification
if obj.contains_key("anyOf") || obj.contains_key("oneOf") { if (obj.contains_key("anyOf") || obj.contains_key("oneOf"))
if let Some(simplified) = Self::try_simplify_union(&obj, defs, strategy, ref_stack) { && let Some(simplified) = Self::try_simplify_union(&obj, defs, strategy, ref_stack)
return simplified; {
} return simplified;
} }
// Build cleaned object // Build cleaned object
@ -244,13 +245,13 @@ impl SchemaCleanr {
return Self::preserve_meta(obj, Value::Object(Map::new())); return Self::preserve_meta(obj, Value::Object(Map::new()));
} }
if let Some(def_name) = Self::parse_local_ref(ref_value) { if let Some(def_name) = Self::parse_local_ref(ref_value)
if let Some(definition) = defs.get(def_name.as_str()) { && let Some(definition) = defs.get(def_name.as_str())
ref_stack.insert(ref_value.to_string()); {
let cleaned = Self::clean_with_defs(definition.clone(), defs, strategy, ref_stack); ref_stack.insert(ref_value.to_string());
ref_stack.remove(ref_value); let cleaned = Self::clean_with_defs(definition.clone(), defs, strategy, ref_stack);
return Self::preserve_meta(obj, cleaned); ref_stack.remove(ref_value);
} return Self::preserve_meta(obj, cleaned);
} }
tracing::warn!("Cannot resolve $ref: {}", ref_value); tracing::warn!("Cannot resolve $ref: {}", ref_value);
@ -342,15 +343,16 @@ impl SchemaCleanr {
if let Some(Value::Null) = obj.get("const") { if let Some(Value::Null) = obj.get("const") {
return true; return true;
} }
if let Some(Value::Array(arr)) = obj.get("enum") { if let Some(Value::Array(arr)) = obj.get("enum")
if arr.len() == 1 && matches!(arr[0], Value::Null) { && arr.len() == 1
return true; && matches!(arr[0], Value::Null)
} {
return true;
} }
if let Some(Value::String(t)) = obj.get("type") { if let Some(Value::String(t)) = obj.get("type")
if t == "null" { && t == "null"
return true; {
} return true;
} }
} }
false false

View File

@ -211,10 +211,8 @@ impl Tool for SkillManageTool {
Err(err) => return Ok(error_result(&err)), Err(err) => return Ok(error_result(&err)),
} }
} }
if reload { if reload && let Err(err) = self.skills.reload() {
if let Err(err) = self.skills.reload() { return Ok(error_result(&err));
return Ok(error_result(&err));
}
} }
json!({ json!({

View File

@ -1,8 +1,8 @@
use parking_lot::RwLock;
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
use std::fs; use std::fs;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::Arc; use std::sync::Arc;
use parking_lot::RwLock;
use std::time::Duration; use std::time::Duration;
use async_trait::async_trait; use async_trait::async_trait;
@ -535,7 +535,11 @@ impl DefaultSubAgentRuntime {
let inherited = session let inherited = session
.parent_topic_id .parent_topic_id
.as_deref() .as_deref()
.and_then(|tid| self.topic_model_selections.as_ref().and_then(|s| s.get(tid))) .and_then(|tid| {
self.topic_model_selections
.as_ref()
.and_then(|s| s.get(tid))
})
.or_else(|| { .or_else(|| {
self.model_selections self.model_selections
.as_ref() .as_ref()
@ -877,10 +881,10 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
// 2. 校验父智能体的子代理策略(白/黑名单),再查找子代理定义。 // 2. 校验父智能体的子代理策略(白/黑名单),再查找子代理定义。
// 与 find_subagent_def 的"def 不可用即拒绝"安全范式一致:策略不通过即拒绝, // 与 find_subagent_def 的"def 不可用即拒绝"安全范式一致:策略不通过即拒绝,
// 防止 LLM 通过选择被禁子代理绕过限制。 // 防止 LLM 通过选择被禁子代理绕过限制。
if let Some(cap) = &parent_context.parent_capability { if let Some(cap) = &parent_context.parent_capability
if let Err(msg) = cap.check_subagent_allowed(&task.subagent_type.name) { && let Err(msg) = cap.check_subagent_allowed(&task.subagent_type.name)
return Err(TaskError::InvalidArguments(msg)); {
} return Err(TaskError::InvalidArguments(msg));
} }
// 3. 查找子代理定义 // 3. 查找子代理定义
@ -1097,8 +1101,7 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
.unwrap_or_default(), .unwrap_or_default(),
}; };
let _ = sub_done_sender.send(result).await; let _ = sub_done_sender.send(result).await;
let _ = let _ = store.update_pending_subagent_status(&task_id_for_spawn, "failed");
store.update_pending_subagent_status(&task_id_for_spawn, "failed");
// _registry_guard drop 时清理 registry 条目 // _registry_guard drop 时清理 registry 条目
return; return;
} }
@ -1140,11 +1143,7 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
String::new(), String::new(),
"cancelled".to_string(), "cancelled".to_string(),
), ),
Err(e) => ( Err(e) => (SubagentStatus::Failed, String::new(), e.to_string()),
SubagentStatus::Failed,
String::new(),
e.to_string(),
),
}; };
// 查询同 topic 下仍未完成的子代理列表 // 查询同 topic 下仍未完成的子代理列表
@ -1181,7 +1180,8 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
SubagentStatus::Timeout => "timeout", SubagentStatus::Timeout => "timeout",
SubagentStatus::Cancelled => "cancelled", SubagentStatus::Cancelled => "cancelled",
}; };
if let Err(e) = store.update_pending_subagent_status(&task_id_for_spawn, status_str) { if let Err(e) = store.update_pending_subagent_status(&task_id_for_spawn, status_str)
{
tracing::warn!( tracing::warn!(
error = %e, error = %e,
task_id = %task_id_for_spawn, task_id = %task_id_for_spawn,
@ -1211,7 +1211,13 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
if let Err(e) = task_repository.save_task_session(&session_done).await { if let Err(e) = task_repository.save_task_session(&session_done).await {
tracing::warn!(error = %e, task_id = %task_id_for_spawn, "Failed to save failed session"); tracing::warn!(error = %e, task_id = %task_id_for_spawn, "Failed to save failed session");
} }
publish_subagent_error(&bus, &session_done, &e.to_string(), &trace_id_owned).await; publish_subagent_error(
&bus,
&session_done,
&e.to_string(),
&trace_id_owned,
)
.await;
} }
} }
// _registry_guard 在此 drop确定性清理 cancel_registry 条目 // _registry_guard 在此 drop确定性清理 cancel_registry 条目
@ -1236,7 +1242,9 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
// ===== 同步路径(子代理嵌套或无 sub_done_sender ===== // ===== 同步路径(子代理嵌套或无 sub_done_sender =====
// 9. 执行任务并处理结果 // 9. 执行任务并处理结果
let result = self.execute_task(agent, &session, &def, task.prompt.clone()).await; let result = self
.execute_task(agent, &session, &def, task.prompt.clone())
.await;
match result { match result {
Ok(tool_result) => { Ok(tool_result) => {
@ -1303,10 +1311,10 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
// 4.1 校验父智能体的子代理策略(白/黑名单)。 // 4.1 校验父智能体的子代理策略(白/黑名单)。
// 安全要求:与 spawn 一致,防止 resume 绕过白名单。若用户切换到不允许 // 安全要求:与 spawn 一致,防止 resume 绕过白名单。若用户切换到不允许
// 该子代理的专家resume 应失败(与 def 被删除即失败的安全语义一致)。 // 该子代理的专家resume 应失败(与 def 被删除即失败的安全语义一致)。
if let Some(cap) = &parent_context.parent_capability { if let Some(cap) = &parent_context.parent_capability
if let Err(msg) = cap.check_subagent_allowed(&session.subagent_type) { && let Err(msg) = cap.check_subagent_allowed(&session.subagent_type)
return Err(TaskError::InvalidArguments(msg)); {
} return Err(TaskError::InvalidArguments(msg));
} }
// 4.2 重新解析 def 以应用工具过滤。 // 4.2 重新解析 def 以应用工具过滤。
@ -1406,10 +1414,11 @@ impl SubAgentRuntime for DefaultSubAgentRuntime {
// token 不在 registry 中(可能已完成但 DB 状态未更新,或进程重启后丢失) // token 不在 registry 中(可能已完成但 DB 状态未更新,或进程重启后丢失)
// 不变量 1条件 UPDATE仅在 status='running' 时转为 cancelled // 不变量 1条件 UPDATE仅在 status='running' 时转为 cancelled
// 避免 spawn 已完成的终态被覆盖completed → cancelled 是非法转换) // 避免 spawn 已完成的终态被覆盖completed → cancelled 是非法转换)
match self match self.store.try_update_pending_subagent_status(
.store &record.task_id,
.try_update_pending_subagent_status(&record.task_id, "running", "cancelled") "running",
{ "cancelled",
) {
Ok(true) => { Ok(true) => {
tracing::info!( tracing::info!(
task_id = %record.task_id, task_id = %record.task_id,
@ -1753,24 +1762,15 @@ impl SubagentRuntime {
/// (生产环境两者相同,但测试场景使用临时目录时必须用 `self.cwd`)。 /// (生产环境两者相同,但测试场景使用临时目录时必须用 `self.cwd`)。
pub fn reload(&self) -> Result<(), String> { pub fn reload(&self) -> Result<(), String> {
let new_catalog = SubagentCatalog::discover_with_cwd(&self.config, &self.cwd); let new_catalog = SubagentCatalog::discover_with_cwd(&self.config, &self.cwd);
let mut guard = self let mut guard = self.catalog.write();
.catalog
.write()
;
*guard = new_catalog; *guard = new_catalog;
Ok(()) Ok(())
} }
/// 列出所有子代理(含禁用项),带 disabled_in_scopes /// 列出所有子代理(含禁用项),带 disabled_in_scopes
pub fn list_with_status(&self) -> Vec<SubagentWithStatus> { pub fn list_with_status(&self) -> Vec<SubagentWithStatus> {
let state = self let state = self.disable_state.read();
.disable_state let catalog = self.catalog.read();
.read()
;
let catalog = self
.catalog
.read()
;
let mut items: Vec<SubagentWithStatus> = catalog let mut items: Vec<SubagentWithStatus> = catalog
.all() .all()
.iter() .iter()
@ -1794,14 +1794,8 @@ impl SubagentRuntime {
/// 可用子代理名称(过滤禁用项) /// 可用子代理名称(过滤禁用项)
pub fn available_names(&self) -> Vec<String> { pub fn available_names(&self) -> Vec<String> {
let state = self let state = self.disable_state.read();
.disable_state let catalog = self.catalog.read();
.read()
;
let catalog = self
.catalog
.read()
;
catalog catalog
.names() .names()
.into_iter() .into_iter()
@ -1811,30 +1805,17 @@ impl SubagentRuntime {
/// 查找可用子代理(过滤禁用项) /// 查找可用子代理(过滤禁用项)
pub fn find_available(&self, name: &str) -> Option<SubagentDef> { pub fn find_available(&self, name: &str) -> Option<SubagentDef> {
let state = self let state = self.disable_state.read();
.disable_state
.read()
;
if state.is_disabled(name) { if state.is_disabled(name) {
return None; return None;
} }
self.catalog self.catalog.read().find(name).cloned()
.read()
.find(name)
.cloned()
} }
/// 生成过滤后的系统索引提示词 /// 生成过滤后的系统索引提示词
pub fn system_index_prompt_filtered(&self) -> Option<String> { pub fn system_index_prompt_filtered(&self) -> Option<String> {
let state = self let state = self.disable_state.read();
.disable_state let catalog = self.catalog.read();
.read()
;
let catalog = self
.catalog
.read()
;
let available_defs: Vec<&SubagentDef> = catalog let available_defs: Vec<&SubagentDef> = catalog
.all() .all()
.into_iter() .into_iter()
@ -1872,14 +1853,8 @@ impl SubagentRuntime {
allowed: Option<&[String]>, allowed: Option<&[String]>,
denied: &[String], denied: &[String],
) -> Option<String> { ) -> Option<String> {
let state = self let state = self.disable_state.read();
.disable_state let catalog = self.catalog.read();
.read()
;
let catalog = self
.catalog
.read()
;
let available_defs: Vec<&SubagentDef> = catalog let available_defs: Vec<&SubagentDef> = catalog
.all() .all()
.into_iter() .into_iter()
@ -1942,13 +1917,7 @@ impl SubagentRuntime {
enabled: bool, enabled: bool,
) -> Result<SubagentAvailabilityChange, String> { ) -> Result<SubagentAvailabilityChange, String> {
// 校验子代理存在 // 校验子代理存在
if self if self.catalog.read().find(name).is_none() {
.catalog
.read()
.find(name)
.is_none()
{
return Err(format!("subagent '{}' not found", name)); return Err(format!("subagent '{}' not found", name));
} }
@ -1969,10 +1938,7 @@ impl SubagentRuntime {
// 更新内存中的 disable_state // 更新内存中的 disable_state
{ {
let mut state = self let mut state = self.disable_state.write();
.disable_state
.write()
;
match scope { match scope {
SubagentScope::User => { SubagentScope::User => {
if enabled { if enabled {
@ -1992,10 +1958,7 @@ impl SubagentRuntime {
} }
// 计算新的 disabled_in_scopes // 计算新的 disabled_in_scopes
let state = self let state = self.disable_state.read();
.disable_state
.read()
;
let disabled_in_scopes = state.disabled_scopes_for(name); let disabled_in_scopes = state.disabled_scopes_for(name);
Ok(SubagentAvailabilityChange { Ok(SubagentAvailabilityChange {
@ -2023,10 +1986,7 @@ impl SubagentRuntime {
reload: bool, reload: bool,
) -> Result<SubagentDef, String> { ) -> Result<SubagentDef, String> {
let def = { let def = {
let catalog = self let catalog = self.catalog.read();
.catalog
.read()
;
catalog catalog
.find(name) .find(name)
.ok_or_else(|| format!("subagent '{}' not found", name))? .ok_or_else(|| format!("subagent '{}' not found", name))?
@ -2089,10 +2049,7 @@ impl SubagentRuntime {
) -> Result<SubagentDef, String> { ) -> Result<SubagentDef, String> {
validate_subagent_name(name)?; validate_subagent_name(name)?;
{ {
let catalog = self let catalog = self.catalog.read();
.catalog
.read()
;
if catalog.find(name).is_some() { if catalog.find(name).is_some() {
return Err(format!("subagent '{}' already exists", name)); return Err(format!("subagent '{}' already exists", name));
} }
@ -2136,17 +2093,10 @@ impl SubagentRuntime {
/// 对齐 `ExpertRuntime::delete_expert`。 /// 对齐 `ExpertRuntime::delete_expert`。
/// - builtin 子代理path 为 None禁止删除。 /// - builtin 子代理path 为 None禁止删除。
/// - 仅当目录内除 SUBAGENT.md 外无其他文件时才删除目录,避免误删用户附件。 /// - 仅当目录内除 SUBAGENT.md 外无其他文件时才删除目录,避免误删用户附件。
pub fn delete_subagent( pub fn delete_subagent(&self, name: &str, reload: bool) -> Result<PathBuf, String> {
&self,
name: &str,
reload: bool,
) -> Result<PathBuf, String> {
validate_subagent_name(name)?; validate_subagent_name(name)?;
let path = { let path = {
let catalog = self let catalog = self.catalog.read();
.catalog
.read()
;
let def = catalog let def = catalog
.find(name) .find(name)
.ok_or_else(|| format!("subagent '{}' not found", name))?; .ok_or_else(|| format!("subagent '{}' not found", name))?;
@ -2201,11 +2151,7 @@ fn validate_subagent_name(name: &str) -> Result<(), String> {
/// 获取指定 scope 下某子代理的 SUBAGENT.md 路径。 /// 获取指定 scope 下某子代理的 SUBAGENT.md 路径。
/// 对齐 `expert_file_path`。 /// 对齐 `expert_file_path`。
fn subagent_file_path( fn subagent_file_path(scope: SubagentScope, name: &str, cwd: &Path) -> Result<PathBuf, String> {
scope: SubagentScope,
name: &str,
cwd: &Path,
) -> Result<PathBuf, String> {
let root = match scope { let root = match scope {
SubagentScope::User => dirs::home_dir() SubagentScope::User => dirs::home_dir()
.map(|p| p.join(".picobot").join("subagents")) .map(|p| p.join(".picobot").join("subagents"))
@ -2632,7 +2578,7 @@ mod tests {
// 禁用后 prompt 不应包含 general无可用子代理时返回 None // 禁用后 prompt 不应包含 general无可用子代理时返回 None
let prompt = runtime.system_index_prompt_filtered(); let prompt = runtime.system_index_prompt_filtered();
assert!(prompt.map_or(true, |p| !p.contains("<name>general</name>"))); assert!(prompt.is_none_or(|p| !p.contains("<name>general</name>")));
} }
#[test] #[test]
@ -3142,10 +3088,7 @@ mod tests {
let item = items.iter().find(|i| i.name == "demo-create").unwrap(); let item = items.iter().find(|i| i.name == "demo-create").unwrap();
assert_eq!(item.description, "demo create agent"); assert_eq!(item.description, "demo create agent");
assert_eq!(item.body.as_deref(), Some("demo body content")); assert_eq!(item.body.as_deref(), Some("demo body content"));
assert_eq!( assert_eq!(item.capability.denied_skills, vec!["skill_x".to_string()]);
item.capability.denied_skills,
vec!["skill_x".to_string()]
);
} }
#[test] #[test]
@ -3288,7 +3231,8 @@ mod tests {
"directory should be preserved when it has other files" "directory should be preserved when it has other files"
); );
assert!( assert!(
!temp.path() !temp
.path()
.join(".picobot") .join(".picobot")
.join("subagents") .join("subagents")
.join("mixed") .join("mixed")

View File

@ -103,7 +103,7 @@ impl Tool for TaskTool {
// 2. 验证描述长度 // 2. 验证描述长度
let word_count = task_args.description.split_whitespace().count(); let word_count = task_args.description.split_whitespace().count();
if task_args.description.len() > 50 || word_count > 7 || word_count < 1 { if task_args.description.len() > 50 || !(1..=7).contains(&word_count) {
return Ok(ToolResult { return Ok(ToolResult {
success: false, success: false,
output: String::new(), output: String::new(),
@ -136,17 +136,17 @@ impl Tool for TaskTool {
// 4. 深度校验仅对嵌套场景生效None = 不限制) // 4. 深度校验仅对嵌套场景生效None = 不限制)
// Some(N) 表示允许最多 N 层嵌套depth=1 的 agent 可创建 depth=2但 depth=2 不能再创建 // Some(N) 表示允许最多 N 层嵌套depth=1 的 agent 可创建 depth=2但 depth=2 不能再创建
if let Some(max_depth) = self.max_nesting_depth { if let Some(max_depth) = self.max_nesting_depth
if context.nesting_depth > max_depth { && context.nesting_depth > max_depth
return Ok(ToolResult { {
success: false, return Ok(ToolResult {
output: String::new(), success: false,
error: Some(format!( output: String::new(),
"Cannot create nested subagent: max nesting depth ({}) reached", error: Some(format!(
max_depth "Cannot create nested subagent: max nesting depth ({}) reached",
)), max_depth
}); )),
} });
} }
// 5. 执行任务 // 5. 执行任务

View File

@ -8,8 +8,10 @@ use crate::utils::current_timestamp;
/// 子代理会话状态 /// 子代理会话状态
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")] #[serde(rename_all = "lowercase")]
#[derive(Default)]
pub enum TaskSessionState { pub enum TaskSessionState {
/// 正在执行 /// 正在执行
#[default]
Running, Running,
/// 已完成 /// 已完成
Completed, Completed,
@ -23,12 +25,6 @@ pub enum TaskSessionState {
Unknown, Unknown,
} }
impl Default for TaskSessionState {
fn default() -> Self {
Self::Running
}
}
impl TaskSessionState { impl TaskSessionState {
pub fn as_str(&self) -> &'static str { pub fn as_str(&self) -> &'static str {
match self { match self {

View File

@ -88,10 +88,10 @@ impl Tool for TodoReadTool {
// 2. 读锁查内存 // 2. 读锁查内存
{ {
let guard = self.state.read().await; let guard = self.state.read().await;
if let Some(items) = guard.get(&scope_key) { if let Some(items) = guard.get(&scope_key)
if !items.is_empty() { && !items.is_empty()
return Ok(success_result(items, &scope_key, "memory")); {
} return Ok(success_result(items, &scope_key, "memory"));
} }
} }

View File

@ -1,5 +1,5 @@
use std::time::Duration;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait; use async_trait::async_trait;
use tokio::sync::{mpsc, watch}; use tokio::sync::{mpsc, watch};
@ -58,11 +58,7 @@ pub trait WaitCoordinator: Send + Sync + 'static {
/// select! 立即返回 `WaitEvent::Cancelled`,并完成完整的状态清理 /// select! 立即返回 `WaitEvent::Cancelled`,并完成完整的状态清理
///(重获取锁、回填 guard、清除 is_waiting、归还 receiver ///(重获取锁、回填 guard、清除 is_waiting、归还 receiver
/// 为 None 时退化为不检查取消(向后兼容,子代理场景)。 /// 为 None 时退化为不检查取消(向后兼容,子代理场景)。
async fn wait( async fn wait(&self, timeout: Duration, cancel_rx: Option<watch::Receiver<()>>) -> WaitEvent;
&self,
timeout: Duration,
cancel_rx: Option<watch::Receiver<()>>,
) -> WaitEvent;
} }
#[derive(Clone, Default)] #[derive(Clone, Default)]

View File

@ -156,9 +156,7 @@ impl Tool for WaitForSubagentsTool {
// 传入 cancel_rx 使 /stop 命令能立即中断等待。 // 传入 cancel_rx 使 /stop 命令能立即中断等待。
// coordinator 在 select! 中以 biased 优先级处理: // coordinator 在 select! 中以 biased 优先级处理:
// 子代理结果 > 用户消息 > 取消信号 > 超时 // 子代理结果 > 用户消息 > 取消信号 > 超时
let event = coordinator let event = coordinator.wait(timeout, context.cancel_rx.clone()).await;
.wait(timeout, context.cancel_rx.clone())
.await;
// 5. 格式化返回结果 // 5. 格式化返回结果
let output = match event { let output = match event {

View File

@ -18,7 +18,7 @@ fn test_message_special_characters() {
/// Test that multi-line system prompt is preserved /// Test that multi-line system prompt is preserved
#[test] #[test]
fn test_multiline_system_prompt() { fn test_multiline_system_prompt() {
let messages = vec![ let messages = [
Message::system( Message::system(
"You are a helpful assistant.\n\nFollow these rules:\n1. Be kind\n2. Be accurate", "You are a helpful assistant.\n\nFollow these rules:\n1. Be kind\n2. Be accurate",
), ),