fix: 修复后台线程调用 UI 导致的 macOS 崩溃问题

修复了在 tokio 后台线程中调用 macOS UI 操作导致的崩溃:
- 使用 run_on_main_thread() 将更新窗口创建操作调度到主线程
- 解决 "Must only be used from the main thread" 断言失败
- 修复测试代码中过时的 iflow 字段引用
- 更新版本号到 0.48.0
This commit is contained in:
coso
2026-01-22 12:37:23 +08:00
parent 2989ba2582
commit 2e41683134
71 changed files with 2542 additions and 2308 deletions
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.47.4",
"version": "0.48.0",
"type": "module",
"repository": {
"type": "git",
@@ -498,11 +498,13 @@ pub fn get_default_check_model(provider_type: PoolProviderType) -> &'static str
PoolProviderType::OpenAI => "gpt-3.5-turbo",
// 使用 claude-sonnet-4-5-20250929,兼容更多代理服务器
PoolProviderType::Claude => "claude-sonnet-4-5-20250929",
PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929",
// Anthropic 兼容格式使用相同的健康检查模型
PoolProviderType::AnthropicCompatible => "claude-sonnet-4-5-20250929",
PoolProviderType::Antigravity => "gemini-3-pro-preview",
PoolProviderType::Vertex => "gemini-2.0-flash",
PoolProviderType::GeminiApiKey => "gemini-2.5-flash",
PoolProviderType::Codex => "gpt-4o-mini",
PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929",
// API Key Provider 类型
PoolProviderType::Anthropic => "claude-sonnet-4-5-20250929",
PoolProviderType::AzureOpenai => "gpt-4o-mini",
@@ -13,13 +13,16 @@ pub enum ProviderType {
#[serde(rename = "openai")]
OpenAI,
Claude,
#[serde(rename = "claude_oauth")]
ClaudeOAuth,
/// Anthropic 兼容格式(支持 system 数组格式等变体)
#[serde(rename = "anthropic_compatible")]
AnthropicCompatible,
Antigravity,
Vertex,
#[serde(rename = "gemini_api_key")]
GeminiApiKey,
Codex,
#[serde(rename = "claude_oauth")]
ClaudeOAuth,
// API Key Provider 类型
Anthropic,
#[serde(rename = "azure_openai")]
@@ -36,11 +39,12 @@ impl std::fmt::Display for ProviderType {
ProviderType::Gemini => write!(f, "gemini"),
ProviderType::OpenAI => write!(f, "openai"),
ProviderType::Claude => write!(f, "claude"),
ProviderType::ClaudeOAuth => write!(f, "claude_oauth"),
ProviderType::AnthropicCompatible => write!(f, "anthropic_compatible"),
ProviderType::Antigravity => write!(f, "antigravity"),
ProviderType::Vertex => write!(f, "vertex"),
ProviderType::GeminiApiKey => write!(f, "gemini_api_key"),
ProviderType::Codex => write!(f, "codex"),
ProviderType::ClaudeOAuth => write!(f, "claude_oauth"),
ProviderType::Anthropic => write!(f, "anthropic"),
ProviderType::AzureOpenai => write!(f, "azure_openai"),
ProviderType::AwsBedrock => write!(f, "aws_bedrock"),
@@ -58,11 +62,14 @@ impl std::str::FromStr for ProviderType {
"gemini" => Ok(ProviderType::Gemini),
"openai" => Ok(ProviderType::OpenAI),
"claude" => Ok(ProviderType::Claude),
"claude_oauth" => Ok(ProviderType::ClaudeOAuth),
"anthropic_compatible" | "anthropic-compatible" => {
Ok(ProviderType::AnthropicCompatible)
}
"antigravity" => Ok(ProviderType::Antigravity),
"vertex" => Ok(ProviderType::Vertex),
"gemini_api_key" => Ok(ProviderType::GeminiApiKey),
"codex" => Ok(ProviderType::Codex),
"claude_oauth" => Ok(ProviderType::ClaudeOAuth),
"anthropic" => Ok(ProviderType::Anthropic),
"azure_openai" | "azure-openai" => Ok(ProviderType::AzureOpenai),
"aws_bedrock" | "aws-bedrock" => Ok(ProviderType::AwsBedrock),
@@ -5,34 +5,6 @@
"name": "Anthropic"
},
"models": [
{
"id": "claude-opus-4-5",
"name": "Claude Opus 4.5 (latest)",
"family": "claude-opus",
"tier": "max",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 15,
"output": 75,
"cache_read": 1.5,
"cache_write": 18.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 32000
},
"status": "active",
"release_date": "2025-02-24",
"is_latest": true
},
{
"id": "claude-opus-4-5-20251101",
"name": "Claude Opus 4.5",
@@ -61,34 +33,6 @@
"release_date": "2025-02-24",
"is_latest": false
},
{
"id": "claude-sonnet-4-5",
"name": "Claude Sonnet 4.5 (latest)",
"family": "claude-sonnet",
"tier": "pro",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 3,
"output": 15,
"cache_read": 0.3,
"cache_write": 3.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 64000
},
"status": "active",
"release_date": "2025-09-29",
"is_latest": true
},
{
"id": "claude-sonnet-4-5-20250929",
"name": "Claude Sonnet 4.5",
@@ -172,174 +116,6 @@
"status": "active",
"release_date": "2025-10-01",
"is_latest": false
},
{
"id": "claude-opus-4-1",
"name": "Claude Opus 4.1 (latest)",
"family": "claude-opus",
"tier": "max",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 15,
"output": 75,
"cache_read": 1.5,
"cache_write": 18.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 32000
},
"status": "active",
"release_date": "2025-08-05",
"is_latest": false
},
{
"id": "claude-opus-4-1-20250805",
"name": "Claude Opus 4.1",
"family": "claude-opus",
"tier": "max",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 15,
"output": 75,
"cache_read": 1.5,
"cache_write": 18.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 32000
},
"status": "active",
"release_date": "2025-08-05",
"is_latest": false
},
{
"id": "claude-opus-4-0",
"name": "Claude Opus 4 (latest)",
"family": "claude-opus",
"tier": "max",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 15,
"output": 75,
"cache_read": 1.5,
"cache_write": 18.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 32000
},
"status": "active",
"release_date": "2025-05-14",
"is_latest": false
},
{
"id": "claude-opus-4-20250514",
"name": "Claude Opus 4",
"family": "claude-opus",
"tier": "max",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 15,
"output": 75,
"cache_read": 1.5,
"cache_write": 18.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 32000
},
"status": "active",
"release_date": "2025-05-14",
"is_latest": false
},
{
"id": "claude-sonnet-4-0",
"name": "Claude Sonnet 4 (latest)",
"family": "claude-sonnet",
"tier": "pro",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 3,
"output": 15,
"cache_read": 0.3,
"cache_write": 3.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 64000
},
"status": "active",
"release_date": "2025-05-14",
"is_latest": false
},
{
"id": "claude-sonnet-4-20250514",
"name": "Claude Sonnet 4",
"family": "claude-sonnet",
"tier": "pro",
"capabilities": {
"vision": true,
"tools": true,
"streaming": true,
"json_mode": true,
"function_calling": true,
"reasoning": true
},
"pricing": {
"input": 3,
"output": 15,
"cache_read": 0.3,
"cache_write": 3.75,
"currency": "USD"
},
"limits": {
"context": 200000,
"max_output": 64000
},
"status": "active",
"release_date": "2025-05-14",
"is_latest": false
}
],
"updated_at": "2026-01-12T00:00:00.000Z",
+67 -29
View File
@@ -100,33 +100,9 @@ impl NativeAgent {
/// 获取 API 请求的有效 base_url
///
/// 对于自定义 Provider(如 moonshot),返回 `{base_url}/api/provider/{provider_id}`
/// 对于内置 Provider,返回原始 base_url
/// 所有 Provider 都使用标准路由,不再使用 Amp CLI 路由前缀
fn get_effective_base_url(&self) -> String {
if let Some(ref pid) = self.provider_id {
// 检查 provider_id 是否是已知的内置类型
let is_builtin = matches!(
pid.to_lowercase().as_str(),
"openai"
| "claude"
| "anthropic"
| "gemini"
| "kiro"
| "qwen"
| "codex"
| "antigravity"
| "iflow"
);
if is_builtin {
self.base_url.clone()
} else {
// 自定义 Provider,使用 provider 特定路由
// 例如:http://127.0.0.1:8999/api/provider/moonshot
format!("{}/api/provider/{}", self.base_url, pid)
}
} else {
self.base_url.clone()
}
self.base_url.clone()
}
/// 检查是否是自定义 Provider
@@ -879,8 +855,56 @@ impl NativeAgentState {
})
}
/// 创建临时 Agent 用于异步操作(支持根据模型名称动态选择协议)
fn create_temp_agent_with_model(&self, model: &str) -> Result<NativeAgent, String> {
let guard = self.agent.read();
let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?;
let client = Client::builder()
.timeout(Duration::from_secs(300))
.connect_timeout(Duration::from_secs(30))
.no_proxy()
.build()
.map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?;
// 如果有自定义 provider_id,尝试从模型名称推断协议类型
let provider_type = if let Some(provider_id) = &agent.provider_id {
ProviderType::from_provider_and_model(provider_id, model)
} else {
agent.provider_type
};
let protocol = create_protocol(provider_type);
info!(
"[NativeAgent] 创建临时 Agent: model={}, provider_type={:?}, provider_id={:?}, protocol_endpoint={}",
model,
provider_type,
agent.provider_id,
protocol.endpoint()
);
Ok(NativeAgent {
client,
base_url: agent.base_url.clone(),
api_key: agent.api_key.clone(),
sessions: agent.sessions.clone(),
config: agent.config.clone(),
provider_type,
protocol,
provider_id: agent.provider_id.clone(),
})
}
pub async fn chat(&self, request: NativeChatRequest) -> Result<NativeChatResponse, String> {
let temp_agent = self.create_temp_agent()?;
let model = request.model.clone().unwrap_or_else(|| {
self.agent
.read()
.as_ref()
.map(|a| a.config.model.clone())
.unwrap_or_default()
});
let temp_agent = self.create_temp_agent_with_model(&model)?;
temp_agent.chat(request).await
}
@@ -889,7 +913,14 @@ impl NativeAgentState {
request: NativeChatRequest,
tx: mpsc::Sender<StreamEvent>,
) -> Result<StreamResult, String> {
let temp_agent = self.create_temp_agent()?;
let model = request.model.clone().unwrap_or_else(|| {
self.agent
.read()
.as_ref()
.map(|a| a.config.model.clone())
.unwrap_or_default()
});
let temp_agent = self.create_temp_agent_with_model(&model)?;
temp_agent.chat_stream(request, None, tx).await
}
@@ -899,7 +930,14 @@ impl NativeAgentState {
tx: mpsc::Sender<StreamEvent>,
tool_loop_engine: &ToolLoopEngine,
) -> Result<StreamResult, String> {
let temp_agent = self.create_temp_agent()?;
let model = request.model.clone().unwrap_or_else(|| {
self.agent
.read()
.as_ref()
.map(|a| a.config.model.clone())
.unwrap_or_default()
});
let temp_agent = self.create_temp_agent_with_model(&model)?;
temp_agent
.chat_stream_with_tools(request, tx, tool_loop_engine)
.await
@@ -14,7 +14,6 @@ use async_trait::async_trait;
use futures::StreamExt;
use reqwest::Client;
use serde::Serialize;
use std::collections::HashMap;
use tokio::sync::mpsc;
use tracing::{debug, error, info, warn};
@@ -532,11 +531,6 @@ impl Protocol for AnthropicProtocol {
.header("Content-Type", "application/json")
.header("anthropic-version", "2023-06-01");
// 添加 X-Provider-Id header 用于精确路由
if let Some(pid) = provider_id {
req_builder = req_builder.header("X-Provider-Id", pid);
}
let response = req_builder
.json(&request)
.send()
@@ -646,11 +640,6 @@ impl Protocol for AnthropicProtocol {
.header("Content-Type", "application/json")
.header("anthropic-version", "2023-06-01");
// 添加 X-Provider-Id header 用于精确路由
if let Some(pid) = provider_id {
req_builder = req_builder.header("X-Provider-Id", pid);
}
let response = req_builder
.json(&request)
.send()
+1 -1
View File
@@ -109,7 +109,7 @@ pub struct ToolLoopConfig {
impl Default for ToolLoopConfig {
fn default() -> Self {
Self {
max_iterations: 25, // 默认最大 25 次迭代
max_iterations: 50, // 默认最大 25 次迭代
}
}
}
+35
View File
@@ -52,6 +52,41 @@ impl ProviderType {
}
}
/// 从 provider 字符串和模型名称推断 provider 类型
///
/// 对于自定义 Provider ID(如 custom-xxx),尝试从模型名称推断协议类型
pub fn from_provider_and_model(provider: &str, model: &str) -> Self {
// 首先检查是否是自定义 Provider ID(以 custom- 开头)
if provider.starts_with("custom-") {
// 自定义 Provider 使用 Anthropic 兼容协议(Anthropic Compatible)
return Self::AnthropicCompatible;
}
// 对于其他 Provider,尝试直接解析
let provider_type = Self::from_str(provider);
// 如果能被识别,直接返回
if !matches!(provider_type, Self::OpenAI) || provider.eq_ignore_ascii_case("openai") {
return provider_type;
}
// 对于标准 Provider (kiro/openai/claude/gemini 等),尝试从模型名称推断协议类型
let model_lower = model.to_lowercase();
// Claude 模型使用 Anthropic 协议
if model_lower.starts_with("claude-") || model_lower.starts_with("anthropic-") {
return Self::Claude;
}
// Gemini 模型
if model_lower.starts_with("gemini") || model_lower.contains("gemini") {
return Self::Gemini;
}
// 默认使用 OpenAI 协议
Self::OpenAI
}
/// 获取 API 端点路径
pub fn endpoint(&self) -> &'static str {
match self {
+2
View File
@@ -89,6 +89,8 @@ pub async fn check_api_compatibility(
("claude-sonnet-4-5", "basic"),
("claude-sonnet-4-5", "tool_call"),
],
// Anthropic 兼容格式 - 使用 Claude 相同的测试
ProviderType::AnthropicCompatible => vec![],
ProviderType::OpenAI | ProviderType::Claude => vec![],
// API Key Provider 类型 - 暂不支持自动测试
ProviderType::Anthropic
+59 -2
View File
@@ -26,6 +26,7 @@ pub struct NativeAgentStatus {
pub async fn native_agent_init(
agent_state: State<'_, NativeAgentState>,
app_state: State<'_, AppState>,
db: State<'_, crate::database::DbConnection>,
) -> Result<NativeAgentStatus, String> {
tracing::info!("[NativeAgent] 初始化 Agent");
@@ -48,7 +49,63 @@ pub async fn native_agent_init(
let api_key = api_key.ok_or_else(|| "ProxyCast API Server 未配置 API Key".to_string())?;
let base_url = get_local_url(&host, port);
let provider_type = ProviderType::from_str(&default_provider);
// 对于自定义 Provider ID(如 custom-xxx),使用 Anthropic 兼容协议
let provider_type = if default_provider.starts_with("custom-") {
tracing::info!(
"[NativeAgent] 自定义 Provider ID '{}',使用 Anthropic 兼容协议",
default_provider
);
ProviderType::AnthropicCompatible
} else {
// 对于标准 Provider,从数据库查询类型
match crate::database::dao::api_key_provider::ApiKeyProviderDao::get_provider_by_id(
&*db.lock().map_err(|e| e.to_string())?,
&default_provider,
) {
Ok(Some(provider)) => {
// 从数据库的 provider_type 转换为 ProviderType
match provider.provider_type {
crate::database::dao::api_key_provider::ApiProviderType::Anthropic |
crate::database::dao::api_key_provider::ApiProviderType::AnthropicCompatible => {
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {:?} (Anthropic)",
provider.provider_type
);
ProviderType::Claude
}
crate::database::dao::api_key_provider::ApiProviderType::Gemini => {
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {:?} (Gemini)",
provider.provider_type
);
ProviderType::Gemini
}
_ => {
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {:?} (OpenAI)",
provider.provider_type
);
ProviderType::OpenAI
}
}
}
Ok(None) => {
tracing::warn!(
"[NativeAgent] 数据库中未找到 Provider '{}',使用字符串解析",
default_provider
);
ProviderType::from_str(&default_provider)
}
Err(e) => {
tracing::warn!(
"[NativeAgent] 从数据库查询 Provider 失败: {},使用字符串解析",
e
);
ProviderType::from_str(&default_provider)
}
}
};
tracing::info!(
"[NativeAgent] 初始化 Agent: base_url={}, provider={:?}, use_default_prompt={}",
@@ -62,7 +119,7 @@ pub async fn native_agent_init(
base_url.clone(),
api_key,
provider_type,
Some(default_provider),
Some(default_provider.to_string()),
&agent_config,
)?;
+10 -6
View File
@@ -213,12 +213,16 @@ pub async fn start_background_update_check(
});
if should_notify {
// 打开独立的更新提醒窗口
if let Err(e) =
update_window::open_update_window(&app_handle_clone, &result)
{
tracing::error!("[更新检查] 打开更新窗口失败: {}", e);
}
// 打开独立的更新提醒窗口 - 必须在主线程执行
let app_handle_for_ui = app_handle_clone.clone();
let result_clone = result.clone();
let _ = app_handle_clone.run_on_main_thread(move || {
if let Err(e) =
update_window::open_update_window(&app_handle_for_ui, &result_clone)
{
tracing::error!("[更新检查] 打开更新窗口失败: {}", e);
}
});
}
}
}
+12 -98
View File
@@ -1228,7 +1228,6 @@ fn arb_credential_pool_config() -> impl Strategy<Value = CredentialPoolConfig> {
gemini_api_keys: vec![],
vertex_api_keys: vec![],
codex: vec![],
iflow: vec![],
},
)
}
@@ -2223,30 +2222,6 @@ fn arb_vertex_api_key_entry() -> impl Strategy<Value = crate::config::VertexApiK
})
}
/// 生成随机的 iFlow 凭证条目
fn arb_iflow_credential_entry() -> impl Strategy<Value = crate::config::IFlowCredentialEntry> {
(
"[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s),
proptest::option::of("[a-z]+/iflow-token-[0-9]{1,5}\\.json".prop_map(|s| s)),
prop_oneof![Just("oauth".to_string()), Just("cookie".to_string())],
proptest::option::of("[a-zA-Z0-9=;]+".prop_map(|s| s)),
proptest::option::of("http://proxy\\.[a-z]+\\.com:[0-9]{4}".prop_map(|s| s)),
any::<bool>(),
)
.prop_map(
|(id, token_file, auth_type, cookies, proxy_url, disabled)| {
crate::config::IFlowCredentialEntry {
id,
token_file,
auth_type,
cookies,
proxy_url,
disabled,
}
},
)
}
/// 生成包含新 Provider 凭证的凭证池配置
fn arb_extended_credential_pool_config() -> impl Strategy<Value = CredentialPoolConfig> {
(
@@ -2257,30 +2232,20 @@ fn arb_extended_credential_pool_config() -> impl Strategy<Value = CredentialPool
proptest::collection::vec(arb_api_key_entry(), 0..3),
proptest::collection::vec(arb_gemini_api_key_entry(), 0..3),
proptest::collection::vec(arb_vertex_api_key_entry(), 0..3),
proptest::collection::vec(arb_oauth_credential_entry(), 0..3),
proptest::collection::vec(arb_iflow_credential_entry(), 0..3),
proptest::collection::vec(arb_credential_entry(), 0..3),
)
.prop_map(
|(
kiro,
gemini,
qwen,
openai,
claude,
gemini_api_keys,
vertex_api_keys,
codex,
iflow,
)| CredentialPoolConfig {
kiro,
gemini,
qwen,
openai,
claude,
gemini_api_keys,
vertex_api_keys,
codex,
iflow,
|(kiro, gemini, qwen, openai, claude, gemini_api_keys, vertex_api_keys, codex)| {
CredentialPoolConfig {
kiro,
gemini,
qwen,
openai,
claude,
gemini_api_keys,
vertex_api_keys,
codex,
}
},
)
}
@@ -2318,16 +2283,6 @@ proptest! {
parsed.credential_pool.gemini.len(),
"Gemini OAuth 凭证数量往返不一致"
);
prop_assert_eq!(
pool.codex.len(),
parsed.credential_pool.codex.len(),
"Codex OAuth 凭证数量往返不一致"
);
prop_assert_eq!(
pool.iflow.len(),
parsed.credential_pool.iflow.len(),
"iFlow 凭证数量往返不一致"
);
// 验证 Gemini API Key 多账号配置往返一致性
prop_assert_eq!(
@@ -2343,44 +2298,6 @@ proptest! {
"Vertex AI 凭证数量往返不一致"
);
// 验证每个 Codex OAuth 凭证的详细内容
for (original, parsed_entry) in pool.codex.iter().zip(parsed.credential_pool.codex.iter()) {
prop_assert_eq!(
&original.id,
&parsed_entry.id,
"Codex 凭证 ID 往返不一致"
);
prop_assert_eq!(
&original.token_file,
&parsed_entry.token_file,
"Codex Token 文件路径往返不一致"
);
prop_assert_eq!(
original.disabled,
parsed_entry.disabled,
"Codex 禁用状态往返不一致"
);
prop_assert_eq!(
&original.proxy_url,
&parsed_entry.proxy_url,
"Codex 代理 URL 往返不一致"
);
}
// 验证每个 iFlow 凭证的详细内容
for (original, parsed_entry) in pool.iflow.iter().zip(parsed.credential_pool.iflow.iter()) {
prop_assert_eq!(
&original.id,
&parsed_entry.id,
"iFlow 凭证 ID 往返不一致"
);
prop_assert_eq!(
&original.auth_type,
&parsed_entry.auth_type,
"iFlow 认证类型往返不一致"
);
}
// 验证每个 Gemini API Key 的详细内容
for (original, parsed_entry) in pool.gemini_api_keys.iter().zip(parsed.credential_pool.gemini_api_keys.iter()) {
prop_assert_eq!(
@@ -2431,7 +2348,6 @@ fn arb_provider_name() -> impl Strategy<Value = String> {
Just("openai".to_string()),
Just("claude".to_string()),
Just("codex".to_string()),
Just("iflow".to_string()),
]
}
@@ -2581,7 +2497,6 @@ fn arb_valid_provider_type() -> impl Strategy<Value = String> {
Just("gemini_api_key".to_string()),
Just("codex".to_string()),
Just("claude_oauth".to_string()),
Just("iflow".to_string()),
]
}
@@ -2601,7 +2516,6 @@ fn arb_invalid_provider_type() -> impl Strategy<Value = String> {
| "gemini_api_key"
| "codex"
| "claude_oauth"
| "iflow"
)
})
}
+3 -1
View File
@@ -57,11 +57,13 @@ impl ProtocolSelector {
PoolProviderType::Gemini => Protocol::Gemini,
PoolProviderType::OpenAI => Protocol::OpenAI,
PoolProviderType::Claude => Protocol::Anthropic,
PoolProviderType::ClaudeOAuth => Protocol::Anthropic, // Claude OAuth uses Anthropic protocol
// Anthropic 兼容格式使用 Anthropic 协议
PoolProviderType::AnthropicCompatible => Protocol::Anthropic,
PoolProviderType::Antigravity => Protocol::Antigravity,
PoolProviderType::Vertex => Protocol::Gemini, // Vertex AI uses Gemini protocol
PoolProviderType::GeminiApiKey => Protocol::Gemini, // Gemini API Key uses Gemini protocol
PoolProviderType::Codex => Protocol::OpenAI, // Codex uses OpenAI protocol
PoolProviderType::ClaudeOAuth => Protocol::Anthropic, // Claude OAuth uses Anthropic protocol
// API Key Provider 类型
PoolProviderType::Anthropic => Protocol::Anthropic,
PoolProviderType::AzureOpenai => Protocol::OpenAI,
+6
View File
@@ -368,6 +368,12 @@ impl CredentialSyncService {
"Claude OAuth 凭证暂不支持同步到配置".to_string(),
));
}
// Anthropic 兼容格式 - 不支持同步到配置
PoolProviderType::AnthropicCompatible => {
return Err(SyncError::InvalidCredentialType(
"Anthropic Compatible 凭证暂不支持同步到配置".to_string(),
));
}
// API Key Provider 类型 - 不支持同步到配置
PoolProviderType::Anthropic
| PoolProviderType::AzureOpenai
+11 -3
View File
@@ -321,12 +321,18 @@ impl ProviderCredential {
/// 检查凭证是否适用于指定的客户端类型
///
/// 某些凭证可能有使用限制,例如 Claude Code 专用凭证只能用于 Claude Code 客户端
pub fn is_compatible_with_client(&self, client_type: Option<&crate::server::client_detector::ClientType>) -> bool {
pub fn is_compatible_with_client(
&self,
client_type: Option<&crate::server::client_detector::ClientType>,
) -> bool {
// 检查是否是 Claude Code 专用凭证
if let Some(error_msg) = &self.last_error_message {
if error_msg.contains("only authorized for use with Claude Code") {
// 这是 Claude Code 专用凭证,只能用于 Claude Code 客户端
return matches!(client_type, Some(crate::server::client_detector::ClientType::ClaudeCode));
return matches!(
client_type,
Some(crate::server::client_detector::ClientType::ClaudeCode)
);
}
}
@@ -514,11 +520,13 @@ pub fn get_default_check_model(provider_type: PoolProviderType) -> &'static str
PoolProviderType::OpenAI => "gpt-3.5-turbo",
// 使用 claude-sonnet-4-5-20250929,兼容更多代理服务器
PoolProviderType::Claude => "claude-sonnet-4-5-20250929",
PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929",
// Anthropic 兼容格式使用相同的健康检查模型
PoolProviderType::AnthropicCompatible => "claude-sonnet-4-5-20250929",
PoolProviderType::Antigravity => "gemini-3-pro-preview",
PoolProviderType::Vertex => "gemini-2.0-flash",
PoolProviderType::GeminiApiKey => "gemini-2.5-flash",
PoolProviderType::Codex => "gpt-4o-mini",
PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929",
// API Key Provider 类型
PoolProviderType::Anthropic => "claude-sonnet-4-5-20250929",
PoolProviderType::AzureOpenai => "gpt-4o-mini",
+4 -3
View File
@@ -281,16 +281,17 @@ impl ClaudeCustomProvider {
if !status.is_success() {
let body = resp.text().await.unwrap_or_default();
// 检查是否是 Claude Code 专用凭证限制错误
if body.contains("only authorized for use with Claude Code") {
return Err(format!(
"凭证限制错误: 当前 Claude 凭证只能用于 Claude Code,不能用于通用 API 调用。\
请使用通用的 Claude API Key 或 Anthropic API Key。\
错误详情: {status} - {body}"
).into());
)
.into());
}
return Err(format!("Claude API error: {status} - {body}").into());
}
+47 -22
View File
@@ -743,7 +743,12 @@ pub async fn chat_completions(
);
let cred = state
.pool_service
.select_credential_with_client_check(db, explicit_provider_id, Some(&request.model), Some(&client_type))
.select_credential_with_client_check(
db,
explicit_provider_id,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
@@ -781,7 +786,12 @@ pub async fn chat_completions(
);
let cred = state
.pool_service
.select_credential_with_client_check(db, &selected_provider, Some(&request.model), Some(&client_type))
.select_credential_with_client_check(
db,
&selected_provider,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
@@ -825,12 +835,16 @@ pub async fn chat_completions(
provider_id_lower
);
match state.api_key_service.get_fallback_credential(
db,
&crate::models::provider_pool_model::PoolProviderType::OpenAI,
Some(&provider_id_lower),
Some(&client_type),
).await {
match state
.api_key_service
.get_fallback_credential(
db,
&crate::models::provider_pool_model::PoolProviderType::OpenAI,
Some(&provider_id_lower),
Some(&client_type),
)
.await
{
Ok(Some(cred)) => {
eprintln!(
"[CHAT_COMPLETIONS] 通过 provider_id '{}' 找到凭证: name={:?}",
@@ -1932,7 +1946,12 @@ pub async fn anthropic_messages(
);
let cred = state
.pool_service
.select_credential_with_client_check(db, explicit_provider_id, Some(&request.model), Some(&client_type))
.select_credential_with_client_check(
db,
explicit_provider_id,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
@@ -1969,7 +1988,12 @@ pub async fn anthropic_messages(
);
let cred = state
.pool_service
.select_credential_with_client_check(db, &selected_provider, Some(&request.model), Some(&client_type))
.select_credential_with_client_check(
db,
&selected_provider,
Some(&request.model),
Some(&client_type),
)
.ok()
.flatten();
@@ -2008,12 +2032,16 @@ pub async fn anthropic_messages(
selected_provider
);
match state.api_key_service.get_fallback_credential(
db,
&crate::models::provider_pool_model::PoolProviderType::Anthropic,
Some(&selected_provider),
Some(&client_type),
).await {
match state
.api_key_service
.get_fallback_credential(
db,
&crate::models::provider_pool_model::PoolProviderType::Anthropic,
Some(&selected_provider),
Some(&client_type),
)
.await
{
Ok(Some(cred)) => {
eprintln!(
"[ANTHROPIC_MESSAGES] 通过 provider_id '{}' 找到凭证: name={:?}",
@@ -2053,12 +2081,9 @@ pub async fn anthropic_messages(
// 启动 Flow 捕获
let llm_request = build_llm_request_from_anthropic(&request, "/v1/messages", &headers);
// 尝试将 selected_provider 解析为 ProviderType
// 如果是自定义 provider ID,则使用 OpenAI 作为默认值
// 使用实际的 provider ID 构建 Flow Metadata
let provider_type = selected_provider
.parse::<ProviderType>()
.unwrap_or(ProviderType::OpenAI);
// 使用凭证的实际 provider_type(支持自定义 Provider)
// 对于自定义 Provider ID,凭证的 provider_type 已通过数据库查询正确设置
let provider_type = cred.provider_type;
// 从凭证名称中提取 Provider 显示名称
// 凭证名称格式:Some("[降级] DeepSeek") 或 Some("DeepSeek")
@@ -416,6 +416,25 @@ pub async fn management_add_credential(
);
}
}
// Anthropic 兼容格式 - 使用 ClaudeKey(与 Anthropic 相同)
PoolProviderType::AnthropicCompatible => {
if let Some(api_key) = request.api_key {
CredentialData::ClaudeKey {
api_key,
base_url: request.base_url,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "API key is required for Anthropic Compatible provider"
.to_string(),
id: None,
}),
);
}
}
// API Key Provider 类型 - 不支持通过此接口添加凭证
PoolProviderType::AzureOpenai | PoolProviderType::AwsBedrock | PoolProviderType::Ollama => {
return (
-356
View File
@@ -1016,8 +1016,6 @@ async fn run_server(
"/v1/images/generations",
post(handlers::handle_image_generation),
)
// Gemini 原生协议路由
.route("/v1/gemini/{*path}", post(gemini_generate_content))
// WebSocket 路由
.route("/v1/ws", get(handlers::ws_upgrade_handler))
.route("/ws", get(handlers::ws_upgrade_handler))
@@ -1030,28 +1028,6 @@ async fn run_server(
"/{selector}/v1/chat/completions",
post(chat_completions_with_selector),
)
// Amp CLI 路由
.route(
"/api/provider/{provider}/v1/chat/completions",
post(
|State(state): State<AppState>,
Path(provider): Path<String>,
headers: HeaderMap,
Json(mut request): Json<crate::models::openai::ChatCompletionRequest>| async {
amp_chat_completions(State(state), Path(provider), headers, Json(request)).await
}
),
)
// TODO: amp_messages 和 amp_management_proxy 路由暂时禁用
// .route("/api/provider/{provider}/v1/messages", post(amp_messages))
// .route(
// "/api/auth/{*path}",
// axum::routing::any(amp_management_proxy_auth),
// )
// .route(
// "/api/user/{*path}",
// axum::routing::any(amp_management_proxy_user),
// )
// 管理 API 路由
.merge(management_routes)
// Kiro凭证管理API路由
@@ -1743,338 +1719,6 @@ async fn chat_completions_with_selector(
}
}
// ============ Amp CLI 路由处理 ============
/// Amp CLI chat completions 处理
///
/// 处理 `/api/provider/:provider/v1/chat/completions` 路由
/// 支持模型映射,将不可用模型映射到可用替代
async fn amp_chat_completions(
State(state): State<AppState>,
Path(provider): Path<String>,
headers: HeaderMap,
Json(mut request): Json<ChatCompletionRequest>,
) -> Response {
if let Err(e) = handlers::verify_api_key(&headers, &state.api_key).await {
state.logs.write().await.add(
"warn",
&format!(
"Unauthorized request to /api/provider/{}/v1/chat/completions",
provider
),
);
return e.into_response();
}
// 应用模型映射
let original_model = request.model.clone();
let mapped_model = state.amp_router.apply_model_mapping(&request.model);
if mapped_model != original_model {
state.logs.write().await.add(
"info",
&format!(
"[AMP] Model mapping applied: {} -> {}",
original_model, mapped_model
),
);
request.model = mapped_model;
}
state.logs.write().await.add(
"info",
&format!(
"[AMP] POST /api/provider/{}/v1/chat/completions model={} stream={}",
provider, request.model, request.stream
),
);
// 尝试根据 provider 名称选择凭证
eprintln!(
"[AMP] 开始查找凭证: provider={}, model={}, db={}",
provider,
request.model,
state.db.is_some()
);
let credential = match &state.db {
Some(db) => {
eprintln!(
"[AMP] 使用 select_credential 查找凭证(Provider Pool): provider={}",
provider
);
// 先尝试从 Provider Pool 查找
let pool_cred = if let Ok(Some(cred)) =
state
.pool_service
.select_credential(db, &provider, Some(&request.model))
{
eprintln!("[AMP] select_credential 找到凭证: {:?}", cred.name);
Some(cred)
}
// 然后尝试按名称查找
else if let Ok(Some(cred)) = state.pool_service.get_by_name(db, &provider) {
eprintln!("[AMP] get_by_name 找到凭证: {:?}", cred.name);
Some(cred)
}
// 最后尝试按 UUID 查找
else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &provider) {
eprintln!("[AMP] get_by_uuid 找到凭证: {:?}", cred.name);
Some(cred)
} else {
None
};
// 如果 Provider Pool 中没有找到,尝试从 API Key Provider 查找
if pool_cred.is_none() {
eprintln!(
"[AMP] Provider Pool 中未找到凭证,尝试 API Key Provider: provider={}",
provider
);
match state.api_key_service.get_fallback_credential(
db,
&crate::models::provider_pool_model::PoolProviderType::OpenAI,
Some(&provider),
None, // 没有客户端类型检测
).await {
Ok(Some(cred)) => {
eprintln!(
"[AMP] 通过 provider_id '{}' 找到 API Key Provider 凭证: name={:?}",
provider, cred.name
);
Some(cred)
}
Ok(None) => {
eprintln!("[AMP] 未找到任何凭证 for provider '{}'", provider);
None
}
Err(e) => {
eprintln!("[AMP] 查找 API Key Provider 凭证时出错: {}", e);
None
}
}
} else {
pool_cred
}
}
None => {
eprintln!("[AMP] 数据库未初始化");
None
}
};
match credential {
Some(cred) => {
state.logs.write().await.add(
"info",
&format!(
"[AMP] Using credential: type={} name={:?} uuid={}",
cred.provider_type,
cred.name,
&cred.uuid[..8]
),
);
// 注意:这里没有 Flow 捕获,因为是通过 AMP CLI 路由的请求
handlers::call_provider_openai(&state, &cred, &request, None).await.into_response()
}
None => {
// 不再回退到默认 provider,直接返回错误
state.logs.write().await.add(
"error",
&format!(
"[AMP] No available credentials for provider '{}', refusing to fallback",
provider
),
);
(
StatusCode::SERVICE_UNAVAILABLE,
Json(serde_json::json!({
"error": {
"message": format!("No available credentials for provider '{}'", provider),
"type": "provider_unavailable",
"code": "no_credentials"
}
})),
)
.into_response()
}
}
}
/// Amp CLI 管理代理内部实现
///
/// 处理 `/api/auth/*` 和 `/api/user/*` 路由
/// 将请求代理到上游 URL
///
/// # 参数
/// - `path`: 请求路径(不含 /api/ 前缀,如 "auth/login" 或 "user/profile")
async fn amp_management_proxy_internal(
state: AppState,
path: &str,
headers: HeaderMap,
method: axum::http::Method,
body: axum::body::Bytes,
) -> Response {
let full_path = format!("/api/{}", path);
// 检查是否是管理路由
if !state.amp_router.is_management_route(&full_path) {
state.logs.write().await.add(
"warn",
&format!("[AMP] Invalid management route: {}", full_path),
);
return (
StatusCode::NOT_FOUND,
Json(serde_json::json!({"error": {"message": "Not found"}})),
)
.into_response();
}
// 检查 localhost 限制
if state.amp_router.restrict_management_to_localhost() {
// 从 headers 中获取客户端 IP
let client_ip = headers
.get("x-forwarded-for")
.and_then(|v| v.to_str().ok())
.map(|s| s.split(',').next().unwrap_or("").trim().to_string())
.or_else(|| {
headers
.get("x-real-ip")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
});
if let Some(ip) = &client_ip {
let is_localhost = ip == "127.0.0.1" || ip == "::1" || ip == "localhost";
if !is_localhost {
state.logs.write().await.add(
"warn",
&format!("[AMP] Management proxy blocked from non-localhost: {}", ip),
);
return (
StatusCode::FORBIDDEN,
Json(serde_json::json!({"error": {"message": "Management endpoints are restricted to localhost"}})),
)
.into_response();
}
}
}
// 获取上游 URL
let upstream_url = match state.amp_router.get_management_upstream_path(&full_path) {
Some(url) => url,
None => {
state.logs.write().await.add(
"warn",
"[AMP] No upstream URL configured for management proxy",
);
return (
StatusCode::SERVICE_UNAVAILABLE,
Json(serde_json::json!({"error": {"message": "Upstream URL not configured"}})),
)
.into_response();
}
};
state.logs.write().await.add(
"info",
&format!(
"[AMP] Proxying management request: {} {} -> {}",
method, full_path, upstream_url
),
);
// 创建 HTTP 客户端
let client = reqwest::Client::new();
// 构建请求
let mut request_builder = match method {
axum::http::Method::GET => client.get(&upstream_url),
axum::http::Method::POST => client.post(&upstream_url),
axum::http::Method::PUT => client.put(&upstream_url),
axum::http::Method::DELETE => client.delete(&upstream_url),
axum::http::Method::PATCH => client.patch(&upstream_url),
axum::http::Method::HEAD => client.head(&upstream_url),
axum::http::Method::OPTIONS => client.request(reqwest::Method::OPTIONS, &upstream_url),
_ => {
return (
StatusCode::METHOD_NOT_ALLOWED,
Json(serde_json::json!({"error": {"message": "Method not allowed"}})),
)
.into_response();
}
};
// 复制请求头(排除 host 和 content-length)
for (name, value) in headers.iter() {
let name_str = name.as_str().to_lowercase();
if name_str != "host" && name_str != "content-length" {
if let Ok(value_str) = value.to_str() {
request_builder = request_builder.header(name.as_str(), value_str);
}
}
}
// 添加请求体
if !body.is_empty() {
request_builder = request_builder.body(body.to_vec());
}
// 发送请求
match request_builder.send().await {
Ok(response) => {
let status = response.status();
let response_headers = response.headers().clone();
match response.bytes().await {
Ok(response_body) => {
let mut builder = Response::builder().status(status.as_u16());
// 复制响应头
for (name, value) in response_headers.iter() {
let name_str = name.as_str().to_lowercase();
// 排除 transfer-encoding 和 content-length(axum 会自动处理)
if name_str != "transfer-encoding" && name_str != "content-length" {
builder = builder.header(name.as_str(), value.to_str().unwrap_or(""));
}
}
builder
.body(Body::from(response_body.to_vec()))
.unwrap_or_else(|_| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": "Failed to build response"}})),
)
.into_response()
})
}
Err(e) => {
state.logs.write().await.add(
"error",
&format!("[AMP] Failed to read upstream response: {}", e),
);
(
StatusCode::BAD_GATEWAY,
Json(serde_json::json!({"error": {"message": format!("Failed to read upstream response: {}", e)}})),
)
.into_response()
}
}
}
Err(e) => {
state.logs.write().await.add(
"error",
&format!("[AMP] Failed to proxy request to upstream: {}", e),
);
(
StatusCode::BAD_GATEWAY,
Json(serde_json::json!({"error": {"message": format!("Failed to connect to upstream: {}", e)}})),
)
.into_response()
}
}
}
/// 内部 Anthropic messages 处理 (使用默认 Kiro)
/// 预留:用于内部直接调用 Kiro API
#[allow(dead_code)]
@@ -861,7 +861,10 @@ impl ApiKeyProviderService {
"[get_fallback_credential] 尝试按 provider_id '{}' 查找",
provider_id
);
if let Some(cred) = self.find_by_provider_id(db, provider_id, client_type).await? {
if let Some(cred) = self
.find_by_provider_id(db, provider_id, client_type)
.await?
{
eprintln!(
"[get_fallback_credential] 通过 provider_id '{}' 找到凭证: {:?}",
provider_id, cred.name
@@ -912,6 +915,7 @@ impl ApiKeyProviderService {
// API Key Provider 类型 - 直接映射
PoolProviderType::Anthropic => Some(ApiProviderType::Anthropic),
PoolProviderType::AnthropicCompatible => Some(ApiProviderType::AnthropicCompatible),
PoolProviderType::AzureOpenai => Some(ApiProviderType::AzureOpenai),
PoolProviderType::AwsBedrock => Some(ApiProviderType::AwsBedrock),
PoolProviderType::Ollama => Some(ApiProviderType::Ollama),
@@ -1004,8 +1008,8 @@ impl ApiKeyProviderService {
let conn = db.lock().map_err(|e| e.to_string())?;
// 直接按 provider_id 查找
let provider =
ApiKeyProviderDao::get_provider_by_id(&conn, provider_id).map_err(|e| e.to_string())?;
let provider = ApiKeyProviderDao::get_provider_by_id(&conn, provider_id)
.map_err(|e| e.to_string())?;
let provider = match provider {
Some(p) if p.enabled => {
@@ -1072,14 +1076,20 @@ impl ApiKeyProviderService {
if provider.provider_type == ApiProviderType::Anthropic {
if let Some(client) = client_type {
// 对于 Claude Code 客户端,可以使用任何 Claude 凭证
if matches!(client, crate::server::client_detector::ClientType::ClaudeCode) {
if matches!(
client,
crate::server::client_detector::ClientType::ClaudeCode
) {
selected_key = Some(candidate_key);
break;
}
// 对于其他客户端,需要检查凭证是否是 Claude Code 专用
// 通过发送测试请求来检查
if let Err(e) = self.test_claude_key_compatibility(&api_key, &provider.api_host).await {
if let Err(e) = self
.test_claude_key_compatibility(&api_key, &provider.api_host)
.await
{
if e.contains("CLAUDE_CODE_ONLY") {
eprintln!(
"[find_by_provider_id] API Key {} 是 Claude Code 专用,跳过 (客户端: {:?})",
@@ -1141,6 +1151,15 @@ impl ApiKeyProviderService {
};
(data, PoolProviderType::Claude)
}
ApiProviderType::AnthropicCompatible => {
// Anthropic 兼容格式使用 ClaudeKey(与 Anthropic 相同的凭证数据)
// 但使用 AnthropicCompatible 作为 PoolProviderType,以便使用正确的端点
let data = CredentialData::ClaudeKey {
api_key: api_key.to_string(),
base_url: Some(provider.api_host.clone()),
};
(data, PoolProviderType::AnthropicCompatible)
}
ApiProviderType::Gemini => {
// Gemini 类型使用 GeminiApiKey
let data = CredentialData::GeminiApiKey {
@@ -1289,12 +1308,18 @@ impl ApiKeyProviderService {
let test_model = model_name
.or_else(|| provider.custom_models.first().cloned())
.unwrap_or_else(|| "claude-3-haiku-20240307".to_string());
match self.test_anthropic_connection(&api_key, &provider.api_host, &test_model).await {
match self
.test_anthropic_connection(&api_key, &provider.api_host, &test_model)
.await
{
Ok(models) => Ok(models),
Err(e) if e == "CLAUDE_CODE_ONLY" => {
// Claude Code 专用凭证限制错误,返回特殊错误信息
Err("凭证限制: 当前 Claude 凭证只能用于 Claude Code,不能用于通用 API 调用".to_string())
Err(
"凭证限制: 当前 Claude 凭证只能用于 Claude Code,不能用于通用 API 调用"
.to_string(),
)
}
Err(e) => Err(e),
}
@@ -1444,12 +1469,12 @@ impl ApiKeyProviderService {
Ok(())
} else {
let body = response.text().await.unwrap_or_default();
// 检查是否是 Claude Code 专用凭证限制错误
if body.contains("only authorized for use with Claude Code") {
return Err("CLAUDE_CODE_ONLY".to_string());
}
// 其他错误不影响兼容性判断
Ok(())
}
@@ -1484,12 +1509,12 @@ impl ApiKeyProviderService {
} else {
let status = response.status();
let body = response.text().await.unwrap_or_default();
// 检查是否是 Claude Code 专用凭证限制错误
if body.contains("only authorized for use with Claude Code") {
return Err("CLAUDE_CODE_ONLY".to_string());
}
Err(format!("API 返回错误: {} - {}", status, body))
}
}
@@ -270,8 +270,9 @@ impl ProviderPoolService {
);
credentials.extend(assistant_creds);
} else if pt == PoolProviderType::Claude {
let ai_provider_creds = ProviderPoolDao::get_by_type(&conn, &PoolProviderType::Anthropic)
.map_err(|e| e.to_string())?;
let ai_provider_creds =
ProviderPoolDao::get_by_type(&conn, &PoolProviderType::Anthropic)
.map_err(|e| e.to_string())?;
eprintln!(
"[SELECT_CREDENTIAL] Assistant: adding {} AI Provider credentials",
ai_provider_creds.len()
@@ -398,7 +399,9 @@ impl ProviderPoolService {
);
// Step 1: 尝试从 Provider Pool 选择 (OAuth + API Key)
if let Some(cred) = self.select_credential_with_client_check(db, provider_type, model, client_type)? {
if let Some(cred) =
self.select_credential_with_client_check(db, provider_type, model, client_type)?
{
eprintln!(
"[select_credential_with_fallback] 从 Provider Pool 找到凭证: {:?}",
cred.name
@@ -416,7 +419,10 @@ impl ProviderPoolService {
// 传入 provider_id_hint 支持 60+ Provider
eprintln!("[select_credential_with_fallback] 调用 get_fallback_credential");
if let Some(cred) = api_key_service.get_fallback_credential(db, &pt, provider_id_hint, client_type).await? {
if let Some(cred) = api_key_service
.get_fallback_credential(db, &pt, provider_id_hint, client_type)
.await?
{
eprintln!(
"[select_credential_with_fallback] 智能降级成功: {:?}",
cred.name
@@ -450,7 +456,8 @@ impl ProviderPoolService {
model,
provider_id_hint,
None, // 兼容方法不传递客户端类型
).await
)
.await
}
/// 基于权重分数选择最优凭证
+1 -1
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "ProxyCast",
"version": "0.47.4",
"version": "0.48.0",
"identifier": "com.proxycast.app",
"build": {
"beforeDevCommand": "npm run dev",
@@ -147,7 +147,11 @@ export function DecisionPanel({ request, onSubmit }: DecisionPanelProps) {
};
// 渲染用户问题面板
if (request.actionType === "ask_user" && request.questions && request.questions.length > 0) {
if (
request.actionType === "ask_user" &&
request.questions &&
request.questions.length > 0
) {
const questions = request.questions;
return (
<Card className="border-blue-200 bg-blue-50/50 dark:border-blue-800 dark:bg-blue-950/20">
@@ -108,9 +108,7 @@ export const InputbarTools: React.FC<InputbarToolsProps> = ({
onClick={() => onToolClick?.("canvas")}
className={isCanvasOpen ? "active" : ""}
>
<PanelRight
className={isCanvasOpen ? "text-primary" : ""}
/>
<PanelRight className={isCanvasOpen ? "text-primary" : ""} />
</ToolButton>
</TooltipTrigger>
<TooltipContent side="top">
@@ -63,23 +63,23 @@ export const MessageList: React.FC<MessageListProps> = ({
const handleScroll = () => {
const { scrollTop, scrollHeight, clientHeight } = container;
const isAtBottom = scrollHeight - scrollTop - clientHeight < 50; // 50px 容差
setIsUserScrolling(true);
setShouldAutoScroll(isAtBottom);
// 清除之前的定时器
clearTimeout(scrollTimeout);
// 500ms 后认为用户停止滚动
scrollTimeout = setTimeout(() => {
setIsUserScrolling(false);
}, 500);
};
container.addEventListener('scroll', handleScroll, { passive: true });
container.addEventListener("scroll", handleScroll, { passive: true });
return () => {
container.removeEventListener('scroll', handleScroll);
container.removeEventListener("scroll", handleScroll);
clearTimeout(scrollTimeout);
};
}, []);
@@ -363,7 +363,12 @@ export const StreamingRenderer: React.FC<StreamingRendererProps> = memo(
const result = parseAIResponse(visibleText, isStreaming);
// 添加调试日志
if (result.hasWriteFile) {
console.log("[StreamingRenderer] 检测到 write_file:", result.parts.filter(p => p.type === "write_file" || p.type === "pending_write_file"));
console.log(
"[StreamingRenderer] 检测到 write_file:",
result.parts.filter(
(p) => p.type === "write_file" || p.type === "pending_write_file",
),
);
}
return result;
}, [visibleText, isStreaming]);
@@ -409,10 +414,13 @@ export const StreamingRenderer: React.FC<StreamingRendererProps> = memo(
// 判断是否有可见内容
const hasVisibleContent = useInterleavedMode
? contentParts.some(
(part) =>
(part) =>
(part.type === "text" && part.text.length > 0) ||
(part.type === "thinking" && part.text.length > 0),
) || (isStreaming && (content.length > 0 || (externalThinking && externalThinking.length > 0)))
) ||
(isStreaming &&
(content.length > 0 ||
(externalThinking && externalThinking.length > 0)))
: visibleText.length > 0;
// 交错显示模式:按顺序渲染 contentParts
@@ -433,7 +441,14 @@ export const StreamingRenderer: React.FC<StreamingRendererProps> = memo(
// 添加调试日志
if (partParsed.hasWriteFile) {
console.log("[StreamingRenderer] 交错模式检测到 write_file:", partParsed.parts.filter(p => p.type === "write_file" || p.type === "pending_write_file"));
console.log(
"[StreamingRenderer] 交错模式检测到 write_file:",
partParsed.parts.filter(
(p) =>
p.type === "write_file" ||
p.type === "pending_write_file",
),
);
}
// 处理文件写入回调
@@ -627,10 +627,12 @@ export function useAgentChat(options: UseAgentChatOptions = {}) {
// 如果是写入文件工具,立即调用 onWriteFile 展开右边栏
const toolName = data.tool_name.toLowerCase();
console.log(`[Tool Start] 工具名称: ${data.tool_name}, 小写: ${toolName}`);
console.log(
`[Tool Start] 工具名称: ${data.tool_name}, 小写: ${toolName}`,
);
console.log(`[Tool Start] 工具参数: ${data.arguments}`);
console.log(`[Tool Start] onWriteFile 回调存在: ${!!onWriteFile}`);
if (toolName.includes("write") || toolName.includes("create")) {
console.log(`[Tool Start] 匹配到文件写入工具: ${data.tool_name}`);
try {
@@ -638,12 +640,16 @@ export function useAgentChat(options: UseAgentChatOptions = {}) {
console.log(`[Tool Start] 解析后的参数:`, args);
const filePath = args.path || args.file_path || args.filePath;
const content = args.content || args.text || "";
console.log(`[Tool Start] 文件路径: ${filePath}, 内容长度: ${content.length}`);
console.log(
`[Tool Start] 文件路径: ${filePath}, 内容长度: ${content.length}`,
);
if (filePath && content && onWriteFile) {
console.log(`[Tool Start] 触发文件写入: ${filePath}`);
onWriteFile(content, filePath);
} else {
console.log(`[Tool Start] 文件写入条件不满足: filePath=${!!filePath}, content=${!!content}, onWriteFile=${!!onWriteFile}`);
console.log(
`[Tool Start] 文件写入条件不满足: filePath=${!!filePath}, content=${!!content}, onWriteFile=${!!onWriteFile}`,
);
}
} catch (e) {
console.warn("[Tool Start] 解析工具参数失败:", e);
@@ -683,11 +689,16 @@ export function useAgentChat(options: UseAgentChatOptions = {}) {
case "action_required": {
// 权限确认请求 - 添加到权限请求列表和 contentParts
console.log(`[Action Required] ${data.action_type} (${data.request_id})`);
console.log(
`[Action Required] ${data.action_type} (${data.request_id})`,
);
const actionRequired: ActionRequired = {
requestId: data.request_id,
actionType: data.action_type as "tool_confirmation" | "ask_user" | "elicitation",
actionType: data.action_type as
| "tool_confirmation"
| "ask_user"
| "elicitation",
toolName: data.tool_name,
arguments: data.arguments,
prompt: data.prompt,
@@ -712,7 +723,10 @@ export function useAgentChat(options: UseAgentChatOptions = {}) {
return {
...msg,
actionRequests: [...(msg.actionRequests || []), actionRequired],
actionRequests: [
...(msg.actionRequests || []),
actionRequired,
],
// 添加到 contentParts,支持交错显示
contentParts: addActionRequiredToParts(
msg.contentParts || [],
@@ -958,7 +972,7 @@ export function useAgentChat(options: UseAgentChatOptions = {}) {
confirmed: response.confirmed,
response: response.response,
});
// 移除已处理的权限请求
setMessages((prev) =>
prev.map((msg) => ({
+61 -8
View File
@@ -475,7 +475,8 @@ export function AgentChatPage({
setGeneralCanvasState((prev) => ({
...prev,
isOpen: !prev.isOpen,
contentType: prev.contentType === "empty" ? "markdown" : prev.contentType,
contentType:
prev.contentType === "empty" ? "markdown" : prev.contentType,
content: prev.content || "# 新文档\n\n在这里开始编写内容...",
}));
setLayoutMode((prev) => (prev === "chat" ? "chat-canvas" : "chat"));
@@ -490,7 +491,8 @@ export function AgentChatPage({
createInitialCanvasState(
mappedTheme,
"# 新文档\n\n在这里开始编写内容...",
) || createInitialDocumentState("# 新文档\n\n在这里开始编写内容...");
) ||
createInitialDocumentState("# 新文档\n\n在这里开始编写内容...");
setCanvasState(initialState);
}
return "chat-canvas";
@@ -521,11 +523,39 @@ export function AgentChatPage({
// General 主题使用专门的画布处理
if (activeTheme === "general") {
const ext = fileName.split(".").pop()?.toLowerCase() || "";
const isCode = ["js", "ts", "tsx", "jsx", "py", "rs", "go", "java", "c", "cpp", "h", "css", "scss", "json", "yaml", "yml", "toml", "xml", "html", "sql", "sh", "bash"].includes(ext);
const isCode = [
"js",
"ts",
"tsx",
"jsx",
"py",
"rs",
"go",
"java",
"c",
"cpp",
"h",
"css",
"scss",
"json",
"yaml",
"yml",
"toml",
"xml",
"html",
"sql",
"sh",
"bash",
].includes(ext);
const isMd = ["md", "markdown"].includes(ext);
console.log("[AgentChatPage] General 主题文件写入:", fileName, "类型:", isCode ? "code" : isMd ? "markdown" : "file");
console.log(
"[AgentChatPage] General 主题文件写入:",
fileName,
"类型:",
isCode ? "code" : isMd ? "markdown" : "file",
);
setGeneralCanvasState({
isOpen: true,
contentType: isCode ? "code" : isMd ? "markdown" : "file",
@@ -678,9 +708,32 @@ export function AgentChatPage({
// General 主题使用专门的画布
if (activeTheme === "general") {
const ext = fileName.split(".").pop()?.toLowerCase() || "";
const isCode = ["js", "ts", "tsx", "jsx", "py", "rs", "go", "java", "c", "cpp", "h", "css", "scss", "json", "yaml", "yml", "toml", "xml", "html", "sql", "sh", "bash"].includes(ext);
const isCode = [
"js",
"ts",
"tsx",
"jsx",
"py",
"rs",
"go",
"java",
"c",
"cpp",
"h",
"css",
"scss",
"json",
"yaml",
"yml",
"toml",
"xml",
"html",
"sql",
"sh",
"bash",
].includes(ext);
const isMd = ["md", "markdown"].includes(ext);
setGeneralCanvasState({
isOpen: true,
contentType: isCode ? "code" : isMd ? "markdown" : "file",
+1 -5
View File
@@ -131,7 +131,6 @@ export const PROVIDER_CONFIG: Record<
"claude-opus-4-5-20251101",
"claude-sonnet-4-5-20250929",
"claude-sonnet-4-20250514",
],
},
openai: {
@@ -153,10 +152,7 @@ export const PROVIDER_CONFIG: Record<
},
gemini: {
label: "Gemini",
models: [
"gemini-3-pro-preview",
"gemini-3-flash-preview",
],
models: ["gemini-3-pro-preview", "gemini-3-flash-preview"],
},
qwen: {
label: "通义千问",
@@ -80,6 +80,7 @@ const getProviderApiType = (provider: string): ApiType => {
// Anthropic 类型
if (
p === "anthropic" ||
p === "anthropic-compatible" ||
p === "claude" ||
p === "claude_oauth" ||
p === "kiro"
+40 -21
View File
@@ -6,13 +6,13 @@
* @requirements 3.1, 3.5, 9.4
*/
import React, { useState, useCallback, useEffect, useRef } from 'react';
import { ChatPanel } from './chat/ChatPanel';
import { CanvasPanel } from './canvas/CanvasPanel';
import { ErrorBoundary } from './chat/ErrorBoundary';
import { useGeneralChatStore } from './store/useGeneralChatStore';
import type { CanvasState, GeneralChatPageProps } from './types';
import { DEFAULT_CANVAS_STATE } from './types';
import React, { useState, useCallback, useEffect, useRef } from "react";
import { ChatPanel } from "./chat/ChatPanel";
import { CanvasPanel } from "./canvas/CanvasPanel";
import { ErrorBoundary } from "./chat/ErrorBoundary";
import { useGeneralChatStore } from "./store/useGeneralChatStore";
import type { CanvasState, GeneralChatPageProps } from "./types";
import { DEFAULT_CANVAS_STATE } from "./types";
/**
* 通用对话主页面
@@ -26,15 +26,12 @@ export const GeneralChatPage: React.FC<GeneralChatPageProps> = ({
initialSessionId,
onNavigate,
}) => {
const {
currentSessionId,
selectSession,
sessions,
createSession,
} = useGeneralChatStore();
const { currentSessionId, selectSession, sessions, createSession } =
useGeneralChatStore();
// 画布状态
const [canvasState, setCanvasState] = useState<CanvasState>(DEFAULT_CANVAS_STATE);
const [canvasState, setCanvasState] =
useState<CanvasState>(DEFAULT_CANVAS_STATE);
// 使用 ref 防止重复创建会话
const sessionCreatedRef = useRef(false);
@@ -43,12 +40,22 @@ export const GeneralChatPage: React.FC<GeneralChatPageProps> = ({
useEffect(() => {
if (initialSessionId) {
selectSession(initialSessionId);
} else if (sessions.length === 0 && !currentSessionId && !sessionCreatedRef.current) {
} else if (
sessions.length === 0 &&
!currentSessionId &&
!sessionCreatedRef.current
) {
// 如果没有会话,创建一个新会话(只创建一次)
sessionCreatedRef.current = true;
createSession();
}
}, [initialSessionId, selectSession, sessions.length, currentSessionId, createSession]);
}, [
initialSessionId,
selectSession,
sessions.length,
currentSessionId,
createSession,
]);
// 打开画布
const handleOpenCanvas = useCallback((state: CanvasState) => {
@@ -62,7 +69,7 @@ export const GeneralChatPage: React.FC<GeneralChatPageProps> = ({
// 画布内容变更
const handleCanvasContentChange = useCallback((content: string) => {
setCanvasState(prev => ({ ...prev, content }));
setCanvasState((prev) => ({ ...prev, content }));
}, []);
return (
@@ -72,8 +79,14 @@ export const GeneralChatPage: React.FC<GeneralChatPageProps> = ({
<ErrorBoundary
componentName="ChatPanel"
onError={(error, errorInfo) => {
console.error('[GeneralChatPage] ChatPanel 渲染错误:', error.message);
console.error('[GeneralChatPage] 组件堆栈:', errorInfo.componentStack);
console.error(
"[GeneralChatPage] ChatPanel 渲染错误:",
error.message,
);
console.error(
"[GeneralChatPage] 组件堆栈:",
errorInfo.componentStack,
);
}}
>
<ChatPanel
@@ -90,8 +103,14 @@ export const GeneralChatPage: React.FC<GeneralChatPageProps> = ({
<ErrorBoundary
componentName="CanvasPanel"
onError={(error, errorInfo) => {
console.error('[GeneralChatPage] CanvasPanel 渲染错误:', error.message);
console.error('[GeneralChatPage] 组件堆栈:', errorInfo.componentStack);
console.error(
"[GeneralChatPage] CanvasPanel 渲染错误:",
error.message,
);
console.error(
"[GeneralChatPage] 组件堆栈:",
errorInfo.componentStack,
);
}}
>
<CanvasPanel
@@ -6,10 +6,10 @@
* @requirements 3.3, 4.4
*/
import React, { useState } from 'react';
import type { CanvasState } from '../types';
import { CodePreview } from './CodePreview';
import { MarkdownPreview } from './MarkdownPreview';
import React, { useState } from "react";
import type { CanvasState } from "../types";
import { CodePreview } from "./CodePreview";
import { MarkdownPreview } from "./MarkdownPreview";
interface CanvasPanelProps {
/** 画布状态 */
@@ -39,11 +39,14 @@ export const CanvasPanel: React.FC<CanvasPanelProps> = ({
// 下载内容
const handleDownload = () => {
const filename = state.filename ||
(state.contentType === 'code' ? `code.${state.language || 'txt'}` : 'content.md');
const blob = new Blob([state.content], { type: 'text/plain' });
const filename =
state.filename ||
(state.contentType === "code"
? `code.${state.language || "txt"}`
: "content.md");
const blob = new Blob([state.content], { type: "text/plain" });
const url = URL.createObjectURL(blob);
const a = document.createElement('a');
const a = document.createElement("a");
a.href = url;
a.download = filename;
a.click();
@@ -60,7 +63,8 @@ export const CanvasPanel: React.FC<CanvasPanelProps> = ({
<div className="flex items-center justify-between px-4 py-2 border-b border-border bg-muted/50">
<div className="flex items-center gap-2">
<span className="text-sm font-medium text-foreground">
{state.filename || (state.contentType === 'code' ? '代码预览' : '内容预览')}
{state.filename ||
(state.contentType === "code" ? "代码预览" : "内容预览")}
</span>
{state.language && (
<span className="text-xs px-2 py-0.5 bg-muted rounded text-muted-foreground">
@@ -73,15 +77,35 @@ export const CanvasPanel: React.FC<CanvasPanelProps> = ({
<button
onClick={handleCopy}
className="p-1.5 text-muted-foreground hover:text-foreground hover:bg-muted rounded transition-colors"
title={copied ? '已复制' : '复制'}
title={copied ? "已复制" : "复制"}
>
{copied ? (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M5 13l4 4L19 7" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M5 13l4 4L19 7"
/>
</svg>
) : (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z"
/>
</svg>
)}
</button>
@@ -91,8 +115,18 @@ export const CanvasPanel: React.FC<CanvasPanelProps> = ({
className="p-1.5 text-muted-foreground hover:text-foreground hover:bg-muted rounded transition-colors"
title="下载"
>
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M4 16v1a3 3 0 003 3h10a3 3 0 003-3v-1m-4-4l-4 4m0 0l-4-4m4 4V4" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M4 16v1a3 3 0 003 3h10a3 3 0 003-3v-1m-4-4l-4 4m0 0l-4-4m4 4V4"
/>
</svg>
</button>
{/* 关闭按钮 */}
@@ -101,8 +135,18 @@ export const CanvasPanel: React.FC<CanvasPanelProps> = ({
className="p-1.5 text-muted-foreground hover:text-foreground hover:bg-muted rounded transition-colors"
title="关闭"
>
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M6 18L18 6M6 6l12 12" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M6 18L18 6M6 6l12 12"
/>
</svg>
</button>
</div>
@@ -110,14 +154,14 @@ export const CanvasPanel: React.FC<CanvasPanelProps> = ({
{/* 内容区域 */}
<div className="flex-1 overflow-auto">
{state.contentType === 'code' ? (
{state.contentType === "code" ? (
<CodePreview
code={state.content}
language={state.language || 'plaintext'}
language={state.language || "plaintext"}
isEditing={state.isEditing}
onContentChange={onContentChange}
/>
) : state.contentType === 'markdown' ? (
) : state.contentType === "markdown" ? (
<MarkdownPreview
content={state.content}
isEditing={state.isEditing}
@@ -6,7 +6,7 @@
* @requirements 4.1, 4.4
*/
import React, { useMemo } from 'react';
import React, { useMemo } from "react";
interface CodePreviewProps {
/** 代码内容 */
@@ -29,7 +29,7 @@ export const CodePreview: React.FC<CodePreviewProps> = ({
onContentChange,
}) => {
// 计算行号
const lines = useMemo(() => code.split('\n'), [code]);
const lines = useMemo(() => code.split("\n"), [code]);
const lineCount = lines.length;
const lineNumberWidth = String(lineCount).length * 10 + 20;
@@ -6,7 +6,7 @@
* @requirements 4.5
*/
import React, { useState } from 'react';
import React, { useState } from "react";
interface MarkdownPreviewProps {
/** Markdown 内容 */
@@ -24,27 +24,39 @@ interface MarkdownPreviewProps {
const renderMarkdown = (content: string): string => {
let html = content
// 转义 HTML
.replace(/&/g, '&amp;')
.replace(/</g, '&lt;')
.replace(/>/g, '&gt;')
.replace(/&/g, "&amp;")
.replace(/</g, "&lt;")
.replace(/>/g, "&gt;")
// 标题
.replace(/^### (.*$)/gm, '<h3 class="text-lg font-semibold mt-4 mb-2">$1</h3>')
.replace(/^## (.*$)/gm, '<h2 class="text-xl font-semibold mt-4 mb-2">$1</h2>')
.replace(
/^### (.*$)/gm,
'<h3 class="text-lg font-semibold mt-4 mb-2">$1</h3>',
)
.replace(
/^## (.*$)/gm,
'<h2 class="text-xl font-semibold mt-4 mb-2">$1</h2>',
)
.replace(/^# (.*$)/gm, '<h1 class="text-2xl font-bold mt-4 mb-2">$1</h1>')
// 粗体和斜体
.replace(/\*\*\*(.*?)\*\*\*/g, '<strong><em>$1</em></strong>')
.replace(/\*\*(.*?)\*\*/g, '<strong>$1</strong>')
.replace(/\*(.*?)\*/g, '<em>$1</em>')
.replace(/\*\*\*(.*?)\*\*\*/g, "<strong><em>$1</em></strong>")
.replace(/\*\*(.*?)\*\*/g, "<strong>$1</strong>")
.replace(/\*(.*?)\*/g, "<em>$1</em>")
// 行内代码
.replace(/`([^`]+)`/g, '<code class="px-1 py-0.5 bg-ink-100 rounded text-sm font-mono">$1</code>')
.replace(
/`([^`]+)`/g,
'<code class="px-1 py-0.5 bg-ink-100 rounded text-sm font-mono">$1</code>',
)
// 链接
.replace(/\[([^\]]+)\]\(([^)]+)\)/g, '<a href="$2" class="text-accent hover:underline" target="_blank">$1</a>')
.replace(
/\[([^\]]+)\]\(([^)]+)\)/g,
'<a href="$2" class="text-accent hover:underline" target="_blank">$1</a>',
)
// 列表
.replace(/^\s*[-*]\s+(.*$)/gm, '<li class="ml-4">$1</li>')
// 段落
.replace(/\n\n/g, '</p><p class="mb-2">')
// 换行
.replace(/\n/g, '<br/>');
.replace(/\n/g, "<br/>");
return `<p class="mb-2">${html}</p>`;
};
@@ -57,7 +69,7 @@ export const MarkdownPreview: React.FC<MarkdownPreviewProps> = ({
isEditing = false,
onContentChange,
}) => {
const [viewMode, setViewMode] = useState<'preview' | 'source'>('preview');
const [viewMode, setViewMode] = useState<"preview" | "source">("preview");
if (isEditing) {
return (
@@ -65,21 +77,21 @@ export const MarkdownPreview: React.FC<MarkdownPreviewProps> = ({
{/* 模式切换 */}
<div className="flex items-center gap-2 px-4 py-2 border-b border-ink-200 bg-ink-50">
<button
onClick={() => setViewMode('preview')}
onClick={() => setViewMode("preview")}
className={`px-3 py-1 text-sm rounded transition-colors ${
viewMode === 'preview'
? 'bg-accent text-white'
: 'text-ink-600 hover:bg-ink-100'
viewMode === "preview"
? "bg-accent text-white"
: "text-ink-600 hover:bg-ink-100"
}`}
>
预览
</button>
<button
onClick={() => setViewMode('source')}
onClick={() => setViewMode("source")}
className={`px-3 py-1 text-sm rounded transition-colors ${
viewMode === 'source'
? 'bg-accent text-white'
: 'text-ink-600 hover:bg-ink-100'
viewMode === "source"
? "bg-accent text-white"
: "text-ink-600 hover:bg-ink-100"
}`}
>
源码
@@ -87,7 +99,7 @@ export const MarkdownPreview: React.FC<MarkdownPreviewProps> = ({
</div>
{/* 内容区域 */}
<div className="flex-1 overflow-auto">
{viewMode === 'source' ? (
{viewMode === "source" ? (
<textarea
value={content}
onChange={(e) => onContentChange?.(e.target.value)}
+3 -3
View File
@@ -4,6 +4,6 @@
* @module components/general-chat/canvas
*/
export { CanvasPanel } from './CanvasPanel';
export { CodePreview } from './CodePreview';
export { MarkdownPreview } from './MarkdownPreview';
export { CanvasPanel } from "./CanvasPanel";
export { CodePreview } from "./CodePreview";
export { MarkdownPreview } from "./MarkdownPreview";
@@ -6,11 +6,11 @@
* @requirements 6.2, 6.7, 2.6, 9.2, 9.3, 9.5
*/
import React, { useState, useMemo } from 'react';
import type { Message, ContentBlock } from '../types';
import { CodeBlock } from './CodeBlock';
import { ErrorDisplay } from './ErrorDisplay';
import { ImageMessage } from './ImageMessage';
import React, { useState, useMemo } from "react";
import type { Message, ContentBlock } from "../types";
import { CodeBlock } from "./CodeBlock";
import { ErrorDisplay } from "./ErrorDisplay";
import { ImageMessage } from "./ImageMessage";
interface AssistantMessageProps {
/** 消息数据 */
@@ -45,15 +45,15 @@ const parseContent = (content: string): ContentBlock[] => {
if (match.index > lastIndex) {
const text = content.slice(lastIndex, match.index).trim();
if (text) {
blocks.push({ type: 'text', content: text });
blocks.push({ type: "text", content: text });
}
}
// 添加代码块
blocks.push({
type: 'code',
type: "code",
content: match[2].trim(),
language: match[1] || 'plaintext',
language: match[1] || "plaintext",
});
lastIndex = match.index + match[0].length;
@@ -63,13 +63,13 @@ const parseContent = (content: string): ContentBlock[] => {
if (lastIndex < content.length) {
const text = content.slice(lastIndex).trim();
if (text) {
blocks.push({ type: 'text', content: text });
blocks.push({ type: "text", content: text });
}
}
// 如果没有解析出任何块,返回整个内容作为文本
if (blocks.length === 0) {
blocks.push({ type: 'text', content });
blocks.push({ type: "text", content });
}
return blocks;
@@ -92,17 +92,23 @@ export const AssistantMessage: React.FC<AssistantMessageProps> = ({
const [copied, setCopied] = useState(false);
// 判断是否为错误状态
const isError = message.status === 'error' && message.error;
const isError = message.status === "error" && message.error;
// 使用流式内容或消息内容
const displayContent = isStreaming && streamingContent ? streamingContent : message.content;
const displayContent =
isStreaming && streamingContent ? streamingContent : message.content;
// 解析内容块(仅用于文本内容的 Markdown 解析)
const parsedContentBlocks = useMemo(() => parseContent(displayContent), [displayContent]);
const parsedContentBlocks = useMemo(
() => parseContent(displayContent),
[displayContent],
);
// 合并消息中的图片块和解析出的内容块
const allContentBlocks = useMemo(() => {
const imageBlocks = message.blocks.filter(block => block.type === 'image');
const imageBlocks = message.blocks.filter(
(block) => block.type === "image",
);
return [...imageBlocks, ...parsedContentBlocks];
}, [message.blocks, parsedContentBlocks]);
@@ -121,9 +127,21 @@ export const AssistantMessage: React.FC<AssistantMessageProps> = ({
<div className="max-w-[85%] flex flex-col gap-1">
{/* 头像和标签 */}
<div className="flex items-center gap-2 mb-1">
<div className={`w-6 h-6 rounded-full flex items-center justify-center ${isError ? 'bg-red-100' : 'bg-accent/10'}`}>
<svg className={`w-4 h-4 ${isError ? 'text-red-500' : 'text-accent'}`} fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M9.75 17L9 20l-1 1h8l-1-1-.75-3M3 13h18M5 17h14a2 2 0 002-2V5a2 2 0 00-2-2H5a2 2 0 00-2 2v10a2 2 0 002 2z" />
<div
className={`w-6 h-6 rounded-full flex items-center justify-center ${isError ? "bg-red-100" : "bg-accent/10"}`}
>
<svg
className={`w-4 h-4 ${isError ? "text-red-500" : "text-accent"}`}
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M9.75 17L9 20l-1 1h8l-1-1-.75-3M3 13h18M5 17h14a2 2 0 002-2V5a2 2 0 00-2-2H5a2 2 0 00-2 2v10a2 2 0 002 2z"
/>
</svg>
</div>
<span className="text-xs text-ink-500">AI 助手</span>
@@ -151,19 +169,23 @@ export const AssistantMessage: React.FC<AssistantMessageProps> = ({
)}
{/* 消息内容(非错误状态或有部分内容时显示) */}
{(!isError || displayContent || message.blocks.some(block => block.type === 'image')) && (
<div className={`bg-surface-secondary rounded-2xl rounded-tl-sm px-4 py-3 ${isError ? 'opacity-60' : ''}`}>
{(!isError ||
displayContent ||
message.blocks.some((block) => block.type === "image")) && (
<div
className={`bg-surface-secondary rounded-2xl rounded-tl-sm px-4 py-3 ${isError ? "opacity-60" : ""}`}
>
<div className="prose prose-sm max-w-none">
{allContentBlocks.map((block, index) => (
<div key={index} className="mb-2 last:mb-0">
{block.type === 'code' ? (
{block.type === "code" ? (
<CodeBlock
code={block.content}
language={block.language || 'plaintext'}
language={block.language || "plaintext"}
onCopy={() => onCopy(block.content)}
onOpenInCanvas={() => onOpenInCanvas(block)}
/>
) : block.type === 'image' ? (
) : block.type === "image" ? (
<ImageMessage block={block} />
) : (
<p className="text-sm text-ink-800 whitespace-pre-wrap break-words">
@@ -182,15 +204,35 @@ export const AssistantMessage: React.FC<AssistantMessageProps> = ({
<button
onClick={handleCopy}
className="p-1 text-ink-400 hover:text-ink-600 transition-colors"
title={copied ? '已复制' : '复制'}
title={copied ? "已复制" : "复制"}
>
{copied ? (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M5 13l4 4L19 7" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M5 13l4 4L19 7"
/>
</svg>
) : (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z"
/>
</svg>
)}
</button>
@@ -200,8 +242,18 @@ export const AssistantMessage: React.FC<AssistantMessageProps> = ({
className="p-1 text-ink-400 hover:text-ink-600 transition-colors"
title="重新生成"
>
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15"
/>
</svg>
</button>
)}
@@ -211,9 +263,7 @@ export const AssistantMessage: React.FC<AssistantMessageProps> = ({
{/* 元数据 */}
{message.metadata && !isStreaming && !isError && (
<div className="flex items-center gap-2 px-1 text-xs text-ink-400">
{message.metadata.model && (
<span>{message.metadata.model}</span>
)}
{message.metadata.model && <span>{message.metadata.model}</span>}
{message.metadata.tokens && (
<span>· {message.metadata.tokens} tokens</span>
)}
+114 -75
View File
@@ -6,17 +6,17 @@
* @requirements 3.1, 3.6, 5.4, 5.5, 2.6, 9.5, 10.2
*/
import React, { useState, useMemo, useCallback } from 'react';
import { Settings, AlertTriangle, Loader2 } from 'lucide-react';
import { useGeneralChatStore } from '../store/useGeneralChatStore';
import type { CanvasState, ContentBlock } from '../types';
import { DEFAULT_PAGINATION_STATE } from '../types';
import { MessageList } from './MessageList';
import { Inputbar } from '@/components/agent/chat/components/Inputbar';
import { CompactModelSelector } from './CompactModelSelector';
import { WorkflowStatusPanel } from '../components/WorkflowStatusPanel';
import { useConfiguredProviders } from '@/hooks/useConfiguredProviders';
import type { MessageImage } from '@/components/agent/chat/types';
import React, { useState, useMemo, useCallback } from "react";
import { Settings, AlertTriangle, Loader2 } from "lucide-react";
import { useGeneralChatStore } from "../store/useGeneralChatStore";
import type { CanvasState, ContentBlock } from "../types";
import { DEFAULT_PAGINATION_STATE } from "../types";
import { MessageList } from "./MessageList";
import { Inputbar } from "@/components/agent/chat/components/Inputbar";
import { CompactModelSelector } from "./CompactModelSelector";
import { WorkflowStatusPanel } from "../components/WorkflowStatusPanel";
import { useConfiguredProviders } from "@/hooks/useConfiguredProviders";
import type { MessageImage } from "@/components/agent/chat/types";
interface ChatPanelProps {
/** 当前会话 ID */
@@ -66,7 +66,9 @@ interface NoProviderPromptProps {
/**
* 无 Provider 提示组件
*/
const NoProviderPrompt: React.FC<NoProviderPromptProps> = ({ onNavigateToConfig }) => (
const NoProviderPrompt: React.FC<NoProviderPromptProps> = ({
onNavigateToConfig,
}) => (
<div className="flex flex-col items-center justify-center h-full text-center px-8">
<div className="w-20 h-20 mb-6 rounded-full bg-amber-100 flex items-center justify-center">
<AlertTriangle className="w-10 h-10 text-amber-600" />
@@ -75,8 +77,8 @@ const NoProviderPrompt: React.FC<NoProviderPromptProps> = ({ onNavigateToConfig
尚未配置 AI Provider
</h3>
<p className="text-sm text-ink-500 max-w-md mb-6 leading-relaxed">
要开始使用 AI 对话功能,您需要先配置至少一个 Provider 凭证。
支持 Kiro、Gemini、OpenAI、Claude 等多种 Provider。
要开始使用 AI 对话功能,您需要先配置至少一个 Provider 凭证。 支持
Kiro、Gemini、OpenAI、Claude 等多种 Provider。
</p>
{onNavigateToConfig && (
<button
@@ -102,9 +104,11 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
onOpenCanvas,
onNavigate,
}) => {
const [input, setInput] = useState('');
const [retryingMessageId, setRetryingMessageId] = useState<string | null>(null);
const [input, setInput] = useState("");
const [retryingMessageId, setRetryingMessageId] = useState<string | null>(
null,
);
// 直接从 store 获取状态
const messages = useGeneralChatStore((state) => state.messages);
const streaming = useGeneralChatStore((state) => state.streaming);
@@ -113,12 +117,20 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
const sendMessage = useGeneralChatStore((state) => state.sendMessage);
const stopGeneration = useGeneralChatStore((state) => state.stopGeneration);
const retryMessage = useGeneralChatStore((state) => state.retryMessage);
const loadMoreMessages = useGeneralChatStore((state) => state.loadMoreMessages);
const initializeWorkflow = useGeneralChatStore((state) => state.initializeWorkflow);
const getWorkflowManager = useGeneralChatStore((state) => state.getWorkflowManager);
const loadMoreMessages = useGeneralChatStore(
(state) => state.loadMoreMessages,
);
const initializeWorkflow = useGeneralChatStore(
(state) => state.initializeWorkflow,
);
const getWorkflowManager = useGeneralChatStore(
(state) => state.getWorkflowManager,
);
// 获取分页状态
const paginationState = sessionId ? pagination[sessionId] || DEFAULT_PAGINATION_STATE : DEFAULT_PAGINATION_STATE;
const paginationState = sessionId
? pagination[sessionId] || DEFAULT_PAGINATION_STATE
: DEFAULT_PAGINATION_STATE;
const { hasMoreMessages, isLoadingMore } = paginationState;
// 直接使用 useConfiguredProviders 检查是否有可用的 Provider
@@ -131,14 +143,16 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
// 获取当前会话的消息
const currentMessages = useMemo(
() => (sessionId ? messages[sessionId] || [] : []),
[sessionId, messages]
[sessionId, messages],
);
// 获取工作流状态
const workflowManager = sessionId ? getWorkflowManager(sessionId) : null;
const isWorkflowInitialized = workflowManager !== null;
const messageCount = currentMessages.length;
const visualOperationCount = currentMessages.filter(m => m.images && m.images.length > 0).length;
const visualOperationCount = currentMessages.filter(
(m) => m.images && m.images.length > 0,
).length;
// 处理加载更多消息
const handleLoadMore = useCallback(() => {
@@ -148,11 +162,14 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
}, [sessionId, hasMoreMessages, isLoadingMore, loadMoreMessages]);
// 处理初始化工作流
const handleInitializeWorkflow = useCallback(async (projectName: string, goal: string) => {
if (sessionId) {
await initializeWorkflow(sessionId, projectName, goal);
}
}, [sessionId, initializeWorkflow]);
const handleInitializeWorkflow = useCallback(
async (projectName: string, goal: string) => {
if (sessionId) {
await initializeWorkflow(sessionId, projectName, goal);
}
},
[sessionId, initializeWorkflow],
);
// 处理结束工作流
const handleFinalizeWorkflow = useCallback(async () => {
@@ -160,41 +177,53 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
try {
await workflowManager.finalizeWorkflow();
} catch (error) {
console.error('结束工作流失败:', error);
console.error("结束工作流失败:", error);
}
}
}, [sessionId, workflowManager]);
// 处理导航到 Provider 配置页面
const handleNavigateToProviderConfig = useCallback(() => {
onNavigate?.('provider-pool');
onNavigate?.("provider-pool");
}, [onNavigate]);
// 处理发送消息
const handleSend = useCallback(async (images?: MessageImage[], _webSearch?: boolean, _thinking?: boolean) => {
if (!sessionId || (!input.trim() && (!images || images.length === 0))) return;
const content = input.trim();
setInput('');
// 将 MessageImage 转换为 File 对象(用于 store)
let files: File[] | undefined;
if (images && images.length > 0) {
files = images.map((img, index) => {
// 将 base64 转换回 Blob,然后创建 File
const byteCharacters = atob(img.data);
const byteNumbers = new Array(byteCharacters.length);
for (let i = 0; i < byteCharacters.length; i++) {
byteNumbers[i] = byteCharacters.charCodeAt(i);
}
const byteArray = new Uint8Array(byteNumbers);
const blob = new Blob([byteArray], { type: img.mediaType });
return new File([blob], `image_${index}.${img.mediaType.split('/')[1]}`, { type: img.mediaType });
});
}
await sendMessage(content || '请分析这张图片', files);
}, [sessionId, input, sendMessage]);
const handleSend = useCallback(
async (
images?: MessageImage[],
_webSearch?: boolean,
_thinking?: boolean,
) => {
if (!sessionId || (!input.trim() && (!images || images.length === 0)))
return;
const content = input.trim();
setInput("");
// 将 MessageImage 转换为 File 对象(用于 store)
let files: File[] | undefined;
if (images && images.length > 0) {
files = images.map((img, index) => {
// 将 base64 转换回 Blob,然后创建 File
const byteCharacters = atob(img.data);
const byteNumbers = new Array(byteCharacters.length);
for (let i = 0; i < byteCharacters.length; i++) {
byteNumbers[i] = byteCharacters.charCodeAt(i);
}
const byteArray = new Uint8Array(byteNumbers);
const blob = new Blob([byteArray], { type: img.mediaType });
return new File(
[blob],
`image_${index}.${img.mediaType.split("/")[1]}`,
{ type: img.mediaType },
);
});
}
await sendMessage(content || "请分析这张图片", files);
},
[sessionId, input, sendMessage],
);
// 处理停止生成
const handleStop = useCallback(() => {
@@ -207,26 +236,32 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
}, []);
// 处理在画布中打开
const handleOpenInCanvas = useCallback((block: ContentBlock) => {
onOpenCanvas({
isOpen: true,
contentType: block.type === 'code' ? 'code' : 'markdown',
content: block.content,
language: block.language,
filename: block.filename,
isEditing: false,
});
}, [onOpenCanvas]);
const handleOpenInCanvas = useCallback(
(block: ContentBlock) => {
onOpenCanvas({
isOpen: true,
contentType: block.type === "code" ? "code" : "markdown",
content: block.content,
language: block.language,
filename: block.filename,
isEditing: false,
});
},
[onOpenCanvas],
);
// 处理重试消息
const handleRetry = useCallback(async (messageId: string) => {
setRetryingMessageId(messageId);
try {
await retryMessage(messageId);
} finally {
setRetryingMessageId(null);
}
}, [retryMessage]);
const handleRetry = useCallback(
async (messageId: string) => {
setRetryingMessageId(messageId);
try {
await retryMessage(messageId);
} finally {
setRetryingMessageId(null);
}
},
[retryMessage],
);
// Provider 加载中时显示加载状态
if (isProviderLoading) {
@@ -244,7 +279,11 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
if (!hasAvailableProvider) {
return (
<div className="flex flex-col h-full bg-surface">
<NoProviderPrompt onNavigateToConfig={onNavigate ? handleNavigateToProviderConfig : undefined} />
<NoProviderPrompt
onNavigateToConfig={
onNavigate ? handleNavigateToProviderConfig : undefined
}
/>
</div>
);
}
@@ -302,8 +341,8 @@ export const ChatPanel: React.FC<ChatPanelProps> = ({
</div>
{/* 工作流状态面板 */}
<WorkflowStatusPanel
sessionId={sessionId || ''}
<WorkflowStatusPanel
sessionId={sessionId || ""}
isWorkflowActive={workflowEnabled}
isWorkflowInitialized={isWorkflowInitialized}
messageCount={messageCount}
+62 -32
View File
@@ -6,7 +6,7 @@
* @requirements 2.4, 6.3
*/
import React, { useState } from 'react';
import React, { useState } from "react";
interface CodeBlockProps {
/** 代码内容 */
@@ -26,30 +26,30 @@ interface CodeBlockProps {
*/
const getLanguageDisplayName = (lang: string): string => {
const names: Record<string, string> = {
javascript: 'JavaScript',
typescript: 'TypeScript',
python: 'Python',
rust: 'Rust',
go: 'Go',
java: 'Java',
cpp: 'C++',
c: 'C',
csharp: 'C#',
ruby: 'Ruby',
php: 'PHP',
swift: 'Swift',
kotlin: 'Kotlin',
html: 'HTML',
css: 'CSS',
scss: 'SCSS',
json: 'JSON',
yaml: 'YAML',
xml: 'XML',
sql: 'SQL',
bash: 'Bash',
shell: 'Shell',
markdown: 'Markdown',
plaintext: 'Text',
javascript: "JavaScript",
typescript: "TypeScript",
python: "Python",
rust: "Rust",
go: "Go",
java: "Java",
cpp: "C++",
c: "C",
csharp: "C#",
ruby: "Ruby",
php: "PHP",
swift: "Swift",
kotlin: "Kotlin",
html: "HTML",
css: "CSS",
scss: "SCSS",
json: "JSON",
yaml: "YAML",
xml: "XML",
sql: "SQL",
bash: "Bash",
shell: "Shell",
markdown: "Markdown",
plaintext: "Text",
};
return names[lang.toLowerCase()] || lang;
};
@@ -87,15 +87,35 @@ export const CodeBlock: React.FC<CodeBlockProps> = ({
<button
onClick={handleCopy}
className="p-1 text-ink-500 hover:text-ink-700 transition-colors"
title={copied ? '已复制' : '复制代码'}
title={copied ? "已复制" : "复制代码"}
>
{copied ? (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M5 13l4 4L19 7" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M5 13l4 4L19 7"
/>
</svg>
) : (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z"
/>
</svg>
)}
</button>
@@ -105,8 +125,18 @@ export const CodeBlock: React.FC<CodeBlockProps> = ({
className="p-1 text-ink-500 hover:text-ink-700 transition-colors"
title="在画布中打开"
>
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M10 6H6a2 2 0 00-2 2v10a2 2 0 002 2h10a2 2 0 002-2v-4M14 4h6m0 0v6m0-6L10 14" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M10 6H6a2 2 0 00-2 2v10a2 2 0 002 2h10a2 2 0 002-2v-4M14 4h6m0 0v6m0-6L10 14"
/>
</svg>
</button>
</div>
@@ -8,11 +8,11 @@
* @requirements 5.4
*/
import React, { useState, useRef, useEffect } from 'react';
import { ChevronDown, Loader2, AlertCircle, Check } from 'lucide-react';
import { cn } from '@/lib/utils';
import { useProvider } from '../hooks/useProvider';
import { getProviderLabel } from '@/lib/constants/providerMappings';
import React, { useState, useRef, useEffect } from "react";
import { ChevronDown, Loader2, AlertCircle, Check } from "lucide-react";
import { cn } from "@/lib/utils";
import { useProvider } from "../hooks/useProvider";
import { getProviderLabel } from "@/lib/constants/providerMappings";
// ============================================================================
// 类型定义
@@ -38,7 +38,11 @@ interface DropdownMenuProps {
/**
* 下拉菜单容器
*/
const DropdownMenu: React.FC<DropdownMenuProps> = ({ isOpen, onClose, children }) => {
const DropdownMenu: React.FC<DropdownMenuProps> = ({
isOpen,
onClose,
children,
}) => {
const menuRef = useRef<HTMLDivElement>(null);
// 点击外部关闭
@@ -50,11 +54,11 @@ const DropdownMenu: React.FC<DropdownMenuProps> = ({ isOpen, onClose, children }
};
if (isOpen) {
document.addEventListener('mousedown', handleClickOutside);
document.addEventListener("mousedown", handleClickOutside);
}
return () => {
document.removeEventListener('mousedown', handleClickOutside);
document.removeEventListener("mousedown", handleClickOutside);
};
}, [isOpen, onClose]);
@@ -106,20 +110,20 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
// 获取显示文本
const displayText = React.useMemo(() => {
if (isLoading) {
return '加载中...';
return "加载中...";
}
if (!hasAvailableProvider) {
return '未配置 Provider';
return "未配置 Provider";
}
if (!selectedProvider) {
return '选择模型';
return "选择模型";
}
const providerLabel = getProviderLabel(selectedProvider.key);
if (!selectedModelId) {
return providerLabel;
}
// 简化模型名称显示
const shortModelId = selectedModelId.split('/').pop() || selectedModelId;
const shortModelId = selectedModelId.split("/").pop() || selectedModelId;
return `${providerLabel} / ${shortModelId}`;
}, [isLoading, hasAvailableProvider, selectedProvider, selectedModelId]);
@@ -137,7 +141,12 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
// 无 Provider 时显示提示
if (!hasAvailableProvider && !isLoading) {
return (
<div className={cn('flex items-center gap-2 text-sm text-amber-600', className)}>
<div
className={cn(
"flex items-center gap-2 text-sm text-amber-600",
className,
)}
>
<AlertCircle className="h-4 w-4" />
<span>请先配置 Provider 凭证</span>
</div>
@@ -145,18 +154,18 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
}
return (
<div className={cn('relative', className)}>
<div className={cn("relative", className)}>
{/* 触发按钮 */}
<button
type="button"
onClick={() => !disabled && setIsOpen(!isOpen)}
disabled={disabled || isLoading}
className={cn(
'flex items-center gap-2 px-3 py-1.5 text-sm rounded-md transition-colors',
'border border-ink-200 hover:border-ink-300 hover:bg-ink-50',
'focus:outline-none focus:ring-2 focus:ring-accent/50',
disabled && 'opacity-50 cursor-not-allowed',
isOpen && 'border-accent bg-accent/5'
"flex items-center gap-2 px-3 py-1.5 text-sm rounded-md transition-colors",
"border border-ink-200 hover:border-ink-300 hover:bg-ink-50",
"focus:outline-none focus:ring-2 focus:ring-accent/50",
disabled && "opacity-50 cursor-not-allowed",
isOpen && "border-accent bg-accent/5",
)}
>
{isLoading ? (
@@ -169,8 +178,8 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
<span className="truncate max-w-[200px]">{displayText}</span>
<ChevronDown
className={cn(
'h-4 w-4 text-ink-400 transition-transform',
isOpen && 'rotate-180'
"h-4 w-4 text-ink-400 transition-transform",
isOpen && "rotate-180",
)}
/>
</button>
@@ -190,10 +199,10 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
type="button"
onClick={() => handleProviderSelect(provider.key)}
className={cn(
'w-full flex items-center gap-2 px-2 py-1.5 text-xs rounded transition-colors',
"w-full flex items-center gap-2 px-2 py-1.5 text-xs rounded transition-colors",
selectedProvider?.key === provider.key
? 'bg-accent text-white'
: 'hover:bg-ink-100 text-ink-700'
? "bg-accent text-white"
: "hover:bg-ink-100 text-ink-700",
)}
>
<span className="truncate">{provider.label}</span>
@@ -208,7 +217,7 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
<h4 className="text-xs font-medium text-ink-600">
{selectedProvider
? `${getProviderLabel(selectedProvider.key)} 模型`
: '请选择 Provider'}
: "请选择 Provider"}
</h4>
</div>
<div className="flex-1 overflow-y-auto p-1">
@@ -220,7 +229,7 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
availableModelIds.map((modelId) => {
const isSelected = selectedModelId === modelId;
// 简化模型名称显示
const displayName = modelId.split('/').pop() || modelId;
const displayName = modelId.split("/").pop() || modelId;
return (
<button
@@ -228,14 +237,16 @@ export const CompactModelSelector: React.FC<CompactModelSelectorProps> = ({
type="button"
onClick={() => handleModelSelect(modelId)}
className={cn(
'w-full flex items-center justify-between px-2 py-1.5 text-xs rounded transition-colors',
"w-full flex items-center justify-between px-2 py-1.5 text-xs rounded transition-colors",
isSelected
? 'bg-accent/10 text-accent border border-accent/30'
: 'hover:bg-ink-100 text-ink-700 border border-transparent'
? "bg-accent/10 text-accent border border-accent/30"
: "hover:bg-ink-100 text-ink-700 border border-transparent",
)}
>
<span className="truncate">{displayName}</span>
{isSelected && <Check className="h-3 w-3 flex-shrink-0" />}
{isSelected && (
<Check className="h-3 w-3 flex-shrink-0" />
)}
</button>
);
})
@@ -9,8 +9,8 @@
* @requirements 9.4
*/
import React, { Component, ErrorInfo, ReactNode } from 'react';
import { AlertTriangle, RefreshCw, Home, Bug } from 'lucide-react';
import React, { Component, ErrorInfo, ReactNode } from "react";
import { AlertTriangle, RefreshCw, Home, Bug } from "lucide-react";
// ============================================================================
// 类型定义
@@ -55,7 +55,10 @@ interface ErrorBoundaryState {
*
* @requirements 9.4
*/
export class ErrorBoundary extends Component<ErrorBoundaryProps, ErrorBoundaryState> {
export class ErrorBoundary extends Component<
ErrorBoundaryProps,
ErrorBoundaryState
> {
constructor(props: ErrorBoundaryProps) {
super(props);
this.state = {
@@ -87,7 +90,7 @@ export class ErrorBoundary extends Component<ErrorBoundaryProps, ErrorBoundarySt
this.setState({ errorInfo });
// 记录错误日志
const logContext = componentName ? `[${componentName}]` : '[ErrorBoundary]';
const logContext = componentName ? `[${componentName}]` : "[ErrorBoundary]";
console.error(`${logContext} 捕获到渲染错误:`, error);
console.error(`${logContext} 组件堆栈:`, errorInfo.componentStack);
@@ -116,7 +119,7 @@ export class ErrorBoundary extends Component<ErrorBoundaryProps, ErrorBoundarySt
componentName: this.props.componentName,
};
console.info('[ErrorBoundary] 错误报告:', errorReport);
console.info("[ErrorBoundary] 错误报告:", errorReport);
}
/**
@@ -141,7 +144,7 @@ export class ErrorBoundary extends Component<ErrorBoundaryProps, ErrorBoundarySt
* 返回首页
*/
private handleGoHome = (): void => {
window.location.href = '/';
window.location.href = "/";
};
render(): ReactNode {
@@ -204,9 +207,7 @@ const DefaultFallbackUI: React.FC<DefaultFallbackUIProps> = ({
</div>
{/* 错误标题 */}
<h2 className="text-xl font-semibold text-ink-900 mb-3">
内容渲染失败
</h2>
<h2 className="text-xl font-semibold text-ink-900 mb-3">内容渲染失败</h2>
{/* 错误描述 */}
<p className="text-sm text-ink-500 text-center max-w-md mb-6 leading-relaxed">
@@ -252,7 +253,7 @@ const DefaultFallbackUI: React.FC<DefaultFallbackUIProps> = ({
className="flex items-center gap-2 text-xs text-ink-400 hover:text-ink-600 transition-colors mx-auto"
>
<Bug className="w-3 h-3" />
{showDetails ? '隐藏错误详情' : '查看错误详情'}
{showDetails ? "隐藏错误详情" : "查看错误详情"}
</button>
{showDetails && (
@@ -264,12 +265,16 @@ const DefaultFallbackUI: React.FC<DefaultFallbackUIProps> = ({
<div className="mb-3">
<p className="text-xs font-medium text-ink-600 mb-1">错误信息:</p>
<p className="text-xs text-red-600 font-mono break-all">{error.message}</p>
<p className="text-xs text-red-600 font-mono break-all">
{error.message}
</p>
</div>
{error.stack && (
<div className="mb-3">
<p className="text-xs font-medium text-ink-600 mb-1">堆栈跟踪:</p>
<p className="text-xs font-medium text-ink-600 mb-1">
堆栈跟踪:
</p>
<pre className="text-xs text-ink-500 font-mono overflow-x-auto whitespace-pre-wrap break-all max-h-32 overflow-y-auto">
{error.stack}
</pre>
@@ -278,7 +283,9 @@ const DefaultFallbackUI: React.FC<DefaultFallbackUIProps> = ({
{errorInfo?.componentStack && (
<div>
<p className="text-xs font-medium text-ink-600 mb-1">组件堆栈:</p>
<p className="text-xs font-medium text-ink-600 mb-1">
组件堆栈:
</p>
<pre className="text-xs text-ink-500 font-mono overflow-x-auto whitespace-pre-wrap break-all max-h-32 overflow-y-auto">
{errorInfo.componentStack}
</pre>
@@ -8,9 +8,16 @@
* @requirements 2.6, 9.2, 9.3, 9.5
*/
import React, { useState, useEffect } from 'react';
import { AlertCircle, RefreshCw, WifiOff, Clock, AlertTriangle, XCircle } from 'lucide-react';
import type { ErrorInfo, ErrorCode } from '../types';
import React, { useState, useEffect } from "react";
import {
AlertCircle,
RefreshCw,
WifiOff,
Clock,
AlertTriangle,
XCircle,
} from "lucide-react";
import type { ErrorInfo, ErrorCode } from "../types";
interface ErrorDisplayProps {
/** 错误信息 */
@@ -25,22 +32,22 @@ interface ErrorDisplayProps {
* 根据错误代码获取对应的图标
*/
const getErrorIcon = (code: ErrorCode): React.ReactNode => {
const iconClass = 'w-5 h-5';
const iconClass = "w-5 h-5";
switch (code) {
case 'NETWORK_ERROR':
case "NETWORK_ERROR":
return <WifiOff className={iconClass} />;
case 'TIMEOUT':
case "TIMEOUT":
return <Clock className={iconClass} />;
case 'RATE_LIMIT':
case "RATE_LIMIT":
return <AlertTriangle className={iconClass} />;
case 'TOKEN_LIMIT':
case "TOKEN_LIMIT":
return <XCircle className={iconClass} />;
case 'AUTH_ERROR':
case "AUTH_ERROR":
return <AlertCircle className={iconClass} />;
case 'SERVER_ERROR':
case 'PROVIDER_ERROR':
case 'UNKNOWN_ERROR':
case "SERVER_ERROR":
case "PROVIDER_ERROR":
case "UNKNOWN_ERROR":
default:
return <AlertCircle className={iconClass} />;
}
@@ -51,19 +58,19 @@ const getErrorIcon = (code: ErrorCode): React.ReactNode => {
*/
const getErrorColorClass = (code: ErrorCode): string => {
switch (code) {
case 'NETWORK_ERROR':
case 'TIMEOUT':
return 'bg-amber-50 border-amber-200 text-amber-800';
case 'RATE_LIMIT':
return 'bg-orange-50 border-orange-200 text-orange-800';
case 'TOKEN_LIMIT':
case 'AUTH_ERROR':
return 'bg-red-50 border-red-200 text-red-800';
case 'SERVER_ERROR':
case 'PROVIDER_ERROR':
case 'UNKNOWN_ERROR':
case "NETWORK_ERROR":
case "TIMEOUT":
return "bg-amber-50 border-amber-200 text-amber-800";
case "RATE_LIMIT":
return "bg-orange-50 border-orange-200 text-orange-800";
case "TOKEN_LIMIT":
case "AUTH_ERROR":
return "bg-red-50 border-red-200 text-red-800";
case "SERVER_ERROR":
case "PROVIDER_ERROR":
case "UNKNOWN_ERROR":
default:
return 'bg-red-50 border-red-200 text-red-700';
return "bg-red-50 border-red-200 text-red-700";
}
};
@@ -81,9 +88,13 @@ export const ErrorDisplay: React.FC<ErrorDisplayProps> = ({
// 处理 rate limit 倒计时
useEffect(() => {
if (error.code === 'RATE_LIMIT' && error.retryAfter && error.retryAfter > 0) {
if (
error.code === "RATE_LIMIT" &&
error.retryAfter &&
error.retryAfter > 0
) {
setCountdown(error.retryAfter);
const timer = setInterval(() => {
setCountdown((prev) => {
if (prev === null || prev <= 1) {
@@ -99,27 +110,24 @@ export const ErrorDisplay: React.FC<ErrorDisplayProps> = ({
}, [error.code, error.retryAfter]);
const colorClass = getErrorColorClass(error.code);
const canRetry = error.retryable && !isRetrying && (countdown === null || countdown === 0);
const canRetry =
error.retryable && !isRetrying && (countdown === null || countdown === 0);
return (
<div className={`flex items-start gap-3 p-3 rounded-lg border ${colorClass}`}>
<div
className={`flex items-start gap-3 p-3 rounded-lg border ${colorClass}`}
>
{/* 错误图标 */}
<div className="flex-shrink-0 mt-0.5">
{getErrorIcon(error.code)}
</div>
<div className="flex-shrink-0 mt-0.5">{getErrorIcon(error.code)}</div>
{/* 错误内容 */}
<div className="flex-1 min-w-0">
{/* 错误消息 */}
<p className="text-sm font-medium">
{error.message}
</p>
<p className="text-sm font-medium">{error.message}</p>
{/* 倒计时提示 */}
{countdown !== null && countdown > 0 && (
<p className="text-xs mt-1 opacity-80">
{countdown} 秒后可重试
</p>
<p className="text-xs mt-1 opacity-80">{countdown} 秒后可重试</p>
)}
{/* 详细信息(可选显示) */}
@@ -143,15 +151,24 @@ export const ErrorDisplay: React.FC<ErrorDisplayProps> = ({
className={`
flex-shrink-0 flex items-center gap-1.5 px-3 py-1.5 rounded-md text-sm font-medium
transition-all duration-200
${canRetry
? 'bg-white/80 hover:bg-white shadow-sm cursor-pointer'
: 'bg-white/40 cursor-not-allowed opacity-50'
${
canRetry
? "bg-white/80 hover:bg-white shadow-sm cursor-pointer"
: "bg-white/40 cursor-not-allowed opacity-50"
}
`}
title={canRetry ? '点击重试' : (countdown ? `${countdown}秒后可重试` : '正在重试...')}
title={
canRetry
? "点击重试"
: countdown
? `${countdown}秒后可重试`
: "正在重试..."
}
>
<RefreshCw className={`w-4 h-4 ${isRetrying ? 'animate-spin' : ''}`} />
{isRetrying ? '重试中...' : '重试'}
<RefreshCw
className={`w-4 h-4 ${isRetrying ? "animate-spin" : ""}`}
/>
{isRetrying ? "重试中..." : "重试"}
</button>
)}
</div>
@@ -6,9 +6,9 @@
* 用于在消息中显示图片内容,支持点击放大预览
*/
import React, { useState, useCallback } from 'react';
import { X, ZoomIn, Download } from 'lucide-react';
import type { ContentBlock } from '../types';
import React, { useState, useCallback } from "react";
import { X, ZoomIn, Download } from "lucide-react";
import type { ContentBlock } from "../types";
interface ImageMessageProps {
/** 图片内容块 */
@@ -38,9 +38,9 @@ const ImagePreviewModal: React.FC<ImagePreviewModalProps> = ({
onClose,
}) => {
const handleDownload = useCallback(() => {
const link = document.createElement('a');
const link = document.createElement("a");
link.href = src;
link.download = filename || 'image.png';
link.download = filename || "image.png";
document.body.appendChild(link);
link.click();
document.body.removeChild(link);
@@ -51,22 +51,22 @@ const ImagePreviewModal: React.FC<ImagePreviewModalProps> = ({
return (
<div className="fixed inset-0 z-50 flex items-center justify-center bg-black/80 backdrop-blur-sm">
{/* 背景遮罩 */}
<div
className="absolute inset-0"
<div
className="absolute inset-0"
onClick={onClose}
role="button"
tabIndex={0}
onKeyDown={(e) => {
if (e.key === 'Escape') onClose();
if (e.key === "Escape") onClose();
}}
/>
{/* 图片容器 */}
<div className="relative max-w-[90vw] max-h-[90vh] bg-white rounded-lg shadow-2xl">
{/* 工具栏 */}
<div className="absolute top-0 left-0 right-0 z-10 flex items-center justify-between p-4 bg-gradient-to-b from-black/50 to-transparent">
<div className="text-white text-sm font-medium">
{filename || '图片预览'}
{filename || "图片预览"}
</div>
<div className="flex items-center gap-2">
<button
@@ -87,13 +87,13 @@ const ImagePreviewModal: React.FC<ImagePreviewModalProps> = ({
</button>
</div>
</div>
{/* 图片 */}
<img
src={src}
alt={filename || '图片'}
alt={filename || "图片"}
className="max-w-full max-h-[90vh] object-contain rounded-lg"
style={{ minWidth: '300px', minHeight: '200px' }}
style={{ minWidth: "300px", minHeight: "200px" }}
/>
</div>
</div>
@@ -103,9 +103,12 @@ const ImagePreviewModal: React.FC<ImagePreviewModalProps> = ({
/**
* 图片消息组件
*/
export const ImageMessage: React.FC<ImageMessageProps> = ({ block, onClick }) => {
export const ImageMessage: React.FC<ImageMessageProps> = ({
block,
onClick,
}) => {
const [isPreviewOpen, setIsPreviewOpen] = useState(false);
const handleImageClick = useCallback(() => {
setIsPreviewOpen(true);
onClick?.();
@@ -116,15 +119,18 @@ export const ImageMessage: React.FC<ImageMessageProps> = ({ block, onClick }) =>
}, []);
// 处理键盘事件
const handleKeyDown = useCallback((e: React.KeyboardEvent) => {
if (e.key === 'Enter' || e.key === ' ') {
e.preventDefault();
handleImageClick();
}
}, [handleImageClick]);
const handleKeyDown = useCallback(
(e: React.KeyboardEvent) => {
if (e.key === "Enter" || e.key === " ") {
e.preventDefault();
handleImageClick();
}
},
[handleImageClick],
);
// 确保是图片类型
if (block.type !== 'image') {
if (block.type !== "image") {
return null;
}
@@ -142,11 +148,11 @@ export const ImageMessage: React.FC<ImageMessageProps> = ({ block, onClick }) =>
>
<img
src={block.content}
alt={block.filename || '图片'}
alt={block.filename || "图片"}
className="w-full h-auto max-h-64 object-cover transition-transform duration-200 group-hover:scale-105"
loading="lazy"
/>
{/* 悬浮遮罩 */}
<div className="absolute inset-0 bg-black/0 group-hover:bg-black/20 transition-colors duration-200 flex items-center justify-center">
<div className="opacity-0 group-hover:opacity-100 transition-opacity duration-200 bg-white/90 rounded-full p-2">
@@ -154,7 +160,7 @@ export const ImageMessage: React.FC<ImageMessageProps> = ({ block, onClick }) =>
</div>
</div>
</div>
{/* 文件名标签 */}
{block.filename && (
<div className="mt-2 text-xs text-ink-500 truncate">
@@ -174,4 +180,4 @@ export const ImageMessage: React.FC<ImageMessageProps> = ({ block, onClick }) =>
);
};
export default ImageMessage;
export default ImageMessage;
@@ -6,11 +6,11 @@
* @requirements 6.1, 6.2, 6.7, 2.6, 9.5
*/
import React from 'react';
import React from "react";
import type { Message, ContentBlock } from '../types';
import { UserMessage } from './UserMessage';
import { AssistantMessage } from './AssistantMessage';
import type { Message, ContentBlock } from "../types";
import { UserMessage } from "./UserMessage";
import { AssistantMessage } from "./AssistantMessage";
interface MessageItemProps {
/** 消息数据 */
@@ -44,16 +44,11 @@ export const MessageItem: React.FC<MessageItemProps> = ({
onRetry,
isRetrying,
}) => {
if (message.role === 'user') {
return (
<UserMessage
message={message}
onCopy={onCopy}
/>
);
if (message.role === "user") {
return <UserMessage message={message} onCopy={onCopy} />;
}
if (message.role === 'assistant') {
if (message.role === "assistant") {
return (
<AssistantMessage
message={message}
@@ -10,18 +10,18 @@
* - 支持动态高度消息(聊天消息高度不固定)
* - 保持自动滚动到底部功能
* - 流式消息更新时滚动行为正常
*
*
* 分页加载实现说明:
* - 滚动到顶部时自动加载更多历史消息
* - 每次加载 20-50 条消息(可配置)
* - 与虚拟滚动兼容
*/
import React, { useRef, useEffect, useCallback } from 'react';
import { useVirtualizer } from '@tanstack/react-virtual';
import React, { useRef, useEffect, useCallback } from "react";
import { useVirtualizer } from "@tanstack/react-virtual";
import type { Message, ContentBlock } from '../types';
import { MessageItem } from './MessageItem';
import type { Message, ContentBlock } from "../types";
import { MessageItem } from "./MessageItem";
/** 虚拟滚动启用阈值 - 超过此数量启用虚拟滚动 */
const VIRTUAL_SCROLL_THRESHOLD = 50;
@@ -66,7 +66,7 @@ interface MessageListProps {
* 当消息数量超过阈值时启用虚拟滚动,
* 确保大量消息时 DOM 元素数量受控。
* 支持滚动到顶部时加载更多历史消息。
*
*
* @requirements 10.1, 10.2
*/
export const MessageList: React.FC<MessageListProps> = ({
@@ -114,12 +114,12 @@ export const MessageList: React.FC<MessageListProps> = ({
// 滚动到底部
const scrollToBottom = useCallback(
(behavior: 'smooth' | 'auto' = 'smooth') => {
(behavior: "smooth" | "auto" = "smooth") => {
if (useVirtualScroll) {
virtualizer.scrollToIndex(messages.length - 1, {
align: 'end',
align: "end",
// 使用类型断言解决 @tanstack/react-virtual 与 DOM ScrollBehavior 类型不兼容问题
behavior: behavior as 'smooth' | 'auto',
behavior: behavior as "smooth" | "auto",
});
} else {
const parent = parentRef.current;
@@ -131,7 +131,7 @@ export const MessageList: React.FC<MessageListProps> = ({
}
}
},
[useVirtualScroll, virtualizer, messages.length]
[useVirtualScroll, virtualizer, messages.length],
);
// 监听滚动事件,更新是否在底部的状态
@@ -141,7 +141,7 @@ export const MessageList: React.FC<MessageListProps> = ({
const handleScroll = () => {
isAtBottomRef.current = checkIsAtBottom();
// 检查是否滚动到顶部,触发加载更多
if (
onLoadMore &&
@@ -153,16 +153,17 @@ export const MessageList: React.FC<MessageListProps> = ({
}
};
parent.addEventListener('scroll', handleScroll, { passive: true });
return () => parent.removeEventListener('scroll', handleScroll);
parent.addEventListener("scroll", handleScroll, { passive: true });
return () => parent.removeEventListener("scroll", handleScroll);
}, [checkIsAtBottom, onLoadMore, hasMoreMessages, isLoadingMore]);
// 新消息到达时自动滚动到底部
useEffect(() => {
const messageCountChanged = messages.length !== prevMessageCountRef.current;
const firstMessageId = messages.length > 0 ? messages[0].id : null;
const firstMessageChanged = firstMessageId !== prevFirstMessageIdRef.current;
const firstMessageChanged =
firstMessageId !== prevFirstMessageIdRef.current;
// 更新引用
const prevCount = prevMessageCountRef.current;
prevMessageCountRef.current = messages.length;
@@ -177,8 +178,8 @@ export const MessageList: React.FC<MessageListProps> = ({
const newMessagesCount = messages.length - prevCount;
requestAnimationFrame(() => {
virtualizer.scrollToIndex(newMessagesCount, {
align: 'start',
behavior: 'auto' as 'smooth' | 'auto',
align: "start",
behavior: "auto" as "smooth" | "auto",
});
});
}
@@ -189,16 +190,22 @@ export const MessageList: React.FC<MessageListProps> = ({
if (messageCountChanged && isAtBottomRef.current) {
// 使用 requestAnimationFrame 确保 DOM 更新后再滚动
requestAnimationFrame(() => {
scrollToBottom('smooth');
scrollToBottom("smooth");
});
}
}, [messages.length, messages, scrollToBottom, useVirtualScroll, virtualizer]);
}, [
messages.length,
messages,
scrollToBottom,
useVirtualScroll,
virtualizer,
]);
// 流式内容更新时保持滚动到底部
useEffect(() => {
if (isStreaming && isAtBottomRef.current) {
requestAnimationFrame(() => {
scrollToBottom('auto');
scrollToBottom("auto");
});
}
}, [isStreaming, partialContent, scrollToBottom]);
@@ -207,7 +214,8 @@ export const MessageList: React.FC<MessageListProps> = ({
const renderMessageItem = useCallback(
(message: Message, index: number) => {
const isLast = index === messages.length - 1;
const isStreamingMessage = isLast && isStreaming && message.role === 'assistant';
const isStreamingMessage =
isLast && isStreaming && message.role === "assistant";
const isRetrying = retryingMessageId === message.id;
return (
@@ -218,7 +226,9 @@ export const MessageList: React.FC<MessageListProps> = ({
streamingContent={isStreamingMessage ? partialContent : undefined}
onCopy={onCopy}
onOpenInCanvas={onOpenInCanvas}
onRegenerate={onRegenerate ? () => onRegenerate(message.id) : undefined}
onRegenerate={
onRegenerate ? () => onRegenerate(message.id) : undefined
}
onRetry={onRetry ? () => onRetry(message.id) : undefined}
isRetrying={isRetrying}
/>
@@ -233,7 +243,7 @@ export const MessageList: React.FC<MessageListProps> = ({
onRegenerate,
onRetry,
retryingMessageId,
]
],
);
// 非虚拟滚动模式(消息数量较少时)
@@ -256,7 +266,9 @@ export const MessageList: React.FC<MessageListProps> = ({
{/* 没有更多消息提示 */}
{!hasMoreMessages && messages.length > 0 && (
<div className="flex justify-center py-2">
<span className="text-xs text-muted-foreground">已加载全部消息</span>
<span className="text-xs text-muted-foreground">
已加载全部消息
</span>
</div>
)}
{messages.map((message, index) => renderMessageItem(message, index))}
@@ -292,7 +304,9 @@ export const MessageList: React.FC<MessageListProps> = ({
{/* 没有更多消息提示 */}
{!hasMoreMessages && messages.length > 0 && !isLoadingMore && (
<div className="absolute top-0 left-0 right-0 flex justify-center py-2 z-10">
<span className="text-xs text-muted-foreground">已加载全部消息</span>
<span className="text-xs text-muted-foreground">
已加载全部消息
</span>
</div>
)}
{/* 虚拟滚动内容 */}
@@ -107,7 +107,7 @@ export function detectLists(content: string): {
* @returns 代码块信息数组
*/
export function detectCodeBlocks(
content: string
content: string,
): Array<{ language: string; hasContent: boolean }> {
const codeBlockRegex = /```(\w+)?\n([\s\S]*?)```/g;
const codeBlocks: Array<{ language: string; hasContent: boolean }> = [];
@@ -206,7 +206,7 @@ const unorderedListArbitrary: fc.Arbitrary<string> = fc
fc
.string({ minLength: 1, maxLength: 30 })
.filter((s) => s.trim().length > 0 && !s.includes("\n")),
{ minLength: 1, maxLength: 5 }
{ minLength: 1, maxLength: 5 },
)
.map((items) => items.map((item) => `- ${item.trim()}`).join("\n"));
@@ -218,10 +218,10 @@ const orderedListArbitrary: fc.Arbitrary<string> = fc
fc
.string({ minLength: 1, maxLength: 30 })
.filter((s) => s.trim().length > 0 && !s.includes("\n")),
{ minLength: 1, maxLength: 5 }
{ minLength: 1, maxLength: 5 },
)
.map((items) =>
items.map((item, index) => `${index + 1}. ${item.trim()}`).join("\n")
items.map((item, index) => `${index + 1}. ${item.trim()}`).join("\n"),
);
/**
@@ -237,7 +237,7 @@ const codeBlockArbitrary: fc.Arbitrary<{ markdown: string; language: string }> =
"rust",
"go",
"java",
"plaintext"
"plaintext",
),
code: fc
.string({ minLength: 1, maxLength: 100 })
@@ -263,9 +263,9 @@ const tableArbitrary: fc.Arbitrary<string> = fc
fc
.string({ minLength: 1, maxLength: 10 })
.filter((s) => s.trim().length > 0 && !s.includes("|")),
{ minLength: cols, maxLength: cols }
{ minLength: cols, maxLength: cols },
),
{ minLength: rows + 1, maxLength: rows + 1 }
{ minLength: rows + 1, maxLength: rows + 1 },
)
.map((data) => {
const header = `| ${data[0].map((c) => c.trim()).join(" | ")} |`;
@@ -274,7 +274,7 @@ const tableArbitrary: fc.Arbitrary<string> = fc
.slice(1)
.map((row) => `| ${row.map((c) => c.trim()).join(" | ")} |`);
return [header, separator, ...bodyRows].join("\n");
})
}),
);
/**
@@ -293,9 +293,9 @@ const mathFormulaArbitrary: fc.Arbitrary<{
.constantFrom(
"\\sum_{i=1}^{n} x_i",
"\\int_0^\\infty e^{-x} dx",
"\\begin{matrix} a & b \\\\ c & d \\end{matrix}"
"\\begin{matrix} a & b \\\\ c & d \\end{matrix}",
)
.map((formula) => ({ markdown: `$$${formula}$$`, isBlock: true }))
.map((formula) => ({ markdown: `$$${formula}$$`, isBlock: true })),
);
/**
@@ -316,7 +316,7 @@ const formattedTextArbitrary: fc.Arbitrary<{
!s.includes("*") &&
!s.includes("_") &&
!s.includes("[") &&
!s.includes("]")
!s.includes("]"),
),
format: fc.constantFrom("bold", "italic", "link", "boldItalic"),
})
@@ -382,23 +382,24 @@ describe("消息渲染属性测试", () => {
* **Validates: Requirements 2.3**
*/
test.prop(
[fc.integer({ min: 1, max: 6 }).chain((level) => headingArbitrary(level))],
{ numRuns: 100 }
)(
"对于任意 Markdown 标题,应正确检测标题级别",
(heading: string) => {
const headings = detectHeadings(heading);
[
fc
.integer({ min: 1, max: 6 })
.chain((level) => headingArbitrary(level)),
],
{ numRuns: 100 },
)("对于任意 Markdown 标题,应正确检测标题级别", (heading: string) => {
const headings = detectHeadings(heading);
// 应该检测到至少一个标题
expect(headings.length).toBeGreaterThanOrEqual(1);
// 应该检测到至少一个标题
expect(headings.length).toBeGreaterThanOrEqual(1);
// 标题级别应在 1-6 范围内
headings.forEach((level) => {
expect(level).toBeGreaterThanOrEqual(1);
expect(level).toBeLessThanOrEqual(6);
});
}
);
// 标题级别应在 1-6 范围内
headings.forEach((level) => {
expect(level).toBeGreaterThanOrEqual(1);
expect(level).toBeLessThanOrEqual(6);
});
});
/**
* 7.2 无序列表检测测试
@@ -410,7 +411,7 @@ describe("消息渲染属性测试", () => {
(list: string) => {
const { hasUnorderedList } = detectLists(list);
expect(hasUnorderedList).toBe(true);
}
},
);
/**
@@ -423,7 +424,7 @@ describe("消息渲染属性测试", () => {
(list: string) => {
const { hasOrderedList } = detectLists(list);
expect(hasOrderedList).toBe(true);
}
},
);
/**
@@ -445,7 +446,7 @@ describe("消息渲染属性测试", () => {
const codeBlock = blocks.find((b) => b.type === "code");
expect(codeBlock).toBeDefined();
expect(codeBlock?.language).toBe(language);
}
},
);
/**
@@ -458,7 +459,7 @@ describe("消息渲染属性测试", () => {
(table: string) => {
const hasTable = detectTables(table);
expect(hasTable).toBe(true);
}
},
);
/**
@@ -476,7 +477,7 @@ describe("消息渲染属性测试", () => {
} else {
expect(hasInlineMath).toBe(true);
}
}
},
);
/**
@@ -508,7 +509,7 @@ describe("消息渲染属性测试", () => {
if (hasLink) {
expect(detected.hasLink).toBe(true);
}
}
},
);
/**
@@ -531,35 +532,32 @@ describe("消息渲染属性测试", () => {
fc
.string({ minLength: 1, maxLength: 50 })
.filter((s) => s.trim().length > 0 && !s.includes("```")),
codeBlockArbitrary
codeBlockArbitrary,
)
.map(([text, code]) => `${text}\n\n${code.markdown}`)
.map(([text, code]) => `${text}\n\n${code.markdown}`),
),
],
{ numRuns: 100 }
)(
"解析后的内容块应包含原始内容的所有非空白部分",
(content: string) => {
const blocks = parseContent(content);
{ numRuns: 100 },
)("解析后的内容块应包含原始内容的所有非空白部分", (content: string) => {
const blocks = parseContent(content);
// 应该至少有一个内容块
expect(blocks.length).toBeGreaterThanOrEqual(1);
// 应该至少有一个内容块
expect(blocks.length).toBeGreaterThanOrEqual(1);
// 每个块都应该有类型和内容
blocks.forEach((block) => {
expect(block.type).toBeDefined();
expect(["text", "code", "image", "file"]).toContain(block.type);
// 内容可以为空字符串,但类型必须正确
expect(typeof block.content).toBe("string");
});
// 每个块都应该有类型和内容
blocks.forEach((block) => {
expect(block.type).toBeDefined();
expect(["text", "code", "image", "file"]).toContain(block.type);
// 内容可以为空字符串,但类型必须正确
expect(typeof block.content).toBe("string");
});
// 如果原始内容包含代码块,解析结果应包含代码块
if (content.includes("```")) {
const hasCodeBlock = blocks.some((b) => b.type === "code");
expect(hasCodeBlock).toBe(true);
}
// 如果原始内容包含代码块,解析结果应包含代码块
if (content.includes("```")) {
const hasCodeBlock = blocks.some((b) => b.type === "code");
expect(hasCodeBlock).toBe(true);
}
);
});
/**
* 7.9 代码块语言标签正确性测试
@@ -578,10 +576,10 @@ describe("消息渲染属性测试", () => {
"cpp",
"c",
"ruby",
"php"
"php",
),
],
{ numRuns: 100 }
{ numRuns: 100 },
)(
"代码块的语言标签应与原始 Markdown 中指定的语言一致",
(language: string) => {
@@ -591,7 +589,7 @@ describe("消息渲染属性测试", () => {
const codeBlock = blocks.find((b) => b.type === "code");
expect(codeBlock).toBeDefined();
expect(codeBlock?.language).toBe(language);
}
},
);
/**
@@ -605,17 +603,14 @@ describe("消息渲染属性测试", () => {
.string({ minLength: 1, maxLength: 50 })
.filter((s) => s.trim().length > 0 && !s.includes("```")),
],
{ numRuns: 100 }
)(
"没有语言标签的代码块应默认为 plaintext",
(code: string) => {
const markdown = `\`\`\`\n${code}\n\`\`\``;
const blocks = parseContent(markdown);
{ numRuns: 100 },
)("没有语言标签的代码块应默认为 plaintext", (code: string) => {
const markdown = `\`\`\`\n${code}\n\`\`\``;
const blocks = parseContent(markdown);
const codeBlock = blocks.find((b) => b.type === "code");
expect(codeBlock).toBeDefined();
expect(codeBlock?.language).toBe("plaintext");
}
);
const codeBlock = blocks.find((b) => b.type === "code");
expect(codeBlock).toBeDefined();
expect(codeBlock?.language).toBe("plaintext");
});
});
});
@@ -6,9 +6,9 @@
* @requirements 6.1
*/
import React, { useState } from 'react';
import type { Message } from '../types';
import { ImageMessage } from './ImageMessage';
import React, { useState } from "react";
import type { Message } from "../types";
import { ImageMessage } from "./ImageMessage";
interface UserMessageProps {
/** 消息数据 */
@@ -34,9 +34,10 @@ export const UserMessage: React.FC<UserMessageProps> = ({
};
// 分离文本和图片内容
const textBlocks = message.blocks.filter(block => block.type === 'text');
const imageBlocks = message.blocks.filter(block => block.type === 'image');
const hasText = textBlocks.length > 0 && textBlocks.some(block => block.content.trim());
const textBlocks = message.blocks.filter((block) => block.type === "text");
const imageBlocks = message.blocks.filter((block) => block.type === "image");
const hasText =
textBlocks.length > 0 && textBlocks.some((block) => block.content.trim());
const hasImages = imageBlocks.length > 0;
return (
@@ -70,15 +71,35 @@ export const UserMessage: React.FC<UserMessageProps> = ({
<button
onClick={handleCopy}
className="p-1 text-ink-400 hover:text-ink-600 transition-colors"
title={copied ? '已复制' : '复制'}
title={copied ? "已复制" : "复制"}
>
{copied ? (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M5 13l4 4L19 7" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M5 13l4 4L19 7"
/>
</svg>
) : (
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z" />
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M8 16H6a2 2 0 01-2-2V6a2 2 0 012-2h8a2 2 0 012 2v2m-6 12h8a2 2 0 002-2v-8a2 2 0 00-2-2h-8a2 2 0 00-2 2v8a2 2 0 002 2z"
/>
</svg>
)}
</button>
+9 -9
View File
@@ -4,12 +4,12 @@
* @module components/general-chat/chat
*/
export { ChatPanel } from './ChatPanel';
export { MessageList } from './MessageList';
export { MessageItem } from './MessageItem';
export { UserMessage } from './UserMessage';
export { AssistantMessage } from './AssistantMessage';
export { CodeBlock } from './CodeBlock';
export { CompactModelSelector } from './CompactModelSelector';
export { ErrorDisplay } from './ErrorDisplay';
export { ErrorBoundary } from './ErrorBoundary';
export { ChatPanel } from "./ChatPanel";
export { MessageList } from "./MessageList";
export { MessageItem } from "./MessageItem";
export { UserMessage } from "./UserMessage";
export { AssistantMessage } from "./AssistantMessage";
export { CodeBlock } from "./CodeBlock";
export { CompactModelSelector } from "./CompactModelSelector";
export { ErrorDisplay } from "./ErrorDisplay";
export { ErrorBoundary } from "./ErrorBoundary";
@@ -1,19 +1,25 @@
/**
* 工作流状态面板组件
*
*
* 显示三阶段工作流的当前状态、进度和统计信息
*/
import React, { useState, useEffect, useCallback } from 'react';
import {
CheckCircle,
Clock,
import React, { useState, useEffect, useCallback } from "react";
import {
CheckCircle,
Clock,
AlertTriangle,
FileText,
Settings,
} from 'lucide-react';
import { ContextMemoryAPI, type MemoryStats } from '../../../lib/api/contextMemory';
import { ToolHooksAPI, type HookExecutionStats } from '../../../lib/api/toolHooks';
} from "lucide-react";
import {
ContextMemoryAPI,
type MemoryStats,
} from "../../../lib/api/contextMemory";
import {
ToolHooksAPI,
type HookExecutionStats,
} from "../../../lib/api/toolHooks";
export interface WorkflowStatusPanelProps {
sessionId: string;
@@ -44,7 +50,7 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
visualOperationCount,
onInitializeWorkflow,
onFinalizeWorkflow,
className = '',
className = "",
}) => {
const [stats, setStats] = useState<WorkflowStats>({
memoryStats: null,
@@ -54,14 +60,14 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
});
const [showInitDialog, setShowInitDialog] = useState(false);
const [projectName, setProjectName] = useState('');
const [goal, setGoal] = useState('');
const [projectName, setProjectName] = useState("");
const [goal, setGoal] = useState("");
// 加载统计信息
const loadStats = useCallback(async () => {
if (!sessionId) return;
setStats(prev => ({ ...prev, isLoading: true, error: null }));
setStats((prev) => ({ ...prev, isLoading: true, error: null }));
try {
const [memoryStats, hookStats] = await Promise.all([
@@ -76,10 +82,10 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
error: null,
});
} catch (error) {
setStats(prev => ({
setStats((prev) => ({
...prev,
isLoading: false,
error: error instanceof Error ? error.message : '加载统计信息失败',
error: error instanceof Error ? error.message : "加载统计信息失败",
}));
}
}, [sessionId]);
@@ -97,8 +103,8 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
if (projectName.trim() && goal.trim() && onInitializeWorkflow) {
onInitializeWorkflow(projectName.trim(), goal.trim());
setShowInitDialog(false);
setProjectName('');
setGoal('');
setProjectName("");
setGoal("");
}
};
@@ -114,19 +120,21 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
const getStatusText = () => {
if (!isWorkflowInitialized) {
return '工作流未初始化';
return "工作流未初始化";
}
if (!isWorkflowActive) {
return '工作流已暂停';
return "工作流已暂停";
}
if (stats.memoryStats && stats.memoryStats.unresolved_errors > 0) {
return `工作流运行中 (${stats.memoryStats.unresolved_errors} 个错误)`;
}
return '工作流运行正常';
return "工作流运行正常";
};
return (
<div className={`bg-white border border-gray-200 rounded-lg shadow-sm ${className}`}>
<div
className={`bg-white border border-gray-200 rounded-lg shadow-sm ${className}`}
>
{/* 标题栏 */}
<div className="px-4 py-3 border-b border-gray-200 flex items-center justify-between">
<div className="flex items-center space-x-2">
@@ -157,9 +165,11 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
{/* 基本状态 */}
<div className="flex items-center justify-between text-sm">
<span className="text-gray-600">状态:</span>
<span className={`font-medium ${
isWorkflowActive ? 'text-green-600' : 'text-gray-600'
}`}>
<span
className={`font-medium ${
isWorkflowActive ? "text-green-600" : "text-gray-600"
}`}
>
{getStatusText()}
</span>
</div>
@@ -181,28 +191,40 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
<div className="border-t border-gray-100 pt-3">
<div className="flex items-center space-x-1 mb-2">
<FileText className="h-4 w-4 text-gray-400" />
<span className="text-xs font-medium text-gray-700">记忆统计</span>
<span className="text-xs font-medium text-gray-700">
记忆统计
</span>
</div>
<div className="grid grid-cols-2 gap-2 text-xs">
<div className="flex justify-between">
<span className="text-gray-600">活跃记忆:</span>
<span className="font-medium">{stats.memoryStats.active_memories}</span>
<span className="font-medium">
{stats.memoryStats.active_memories}
</span>
</div>
<div className="flex justify-between">
<span className="text-gray-600">已归档:</span>
<span className="font-medium">{stats.memoryStats.archived_memories}</span>
<span className="font-medium">
{stats.memoryStats.archived_memories}
</span>
</div>
<div className="flex justify-between">
<span className="text-gray-600">未解决错误:</span>
<span className={`font-medium ${
stats.memoryStats.unresolved_errors > 0 ? 'text-red-600' : 'text-green-600'
}`}>
<span
className={`font-medium ${
stats.memoryStats.unresolved_errors > 0
? "text-red-600"
: "text-green-600"
}`}
>
{stats.memoryStats.unresolved_errors}
</span>
</div>
<div className="flex justify-between">
<span className="text-gray-600">已解决错误:</span>
<span className="font-medium text-green-600">{stats.memoryStats.resolved_errors}</span>
<span className="font-medium text-green-600">
{stats.memoryStats.resolved_errors}
</span>
</div>
</div>
</div>
@@ -213,17 +235,23 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
<div className="border-t border-gray-100 pt-3">
<div className="flex items-center space-x-1 mb-2">
<Settings className="h-4 w-4 text-gray-400" />
<span className="text-xs font-medium text-gray-700">钩子统计</span>
<span className="text-xs font-medium text-gray-700">
钩子统计
</span>
</div>
<div className="space-y-1">
{Object.entries(stats.hookStats).slice(0, 3).map(([ruleId, stat]) => (
<div key={ruleId} className="flex justify-between text-xs">
<span className="text-gray-600 truncate">{ruleId.split('-')[0]}:</span>
<span className="font-medium">
{stat.success_count}/{stat.execution_count}
</span>
</div>
))}
{Object.entries(stats.hookStats)
.slice(0, 3)
.map(([ruleId, stat]) => (
<div key={ruleId} className="flex justify-between text-xs">
<span className="text-gray-600 truncate">
{ruleId.split("-")[0]}:
</span>
<span className="font-medium">
{stat.success_count}/{stat.execution_count}
</span>
</div>
))}
</div>
</div>
)}
@@ -247,8 +275,10 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
{showInitDialog && (
<div className="fixed inset-0 bg-black bg-opacity-50 flex items-center justify-center z-50">
<div className="bg-white rounded-lg p-6 w-96 max-w-full mx-4">
<h3 className="text-lg font-medium text-gray-900 mb-4">初始化工作流</h3>
<h3 className="text-lg font-medium text-gray-900 mb-4">
初始化工作流
</h3>
<div className="space-y-4">
<div>
<label className="block text-sm font-medium text-gray-700 mb-1">
@@ -262,7 +292,7 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
placeholder="例如:数据分析项目"
/>
</div>
<div>
<label className="block text-sm font-medium text-gray-700 mb-1">
目标描述
@@ -299,4 +329,4 @@ export const WorkflowStatusPanel: React.FC<WorkflowStatusPanelProps> = ({
);
};
export default WorkflowStatusPanel;
export default WorkflowStatusPanel;
+5 -5
View File
@@ -4,8 +4,8 @@
* @module components/general-chat/hooks
*/
export { useStreaming } from './useStreaming';
export { useChat } from './useChat';
export { useSession } from './useSession';
export { useProvider } from './useProvider';
export type { UseProviderResult } from './useProvider';
export { useStreaming } from "./useStreaming";
export { useChat } from "./useChat";
export { useSession } from "./useSession";
export { useProvider } from "./useProvider";
export type { UseProviderResult } from "./useProvider";
+57 -60
View File
@@ -8,10 +8,10 @@
* @requirements 2.1, 5.1, 5.2
*/
import { useCallback } from 'react';
import { invoke } from '@tauri-apps/api/core';
import { useGeneralChatStore } from '../store/useGeneralChatStore';
import type { Message, ProviderConfig } from '../types';
import { useCallback } from "react";
import { invoke } from "@tauri-apps/api/core";
import { useGeneralChatStore } from "../store/useGeneralChatStore";
import type { Message, ProviderConfig } from "../types";
/**
* 发送消息请求参数
@@ -42,65 +42,62 @@ interface UseChatOptions {
* 聊天逻辑 Hook
*/
export const useChat = (options: UseChatOptions) => {
const {
sessionId,
providerConfig,
onMessageSent,
onError,
} = options;
const { sessionId, providerConfig, onMessageSent, onError } = options;
const {
startStreaming,
} = useGeneralChatStore();
const { startStreaming } = useGeneralChatStore();
/**
* 发送消息
*/
const sendMessage = useCallback(async (content: string) => {
if (!sessionId || !content.trim()) {
return;
}
try {
// 构建事件名称
const eventName = `general-chat-stream-${sessionId}`;
// 调用 Tauri 命令发送消息
const request: SendMessageRequest = {
session_id: sessionId,
content: content.trim(),
event_name: eventName,
provider: providerConfig?.providerName,
model: providerConfig?.modelName,
};
const messageId = await invoke<string>('general_chat_send_message', {
request,
});
startStreaming(messageId);
// 消息发送成功
if (onMessageSent) {
const message: Message = {
id: messageId,
sessionId,
role: 'assistant',
content: '',
blocks: [],
status: 'streaming',
createdAt: Date.now(),
};
onMessageSent(message);
const sendMessage = useCallback(
async (content: string) => {
if (!sessionId || !content.trim()) {
return;
}
} catch (error) {
// 停止流式状态
const { stopGeneration } = useGeneralChatStore.getState();
stopGeneration();
const errorMessage = error instanceof Error ? error.message : String(error);
onError?.(errorMessage);
}
}, [sessionId, providerConfig, startStreaming, onMessageSent, onError]);
try {
// 构建事件名称
const eventName = `general-chat-stream-${sessionId}`;
// 调用 Tauri 命令发送消息
const request: SendMessageRequest = {
session_id: sessionId,
content: content.trim(),
event_name: eventName,
provider: providerConfig?.providerName,
model: providerConfig?.modelName,
};
const messageId = await invoke<string>("general_chat_send_message", {
request,
});
startStreaming(messageId);
// 消息发送成功
if (onMessageSent) {
const message: Message = {
id: messageId,
sessionId,
role: "assistant",
content: "",
blocks: [],
status: "streaming",
createdAt: Date.now(),
};
onMessageSent(message);
}
} catch (error) {
// 停止流式状态
const { stopGeneration } = useGeneralChatStore.getState();
stopGeneration();
const errorMessage =
error instanceof Error ? error.message : String(error);
onError?.(errorMessage);
}
},
[sessionId, providerConfig, startStreaming, onMessageSent, onError],
);
/**
* 停止生成
@@ -109,13 +106,13 @@ export const useChat = (options: UseChatOptions) => {
if (!sessionId) return;
try {
await invoke('general_chat_stop_generation', {
await invoke("general_chat_stop_generation", {
sessionId,
});
const { stopGeneration: stopGen } = useGeneralChatStore.getState();
stopGen();
} catch (error) {
console.error('停止生成失败:', error);
console.error("停止生成失败:", error);
}
}, [sessionId]);
@@ -127,7 +124,7 @@ export const useChat = (options: UseChatOptions) => {
// 1. 获取该消息之前的用户消息
// 2. 删除该消息
// 3. 重新发送用户消息
console.log('重新生成消息:', messageId);
console.log("重新生成消息:", messageId);
}, []);
return {
@@ -9,11 +9,14 @@
* @requirements 5.1, 5.2, 5.3
*/
import { useEffect, useCallback, useMemo, useRef } from 'react';
import { useConfiguredProviders, type ConfiguredProvider } from '@/hooks/useConfiguredProviders';
import { useProviderModels } from '@/hooks/useProviderModels';
import { useGeneralChatStore } from '../store/useGeneralChatStore';
import type { ProviderConfig } from '../types';
import { useEffect, useCallback, useMemo, useRef } from "react";
import {
useConfiguredProviders,
type ConfiguredProvider,
} from "@/hooks/useConfiguredProviders";
import { useProviderModels } from "@/hooks/useProviderModels";
import { useGeneralChatStore } from "../store/useGeneralChatStore";
import type { ProviderConfig } from "../types";
// ============================================================================
// 类型定义
@@ -131,21 +134,34 @@ export function useProvider(): UseProviderResult {
if (providersLoading) return;
// 如果没有选中的 Provider,且有可用的 Provider,自动选择第一个
if (!selectedProviderKey && providers.length > 0 && !providerInitializedRef.current) {
if (
!selectedProviderKey &&
providers.length > 0 &&
!providerInitializedRef.current
) {
providerInitializedRef.current = true;
setSelectedProvider(providers[0].key);
return;
}
// 如果选中的 Provider 不在列表中(可能被删除),重新选择
if (selectedProviderKey && !providers.find((p) => p.key === selectedProviderKey)) {
if (
selectedProviderKey &&
!providers.find((p) => p.key === selectedProviderKey)
) {
if (providers.length > 0) {
setSelectedProvider(providers[0].key);
} else {
setSelectedProvider(null);
}
}
}, [providersLoading, providers.length, selectedProviderKey, setSelectedProvider, providers]);
}, [
providersLoading,
providers.length,
selectedProviderKey,
setSelectedProvider,
providers,
]);
// 当 Provider 切换时,重置模型初始化标记
useEffect(() => {
@@ -160,7 +176,11 @@ export function useProvider(): UseProviderResult {
if (modelsLoading) return;
// 如果没有选中的模型,且有可用的模型,自动选择第一个
if (!selectedModelId && availableModelIds.length > 0 && !modelInitializedRef.current) {
if (
!selectedModelId &&
availableModelIds.length > 0 &&
!modelInitializedRef.current
) {
modelInitializedRef.current = true;
setSelectedModel(availableModelIds[0]);
return;
@@ -190,7 +210,7 @@ export function useProvider(): UseProviderResult {
setSelectedModel(null);
}
},
[providers, setSelectedProvider, setSelectedModel]
[providers, setSelectedProvider, setSelectedModel],
);
/**
@@ -202,7 +222,7 @@ export function useProvider(): UseProviderResult {
setSelectedModel(modelId);
}
},
[availableModelIds, setSelectedModel]
[availableModelIds, setSelectedModel],
);
/**
@@ -216,7 +236,9 @@ export function useProvider(): UseProviderResult {
return false;
}
const currentIndex = providers.findIndex((p) => p.key === selectedProviderKey);
const currentIndex = providers.findIndex(
(p) => p.key === selectedProviderKey,
);
const nextIndex = (currentIndex + 1) % providers.length;
const nextProvider = providers[nextIndex];
+90 -69
View File
@@ -8,10 +8,10 @@
* @requirements 1.2, 1.5
*/
import { useCallback, useEffect } from 'react';
import { invoke } from '@tauri-apps/api/core';
import { useGeneralChatStore } from '../store/useGeneralChatStore';
import type { Session } from '../types';
import { useCallback, useEffect } from "react";
import { invoke } from "@tauri-apps/api/core";
import { useGeneralChatStore } from "../store/useGeneralChatStore";
import type { Session } from "../types";
/**
* 后端会话数据结构
@@ -58,10 +58,7 @@ interface UseSessionOptions {
* 会话管理 Hook
*/
export const useSession = (options: UseSessionOptions = {}) => {
const {
autoLoad = true,
onSessionChange,
} = options;
const { autoLoad = true, onSessionChange } = options;
const {
sessions,
@@ -78,93 +75,117 @@ export const useSession = (options: UseSessionOptions = {}) => {
*/
const loadSessions = useCallback(async () => {
try {
const backendSessions = await invoke<BackendSession[]>('general_chat_list_sessions');
const backendSessions = await invoke<BackendSession[]>(
"general_chat_list_sessions",
);
const frontendSessions = backendSessions.map(convertSession);
setSessions(frontendSessions);
} catch (error) {
console.error('加载会话列表失败:', error);
console.error("加载会话列表失败:", error);
}
}, [setSessions]);
/**
* 创建新会话
*/
const createSession = useCallback(async (name?: string): Promise<string | null> => {
try {
const _session = await invoke<BackendSession>('general_chat_create_session', {
name: name || undefined,
metadata: undefined,
});
// 使用 store 的 createSession 方法,它会自动添加会话并设置为当前会话
const sessionId = await createNewSession();
onSessionChange?.(sessionId);
return sessionId;
} catch (error) {
console.error('创建会话失败:', error);
return null;
}
}, [createNewSession, onSessionChange]);
const createSession = useCallback(
async (name?: string): Promise<string | null> => {
try {
const _session = await invoke<BackendSession>(
"general_chat_create_session",
{
name: name || undefined,
metadata: undefined,
},
);
// 使用 store 的 createSession 方法,它会自动添加会话并设置为当前会话
const sessionId = await createNewSession();
onSessionChange?.(sessionId);
return sessionId;
} catch (error) {
console.error("创建会话失败:", error);
return null;
}
},
[createNewSession, onSessionChange],
);
/**
* 切换会话
*/
const switchSession = useCallback(async (sessionId: string) => {
try {
// 加载会话详情
const detail = await invoke<BackendSessionDetail>('general_chat_get_session', {
sessionId,
messageLimit: 50,
});
// 更新会话消息数量
updateSession(sessionId, { messageCount: detail.message_count });
// 切换当前会话
selectSession(sessionId);
onSessionChange?.(sessionId);
} catch (error) {
console.error('切换会话失败:', error);
}
}, [selectSession, updateSession, onSessionChange]);
const switchSession = useCallback(
async (sessionId: string) => {
try {
// 加载会话详情
const detail = await invoke<BackendSessionDetail>(
"general_chat_get_session",
{
sessionId,
messageLimit: 50,
},
);
// 更新会话消息数量
updateSession(sessionId, { messageCount: detail.message_count });
// 切换当前会话
selectSession(sessionId);
onSessionChange?.(sessionId);
} catch (error) {
console.error("切换会话失败:", error);
}
},
[selectSession, updateSession, onSessionChange],
);
/**
* 删除会话
*/
const deleteSession = useCallback(async (sessionId: string) => {
try {
await invoke('general_chat_delete_session', { sessionId });
// 使用 store 的 deleteSession 方法,它会自动处理会话切换逻辑
await useGeneralChatStore.getState().deleteSession(sessionId);
// 获取新的当前会话 ID 并触发回调
const newCurrentId = useGeneralChatStore.getState().currentSessionId;
onSessionChange?.(newCurrentId);
} catch (error) {
console.error('删除会话失败:', error);
}
}, [onSessionChange]);
const deleteSession = useCallback(
async (sessionId: string) => {
try {
await invoke("general_chat_delete_session", { sessionId });
// 使用 store 的 deleteSession 方法,它会自动处理会话切换逻辑
await useGeneralChatStore.getState().deleteSession(sessionId);
// 获取新的当前会话 ID 并触发回调
const newCurrentId = useGeneralChatStore.getState().currentSessionId;
onSessionChange?.(newCurrentId);
} catch (error) {
console.error("删除会话失败:", error);
}
},
[onSessionChange],
);
/**
* 重命名会话
*/
const renameSession = useCallback(async (sessionId: string, name: string) => {
try {
await invoke('general_chat_rename_session', { sessionId, name });
updateSession(sessionId, { name });
} catch (error) {
console.error('重命名会话失败:', error);
}
}, [updateSession]);
const renameSession = useCallback(
async (sessionId: string, name: string) => {
try {
await invoke("general_chat_rename_session", { sessionId, name });
updateSession(sessionId, { name });
} catch (error) {
console.error("重命名会话失败:", error);
}
},
[updateSession],
);
/**
* 自动生成会话标题
* 基于第一条用户消息生成
*/
const generateTitle = useCallback(async (sessionId: string, firstMessage: string) => {
// 简单实现:截取前 20 个字符作为标题
const title = firstMessage.slice(0, 20) + (firstMessage.length > 20 ? '...' : '');
await renameSession(sessionId, title);
}, [renameSession]);
const generateTitle = useCallback(
async (sessionId: string, firstMessage: string) => {
// 简单实现:截取前 20 个字符作为标题
const title =
firstMessage.slice(0, 20) + (firstMessage.length > 20 ? "..." : "");
await renameSession(sessionId, title);
},
[renameSession],
);
// 自动加载会话列表
useEffect(() => {
@@ -8,15 +8,15 @@
* @requirements 2.2, 2.5
*/
import { useEffect, useCallback, useRef } from 'react';
import { listen, type UnlistenFn } from '@tauri-apps/api/event';
import { useGeneralChatStore } from '../store/useGeneralChatStore';
import { useEffect, useCallback, useRef } from "react";
import { listen, type UnlistenFn } from "@tauri-apps/api/event";
import { useGeneralChatStore } from "../store/useGeneralChatStore";
/**
* 流式事件类型
*/
interface StreamEvent {
type: 'start' | 'delta' | 'done' | 'error';
type: "start" | "delta" | "done" | "error";
message_id?: string;
content?: string;
message?: string;
@@ -48,57 +48,57 @@ interface UseStreamingOptions {
export const useStreaming = (options: UseStreamingOptions) => {
const {
sessionId,
eventName = 'general-chat-stream',
eventName = "general-chat-stream",
onStart,
onDelta,
onDone,
onError,
} = options;
const {
startStreaming,
appendStreamingContent,
} = useGeneralChatStore();
const { startStreaming, appendStreamingContent } = useGeneralChatStore();
const unlistenRef = useRef<UnlistenFn | null>(null);
const contentRef = useRef<string>('');
const contentRef = useRef<string>("");
// 处理流式事件
const handleStreamEvent = useCallback((event: { payload: StreamEvent }) => {
const { type, message_id, content, message } = event.payload;
const handleStreamEvent = useCallback(
(event: { payload: StreamEvent }) => {
const { type, message_id, content, message } = event.payload;
switch (type) {
case 'start':
contentRef.current = '';
startStreaming(message_id || '');
onStart?.(message_id || '');
break;
switch (type) {
case "start":
contentRef.current = "";
startStreaming(message_id || "");
onStart?.(message_id || "");
break;
case 'delta':
if (content) {
contentRef.current += content;
appendStreamingContent(content);
onDelta?.(content);
case "delta":
if (content) {
contentRef.current += content;
appendStreamingContent(content);
onDelta?.(content);
}
break;
case "done": {
const { finalizeMessage } = useGeneralChatStore.getState();
finalizeMessage();
onDone?.(message_id || "", contentRef.current);
contentRef.current = "";
break;
}
break;
case 'done': {
const { finalizeMessage } = useGeneralChatStore.getState();
finalizeMessage();
onDone?.(message_id || '', contentRef.current);
contentRef.current = '';
break;
case "error": {
const { stopGeneration: stopGen } = useGeneralChatStore.getState();
stopGen();
onError?.(message || "未知错误");
contentRef.current = "";
break;
}
}
case 'error': {
const { stopGeneration: stopGen } = useGeneralChatStore.getState();
stopGen();
onError?.(message || '未知错误');
contentRef.current = '';
break;
}
}
}, [startStreaming, appendStreamingContent, onStart, onDelta, onDone, onError]);
},
[startStreaming, appendStreamingContent, onStart, onDelta, onDone, onError],
);
// 设置事件监听
useEffect(() => {
@@ -112,7 +112,10 @@ export const useStreaming = (options: UseStreamingOptions) => {
// 设置新的监听器
const eventKey = `${eventName}-${sessionId}`;
unlistenRef.current = await listen<StreamEvent>(eventKey, handleStreamEvent);
unlistenRef.current = await listen<StreamEvent>(
eventKey,
handleStreamEvent,
);
};
setupListener();
@@ -130,7 +133,7 @@ export const useStreaming = (options: UseStreamingOptions) => {
const stopGeneration = useCallback(() => {
const { stopGeneration: stopGen } = useGeneralChatStore.getState();
stopGen();
contentRef.current = '';
contentRef.current = "";
}, []);
return {
@@ -1,13 +1,13 @@
/**
* 工作流集成 Hook
*
*
* 为通用对话功能集成三阶段工作流,实现自动化上下文管理
*/
import { useCallback, useEffect, useRef } from 'react';
import { useThreeStageWorkflow } from '../../../hooks/useThreeStageWorkflow';
import { ToolHooksAPI } from '../../../lib/api/toolHooks';
import type { Message } from '../types';
import { useCallback, useEffect, useRef } from "react";
import { useThreeStageWorkflow } from "../../../hooks/useThreeStageWorkflow";
import { ToolHooksAPI } from "../../../lib/api/toolHooks";
import type { Message } from "../types";
export interface WorkflowIntegrationOptions {
sessionId: string;
@@ -25,9 +25,20 @@ export interface WorkflowIntegrationState {
export interface WorkflowIntegrationActions {
initializeWorkflow: (projectName: string, goal: string) => Promise<void>;
handlePreMessage: (content: string, messageType: 'user' | 'assistant') => Promise<string | null>;
handlePostMessage: (message: Message, isError?: boolean) => Promise<string | null>;
handleToolUse: (toolName: string, toolParams: any, toolResult: string, error?: string) => Promise<void>;
handlePreMessage: (
content: string,
messageType: "user" | "assistant",
) => Promise<string | null>;
handlePostMessage: (
message: Message,
isError?: boolean,
) => Promise<string | null>;
handleToolUse: (
toolName: string,
toolParams: any,
toolResult: string,
error?: string,
) => Promise<void>;
finalizeSession: () => Promise<string | null>;
toggleWorkflow: () => void;
}
@@ -36,10 +47,14 @@ export interface WorkflowIntegrationActions {
* 工作流集成 Hook
*/
export function useWorkflowIntegration(
options: WorkflowIntegrationOptions
options: WorkflowIntegrationOptions,
): [WorkflowIntegrationState, WorkflowIntegrationActions] {
const { sessionId, enableAutoWorkflow = true, workflowThreshold = 5 } = options;
const {
sessionId,
enableAutoWorkflow = true,
workflowThreshold = 5,
} = options;
const [workflowState, workflowActions] = useThreeStageWorkflow({
sessionId,
autoInitialize: false,
@@ -49,202 +64,207 @@ export function useWorkflowIntegration(
const visualOperationCountRef = useRef(0);
const isWorkflowActiveRef = useRef(false);
const initializeWorkflow = useCallback(async (projectName: string, goal: string) => {
try {
await workflowActions.initializeWorkflow({
sessionId,
projectName,
goal,
phases: [
{
number: 1,
name: '需求理解',
status: 'in_progress',
tasks: [
'理解用户需求和目标',
'识别关键约束条件',
'记录重要信息到 findings.md',
],
},
{
number: 2,
name: '方案制定',
status: 'pending',
tasks: [
'分析可行的解决方案',
'制定执行计划',
'记录关键决策',
],
},
{
number: 3,
name: '执行实施',
status: 'pending',
tasks: [
'按计划执行任务',
'监控执行进度',
'记录执行结果',
],
},
{
number: 4,
name: '验证完善',
status: 'pending',
tasks: [
'验证结果质量',
'完善不足之处',
'总结经验教训',
],
},
],
});
isWorkflowActiveRef.current = true;
// 触发会话开始钩子
await ToolHooksAPI.triggerSessionStart(sessionId, {
project_name: projectName,
goal,
auto_initialized: enableAutoWorkflow.toString(),
});
} catch (error) {
console.error('初始化工作流失败:', error);
throw error;
}
}, [sessionId, workflowActions, enableAutoWorkflow]);
const initializeWorkflow = useCallback(
async (projectName: string, goal: string) => {
try {
await workflowActions.initializeWorkflow({
sessionId,
projectName,
goal,
phases: [
{
number: 1,
name: "需求理解",
status: "in_progress",
tasks: [
"理解用户需求和目标",
"识别关键约束条件",
"记录重要信息到 findings.md",
],
},
{
number: 2,
name: "方案制定",
status: "pending",
tasks: ["分析可行的解决方案", "制定执行计划", "记录关键决策"],
},
{
number: 3,
name: "执行实施",
status: "pending",
tasks: ["按计划执行任务", "监控执行进度", "记录执行结果"],
},
{
number: 4,
name: "验证完善",
status: "pending",
tasks: ["验证结果质量", "完善不足之处", "总结经验教训"],
},
],
});
isWorkflowActiveRef.current = true;
// 触发会话开始钩子
await ToolHooksAPI.triggerSessionStart(sessionId, {
project_name: projectName,
goal,
auto_initialized: enableAutoWorkflow.toString(),
});
} catch (error) {
console.error("初始化工作流失败:", error);
throw error;
}
},
[sessionId, workflowActions, enableAutoWorkflow],
);
// 自动启用工作流检查
useEffect(() => {
if (enableAutoWorkflow &&
messageCountRef.current >= workflowThreshold &&
!workflowState.isInitialized &&
!isWorkflowActiveRef.current) {
if (
enableAutoWorkflow &&
messageCountRef.current >= workflowThreshold &&
!workflowState.isInitialized &&
!isWorkflowActiveRef.current
) {
// 自动初始化工作流
initializeWorkflow('通用对话任务', '协助用户完成复杂任务');
initializeWorkflow("通用对话任务", "协助用户完成复杂任务");
}
}, [enableAutoWorkflow, workflowThreshold, workflowState.isInitialized, initializeWorkflow]);
}, [
enableAutoWorkflow,
workflowThreshold,
workflowState.isInitialized,
initializeWorkflow,
]);
const handlePreMessage = useCallback(async (
content: string,
messageType: 'user' | 'assistant'
): Promise<string | null> => {
if (!workflowState.isInitialized || !isWorkflowActiveRef.current) {
return null;
}
const handlePreMessage = useCallback(
async (
content: string,
messageType: "user" | "assistant",
): Promise<string | null> => {
if (!workflowState.isInitialized || !isWorkflowActiveRef.current) {
return null;
}
messageCountRef.current++;
messageCountRef.current++;
try {
// 只对用户消息执行 Pre-Action
if (messageType === 'user') {
try {
// 只对用户消息执行 Pre-Action
if (messageType === "user") {
const context = {
sessionId,
actionType: "user_message",
actionDescription: `用户消息: ${content.substring(0, 100)}${content.length > 100 ? "..." : ""}`,
messageCount: messageCountRef.current,
};
const preActionResult = await workflowActions.preAction(context);
return preActionResult;
}
return null;
} catch (error) {
console.error("Pre-Message 处理失败:", error);
return null;
}
},
[workflowState.isInitialized, workflowActions, sessionId],
);
const handlePostMessage = useCallback(
async (
message: Message,
isError: boolean = false,
): Promise<string | null> => {
if (!workflowState.isInitialized || !isWorkflowActiveRef.current) {
return null;
}
try {
const context = {
sessionId,
actionType: 'user_message',
actionDescription: `用户消息: ${content.substring(0, 100)}${content.length > 100 ? '...' : ''}`,
actionType:
message.role === "user" ? "user_message" : "assistant_response",
actionDescription: `${message.role} 消息处理`,
messageCount: messageCountRef.current,
};
const preActionResult = await workflowActions.preAction(context);
return preActionResult;
}
return null;
} catch (error) {
console.error('Pre-Message 处理失败:', error);
return null;
}
}, [workflowState.isInitialized, workflowActions, sessionId]);
const postActionResult = await workflowActions.postAction(
context,
message.content,
isError ? "消息处理出现错误" : undefined,
);
const handlePostMessage = useCallback(async (
message: Message,
isError: boolean = false
): Promise<string | null> => {
if (!workflowState.isInitialized || !isWorkflowActiveRef.current) {
return null;
}
// 检测视觉操作
if (message.images && message.images.length > 0) {
visualOperationCountRef.current++;
try {
const context = {
sessionId,
actionType: message.role === 'user' ? 'user_message' : 'assistant_response',
actionDescription: `${message.role} 消息处理`,
messageCount: messageCountRef.current,
};
const postActionResult = await workflowActions.postAction(
context,
message.content,
isError ? '消息处理出现错误' : undefined
);
// 检测视觉操作
if (message.images && message.images.length > 0) {
visualOperationCountRef.current++;
// 应用 2-Action 规则
if (visualOperationCountRef.current >= 2) {
await workflowActions.recordFinding(
'视觉内容分析',
`处理了 ${message.images.length} 个图片,内容: ${message.content.substring(0, 200)}`,
['视觉操作', '2-Action规则']
);
visualOperationCountRef.current = 0;
// 应用 2-Action 规则
if (visualOperationCountRef.current >= 2) {
await workflowActions.recordFinding(
"视觉内容分析",
`处理了 ${message.images.length} 个图片,内容: ${message.content.substring(0, 200)}`,
["视觉操作", "2-Action规则"],
);
visualOperationCountRef.current = 0;
}
}
return postActionResult;
} catch (error) {
console.error("Post-Message 处理失败:", error);
return null;
}
},
[workflowState.isInitialized, workflowActions, sessionId],
);
const handleToolUse = useCallback(
async (
toolName: string,
toolParams: any,
toolResult: string,
error?: string,
) => {
if (!workflowState.isInitialized || !isWorkflowActiveRef.current) {
return;
}
return postActionResult;
} catch (error) {
console.error('Post-Message 处理失败:', error);
return null;
}
}, [workflowState.isInitialized, workflowActions, sessionId]);
try {
// 触发工具使用钩子
if (error) {
await ToolHooksAPI.triggerPostToolUse(
sessionId,
toolName,
toolResult,
toolParams,
`工具 ${toolName} 执行失败`,
messageCountRef.current,
error,
);
} else {
await ToolHooksAPI.triggerPostToolUse(
sessionId,
toolName,
toolResult,
toolParams,
`工具 ${toolName} 执行成功`,
messageCountRef.current,
);
}
const handleToolUse = useCallback(async (
toolName: string,
toolParams: any,
toolResult: string,
error?: string
) => {
if (!workflowState.isInitialized || !isWorkflowActiveRef.current) {
return;
}
try {
// 触发工具使用钩子
if (error) {
await ToolHooksAPI.triggerPostToolUse(
sessionId,
toolName,
toolResult,
toolParams,
`工具 ${toolName} 执行失败`,
messageCountRef.current,
error
);
} else {
await ToolHooksAPI.triggerPostToolUse(
sessionId,
toolName,
toolResult,
toolParams,
`工具 ${toolName} 执行成功`,
messageCountRef.current
// 记录工具使用进度
await workflowActions.recordFinding(
`工具使用: ${toolName}`,
`工具: ${toolName}\n参数: ${JSON.stringify(toolParams, null, 2)}\n结果: ${toolResult.substring(0, 300)}${toolResult.length > 300 ? "..." : ""}${error ? `\n错误: ${error}` : ""}`,
["工具使用", toolName, error ? "错误" : "成功"],
);
} catch (err) {
console.error("工具使用处理失败:", err);
}
// 记录工具使用进度
await workflowActions.recordFinding(
`工具使用: ${toolName}`,
`工具: ${toolName}\n参数: ${JSON.stringify(toolParams, null, 2)}\n结果: ${toolResult.substring(0, 300)}${toolResult.length > 300 ? '...' : ''}${error ? `\n错误: ${error}` : ''}`,
['工具使用', toolName, error ? '错误' : '成功']
);
} catch (err) {
console.error('工具使用处理失败:', err);
}
}, [workflowState.isInitialized, workflowActions, sessionId]);
},
[workflowState.isInitialized, workflowActions, sessionId],
);
const finalizeSession = useCallback(async (): Promise<string | null> => {
if (!workflowState.isInitialized || !isWorkflowActiveRef.current) {
@@ -260,17 +280,17 @@ export function useWorkflowIntegration(
// 检查完成状态
const { isComplete, summary } = await workflowActions.checkCompletion();
// 结束工作流
const finalMessage = await workflowActions.finalizeWorkflow();
isWorkflowActiveRef.current = false;
messageCountRef.current = 0;
visualOperationCountRef.current = 0;
return `${finalMessage}\n\n完成状态: ${isComplete ? '✅ 已完成' : '⏳ 未完成'}\n\n${summary}`;
return `${finalMessage}\n\n完成状态: ${isComplete ? "✅ 已完成" : "⏳ 未完成"}\n\n${summary}`;
} catch (error) {
console.error('结束会话处理失败:', error);
console.error("结束会话处理失败:", error);
return null;
}
}, [workflowState.isInitialized, workflowActions, sessionId]);
@@ -299,4 +319,4 @@ export function useWorkflowIntegration(
return [state, actions];
}
export default useWorkflowIntegration;
export default useWorkflowIntegration;
+14 -7
View File
@@ -7,7 +7,7 @@
*/
// 主页面组件导出
export { default as GeneralChatPage } from './GeneralChatPage';
export { default as GeneralChatPage } from "./GeneralChatPage";
// 类型导出(不触发 react-refresh 警告)
export type {
@@ -28,11 +28,18 @@ export type {
MessageItemProps,
InputBarProps,
CanvasPanelProps,
} from './types';
} from "./types";
// 子模块组件导出
export { ChatPanel, MessageList, MessageItem, UserMessage, AssistantMessage, CodeBlock, ErrorBoundary } from './chat';
export { CanvasPanel, CodePreview, MarkdownPreview } from './canvas';
export { useGeneralChatStore } from './store';
export { useStreaming, useChat, useSession } from './hooks';
export {
ChatPanel,
MessageList,
MessageItem,
UserMessage,
AssistantMessage,
CodeBlock,
ErrorBoundary,
} from "./chat";
export { CanvasPanel, CodePreview, MarkdownPreview } from "./canvas";
export { useGeneralChatStore } from "./store";
export { useStreaming, useChat, useSession } from "./hooks";
+2 -2
View File
@@ -15,6 +15,6 @@ export {
useCanvasState,
useSessions,
useIsStreaming,
} from './useGeneralChatStore';
} from "./useGeneralChatStore";
export type { GeneralChatState } from './useGeneralChatStore';
export type { GeneralChatState } from "./useGeneralChatStore";
@@ -9,8 +9,8 @@
* @requirements 8.1, 8.2
*/
import { create } from 'zustand';
import { persist, createJSONStorage } from 'zustand/middleware';
import { create } from "zustand";
import { persist, createJSONStorage } from "zustand/middleware";
import type {
Session,
Message,
@@ -21,8 +21,8 @@ import type {
ProviderSelectionState,
ErrorInfo,
PaginationState,
} from '../types';
import { ThreeStageWorkflowManager } from '@/lib/workflow/threeStageWorkflow';
} from "../types";
import { ThreeStageWorkflowManager } from "@/lib/workflow/threeStageWorkflow";
import {
DEFAULT_UI_STATE,
DEFAULT_CANVAS_STATE,
@@ -30,7 +30,7 @@ import {
DEFAULT_PROVIDER_SELECTION_STATE,
DEFAULT_PAGINATION_STATE,
parseApiError,
} from '../types';
} from "../types";
// ============================================================================
// Store 状态接口
@@ -110,7 +110,11 @@ export interface GeneralChatState {
// ========== 三阶段工作流操作 ==========
/** 初始化工作流 */
initializeWorkflow: (sessionId: string, projectName: string, goal: string) => Promise<void>;
initializeWorkflow: (
sessionId: string,
projectName: string,
goal: string,
) => Promise<void>;
/** 获取工作流管理器 */
getWorkflowManager: (sessionId: string) => ThreeStageWorkflowManager | null;
/** 启用/禁用工作流 */
@@ -133,7 +137,10 @@ export interface GeneralChatState {
/** 加载更多历史消息 */
loadMoreMessages: (sessionId: string) => Promise<void>;
/** 设置分页状态 */
setPaginationState: (sessionId: string, state: Partial<PaginationState>) => void;
setPaginationState: (
sessionId: string,
state: Partial<PaginationState>,
) => void;
/** 重置分页状态 */
resetPagination: (sessionId: string) => void;
/** 获取分页状态 */
@@ -247,7 +254,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const timestamp = now();
const newSession: Session = {
id,
name: '新对话',
name: "新对话",
createdAt: timestamp,
updatedAt: timestamp,
messageCount: 0,
@@ -309,7 +316,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
renameSession: async (id: string, name: string) => {
set((state) => ({
sessions: state.sessions.map((s) =>
s.id === id ? { ...s, name, updatedAt: now() } : s
s.id === id ? { ...s, name, updatedAt: now() } : s,
),
}));
@@ -324,7 +331,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
updateSession: (id: string, updates: Partial<Session>) => {
set((state) => ({
sessions: state.sessions.map((s) =>
s.id === id ? { ...s, ...updates, updatedAt: now() } : s
s.id === id ? { ...s, ...updates, updatedAt: now() } : s,
),
}));
},
@@ -341,7 +348,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
// 验证:必须有当前会话
if (!currentSessionId) {
console.warn('No current session selected');
console.warn("No current session selected");
return;
}
@@ -355,28 +362,30 @@ export const useGeneralChatStore = create<GeneralChatState>()(
// 将 File 对象转换为 base64 格式
imageData = await Promise.all(
images.map(async (file) => {
return new Promise<{ data: string; media_type: string }>((resolve, reject) => {
const reader = new FileReader();
reader.onload = (e) => {
const result = e.target?.result as string;
if (!result) {
reject(new Error('无法读取文件'));
return;
}
// 提取 base64 数据(去掉 data:image/xxx;base64, 前缀)
const base64Data = result.split(',')[1];
resolve({
data: base64Data,
media_type: file.type,
});
};
reader.onerror = () => reject(new Error('文件读取失败'));
reader.readAsDataURL(file);
});
})
return new Promise<{ data: string; media_type: string }>(
(resolve, reject) => {
const reader = new FileReader();
reader.onload = (e) => {
const result = e.target?.result as string;
if (!result) {
reject(new Error("无法读取文件"));
return;
}
// 提取 base64 数据(去掉 data:image/xxx;base64, 前缀)
const base64Data = result.split(",")[1];
resolve({
data: base64Data,
media_type: file.type,
});
};
reader.onerror = () => reject(new Error("文件读取失败"));
reader.readAsDataURL(file);
},
);
}),
);
} catch (error) {
console.error('图片处理失败:', error);
console.error("图片处理失败:", error);
return;
}
}
@@ -385,17 +394,19 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const userMessage: Message = {
id: messageId,
sessionId: currentSessionId,
role: 'user',
content: content.trim() || '[图片]',
role: "user",
content: content.trim() || "[图片]",
blocks: [
...(content.trim() ? [{ type: 'text' as const, content: content.trim() }] : []),
...(content.trim()
? [{ type: "text" as const, content: content.trim() }]
: []),
...(imageData?.map((img) => ({
type: 'image' as const,
type: "image" as const,
content: `data:${img.media_type};base64,${img.data}`,
mimeType: img.media_type,
})) || []),
],
status: 'complete',
status: "complete",
createdAt: timestamp,
};
@@ -418,10 +429,10 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const assistantMessage: Message = {
id: assistantMessageId,
sessionId: currentSessionId,
role: 'assistant',
content: '',
role: "assistant",
content: "",
blocks: [],
status: 'pending',
status: "pending",
createdAt: timestamp + 1,
};
@@ -429,7 +440,10 @@ export const useGeneralChatStore = create<GeneralChatState>()(
set((state) => ({
messages: {
...state.messages,
[currentSessionId]: [...(state.messages[currentSessionId] || []), assistantMessage],
[currentSessionId]: [
...(state.messages[currentSessionId] || []),
assistantMessage,
],
},
}));
@@ -437,16 +451,23 @@ export const useGeneralChatStore = create<GeneralChatState>()(
get().startStreaming(assistantMessageId);
// 三阶段工作流集成
const { workflowEnabled, workflowManagers, messages: allMessages } = get();
const {
workflowEnabled,
workflowManagers,
messages: allMessages,
} = get();
let workflowManager = workflowManagers[currentSessionId];
// 检查是否应该自动启用工作流
if (!workflowManager && get().shouldAutoEnableWorkflow(currentSessionId)) {
if (
!workflowManager &&
get().shouldAutoEnableWorkflow(currentSessionId)
) {
// 自动初始化工作流
await get().initializeWorkflow(
currentSessionId,
'智能对话任务',
'协助用户完成复杂的对话任务,提供准确和有用的回答'
"智能对话任务",
"协助用户完成复杂的对话任务,提供准确和有用的回答",
);
workflowManager = get().workflowManagers[currentSessionId];
}
@@ -457,12 +478,12 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const messageCount = (allMessages[currentSessionId] || []).length;
const actionContext = {
sessionId: currentSessionId,
actionType: 'send_message',
actionDescription: `发送消息: ${content.substring(0, 100)}${content.length > 100 ? '...' : ''}`,
toolName: 'aster_agent_chat_stream',
actionType: "send_message",
actionDescription: `发送消息: ${content.substring(0, 100)}${content.length > 100 ? "..." : ""}`,
toolName: "aster_agent_chat_stream",
toolParameters: {
message: content.trim() || '请分析这张图片',
hasImages: imageData ? 'true' : 'false',
message: content.trim() || "请分析这张图片",
hasImages: imageData ? "true" : "false",
},
messageCount,
};
@@ -470,16 +491,16 @@ export const useGeneralChatStore = create<GeneralChatState>()(
// Pre-Action: 上下文刷新
await workflowManager.preAction(actionContext);
} catch (error) {
console.warn('工作流 Pre-Action 执行失败:', error);
console.warn("工作流 Pre-Action 执行失败:", error);
}
}
try {
// 调用 Tauri 命令发送消息并开始流式响应
const { invoke } = await import('@tauri-apps/api/core');
await invoke('aster_agent_chat_stream', {
const { invoke } = await import("@tauri-apps/api/core");
await invoke("aster_agent_chat_stream", {
sessionId: currentSessionId,
message: content.trim() || '请分析这张图片',
message: content.trim() || "请分析这张图片",
eventName: `general-chat-stream-${currentSessionId}`,
images: imageData,
});
@@ -487,40 +508,49 @@ export const useGeneralChatStore = create<GeneralChatState>()(
// 如果启用了工作流,执行 Action 阶段
if (workflowManager && workflowEnabled) {
try {
const messageCount = (get().messages[currentSessionId] || []).length;
const messageCount = (get().messages[currentSessionId] || [])
.length;
const actionContext = {
sessionId: currentSessionId,
actionType: 'send_message',
actionDescription: `发送消息: ${content.substring(0, 100)}${content.length > 100 ? '...' : ''}`,
toolName: 'aster_agent_chat_stream',
actionType: "send_message",
actionDescription: `发送消息: ${content.substring(0, 100)}${content.length > 100 ? "..." : ""}`,
toolName: "aster_agent_chat_stream",
messageCount,
};
await workflowManager.executeAction(actionContext, '消息发送成功,等待 AI 响应');
await workflowManager.executeAction(
actionContext,
"消息发送成功,等待 AI 响应",
);
} catch (error) {
console.warn('工作流 Action 执行失败:', error);
console.warn("工作流 Action 执行失败:", error);
}
}
} catch (error) {
console.error('发送消息失败:', error);
console.error("发送消息失败:", error);
// 如果启用了工作流,执行 Post-Action 阶段(错误情况)
if (workflowManager && workflowEnabled) {
try {
const messageCount = (get().messages[currentSessionId] || []).length;
const messageCount = (get().messages[currentSessionId] || [])
.length;
const actionContext = {
sessionId: currentSessionId,
actionType: 'send_message',
actionDescription: `发送消息: ${content.substring(0, 100)}${content.length > 100 ? '...' : ''}`,
actionType: "send_message",
actionDescription: `发送消息: ${content.substring(0, 100)}${content.length > 100 ? "..." : ""}`,
messageCount,
};
await workflowManager.postAction(actionContext, '', error as string);
await workflowManager.postAction(
actionContext,
"",
error as string,
);
} catch (workflowError) {
console.warn('工作流 Post-Action 执行失败:', workflowError);
console.warn("工作流 Post-Action 执行失败:", workflowError);
}
}
// 设置错误状态
get().setMessageError(assistantMessageId, error as string);
}
@@ -538,8 +568,12 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const currentMessages = messages[currentSessionId] || [];
const updatedMessages = currentMessages.map((m) =>
m.id === streaming.currentMessageId
? { ...m, status: 'complete' as const, content: streaming.partialContent }
: m
? {
...m,
status: "complete" as const,
content: streaming.partialContent,
}
: m,
);
set((state) => ({
@@ -572,7 +606,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const updatedMessages = currentMessages.map((m) =>
m.id === streaming.currentMessageId
? { ...m, content: get().streaming.partialContent }
: m
: m,
);
set((state) => ({
@@ -585,7 +619,13 @@ export const useGeneralChatStore = create<GeneralChatState>()(
},
finalizeMessage: async (metadata?: MessageMetadata) => {
const { streaming, currentSessionId, messages, workflowManagers, workflowEnabled } = get();
const {
streaming,
currentSessionId,
messages,
workflowManagers,
workflowEnabled,
} = get();
if (!streaming.currentMessageId || !currentSessionId) {
return;
@@ -596,11 +636,11 @@ export const useGeneralChatStore = create<GeneralChatState>()(
m.id === streaming.currentMessageId
? {
...m,
status: 'complete' as const,
status: "complete" as const,
content: streaming.partialContent,
metadata,
}
: m
: m,
);
set((state) => ({
@@ -618,14 +658,17 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const messageCount = updatedMessages.length;
const actionContext = {
sessionId: currentSessionId,
actionType: 'receive_message',
actionDescription: `接收 AI 响应: ${streaming.partialContent.substring(0, 100)}${streaming.partialContent.length > 100 ? '...' : ''}`,
actionType: "receive_message",
actionDescription: `接收 AI 响应: ${streaming.partialContent.substring(0, 100)}${streaming.partialContent.length > 100 ? "..." : ""}`,
messageCount,
};
await workflowManager.postAction(actionContext, streaming.partialContent);
await workflowManager.postAction(
actionContext,
streaming.partialContent,
);
} catch (error) {
console.warn('工作流 Post-Action 执行失败:', error);
console.warn("工作流 Post-Action 执行失败:", error);
}
}
},
@@ -655,7 +698,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const currentMessages = messages[currentSessionId] || [];
const updatedMessages = currentMessages.map((m) =>
m.id === messageId ? { ...m, ...updates } : m
m.id === messageId ? { ...m, ...updates } : m,
);
set((state) => ({
@@ -671,7 +714,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
streaming: {
isStreaming: true,
currentMessageId: messageId,
partialContent: '',
partialContent: "",
},
});
},
@@ -681,15 +724,14 @@ export const useGeneralChatStore = create<GeneralChatState>()(
if (!currentSessionId) return;
// 如果传入的是字符串,解析为 ErrorInfo
const errorInfo: ErrorInfo = typeof error === 'string'
? parseApiError(error)
: error;
const errorInfo: ErrorInfo =
typeof error === "string" ? parseApiError(error) : error;
const currentMessages = messages[currentSessionId] || [];
const updatedMessages = currentMessages.map((m) =>
m.id === messageId
? { ...m, status: 'error' as const, error: errorInfo }
: m
m.id === messageId
? { ...m, status: "error" as const, error: errorInfo }
: m,
);
set((state) => ({
@@ -707,19 +749,21 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const currentMessages = messages[currentSessionId] || [];
const errorMessage = currentMessages.find((m) => m.id === messageId);
if (!errorMessage || errorMessage.status !== 'error') {
if (!errorMessage || errorMessage.status !== "error") {
return;
}
// 找到错误消息之前的用户消息
const messageIndex = currentMessages.findIndex((m) => m.id === messageId);
const messageIndex = currentMessages.findIndex(
(m) => m.id === messageId,
);
if (messageIndex <= 0) return;
// 查找最近的用户消息
let userMessage: Message | null = null;
for (let i = messageIndex - 1; i >= 0; i--) {
if (currentMessages[i].role === 'user') {
if (currentMessages[i].role === "user") {
userMessage = currentMessages[i];
break;
}
@@ -729,9 +773,14 @@ export const useGeneralChatStore = create<GeneralChatState>()(
// 清除错误状态,将消息状态改为 pending
const updatedMessages = currentMessages.map((m) =>
m.id === messageId
? { ...m, status: 'pending' as const, error: undefined, content: '' }
: m
m.id === messageId
? {
...m,
status: "pending" as const,
error: undefined,
content: "",
}
: m,
);
set((state) => ({
@@ -755,9 +804,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const currentMessages = messages[currentSessionId] || [];
const updatedMessages = currentMessages.map((m) =>
m.id === messageId
? { ...m, error: undefined }
: m
m.id === messageId ? { ...m, error: undefined } : m,
);
set((state) => ({
@@ -772,10 +819,15 @@ export const useGeneralChatStore = create<GeneralChatState>()(
loadMoreMessages: async (sessionId: string) => {
const { messages, pagination } = get();
const currentPagination = pagination[sessionId] || { ...DEFAULT_PAGINATION_STATE };
const currentPagination = pagination[sessionId] || {
...DEFAULT_PAGINATION_STATE,
};
// 如果正在加载或没有更多消息,直接返回
if (currentPagination.isLoadingMore || !currentPagination.hasMoreMessages) {
if (
currentPagination.isLoadingMore ||
!currentPagination.hasMoreMessages
) {
return;
}
@@ -793,23 +845,32 @@ export const useGeneralChatStore = create<GeneralChatState>()(
try {
// 获取当前会话的消息列表
const currentMessages = messages[sessionId] || [];
// 获取最早消息的 ID 用于分页
const oldestMessage = currentMessages.length > 0 ? currentMessages[0] : null;
const oldestMessage =
currentMessages.length > 0 ? currentMessages[0] : null;
const beforeId = oldestMessage?.id || null;
// 调用 Tauri 命令获取更多消息
const { invoke } = await import('@tauri-apps/api/core');
const olderMessages = await invoke<Array<{
id: string;
session_id: string;
role: string;
content: string;
blocks: Array<{ type: string; content: string; language?: string; filename?: string; mime_type?: string }> | null;
status: string;
created_at: number;
metadata: Record<string, unknown> | null;
}>>('general_chat_get_messages', {
const { invoke } = await import("@tauri-apps/api/core");
const olderMessages = await invoke<
Array<{
id: string;
session_id: string;
role: string;
content: string;
blocks: Array<{
type: string;
content: string;
language?: string;
filename?: string;
mime_type?: string;
}> | null;
status: string;
created_at: number;
metadata: Record<string, unknown> | null;
}>
>("general_chat_get_messages", {
sessionId,
limit: currentPagination.pageSize,
beforeId,
@@ -819,32 +880,38 @@ export const useGeneralChatStore = create<GeneralChatState>()(
const convertedMessages: Message[] = olderMessages.map((msg) => ({
id: msg.id,
sessionId: msg.session_id,
role: msg.role as Message['role'],
role: msg.role as Message["role"],
content: msg.content,
blocks: msg.blocks?.map((b) => ({
type: b.type as Message['blocks'][0]['type'],
type: b.type as Message["blocks"][0]["type"],
content: b.content,
language: b.language,
filename: b.filename,
mimeType: b.mime_type,
})) || [{ type: 'text' as const, content: msg.content }],
status: msg.status as Message['status'],
})) || [{ type: "text" as const, content: msg.content }],
status: msg.status as Message["status"],
createdAt: msg.created_at,
metadata: msg.metadata ? {
model: msg.metadata.model as string | undefined,
tokens: msg.metadata.tokens as number | undefined,
duration: msg.metadata.duration as number | undefined,
} : undefined,
metadata: msg.metadata
? {
model: msg.metadata.model as string | undefined,
tokens: msg.metadata.tokens as number | undefined,
duration: msg.metadata.duration as number | undefined,
}
: undefined,
}));
// 判断是否还有更多消息
const hasMore = convertedMessages.length >= currentPagination.pageSize;
const hasMore =
convertedMessages.length >= currentPagination.pageSize;
// 将旧消息添加到列表前面
set((state) => ({
messages: {
...state.messages,
[sessionId]: [...convertedMessages, ...(state.messages[sessionId] || [])],
[sessionId]: [
...convertedMessages,
...(state.messages[sessionId] || []),
],
},
pagination: {
...state.pagination,
@@ -852,12 +919,15 @@ export const useGeneralChatStore = create<GeneralChatState>()(
...currentPagination,
isLoadingMore: false,
hasMoreMessages: hasMore,
oldestMessageId: convertedMessages.length > 0 ? convertedMessages[0].id : currentPagination.oldestMessageId,
oldestMessageId:
convertedMessages.length > 0
? convertedMessages[0].id
: currentPagination.oldestMessageId,
},
},
}));
} catch (error) {
console.error('加载更多消息失败:', error);
console.error("加载更多消息失败:", error);
// 加载失败时重置加载状态
set((state) => ({
pagination: {
@@ -871,12 +941,17 @@ export const useGeneralChatStore = create<GeneralChatState>()(
}
},
setPaginationState: (sessionId: string, state: Partial<PaginationState>) => {
setPaginationState: (
sessionId: string,
state: Partial<PaginationState>,
) => {
set((prev) => ({
pagination: {
...prev.pagination,
[sessionId]: {
...(prev.pagination[sessionId] || { ...DEFAULT_PAGINATION_STATE }),
...(prev.pagination[sessionId] || {
...DEFAULT_PAGINATION_STATE,
}),
...state,
},
},
@@ -1038,9 +1113,13 @@ export const useGeneralChatStore = create<GeneralChatState>()(
// ========== 三阶段工作流操作实现 ==========
initializeWorkflow: async (sessionId: string, projectName: string, goal: string) => {
initializeWorkflow: async (
sessionId: string,
projectName: string,
goal: string,
) => {
const { workflowManagers } = get();
// 如果已存在工作流管理器,直接返回
if (workflowManagers[sessionId]) {
return;
@@ -1048,7 +1127,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
// 创建新的工作流管理器
const workflowManager = new ThreeStageWorkflowManager(sessionId);
// 初始化工作流配置
const workflowConfig = {
sessionId,
@@ -1057,52 +1136,52 @@ export const useGeneralChatStore = create<GeneralChatState>()(
phases: [
{
number: 1,
name: '需求理解与发现',
status: 'in_progress' as const,
name: "需求理解与发现",
status: "in_progress" as const,
tasks: [
'理解用户意图和需求',
'识别约束条件和依赖关系',
'记录发现到 findings.md',
"理解用户意图和需求",
"识别约束条件和依赖关系",
"记录发现到 findings.md",
],
},
{
number: 2,
name: '规划与结构设计',
status: 'pending' as const,
name: "规划与结构设计",
status: "pending" as const,
tasks: [
'定义技术方案和架构',
'创建项目结构(如需要)',
'记录关键决策和理由',
"定义技术方案和架构",
"创建项目结构(如需要)",
"记录关键决策和理由",
],
},
{
number: 3,
name: '实施执行',
status: 'pending' as const,
name: "实施执行",
status: "pending" as const,
tasks: [
'按计划逐步执行',
'执行前先写入文件',
'增量测试并记录结果',
"按计划逐步执行",
"执行前先写入文件",
"增量测试并记录结果",
],
},
{
number: 4,
name: '测试与验证',
status: 'pending' as const,
name: "测试与验证",
status: "pending" as const,
tasks: [
'验证所有需求已满足',
'记录测试结果到 progress.md',
'修复发现的问题并记录解决方案',
"验证所有需求已满足",
"记录测试结果到 progress.md",
"修复发现的问题并记录解决方案",
],
},
{
number: 5,
name: '交付与完成',
status: 'pending' as const,
name: "交付与完成",
status: "pending" as const,
tasks: [
'审查所有输出文件和交付物',
'确保完整性和质量',
'向用户交付最终结果',
"审查所有输出文件和交付物",
"确保完整性和质量",
"向用户交付最终结果",
],
},
],
@@ -1111,7 +1190,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
try {
// 初始化工作流
await workflowManager.initializeWorkflow(workflowConfig);
// 保存到状态中
set((state) => ({
workflowManagers: {
@@ -1122,7 +1201,7 @@ export const useGeneralChatStore = create<GeneralChatState>()(
console.log(`三阶段工作流已为会话 ${sessionId} 初始化`);
} catch (error) {
console.error('工作流初始化失败:', error);
console.error("工作流初始化失败:", error);
throw error;
}
},
@@ -1143,13 +1222,13 @@ export const useGeneralChatStore = create<GeneralChatState>()(
shouldAutoEnableWorkflow: (sessionId: string) => {
const { messages, workflowThreshold, workflowEnabled } = get();
if (!workflowEnabled) return false;
const messageCount = (messages[sessionId] || []).length;
return messageCount >= workflowThreshold;
},
}),
{
name: 'general-chat-storage',
name: "general-chat-storage",
storage: createJSONStorage(() => localStorage),
// 只持久化 UI 状态、当前会话 ID 和 Provider 选择
partialize: (state) => ({
@@ -1160,8 +1239,8 @@ export const useGeneralChatStore = create<GeneralChatState>()(
selectedModelId: state.providerSelection.selectedModelId,
},
}),
}
)
},
),
);
// ============================================================================
@@ -1200,7 +1279,8 @@ export const useUIState = () => useGeneralChatStore((state) => state.ui);
/**
* 获取画布状态
*/
export const useCanvasState = () => useGeneralChatStore((state) => state.canvas);
export const useCanvasState = () =>
useGeneralChatStore((state) => state.canvas);
/**
* 获取会话列表
@@ -1238,7 +1318,7 @@ export const useSelectedModelId = () =>
export const useCurrentPagination = () =>
useGeneralChatStore((state) => {
const { pagination, currentSessionId } = state;
return currentSessionId
return currentSessionId
? pagination[currentSessionId] || { ...DEFAULT_PAGINATION_STATE }
: { ...DEFAULT_PAGINATION_STATE };
});
+78 -56
View File
@@ -38,19 +38,19 @@ export interface Session {
* 消息角色
* @description 标识消息的发送者类型
*/
export type MessageRole = 'user' | 'assistant' | 'system';
export type MessageRole = "user" | "assistant" | "system";
/**
* 消息状态
* @description 标识消息的当前处理状态
*/
export type MessageStatus = 'pending' | 'streaming' | 'complete' | 'error';
export type MessageStatus = "pending" | "streaming" | "complete" | "error";
/**
* 消息内容块类型
* @description 标识内容块的类型
*/
export type ContentBlockType = 'text' | 'code' | 'image' | 'file';
export type ContentBlockType = "text" | "code" | "image" | "file";
/**
* 内容块
@@ -92,14 +92,14 @@ export interface MessageMetadata {
* @requirements 9.2, 9.3
*/
export type ErrorCode =
| 'NETWORK_ERROR' // 网络连接错误
| 'TIMEOUT' // 请求超时
| 'RATE_LIMIT' // 请求频率限制
| 'TOKEN_LIMIT' // Token 数量超限
| 'AUTH_ERROR' // 认证错误
| 'SERVER_ERROR' // 服务器错误
| 'PROVIDER_ERROR' // Provider 错误
| 'UNKNOWN_ERROR'; // 未知错误
| "NETWORK_ERROR" // 网络连接错误
| "TIMEOUT" // 请求超时
| "RATE_LIMIT" // 请求频率限制
| "TOKEN_LIMIT" // Token 数量超限
| "AUTH_ERROR" // 认证错误
| "SERVER_ERROR" // 服务器错误
| "PROVIDER_ERROR" // Provider 错误
| "UNKNOWN_ERROR"; // 未知错误
/**
* 错误信息接口
@@ -129,39 +129,42 @@ export interface ErrorInfo {
*/
export const createErrorInfo = (
code: ErrorCode,
details?: string
details?: string,
): ErrorInfo => {
const errorMessages: Record<ErrorCode, { message: string; retryable: boolean }> = {
const errorMessages: Record<
ErrorCode,
{ message: string; retryable: boolean }
> = {
NETWORK_ERROR: {
message: '网络连接已断开,请检查网络设置',
message: "网络连接已断开,请检查网络设置",
retryable: true,
},
TIMEOUT: {
message: '请求超时,请点击重试',
message: "请求超时,请点击重试",
retryable: true,
},
RATE_LIMIT: {
message: '请求过于频繁,请稍后重试',
message: "请求过于频繁,请稍后重试",
retryable: true,
},
TOKEN_LIMIT: {
message: '对话过长,建议新建会话',
message: "对话过长,建议新建会话",
retryable: false,
},
AUTH_ERROR: {
message: '认证失败,请检查 Provider 配置',
message: "认证失败,请检查 Provider 配置",
retryable: false,
},
SERVER_ERROR: {
message: '服务器错误,请稍后重试',
message: "服务器错误,请稍后重试",
retryable: true,
},
PROVIDER_ERROR: {
message: 'AI 服务暂时不可用,请稍后重试',
message: "AI 服务暂时不可用,请稍后重试",
retryable: true,
},
UNKNOWN_ERROR: {
message: '发生未知错误,请重试',
message: "发生未知错误,请重试",
retryable: true,
},
};
@@ -187,34 +190,53 @@ export const parseApiError = (error: unknown): ErrorInfo => {
const lowerError = errorStr.toLowerCase();
// 根据错误信息判断错误类型
if (lowerError.includes('network') || lowerError.includes('fetch') || lowerError.includes('connection')) {
return createErrorInfo('NETWORK_ERROR', errorStr);
if (
lowerError.includes("network") ||
lowerError.includes("fetch") ||
lowerError.includes("connection")
) {
return createErrorInfo("NETWORK_ERROR", errorStr);
}
if (lowerError.includes('timeout') || lowerError.includes('timed out')) {
return createErrorInfo('TIMEOUT', errorStr);
if (lowerError.includes("timeout") || lowerError.includes("timed out")) {
return createErrorInfo("TIMEOUT", errorStr);
}
if (lowerError.includes('rate limit') || lowerError.includes('429') || lowerError.includes('too many')) {
if (
lowerError.includes("rate limit") ||
lowerError.includes("429") ||
lowerError.includes("too many")
) {
const retryMatch = errorStr.match(/(\d+)\s*(?:seconds?|s)/i);
const info = createErrorInfo('RATE_LIMIT', errorStr);
const info = createErrorInfo("RATE_LIMIT", errorStr);
if (retryMatch) {
info.retryAfter = parseInt(retryMatch[1], 10);
}
return info;
}
if (lowerError.includes('token') && (lowerError.includes('limit') || lowerError.includes('exceed'))) {
return createErrorInfo('TOKEN_LIMIT', errorStr);
if (
lowerError.includes("token") &&
(lowerError.includes("limit") || lowerError.includes("exceed"))
) {
return createErrorInfo("TOKEN_LIMIT", errorStr);
}
if (lowerError.includes('401') || lowerError.includes('unauthorized') || lowerError.includes('auth')) {
return createErrorInfo('AUTH_ERROR', errorStr);
if (
lowerError.includes("401") ||
lowerError.includes("unauthorized") ||
lowerError.includes("auth")
) {
return createErrorInfo("AUTH_ERROR", errorStr);
}
if (lowerError.includes('500') || lowerError.includes('server error') || lowerError.includes('internal')) {
return createErrorInfo('SERVER_ERROR', errorStr);
if (
lowerError.includes("500") ||
lowerError.includes("server error") ||
lowerError.includes("internal")
) {
return createErrorInfo("SERVER_ERROR", errorStr);
}
if (lowerError.includes('provider') || lowerError.includes('model')) {
return createErrorInfo('PROVIDER_ERROR', errorStr);
if (lowerError.includes("provider") || lowerError.includes("model")) {
return createErrorInfo("PROVIDER_ERROR", errorStr);
}
return createErrorInfo('UNKNOWN_ERROR', errorStr);
return createErrorInfo("UNKNOWN_ERROR", errorStr);
};
/**
@@ -252,7 +274,7 @@ export interface Message {
* 画布内容类型
* @description 标识画布中显示的内容类型
*/
export type CanvasContentType = 'code' | 'file' | 'markdown' | 'empty';
export type CanvasContentType = "code" | "file" | "markdown" | "empty";
/**
* 画布状态
@@ -476,11 +498,11 @@ export interface ImagePreviewState {
* 支持的图片格式
*/
export const SUPPORTED_IMAGE_TYPES = [
'image/jpeg',
'image/jpg',
'image/png',
'image/gif',
'image/webp'
"image/jpeg",
"image/jpg",
"image/png",
"image/gif",
"image/webp",
] as const;
/**
@@ -519,17 +541,17 @@ export const isValidImageSize = (file: File): boolean => {
export const fileToImageData = (file: File): Promise<ImageData> => {
return new Promise((resolve, reject) => {
const reader = new FileReader();
reader.onload = (e) => {
const result = e.target?.result as string;
if (!result) {
reject(new Error('无法读取文件'));
reject(new Error("无法读取文件"));
return;
}
// 提取 base64 数据(去掉 data:image/xxx;base64, 前缀)
const base64Data = result.split(',')[1];
const base64Data = result.split(",")[1];
// 创建图片元素获取尺寸
const img = new Image();
img.onload = () => {
@@ -544,18 +566,18 @@ export const fileToImageData = (file: File): Promise<ImageData> => {
};
resolve(imageData);
};
img.onerror = () => {
reject(new Error('无法解析图片'));
reject(new Error("无法解析图片"));
};
img.src = result;
};
reader.onerror = () => {
reject(new Error('文件读取失败'));
reject(new Error("文件读取失败"));
};
reader.readAsDataURL(file);
});
};
@@ -595,8 +617,8 @@ export const DEFAULT_UI_STATE: UIState = {
*/
export const DEFAULT_CANVAS_STATE: CanvasState = {
isOpen: false,
contentType: 'empty',
content: '',
contentType: "empty",
content: "",
isEditing: false,
};
@@ -606,5 +628,5 @@ export const DEFAULT_CANVAS_STATE: CanvasState = {
export const DEFAULT_STREAMING_STATE: StreamingState = {
isStreaming: false,
currentMessageId: null,
partialContent: '',
partialContent: "",
};
-18
View File
@@ -7,8 +7,6 @@ import {
User,
FileCode,
Activity,
Cpu,
Globe,
Minimize2,
Monitor,
Maximize2,
@@ -83,22 +81,6 @@ export const onboardingPlugins: OnboardingPlugin[] = [
downloadUrl:
"https://github.com/aiclientproxy/flow-monitor/releases/latest/download/flow-monitor-plugin.zip",
},
{
id: "machine-id-tool",
name: "机器码管理工具",
description: "查看、修改和管理系统机器码,支持跨平台操作",
icon: Cpu,
downloadUrl:
"https://github.com/aiclientproxy/MachineIdTool/releases/latest/download/machine-id-tool-plugin.zip",
},
{
id: "browser-interception",
name: "浏览器拦截器",
description: "拦截桌面应用的浏览器启动,支持手动复制 URL 到指纹浏览器",
icon: Globe,
downloadUrl:
"https://github.com/aiclientproxy/browser-interception/releases/latest/download/browser-interception-plugin.zip",
},
];
/**
@@ -10,7 +10,6 @@
import React, { useState, useEffect } from "react";
import { AlertCircle, Package, Loader2, ExternalLink } from "lucide-react";
import { safeInvoke } from "@/lib/dev-bridge";
import { BrowserInterceptorTool } from "@/components/tools/browser-interceptor/BrowserInterceptorTool";
import { FlowMonitorPage } from "@/pages";
import { ConfigManagementPage } from "@/components/config/ConfigManagementPage";
import { PluginUIRenderer as DynamicPluginRenderer } from "@/lib/plugin-loader/PluginUIRenderer";
@@ -174,7 +173,6 @@ const builtinPluginComponents: Record<
string,
React.ComponentType<{ onNavigate?: (page: Page) => void }>
> = {
"browser-interception": BrowserInterceptorTool,
"flow-monitor": FlowMonitorPage,
"config-switch": ConfigManagementPage,
};
+38 -22
View File
@@ -9,7 +9,7 @@
*/
import React, { useState, useEffect, useCallback } from "react";
import { Package, Loader2, type LucideIcon } from "lucide-react";
import { Package, Loader2, ArrowLeft, type LucideIcon } from "lucide-react";
import * as LucideIcons from "lucide-react";
import { safeInvoke } from "@/lib/dev-bridge";
import { Badge } from "@/components/ui/badge";
@@ -25,24 +25,8 @@ import { getPluginsForSurface, type PluginUIInfo } from "@/lib/api/pluginUI";
import { PluginInstallDialog } from "@/components/plugins/PluginInstallDialog";
import { ToolCardContextMenu } from "./ToolCardContextMenu";
import { toast } from "sonner";
/**
* 页面类型定义
*
* 支持静态页面和动态插件页面
* - 静态页面: 预定义的页面标识符
* - 动态插件页面: `plugin:${string}` 格式,如 "plugin:machine-id-tool"
*
* _需求: 2.2, 3.2_
*/
type Page =
| "provider-pool"
| "api-server"
| "agent"
| "tools"
| "settings"
| "plugins"
| `plugin:${string}`;
import { ImageAnalysisTool } from "./image-analysis";
import type { Page } from "@/types/page";
interface ToolsPageProps {
/**
@@ -163,7 +147,15 @@ function ToolCard({
/**
* 内置工具列表
*/
const builtinTools: DynamicToolCard[] = [];
const builtinTools: DynamicToolCard[] = [
{
id: "image-analysis",
title: "图像分析",
description: "使用 AI 分析图片内容,支持视觉理解和描述",
icon: "Image",
source: "builtin",
},
];
/**
* 占位工具列表 (敬请期待)
@@ -199,6 +191,7 @@ export function ToolsPage({ onNavigate }: ToolsPageProps) {
const [pluginTools, setPluginTools] = useState<DynamicToolCard[]>([]);
const [loading, setLoading] = useState(true);
const [showInstallDialog, setShowInstallDialog] = useState(false);
const [selectedTool, setSelectedTool] = useState<string | null>(null);
// 加载插件工具和已安装插件列表
const loadPluginTools = useCallback(async () => {
@@ -279,11 +272,18 @@ export function ToolsPage({ onNavigate }: ToolsPageProps) {
// 插件工具: 导航到 plugin:xxx 页面
onNavigate(`plugin:${tool.pluginId}`);
} else {
// 内置工具: 导航到对应页面
onNavigate(tool.id as Page);
// 内置工具: 在当前页面显示工具组件
setSelectedTool(tool.id);
}
};
/**
* 返回工具列表
*/
const handleBackToList = () => {
setSelectedTool(null);
};
/**
* 渲染工具图标
*/
@@ -296,6 +296,22 @@ export function ToolsPage({ onNavigate }: ToolsPageProps) {
);
};
// 如果选中了内置工具,显示该工具
if (selectedTool === "image-analysis") {
return (
<div className="space-y-4">
{/* 返回按钮 */}
<Button variant="ghost" onClick={handleBackToList} className="mb-4">
<ArrowLeft className="w-4 h-4 mr-2" />
返回工具列表
</Button>
{/* 工具组件 */}
<ImageAnalysisTool />
</div>
);
}
return (
<div className="space-y-6">
<div className="flex items-center justify-between">
@@ -5,8 +5,21 @@
*/
import React, { useState, useCallback } from "react";
import { Upload, Image as ImageIcon, Sparkles, Loader2, X, AlertCircle } from "lucide-react";
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card";
import {
Upload,
Image as ImageIcon,
Sparkles,
Loader2,
X,
AlertCircle,
} from "lucide-react";
import {
Card,
CardContent,
CardDescription,
CardHeader,
CardTitle,
} from "@/components/ui/card";
import { Button } from "@/components/ui/button";
import { Badge } from "@/components/ui/badge";
import { toast } from "sonner";
@@ -24,32 +37,35 @@ export function ImageAnalysisTool() {
const [analyzing, setAnalyzing] = useState(false);
const [result, setResult] = useState<AnalysisResult | null>(null);
const handleImageSelect = useCallback((e: React.ChangeEvent<HTMLInputElement>) => {
const file = e.target.files?.[0];
if (!file) return;
const handleImageSelect = useCallback(
(e: React.ChangeEvent<HTMLInputElement>) => {
const file = e.target.files?.[0];
if (!file) return;
if (!file.type.startsWith("image/")) {
toast.error("请选择图片文件");
return;
}
if (!file.type.startsWith("image/")) {
toast.error("请选择图片文件");
return;
}
if (file.size > 10 * 1024 * 1024) {
toast.error("图片大小不能超过 10MB");
return;
}
if (file.size > 10 * 1024 * 1024) {
toast.error("图片大小不能超过 10MB");
return;
}
setSelectedImage(file);
setSelectedImage(file);
// 创建预览
const reader = new FileReader();
reader.onloadend = () => {
setImagePreview(reader.result as string);
};
reader.readAsDataURL(file);
// 创建预览
const reader = new FileReader();
reader.onloadend = () => {
setImagePreview(reader.result as string);
};
reader.readAsDataURL(file);
// 重置结果
setResult(null);
}, []);
// 重置结果
setResult(null);
},
[],
);
const handleClearImage = useCallback(() => {
setSelectedImage(null);
@@ -113,8 +129,12 @@ export function ImageAnalysisTool() {
上传图片,AI 帮您分析内容、识别物体、提取文字
</p>
<div className="flex items-center justify-center gap-2">
<Badge className="bg-primary/10 text-primary border-primary/20">AI 驱动</Badge>
<Badge variant="outline" className="border-muted">支持多模态</Badge>
<Badge className="bg-primary/10 text-primary border-primary/20">
AI 驱动
</Badge>
<Badge variant="outline" className="border-muted">
支持多模态
</Badge>
</div>
</div>
@@ -185,7 +205,7 @@ export function ImageAnalysisTool() {
<div className="flex items-center gap-2">
<span className="font-medium">{selectedImage?.name}</span>
<Badge variant="outline" className="text-xs">
{selectedImage?.type.split('/')[1]?.toUpperCase()}
{selectedImage?.type.split("/")[1]?.toUpperCase()}
</Badge>
</div>
<span className="text-muted-foreground">
@@ -205,9 +225,7 @@ export function ImageAnalysisTool() {
<Sparkles className="w-5 h-5 text-primary" />
分析设置
</CardTitle>
<CardDescription>
描述您想了解的图片内容
</CardDescription>
<CardDescription>描述您想了解的图片内容</CardDescription>
</CardHeader>
<CardContent className="space-y-6">
{/* 输入提示 */}
@@ -265,7 +283,8 @@ export function ImageAnalysisTool() {
<div className="flex items-start gap-2 p-3 rounded-lg bg-muted/50 text-sm text-muted-foreground">
<AlertCircle className="w-4 h-4 mt-0.5 flex-shrink-0" />
<p>
大图片可能导致分析失败,建议使用小于 2MB 的图片以获得最佳体验。
大图片可能导致分析失败,建议使用小于 2MB
的图片以获得最佳体验。
</p>
</div>
</CardContent>
@@ -285,7 +304,9 @@ export function ImageAnalysisTool() {
<CardContent>
{result.error ? (
<div className="p-4 bg-destructive/10 border border-destructive/20 rounded-lg">
<p className="text-destructive whitespace-pre-wrap">{result.error}</p>
<p className="text-destructive whitespace-pre-wrap">
{result.error}
</p>
</div>
) : (
<div className="prose prose-sm max-w-none dark:prose-invert">
@@ -0,0 +1 @@
export { ImageAnalysisTool } from "./ImageAnalysisTool";
+209 -147
View File
@@ -1,14 +1,14 @@
/**
* 三阶段工作流 React Hook
*
*
* 提供在 React 组件中使用三阶段工作流的便捷接口
*/
import { useState, useCallback, useRef, useEffect } from 'react';
import ThreeStageWorkflowManager, {
type WorkflowConfig,
type ActionContext
} from '../lib/workflow/threeStageWorkflow';
import { useState, useCallback, useRef, useEffect } from "react";
import ThreeStageWorkflowManager, {
type WorkflowConfig,
type ActionContext,
} from "../lib/workflow/threeStageWorkflow";
export interface UseThreeStageWorkflowOptions {
sessionId: string;
@@ -29,9 +29,21 @@ export interface WorkflowActions {
initializeWorkflow: (config: WorkflowConfig) => Promise<void>;
preAction: (context: ActionContext) => Promise<string>;
executeAction: (context: ActionContext, result: string) => Promise<void>;
postAction: (context: ActionContext, result: string, error?: string) => Promise<string>;
updatePhaseStatus: (phaseNumber: number, status: 'pending' | 'in_progress' | 'complete', notes?: string) => Promise<void>;
recordFinding: (title: string, content: string, tags?: string[]) => Promise<void>;
postAction: (
context: ActionContext,
result: string,
error?: string,
) => Promise<string>;
updatePhaseStatus: (
phaseNumber: number,
status: "pending" | "in_progress" | "complete",
notes?: string,
) => Promise<void>;
recordFinding: (
title: string,
content: string,
tags?: string[],
) => Promise<void>;
recordDecision: (decision: string, rationale: string) => Promise<void>;
checkCompletion: () => Promise<{ isComplete: boolean; summary: string }>;
finalizeWorkflow: () => Promise<string>;
@@ -42,9 +54,11 @@ export interface WorkflowActions {
/**
* 三阶段工作流 Hook
*/
export function useThreeStageWorkflow(options: UseThreeStageWorkflowOptions): [WorkflowState, WorkflowActions] {
export function useThreeStageWorkflow(
options: UseThreeStageWorkflowOptions,
): [WorkflowState, WorkflowActions] {
const { sessionId, autoInitialize = false, defaultConfig } = options;
const [state, setState] = useState<WorkflowState>({
isInitialized: false,
isLoading: false,
@@ -65,195 +79,241 @@ export function useThreeStageWorkflow(options: UseThreeStageWorkflowOptions): [W
const updateStats = useCallback(async () => {
if (!workflowManagerRef.current) return;
try {
const stats = await workflowManagerRef.current.getSessionStats();
setState(prev => ({
setState((prev) => ({
...prev,
visualOperationCount: stats.visualOperationCount,
errorAttempts: stats.errorAttempts,
}));
} catch (error) {
console.warn('更新统计信息失败:', error);
console.warn("更新统计信息失败:", error);
}
}, []);
const setLoading = useCallback((loading: boolean) => {
setState(prev => ({ ...prev, isLoading: loading }));
setState((prev) => ({ ...prev, isLoading: loading }));
}, []);
const setError = useCallback((error: string | null) => {
setState(prev => ({ ...prev, error }));
setState((prev) => ({ ...prev, error }));
}, []);
const initializeWorkflow = useCallback(async (config: WorkflowConfig) => {
if (!workflowManagerRef.current) return;
setLoading(true);
setError(null);
try {
await workflowManagerRef.current.initializeWorkflow(config);
setState(prev => ({
...prev,
isInitialized: true,
currentPhase: 1,
}));
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : '初始化工作流失败');
} finally {
setLoading(false);
}
}, [setLoading, setError, updateStats]);
const initializeWorkflow = useCallback(
async (config: WorkflowConfig) => {
if (!workflowManagerRef.current) return;
setLoading(true);
setError(null);
try {
await workflowManagerRef.current.initializeWorkflow(config);
setState((prev) => ({
...prev,
isInitialized: true,
currentPhase: 1,
}));
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : "初始化工作流失败");
} finally {
setLoading(false);
}
},
[setLoading, setError, updateStats],
);
// 自动初始化
useEffect(() => {
if (autoInitialize && defaultConfig && !state.isInitialized && workflowManagerRef.current) {
if (
autoInitialize &&
defaultConfig &&
!state.isInitialized &&
workflowManagerRef.current
) {
const config: WorkflowConfig = {
sessionId,
projectName: defaultConfig.projectName || '新项目',
goal: defaultConfig.goal || '待定义目标',
projectName: defaultConfig.projectName || "新项目",
goal: defaultConfig.goal || "待定义目标",
phases: defaultConfig.phases || [
{
number: 1,
name: '需求分析',
status: 'in_progress',
tasks: ['理解用户需求', '识别约束条件', '记录发现'],
name: "需求分析",
status: "in_progress",
tasks: ["理解用户需求", "识别约束条件", "记录发现"],
},
{
number: 2,
name: '方案设计',
status: 'pending',
tasks: ['制定技术方案', '创建项目结构', '记录关键决策'],
name: "方案设计",
status: "pending",
tasks: ["制定技术方案", "创建项目结构", "记录关键决策"],
},
{
number: 3,
name: '实施执行',
status: 'pending',
tasks: ['按计划执行', '增量测试', '记录进展'],
name: "实施执行",
status: "pending",
tasks: ["按计划执行", "增量测试", "记录进展"],
},
],
};
initializeWorkflow(config);
}
}, [autoInitialize, defaultConfig, state.isInitialized, sessionId, initializeWorkflow]);
}, [
autoInitialize,
defaultConfig,
state.isInitialized,
sessionId,
initializeWorkflow,
]);
const preAction = useCallback(async (context: ActionContext): Promise<string> => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
setLoading(true);
setError(null);
try {
const result = await workflowManagerRef.current.preAction(context);
await updateStats();
return result;
} catch (error) {
const errorMessage = error instanceof Error ? error.message : 'Pre-Action 执行失败';
setError(errorMessage);
throw error;
} finally {
setLoading(false);
}
}, [setLoading, setError, updateStats]);
const preAction = useCallback(
async (context: ActionContext): Promise<string> => {
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
const executeAction = useCallback(async (context: ActionContext, result: string) => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
try {
await workflowManagerRef.current.executeAction(context, result);
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : 'Action 执行失败');
throw error;
}
}, [setError, updateStats]);
setLoading(true);
setError(null);
const postAction = useCallback(async (context: ActionContext, result: string, error?: string): Promise<string> => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
setLoading(true);
try {
const message = await workflowManagerRef.current.postAction(context, result, error);
await updateStats();
return message;
} catch (err) {
const errorMessage = err instanceof Error ? err.message : 'Post-Action 执行失败';
setError(errorMessage);
throw err;
} finally {
setLoading(false);
}
}, [setLoading, setError, updateStats]);
try {
const result = await workflowManagerRef.current.preAction(context);
await updateStats();
return result;
} catch (error) {
const errorMessage =
error instanceof Error ? error.message : "Pre-Action 执行失败";
setError(errorMessage);
throw error;
} finally {
setLoading(false);
}
},
[setLoading, setError, updateStats],
);
const updatePhaseStatus = useCallback(async (
phaseNumber: number,
status: 'pending' | 'in_progress' | 'complete',
notes?: string
) => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
try {
await workflowManagerRef.current.updatePhaseStatus(phaseNumber, status, notes);
setState(prev => ({ ...prev, currentPhase: phaseNumber }));
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : '更新阶段状态失败');
throw error;
}
}, [setError, updateStats]);
const executeAction = useCallback(
async (context: ActionContext, result: string) => {
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
const recordFinding = useCallback(async (title: string, content: string, tags: string[] = []) => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
try {
await workflowManagerRef.current.recordFinding(title, content, tags);
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : '记录发现失败');
throw error;
}
}, [setError, updateStats]);
try {
await workflowManagerRef.current.executeAction(context, result);
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : "Action 执行失败");
throw error;
}
},
[setError, updateStats],
);
const recordDecision = useCallback(async (decision: string, rationale: string) => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
try {
await workflowManagerRef.current.recordDecision(decision, rationale);
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : '记录决策失败');
throw error;
}
}, [setError, updateStats]);
const postAction = useCallback(
async (
context: ActionContext,
result: string,
error?: string,
): Promise<string> => {
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
setLoading(true);
try {
const message = await workflowManagerRef.current.postAction(
context,
result,
error,
);
await updateStats();
return message;
} catch (err) {
const errorMessage =
err instanceof Error ? err.message : "Post-Action 执行失败";
setError(errorMessage);
throw err;
} finally {
setLoading(false);
}
},
[setLoading, setError, updateStats],
);
const updatePhaseStatus = useCallback(
async (
phaseNumber: number,
status: "pending" | "in_progress" | "complete",
notes?: string,
) => {
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
try {
await workflowManagerRef.current.updatePhaseStatus(
phaseNumber,
status,
notes,
);
setState((prev) => ({ ...prev, currentPhase: phaseNumber }));
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : "更新阶段状态失败");
throw error;
}
},
[setError, updateStats],
);
const recordFinding = useCallback(
async (title: string, content: string, tags: string[] = []) => {
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
try {
await workflowManagerRef.current.recordFinding(title, content, tags);
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : "记录发现失败");
throw error;
}
},
[setError, updateStats],
);
const recordDecision = useCallback(
async (decision: string, rationale: string) => {
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
try {
await workflowManagerRef.current.recordDecision(decision, rationale);
await updateStats();
} catch (error) {
setError(error instanceof Error ? error.message : "记录决策失败");
throw error;
}
},
[setError, updateStats],
);
const checkCompletion = useCallback(async () => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
try {
const result = await workflowManagerRef.current.checkCompletion();
await updateStats();
return result;
} catch (error) {
setError(error instanceof Error ? error.message : '检查完成状态失败');
setError(error instanceof Error ? error.message : "检查完成状态失败");
throw error;
}
}, [setError, updateStats]);
const finalizeWorkflow = useCallback(async (): Promise<string> => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
setLoading(true);
try {
const result = await workflowManagerRef.current.finalizeWorkflow();
setState(prev => ({ ...prev, isInitialized: false }));
setState((prev) => ({ ...prev, isInitialized: false }));
return result;
} catch (error) {
setError(error instanceof Error ? error.message : '结束工作流失败');
setError(error instanceof Error ? error.message : "结束工作流失败");
throw error;
} finally {
setLoading(false);
@@ -261,12 +321,12 @@ export function useThreeStageWorkflow(options: UseThreeStageWorkflowOptions): [W
}, [setLoading, setError]);
const getSessionStats = useCallback(async () => {
if (!workflowManagerRef.current) throw new Error('工作流未初始化');
if (!workflowManagerRef.current) throw new Error("工作流未初始化");
try {
return await workflowManagerRef.current.getSessionStats();
} catch (error) {
setError(error instanceof Error ? error.message : '获取会话统计失败');
setError(error instanceof Error ? error.message : "获取会话统计失败");
throw error;
}
}, [setError]);
@@ -280,7 +340,9 @@ export function useThreeStageWorkflow(options: UseThreeStageWorkflowOptions): [W
visualOperationCount: 0,
errorAttempts: {},
});
workflowManagerRef.current = sessionId ? new ThreeStageWorkflowManager(sessionId) : null;
workflowManagerRef.current = sessionId
? new ThreeStageWorkflowManager(sessionId)
: null;
}, [sessionId]);
const actions: WorkflowActions = {
@@ -300,4 +362,4 @@ export function useThreeStageWorkflow(options: UseThreeStageWorkflowOptions): [W
return [state, actions];
}
export default useThreeStageWorkflow;
export default useThreeStageWorkflow;
+29 -24
View File
@@ -22,19 +22,19 @@
--radius: 0.5rem;
/* Claude 风格颜色 */
--surface: #FFFFFF;
--surface-secondary: #F5F4F1;
--surface-tertiary: #EFEEE9;
--surface-cream: #FAF9F6;
--ink-900: #1A1915;
--ink-700: #4A4A45;
--surface: #ffffff;
--surface-secondary: #f5f4f1;
--surface-tertiary: #efeee9;
--surface-cream: #faf9f6;
--ink-900: #1a1915;
--ink-700: #4a4a45;
--ink-600: #666661;
--ink-400: #9B9B96;
--claude-accent: #D97757;
--claude-accent-hover: #CC785C;
--claude-accent-subtle: #FDF4F1;
--success: #16A34A;
--info: #2563EB;
--ink-400: #9b9b96;
--claude-accent: #d97757;
--claude-accent-hover: #cc785c;
--claude-accent-subtle: #fdf4f1;
--success: #16a34a;
--info: #2563eb;
}
.dark {
@@ -55,17 +55,17 @@
--border: 217.2 32.6% 17.5%;
/* Claude 风格颜色 - 暗色模式 */
--surface: #1A1915;
--surface-secondary: #2D2D2A;
--surface-tertiary: #3A3A35;
--surface: #1a1915;
--surface-secondary: #2d2d2a;
--surface-tertiary: #3a3a35;
--surface-cream: #242420;
--ink-900: #F5F4F1;
--ink-700: #D1D1CC;
--ink-600: #9B9B96;
--ink-900: #f5f4f1;
--ink-700: #d1d1cc;
--ink-600: #9b9b96;
--ink-400: #666661;
--claude-accent: #E8956F;
--claude-accent-hover: #D97757;
--claude-accent-subtle: #3D2A24;
--claude-accent: #e8956f;
--claude-accent-hover: #d97757;
--claude-accent-subtle: #3d2a24;
}
}
@@ -112,8 +112,12 @@
/* Claude 风格动画 */
@keyframes shimmer {
0% { transform: translateX(-120%); }
100% { transform: translateX(120%); }
0% {
transform: translateX(-120%);
}
100% {
transform: translateX(120%);
}
}
.animate-shimmer {
@@ -166,7 +170,8 @@
}
@keyframes ping {
75%, 100% {
75%,
100% {
transform: scale(2);
opacity: 0;
}
+19 -11
View File
@@ -213,20 +213,28 @@ export function parseStreamEvent(data: unknown): StreamEvent | null {
return {
type: "action_required",
request_id: (event.request_id as string) || "",
action_type: (event.action_type as "tool_confirmation" | "ask_user" | "elicitation") || "tool_confirmation",
action_type:
(event.action_type as
| "tool_confirmation"
| "ask_user"
| "elicitation") || "tool_confirmation",
tool_name: event.tool_name as string | undefined,
arguments: event.arguments as Record<string, unknown> | undefined,
prompt: event.prompt as string | undefined,
questions: event.questions as Array<{
question: string;
header?: string;
options?: Array<{
label: string;
description?: string;
}>;
multiSelect?: boolean;
}> | undefined,
requested_schema: event.requested_schema as Record<string, unknown> | undefined,
questions: event.questions as
| Array<{
question: string;
header?: string;
options?: Array<{
label: string;
description?: string;
}>;
multiSelect?: boolean;
}>
| undefined,
requested_schema: event.requested_schema as
| Record<string, unknown>
| undefined,
};
case "done":
return {
+31 -27
View File
@@ -1,10 +1,10 @@
/**
* 上下文记忆管理 API
*
*
* 基于文件系统的持久化记忆系统,解决 AI Agent 的上下文丢失、目标漂移、错误重复问题
*/
import { invoke } from '@tauri-apps/api/core';
import { invoke } from "@tauri-apps/api/core";
export interface MemoryEntry {
id: string;
@@ -19,7 +19,11 @@ export interface MemoryEntry {
archived: boolean;
}
export type MemoryFileType = 'task_plan' | 'findings' | 'progress' | 'error_log';
export type MemoryFileType =
| "task_plan"
| "findings"
| "progress"
| "error_log";
export interface MemoryStats {
session_id: string;
@@ -60,7 +64,7 @@ export class ContextMemoryAPI {
* 保存记忆条目
*/
static async saveMemoryEntry(request: SaveMemoryRequest): Promise<void> {
return invoke('save_memory_entry', { request });
return invoke("save_memory_entry", { request });
}
/**
@@ -68,9 +72,9 @@ export class ContextMemoryAPI {
*/
static async getSessionMemories(
sessionId: string,
fileType?: MemoryFileType
fileType?: MemoryFileType,
): Promise<MemoryEntry[]> {
return invoke('get_session_memories', {
return invoke("get_session_memories", {
sessionId,
fileType: fileType || null,
});
@@ -80,14 +84,14 @@ export class ContextMemoryAPI {
* 获取记忆上下文(用于 AI 上下文)
*/
static async getMemoryContext(sessionId: string): Promise<string> {
return invoke('get_memory_context', { sessionId });
return invoke("get_memory_context", { sessionId });
}
/**
* 记录错误
*/
static async recordError(request: RecordErrorRequest): Promise<void> {
return invoke('record_error', { request });
return invoke("record_error", { request });
}
/**
@@ -95,9 +99,9 @@ export class ContextMemoryAPI {
*/
static async shouldAvoidOperation(
sessionId: string,
operationDescription: string
operationDescription: string,
): Promise<boolean> {
return invoke('should_avoid_operation', {
return invoke("should_avoid_operation", {
sessionId,
operationDescription,
});
@@ -107,21 +111,21 @@ export class ContextMemoryAPI {
* 标记错误已解决
*/
static async markErrorResolved(request: ResolveErrorRequest): Promise<void> {
return invoke('mark_error_resolved', { request });
return invoke("mark_error_resolved", { request });
}
/**
* 获取记忆统计信息
*/
static async getMemoryStats(sessionId: string): Promise<MemoryStats> {
return invoke('get_memory_stats', { sessionId });
return invoke("get_memory_stats", { sessionId });
}
/**
* 清理过期记忆
*/
static async cleanupExpiredMemories(): Promise<void> {
return invoke('cleanup_expired_memories');
return invoke("cleanup_expired_memories");
}
/**
@@ -131,14 +135,14 @@ export class ContextMemoryAPI {
sessionId: string,
title: string,
content: string,
priority: number = 3
priority: number = 3,
): Promise<void> {
return this.saveMemoryEntry({
session_id: sessionId,
file_type: 'task_plan',
file_type: "task_plan",
title,
content,
tags: ['任务计划'],
tags: ["任务计划"],
priority,
});
}
@@ -151,14 +155,14 @@ export class ContextMemoryAPI {
title: string,
content: string,
tags: string[] = [],
priority: number = 4
priority: number = 4,
): Promise<void> {
return this.saveMemoryEntry({
session_id: sessionId,
file_type: 'findings',
file_type: "findings",
title,
content,
tags: ['发现', ...tags],
tags: ["发现", ...tags],
priority,
});
}
@@ -169,14 +173,14 @@ export class ContextMemoryAPI {
static async logProgress(
sessionId: string,
title: string,
content: string
content: string,
): Promise<void> {
return this.saveMemoryEntry({
session_id: sessionId,
file_type: 'progress',
file_type: "progress",
title,
content,
tags: ['进度'],
tags: ["进度"],
priority: 2,
});
}
@@ -186,15 +190,15 @@ export class ContextMemoryAPI {
*/
static async apply2ActionRule(
sessionId: string,
finding: string
finding: string,
): Promise<void> {
const timestamp = new Date().toLocaleTimeString();
return this.saveFinding(
sessionId,
`2-Action 规则发现 (${timestamp})`,
finding,
['2-Action规则', '自动保存'],
4
["2-Action规则", "自动保存"],
4,
);
}
@@ -205,7 +209,7 @@ export class ContextMemoryAPI {
sessionId: string,
errorDescription: string,
attemptedSolution: string,
operationDescription?: string
operationDescription?: string,
): Promise<{ shouldAvoid: boolean }> {
// 记录错误
await this.recordError({
@@ -223,4 +227,4 @@ export class ContextMemoryAPI {
}
}
export default ContextMemoryAPI;
export default ContextMemoryAPI;
+68 -59
View File
@@ -1,12 +1,16 @@
/**
* 工具钩子管理 API
*
*
* 提供工具执行前后的钩子机制,用于自动化上下文记忆管理
*/
import { invoke } from '@tauri-apps/api/core';
import { invoke } from "@tauri-apps/api/core";
export type HookTrigger = 'session_start' | 'pre_tool_use' | 'post_tool_use' | 'stop';
export type HookTrigger =
| "session_start"
| "pre_tool_use"
| "post_tool_use"
| "stop";
export interface HookRule {
id: string;
@@ -20,7 +24,7 @@ export interface HookRule {
created_at: number;
}
export type HookCondition =
export type HookCondition =
| { tool_name_equals: string }
| { tool_name_contains: string }
| { message_contains: string }
@@ -28,8 +32,15 @@ export type HookCondition =
| { error_count_greater_than: number }
| { custom: { condition_type: string; parameters: Record<string, string> } };
export type HookAction =
| { save_finding: { title: string; content: string; tags: string[]; priority: number } }
export type HookAction =
| {
save_finding: {
title: string;
content: string;
tags: string[];
priority: number;
};
}
| { update_task_plan: { title: string; content: string; priority: number } }
| { log_progress: { title: string; content: string } }
| { record_error: { error_description: string; attempted_solution: string } }
@@ -67,57 +78,62 @@ export class ToolHooksAPI {
* 执行钩子
*/
static async executeHooks(request: ExecuteHooksRequest): Promise<void> {
return invoke('execute_hooks', { request });
return invoke("execute_hooks", { request });
}
/**
* 添加钩子规则
*/
static async addHookRule(rule: HookRule): Promise<void> {
return invoke('add_hook_rule', { rule });
return invoke("add_hook_rule", { rule });
}
/**
* 移除钩子规则
*/
static async removeHookRule(ruleId: string): Promise<void> {
return invoke('remove_hook_rule', { ruleId });
return invoke("remove_hook_rule", { ruleId });
}
/**
* 启用/禁用钩子规则
*/
static async toggleHookRule(ruleId: string, enabled: boolean): Promise<void> {
return invoke('toggle_hook_rule', { ruleId, enabled });
return invoke("toggle_hook_rule", { ruleId, enabled });
}
/**
* 获取所有钩子规则
*/
static async getHookRules(): Promise<HookRule[]> {
return invoke('get_hook_rules');
return invoke("get_hook_rules");
}
/**
* 获取钩子执行统计
*/
static async getHookExecutionStats(): Promise<Record<string, HookExecutionStats>> {
return invoke('get_hook_execution_stats');
static async getHookExecutionStats(): Promise<
Record<string, HookExecutionStats>
> {
return invoke("get_hook_execution_stats");
}
/**
* 清理钩子执行统计
*/
static async clearHookExecutionStats(): Promise<void> {
return invoke('clear_hook_execution_stats');
return invoke("clear_hook_execution_stats");
}
/**
* 触发会话开始钩子
*/
static async triggerSessionStart(sessionId: string, metadata: Record<string, string> = {}): Promise<void> {
static async triggerSessionStart(
sessionId: string,
metadata: Record<string, string> = {},
): Promise<void> {
return this.executeHooks({
trigger: 'session_start',
trigger: "session_start",
context: {
session_id: sessionId,
message_count: 0,
@@ -137,10 +153,10 @@ export class ToolHooksAPI {
toolName: string,
toolParameters: Record<string, string> = {},
messageContent?: string,
messageCount: number = 0
messageCount: number = 0,
): Promise<void> {
return this.executeHooks({
trigger: 'pre_tool_use',
trigger: "pre_tool_use",
context: {
session_id: sessionId,
tool_name: toolName,
@@ -165,10 +181,10 @@ export class ToolHooksAPI {
toolParameters: Record<string, string> = {},
messageContent?: string,
messageCount: number = 0,
errorInfo?: string
errorInfo?: string,
): Promise<void> {
return this.executeHooks({
trigger: 'post_tool_use',
trigger: "post_tool_use",
context: {
session_id: sessionId,
tool_name: toolName,
@@ -180,7 +196,7 @@ export class ToolHooksAPI {
metadata: {
timestamp: new Date().toISOString(),
tool_name: toolName,
has_error: errorInfo ? 'true' : 'false',
has_error: errorInfo ? "true" : "false",
},
},
});
@@ -192,16 +208,16 @@ export class ToolHooksAPI {
static async triggerStop(
sessionId: string,
messageCount: number,
metadata: Record<string, string> = {}
metadata: Record<string, string> = {},
): Promise<void> {
return this.executeHooks({
trigger: 'stop',
trigger: "stop",
context: {
session_id: sessionId,
message_count: messageCount,
metadata: {
timestamp: new Date().toISOString(),
session_end: 'true',
session_end: "true",
...metadata,
},
},
@@ -218,7 +234,7 @@ export class ToolHooksAPI {
trigger: HookTrigger,
conditions: HookCondition[],
actions: HookAction[],
priority: number = 100
priority: number = 100,
): HookRule {
return {
id,
@@ -238,25 +254,22 @@ export class ToolHooksAPI {
*/
static createImportantFindingRule(): HookRule {
return this.createCustomRule(
'important-finding-auto-save',
'重要发现自动保存',
'检测到重要信息时自动保存到 findings.md',
'post_tool_use',
[
{ message_contains: '重要' },
{ message_contains: '发现' },
],
"important-finding-auto-save",
"重要发现自动保存",
"检测到重要信息时自动保存到 findings.md",
"post_tool_use",
[{ message_contains: "重要" }, { message_contains: "发现" }],
[
{
save_finding: {
title: '重要发现 (自动检测)',
content: '检测到重要信息,已自动保存',
tags: ['重要', '自动保存'],
title: "重要发现 (自动检测)",
content: "检测到重要信息,已自动保存",
tags: ["重要", "自动保存"],
priority: 4,
},
},
],
1
1,
);
}
@@ -265,22 +278,20 @@ export class ToolHooksAPI {
*/
static createErrorAutoRecordRule(): HookRule {
return this.createCustomRule(
'error-auto-record',
'错误自动记录',
'检测到错误时自动记录到错误日志',
'post_tool_use',
[
{ message_contains: '错误' },
],
"error-auto-record",
"错误自动记录",
"检测到错误时自动记录到错误日志",
"post_tool_use",
[{ message_contains: "错误" }],
[
{
record_error: {
error_description: '检测到错误',
attempted_solution: '正在尝试解决',
error_description: "检测到错误",
attempted_solution: "正在尝试解决",
},
},
],
1
1,
);
}
@@ -289,24 +300,22 @@ export class ToolHooksAPI {
*/
static create2ActionRule(): HookRule {
return this.createCustomRule(
'2-action-rule',
'2-Action 规则',
'每2次视觉操作后自动保存发现',
'post_tool_use',
[
{ tool_name_contains: 'view' },
],
"2-action-rule",
"2-Action 规则",
"每2次视觉操作后自动保存发现",
"post_tool_use",
[{ tool_name_contains: "view" }],
[
{
save_finding: {
title: '2-Action 规则触发',
content: '视觉操作完成,自动保存发现',
tags: ['2-Action规则', '视觉操作'],
title: "2-Action 规则触发",
content: "视觉操作完成,自动保存发现",
tags: ["2-Action规则", "视觉操作"],
priority: 3,
},
},
],
2
2,
);
}
@@ -330,4 +339,4 @@ export class ToolHooksAPI {
}
}
export default ToolHooksAPI;
export default ToolHooksAPI;
-2
View File
@@ -43,8 +43,6 @@ const PLUGIN_GLOBAL_NAMES: Record<string, string> = {
"gemini-provider": "GeminiProviderPlugin",
"antigravity-provider": "AntigravityProviderPlugin",
"codex-provider": "CodexProviderPlugin",
"machine-id-tool": "MachineIdToolPlugin",
"browser-interception": "BrowserInterceptionPlugin",
"flow-monitor": "FlowMonitorPlugin",
"config-switch": "ConfigSwitchPlugin",
};
+79 -54
View File
@@ -1,19 +1,19 @@
/**
* 三阶段工作流管理器
*
*
* 基于 planning-with-files 的核心机制,实现:
* - Pre-Action → Action → Post-Action 三阶段工作流
* - 自动化上下文工程和错误学习
* - 2-Action 规则和 3次错误协议
*/
import { ContextMemoryAPI } from '../api/contextMemory';
import { ToolHooksAPI } from '../api/toolHooks';
import { ContextMemoryAPI } from "../api/contextMemory";
import { ToolHooksAPI } from "../api/toolHooks";
export interface WorkflowPhase {
number: number;
name: string;
status: 'pending' | 'in_progress' | 'complete';
status: "pending" | "in_progress" | "complete";
tasks: string[];
notes?: string;
}
@@ -62,23 +62,23 @@ export class ThreeStageWorkflowManager {
this.sessionId,
`任务计划: ${config.projectName}`,
taskPlanContent,
5
5,
);
// 创建初始发现记录
await ContextMemoryAPI.saveFinding(
this.sessionId,
'工作流初始化',
"工作流初始化",
`三阶段工作流已初始化\n项目: ${config.projectName}\n目标: ${config.goal}`,
['初始化', '工作流'],
3
["初始化", "工作流"],
3,
);
// 记录初始进度
await ContextMemoryAPI.logProgress(
this.sessionId,
'工作流启动',
`三阶段工作流已启动,共 ${config.phases.length} 个阶段`
"工作流启动",
`三阶段工作流已启动,共 ${config.phases.length} 个阶段`,
);
}
@@ -92,25 +92,27 @@ export class ThreeStageWorkflowManager {
context.toolName || context.actionType,
context.toolParameters || {},
context.actionDescription,
context.messageCount
context.messageCount,
);
// 获取当前记忆上下文
const memoryContext = await ContextMemoryAPI.getMemoryContext(context.sessionId);
const memoryContext = await ContextMemoryAPI.getMemoryContext(
context.sessionId,
);
// 检查是否应该避免该操作(3次错误协议)
const shouldAvoid = await ContextMemoryAPI.shouldAvoidOperation(
context.sessionId,
context.actionDescription
context.actionDescription,
);
if (shouldAvoid) {
const warning = `⚠️ 3次错误协议警告: 该操作已失败3次,建议更换方法\n操作: ${context.actionDescription}`;
await ContextMemoryAPI.recordError({
session_id: context.sessionId,
error_description: `重复失败操作: ${context.actionDescription}`,
attempted_solution: '触发3次错误协议,建议更换方法',
attempted_solution: "触发3次错误协议,建议更换方法",
});
return `${warning}\n\n当前上下文:\n${memoryContext}`;
@@ -119,8 +121,8 @@ export class ThreeStageWorkflowManager {
// 记录上下文刷新
await ContextMemoryAPI.logProgress(
context.sessionId,
'Pre-Action 上下文刷新',
`准备执行: ${context.actionDescription}`
"Pre-Action 上下文刷新",
`准备执行: ${context.actionDescription}`,
);
return `🔄 Pre-Action 上下文刷新完成\n\n准备执行: ${context.actionDescription}\n\n当前记忆上下文:\n${memoryContext}`;
@@ -129,12 +131,15 @@ export class ThreeStageWorkflowManager {
/**
* Action 阶段:执行实际操作
*/
async executeAction(context: ActionContext, actionResult: string): Promise<void> {
async executeAction(
context: ActionContext,
actionResult: string,
): Promise<void> {
// 记录操作执行
await ContextMemoryAPI.logProgress(
context.sessionId,
`执行操作: ${context.actionType}`,
`操作描述: ${context.actionDescription}\n结果: ${actionResult.substring(0, 200)}${actionResult.length > 200 ? '...' : ''}`
`操作描述: ${context.actionDescription}\n结果: ${actionResult.substring(0, 200)}${actionResult.length > 200 ? "..." : ""}`,
);
// 如果是视觉操作,增加计数
@@ -146,8 +151,12 @@ export class ThreeStageWorkflowManager {
/**
* Post-Action 阶段:操作后的状态更新
*/
async postAction(context: ActionContext, actionResult: string, error?: string): Promise<string> {
let message = '📝 Post-Action 状态更新:\n\n';
async postAction(
context: ActionContext,
actionResult: string,
error?: string,
): Promise<string> {
let message = "📝 Post-Action 状态更新:\n\n";
// 处理错误情况
if (error) {
@@ -159,11 +168,11 @@ export class ThreeStageWorkflowManager {
context.sessionId,
error,
`尝试次数: ${attemptCount}`,
context.actionDescription
context.actionDescription,
);
message += `🚨 错误记录 (第${attemptCount}次尝试): ${error}\n`;
if (shouldAvoid) {
message += `⚠️ 已达到3次错误限制,建议更换方法\n`;
}
@@ -177,7 +186,7 @@ export class ThreeStageWorkflowManager {
context.toolParameters || {},
context.actionDescription,
context.messageCount,
error
error,
);
// 应用 2-Action 规则
@@ -201,8 +210,8 @@ export class ThreeStageWorkflowManager {
*/
private async apply2ActionRule(actionResult: string): Promise<void> {
const timestamp = new Date().toLocaleTimeString();
const finding = `2-Action 规则触发 (${timestamp})\n\n最近操作结果:\n${actionResult.substring(0, 500)}${actionResult.length > 500 ? '...' : ''}`;
const finding = `2-Action 规则触发 (${timestamp})\n\n最近操作结果:\n${actionResult.substring(0, 500)}${actionResult.length > 500 ? "..." : ""}`;
await ContextMemoryAPI.apply2ActionRule(this.sessionId, finding);
}
@@ -211,39 +220,43 @@ export class ThreeStageWorkflowManager {
*/
async updatePhaseStatus(
phaseNumber: number,
status: 'pending' | 'in_progress' | 'complete',
notes?: string
status: "pending" | "in_progress" | "complete",
notes?: string,
): Promise<void> {
const statusText = {
pending: '待开始',
in_progress: '进行中',
complete: '已完成',
pending: "待开始",
in_progress: "进行中",
complete: "已完成",
}[status];
await ContextMemoryAPI.saveTaskPlan(
this.sessionId,
`阶段 ${phaseNumber} 状态更新`,
`阶段 ${phaseNumber} 状态已更新为: ${statusText}${notes ? `\n备注: ${notes}` : ''}`,
4
`阶段 ${phaseNumber} 状态已更新为: ${statusText}${notes ? `\n备注: ${notes}` : ""}`,
4,
);
await ContextMemoryAPI.logProgress(
this.sessionId,
`阶段 ${phaseNumber} 状态更新`,
`状态: ${statusText}${notes ? `\n备注: ${notes}` : ''}`
`状态: ${statusText}${notes ? `\n备注: ${notes}` : ""}`,
);
}
/**
* 记录重要发现
*/
async recordFinding(title: string, content: string, tags: string[] = []): Promise<void> {
async recordFinding(
title: string,
content: string,
tags: string[] = [],
): Promise<void> {
await ContextMemoryAPI.saveFinding(
this.sessionId,
title,
content,
['发现', ...tags],
4
["发现", ...tags],
4,
);
}
@@ -255,8 +268,8 @@ export class ThreeStageWorkflowManager {
this.sessionId,
`决策: ${decision}`,
`决策内容: ${decision}\n\n决策理由:\n${rationale}`,
['决策', '重要'],
5
["决策", "重要"],
5,
);
}
@@ -266,19 +279,22 @@ export class ThreeStageWorkflowManager {
async checkCompletion(): Promise<{ isComplete: boolean; summary: string }> {
const stats = await ContextMemoryAPI.getMemoryStats(this.sessionId);
const memories = await ContextMemoryAPI.getSessionMemories(this.sessionId);
// 简单的完成度检查逻辑
const taskPlanMemories = memories.filter(m => m.file_type === 'task_plan');
const hasCompletedPhases = taskPlanMemories.some(m =>
m.content.includes('已完成') || m.content.includes('complete')
const taskPlanMemories = memories.filter(
(m) => m.file_type === "task_plan",
);
const hasCompletedPhases = taskPlanMemories.some(
(m) => m.content.includes("已完成") || m.content.includes("complete"),
);
const summary = `📊 任务完成状态检查:\n\n` +
const summary =
`📊 任务完成状态检查:\n\n` +
`- 活跃记忆: ${stats.active_memories} 个\n` +
`- 未解决错误: ${stats.unresolved_errors} 个\n` +
`- 已解决错误: ${stats.resolved_errors} 个\n` +
`- 是否有已完成阶段: ${hasCompletedPhases ? '是' : '否'}\n\n` +
`${stats.unresolved_errors > 0 ? '⚠️ 仍有未解决的错误需要处理' : '✅ 无未解决错误'}`;
`- 是否有已完成阶段: ${hasCompletedPhases ? "是" : "否"}\n\n` +
`${stats.unresolved_errors > 0 ? "⚠️ 仍有未解决的错误需要处理" : "✅ 无未解决错误"}`;
return {
isComplete: hasCompletedPhases && stats.unresolved_errors === 0,
@@ -300,10 +316,10 @@ export class ThreeStageWorkflowManager {
// 保存会话摘要
await ContextMemoryAPI.saveFinding(
this.sessionId,
'工作流会话摘要',
"工作流会话摘要",
`三阶段工作流已结束\n\n${summary}`,
['摘要', '会话结束'],
5
["摘要", "会话结束"],
5,
);
return `🎉 三阶段工作流已结束\n\n${summary}`;
@@ -320,7 +336,7 @@ export class ThreeStageWorkflowManager {
config.phases.forEach((phase) => {
content += `### 阶段 ${phase.number}: ${phase.name}\n`;
phase.tasks.forEach(task => {
phase.tasks.forEach((task) => {
content += `- [ ] ${task}\n`;
});
content += `- **状态**: ${phase.status}\n\n`;
@@ -352,8 +368,17 @@ export class ThreeStageWorkflowManager {
* 判断是否为视觉操作
*/
private isVisualOperation(actionType: string): boolean {
const visualActions = ['view', 'read', 'browse', 'search', 'screenshot', 'image'];
return visualActions.some(action => actionType.toLowerCase().includes(action));
const visualActions = [
"view",
"read",
"browse",
"search",
"screenshot",
"image",
];
return visualActions.some((action) =>
actionType.toLowerCase().includes(action),
);
}
/**
@@ -365,7 +390,7 @@ export class ThreeStageWorkflowManager {
errorAttempts: Record<string, number>;
}> {
const memoryStats = await ContextMemoryAPI.getMemoryStats(this.sessionId);
return {
memoryStats,
visualOperationCount: this.visualOperationCount,
@@ -374,4 +399,4 @@ export class ThreeStageWorkflowManager {
}
}
export default ThreeStageWorkflowManager;
export default ThreeStageWorkflowManager;
+1
View File
@@ -19,4 +19,5 @@ export type Page =
| "sysinfo"
| "files"
| "web"
| "image-analysis"
| `plugin:${string}`;